Instructions to use BAAI/Brainmu-SpikeCamera with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use BAAI/Brainmu-SpikeCamera with Transformers:
# Load model directly from transformers import SpikeConvFrontend model = SpikeConvFrontend.from_pretrained("BAAI/Brainmu-SpikeCamera", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """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=[] | |