Brainmu-SpikeCamera / src /ui /multi_gpu.py
sunbaby's picture
Upload 69 files
4719196
Raw
History Blame Contribute Delete
3.66 kB
"""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=[]