CoolFace
Apppublic

Oopstom/ReactSeq

sourceHugging Facelgpl-2.1updated 1y agoView on Hugging Face
6likes
inference_engine.py119 linesDownload Raw Back to onmt
1import torch2from onmt.utils.distributed import ErrorHandler, spawned_infer3from onmt.translate.translator import build_translator4from onmt.transforms import get_transforms_cls5from onmt.constants import CorpusTask6from onmt.utils.logging import logger7from onmt.inputters.dynamic_iterator import build_dynamic_dataset_iter8from onmt.inputters.inputter import IterOnDevice9 10 11class InferenceEngine(object):12    """Wrapper Class to run Inference in mulitpocessing with partitioned models.13 14    Args:15        opt: inference options16    """17 18    def __init__(self, opt):19        self.opt = opt20 21        if opt.world_size > 1:22 23            mp = torch.multiprocessing.get_context("spawn")24            # Create a thread to listen for errors in the child processes.25            self.error_queue = mp.SimpleQueue()26            self.error_handler = ErrorHandler(self.error_queue)27            self.queue_instruct = []28            self.queue_result = []29            self.procs = []30 31            print("world_size: ", opt.world_size)32            print("gpu_ranks: ", opt.gpu_ranks)33            print("opt.gpu: ", opt.gpu)34 35            for device_id in range(opt.world_size):36                self.queue_instruct.append(mp.Queue())37                self.queue_result.append(mp.Queue())38                self.procs.append(39                    mp.Process(40                        target=spawned_infer,41                        args=(42                            opt,43                            device_id,44                            self.error_queue,45                            self.queue_instruct[device_id],46                            self.queue_result[device_id],47                        ),48                        daemon=False,49                    )50                )51                self.procs[device_id].start()52                print(" Starting process pid: %d  " % self.procs[device_id].pid)53                self.error_handler.add_child(self.procs[device_id].pid)54        else:55            self.device_id = 056            self.translator = build_translator(57                opt, self.device_id, logger=logger, report_score=True58            )59            self.transforms_cls = get_transforms_cls(opt._all_transform)60 61    def infer_file(self):62        """File inference. Source file must be the opt.src argument"""63        if self.opt.world_size > 1:64            for device_id in range(self.opt.world_size):65                self.queue_instruct[device_id].put(("infer_file", self.opt))66            scores, preds = [], []67            for device_id in range(self.opt.world_size):68                scores.append(self.queue_result[device_id].get())69                preds.append(self.queue_result[device_id].get())70            return scores[0], preds[0]71        else:72            infer_iter = build_dynamic_dataset_iter(73                self.opt,74                self.transforms_cls,75                self.translator.vocabs,76                task=CorpusTask.INFER,77            )78            infer_iter = IterOnDevice(infer_iter, self.device_id)79            scores, preds = self.translator._translate(80                infer_iter,81                infer_iter.transform,82                self.opt.attn_debug,83                self.opt.align_debug,84            )85            return scores, preds86 87    def infer_list(self, src):88        """List of strings inference `src`"""89        if self.opt.world_size > 1:90            for device_id in range(self.opt.world_size):91                self.queue_instruct[device_id].put(("infer_list", src))92            scores, preds = [], []93            for device_id in range(self.opt.world_size):94                scores.append(self.queue_result[device_id].get())95                preds.append(self.queue_result[device_id].get())96            return scores[0], preds[0]97        else:98            infer_iter = build_dynamic_dataset_iter(99                self.opt,100                self.transforms_cls,101                self.translator.vocabs,102                task=CorpusTask.INFER,103                src=src,104            )105            infer_iter = IterOnDevice(infer_iter, self.device_id)106            scores, preds = self.translator._translate(107                infer_iter,108                infer_iter.transform,109                self.opt.attn_debug,110                self.opt.align_debug,111            )112            return scores, preds113 114    def terminate(self):115        if self.opt.world_size > 1:116            for device_id in range(self.opt.world_size):117                self.queue_instruct[device_id].put(("stop"))118                self.procs[device_id].terminate()119