Oopstom/ReactSeq
6
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 