"""Persistent process per GPU; original global sample indices and seeds are preserved.""" import multiprocessing as mp from multiprocessing.connection import wait import os,time,traceback,uuid from pathlib import Path def shard_indices(total,workers): return [list(range(rank,total,workers)) for rank in range(workers)] def worker(conn,gpu,config,log_path): os.environ['CUDA_VISIBLE_DEVICES']=str(gpu) os.environ['OMP_NUM_THREADS']='2' with open(log_path,'a',buffering=1) as log: os.dup2(log.fileno(),1);os.dup2(log.fileno(),2) try: from engine import Engine engine=Engine();engine.load(**config) conn.send({'type':'ready','gpu':gpu}) while True: item=conn.recv() if item['type']=='close':break dat,gt,index,out,prompt=item['args'] row=engine.predict(Path(dat),Path(gt),index,Path(out),prompt) row['gpu_id']=gpu conn.send({'type':'result','row':row}) except EOFError:pass except BaseException as e: traceback.print_exc() try:conn.send({'type':'error','error':f'GPU {gpu}: {e}'}) except Exception:pass finally:conn.close() class MultiGPUEngine: def __init__(self,gpu_ids): if not gpu_ids or len(set(gpu_ids))!=len(gpu_ids):raise ValueError('GPU IDs must be nonempty and unique') self.gpu_ids=list(gpu_ids);self.procs=[];self.conns=[] def load(self,**config): ctx=mp.get_context('spawn') logs=Path(__file__).resolve().parents[1]/'logs' logs.mkdir(parents=True,exist_ok=True) tag=uuid.uuid4().hex[:8] try: for gpu in self.gpu_ids: parent,child=ctx.Pipe() p=ctx.Process(target=worker,args=(child,gpu,config,str(logs/f'ui_gpu{gpu}_{tag}.log')),daemon=True) p.start();child.close();self.procs.append(p);self.conns.append(parent) pending=set(self.conns);deadline=time.monotonic()+1200 while pending: if time.monotonic()>deadline:raise TimeoutError('GPU model loading exceeded 20 minutes') for c in wait(pending,timeout=1): msg=c.recv() if msg.get('type')!='ready':raise RuntimeError(msg.get('error','Invalid model load response')) pending.remove(c) self._check() except BaseException:self.close();raise def _check(self): for gpu,p in zip(self.gpu_ids,self.procs): if not p.is_alive():raise RuntimeError(f'GPU worker {gpu} exited ({p.exitcode})') def predict_many(self,pairs,out,prompt,stop_event): if len(self.conns)!=len(self.gpu_ids):raise RuntimeError('GPU pool is not loaded') queues=[iter(x) for x in shard_indices(len(pairs),len(self.gpu_ids))] active={};seen=set() def dispatch(rank): if stop_event.is_set():return i=next(queues[rank],None) if i is None:return dat,gt=pairs[i];c=self.conns[rank] c.send({'type':'predict','args':(str(dat),str(gt),i,str(out),prompt)}) active[c]=(rank,i) try: for rank in range(len(self.conns)):dispatch(rank) while active: self._check() for c in wait(list(active),timeout=1): rank,expected=active.pop(c);msg=c.recv() if msg.get('type')!='result':raise RuntimeError(msg.get('error','Invalid worker response')) row=msg['row'] if row['index']!=expected or expected in seen:raise RuntimeError('Duplicate or mismatched sample index') seen.add(expected) yield row dispatch(rank) if not stop_event.is_set() and seen!=set(range(len(pairs))):raise RuntimeError('Incomplete GPU result coverage') except BaseException:self.close();raise def close(self): for c in self.conns: try:c.send({'type':'close'}) except Exception:pass for p in self.procs: p.join(timeout=1) if p.is_alive():p.terminate();p.join(timeout=5) for c in self.conns: try:c.close() except Exception:pass self.conns=[];self.procs=[]