File size: 3,663 Bytes
4719196
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
"""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=[]