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
Download src/ui/server.py from BAAI/Brainmu-SpikeCamera: direct link, hf CLI and curl.
- Browser
- Download file 9.49 kB
-
https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/ui/server.py
- Command line
-
hf download hf://BAAI/Brainmu-SpikeCamera/src/ui/server.py
-
curl -L -o server.py https://huggingface.co/BAAI/Brainmu-SpikeCamera/resolve/main/src/ui/server.py
9.49 kB
| #!/usr/bin/env python3 | |
| import csv,json,os,threading,time,traceback,uuid,sys | |
| from pathlib import Path | |
| from http.server import BaseHTTPRequestHandler,ThreadingHTTPServer | |
| from urllib.parse import urlparse,parse_qs | |
| ROOT=Path(__file__).resolve().parents[1]; PROJECT=ROOT | |
| sys.path.insert(0,str(ROOT/'code')) | |
| from project_config import load_config | |
| SETTINGS=load_config()['generation']['inference'] | |
| for key,sub in {'HF_HOME':'huggingface','TORCH_HOME':'torch','TRITON_CACHE_DIR':'triton','TORCHINDUCTOR_CACHE_DIR':'torchinductor','TMPDIR':'tmp'}.items(): | |
| os.environ[key]=str(ROOT/'.cache'/sub);Path(os.environ[key]).mkdir(parents=True,exist_ok=True) | |
| os.environ.update(HF_HUB_OFFLINE='1',TRANSFORMERS_OFFLINE='1',TOKENIZERS_PARALLELISM='false') | |
| DEFAULT_GPU_COUNT=8 | |
| DEFAULTS=dict(model_path=str(ROOT/'checkpoints/Brainmu'),adapter=str(PROJECT/'checkpoints/lora/adapter.safetensors'),frontend=str(PROJECT.parent/'model.safetensors'),spike_dir=str(ROOT/'data/test/spike'),gt_dir=str(ROOT/'data/test/gt')) | |
| lock=threading.RLock(); stop=threading.Event() | |
| state=dict(phase='idle',message='请先加载模型,再加载测试数据。',model_loaded=False,dataset_count=0,requested=0,rows=[],error='',run_dir='',started=None,defaults=DEFAULTS) | |
| engine=None; dataset=[] | |
| def snapshot(): | |
| with lock: | |
| d=dict(state);d['rows']=list(state['rows']) | |
| d['prompt_supported']=True | |
| d['inference_settings']=SETTINGS | |
| d['parallel_supported']=True | |
| d['default_gpu_count']=DEFAULT_GPU_COUNT | |
| d['completed']=len(d['rows']) | |
| d['mean_psnr']=sum(x['psnr_db'] for x in d['rows'])/len(d['rows']) if d['rows'] else None | |
| d['mean_ssim']=sum(x['ssim'] for x in d['rows'])/len(d['rows']) if d['rows'] else None | |
| d['full_1000_complete']=d.get('evaluation_complete',d['phase']=='complete') and d['completed']==d['requested']==1000 | |
| d['elapsed']=round(time.time()-d['started'],1) if d['started'] and d['phase']=='running' else state.get('elapsed',0) | |
| return d | |
| def persist(): | |
| d=snapshot();out=Path(d['run_dir']) | |
| tmp=out/'metrics.json.tmp';tmp.write_text(json.dumps(d,ensure_ascii=False,indent=2));tmp.replace(out/'metrics.json') | |
| fields=['id','index','seed','prompt','gpu_id','psnr_db','ssim','spike','gt','condition','prediction'] | |
| tmp=out/'metrics.csv.tmp' | |
| with tmp.open('w',newline='') as f: | |
| w=csv.DictWriter(f,fieldnames=fields);w.writeheader();w.writerows(sorted(d['rows'],key=lambda r:r['index'])) | |
| tmp.replace(out/'metrics.csv') | |
| def job(fn, persist_run=False): | |
| def wrapped(): | |
| try:fn() | |
| except Exception as e: | |
| traceback.print_exc() | |
| with lock:state.update(phase='error',error=str(e),message='操作失败,请查看错误信息。') | |
| if persist_run: | |
| with lock:state['model_loaded']=False | |
| if state['run_dir']:persist() | |
| threading.Thread(target=wrapped,daemon=True).start() | |
| def action(kind,data): | |
| global engine,dataset | |
| with lock: | |
| busy=state['phase'] in ['loading','validating','running'] | |
| if kind=='stop': | |
| stop.set();state['message']='将在当前样本完成后停止。';return | |
| if busy:raise ValueError('任务正在执行,请稍候。') | |
| if kind=='load': | |
| if state['model_loaded']:return | |
| config={k:str(Path(data.get(k,DEFAULTS[k])).expanduser().resolve()) for k in ['model_path','adapter','frontend']} | |
| for p in [Path(config['model_path'])/n for n in ['ema.safetensors','ae.safetensors','llm_config.json','vit_config.json']]+[Path(config['adapter']),Path(config['adapter']).with_name('adapter_config.json'),Path(config['frontend'])]: | |
| if not p.is_file():raise ValueError('文件不存在:'+str(p)) | |
| gpu_count=int(data.get('gpu_count',DEFAULT_GPU_COUNT)) | |
| if gpu_count not in [1,2,4,8]:raise ValueError('GPU 数量请选择 1、2、4 或 8。') | |
| state.update(phase='loading',gpu_count=gpu_count,gpu_ids=list(range(gpu_count)),message=f'正在向 {gpu_count} 张 GPU 加载 Brainmu 模型…',error='') | |
| def load(): | |
| global engine | |
| from multi_gpu import MultiGPUEngine | |
| candidate=MultiGPUEngine(list(range(gpu_count)));candidate.load(**config);engine=candidate | |
| with lock:state.update(phase='ready',model_loaded=True,model_config=config,message=f'{gpu_count} 卡模型已加载,可以选择测试集。') | |
| job(load) | |
| elif kind=='dataset': | |
| if not state['model_loaded']:raise ValueError('请先加载模型。') | |
| sd=Path(data.get('spike_dir',DEFAULTS['spike_dir'])).expanduser().resolve();gd=Path(data.get('gt_dir',DEFAULTS['gt_dir'])).expanduser().resolve() | |
| dataset=[] | |
| state.update(phase='validating',evaluation_complete=False,dataset_count=0,rows=[],requested=0,run_dir='',started=None,elapsed=0,message='检查 1,000 对 DAT 与 GT…',error='') | |
| def validate(): | |
| global dataset | |
| from PIL import Image | |
| ps=sorted(sd.glob('*.dat')) | |
| if len(ps)!=1000:raise ValueError(f'当前目录有 {len(ps)} 个 DAT,要求恰好 1,000 个。') | |
| pairs=[] | |
| for p in ps: | |
| gt=gd/(p.stem+'.png') | |
| if p.stat().st_size!=512500:raise ValueError('DAT 应为 41×250×400 位:'+p.name) | |
| if not gt.is_file():raise ValueError('缺少对应 GT:'+gt.name) | |
| with Image.open(gt) as im: | |
| if im.size!=(400,250):raise ValueError('GT 应为 400×250:'+gt.name) | |
| im.verify() | |
| pairs.append((p,gt)) | |
| with lock:dataset=pairs;state.update(phase='ready',dataset_count=len(pairs),dataset_config=dict(spike_dir=str(sd),gt_dir=str(gd)),message='已加载 1,000 对 DAT / GT,可以开始测试。') | |
| job(validate) | |
| elif kind=='run': | |
| if not state['model_loaded'] or len(dataset)!=1000:raise ValueError('请先加载模型和 1,000 对数据。') | |
| limit=int(data.get('limit',1000)) | |
| prompt=data.get('prompt',SETTINGS['prompt']) | |
| if not isinstance(prompt,str) or not prompt.strip() or len(prompt)>2000:raise ValueError('Prompt 必须为 1–2000 个字符。') | |
| prompt=prompt.strip() | |
| if limit not in [3,16,1000]:raise ValueError('请选择 3 张、16 张或完整 1,000 张。') | |
| out=ROOT/'outputs/ui_runs'/(time.strftime('%Y%m%d_%H%M%S')+'_'+uuid.uuid4().hex[:6]) | |
| for sub in ['condition','prediction']:(out/sub).mkdir(parents=True) | |
| stop.clear();state.update(phase='running',evaluation_complete=False,run_prompt=prompt,requested=limit,rows=[],run_dir=str(out),started=time.time(),elapsed=0,error='',message='正在重构和计算逐图指标…') | |
| pairs=list(dataset[:limit]) | |
| def run(): | |
| persist() | |
| for row in engine.predict_many(pairs,out,prompt,stop): | |
| with lock: | |
| state['rows'].append(row) | |
| state['message']=f"已完成 {len(state['rows'])} / {limit}:{row['id']}(GPU {row['gpu_id']})" | |
| persist() | |
| with lock: | |
| complete=len(state['rows'])==limit | |
| state.update(phase='complete' if complete else 'stopped',evaluation_complete=complete,message='测试完成。' if complete else '已停止,结果已保存。',elapsed=round(time.time()-state['started'],1)) | |
| persist() | |
| job(run,persist_run=True) | |
| else:raise ValueError('未知操作') | |
| class Handler(BaseHTTPRequestHandler): | |
| def send(self,body,ctype='application/json',status=200,extra=None): | |
| if isinstance(body,(dict,list)):body=json.dumps(body,ensure_ascii=False).encode() | |
| if isinstance(body,str):body=body.encode() | |
| self.send_response(status);self.send_header('Content-Type',ctype);self.send_header('Content-Length',str(len(body)));self.send_header('Cache-Control','no-store') | |
| self.send_header('X-Content-Type-Options','nosniff') | |
| for k,v in (extra or {}).items():self.send_header(k,v) | |
| self.end_headers();self.wfile.write(body) | |
| def do_GET(self): | |
| p=urlparse(self.path) | |
| try: | |
| if p.path=='/':self.send((Path(__file__).parent/'index.html').read_bytes(),'text/html; charset=utf-8') | |
| elif p.path=='/api/state':self.send(snapshot()) | |
| elif p.path=='/image': | |
| q=parse_qs(p.query);i=int(q['i'][0]);kind=q['kind'][0] | |
| if kind not in ['condition','prediction','gt'] or i<0:raise ValueError('Invalid image') | |
| with lock:row=dict(state['rows'][i]) | |
| self.send(Path(row[kind]).read_bytes(),'image/png') | |
| elif p.path=='/download': | |
| q=parse_qs(p.query);ext=q.get('format',['csv'])[0] | |
| if ext not in ['csv','json']:raise ValueError('Invalid format') | |
| with lock:run=state['run_dir'] | |
| if not run:raise ValueError('尚无结果') | |
| self.send((Path(run)/('metrics.'+ext)).read_bytes(),'text/csv; charset=utf-8' if ext=='csv' else 'application/json',extra={'Content-Disposition':f'attachment; filename="metrics.{ext}"'}) | |
| else:self.send({'error':'Not found'},status=404) | |
| except (ValueError,KeyError,IndexError,FileNotFoundError) as e:self.send({'error':str(e)},status=400) | |
| def do_POST(self): | |
| origin=self.headers.get('Origin') | |
| if origin and origin not in ['http://'+self.headers.get('Host','')]: | |
| self.send({'error':'Origin rejected'},status=403);return | |
| try: | |
| if self.headers.get_content_type()!='application/json':raise ValueError('JSON required') | |
| length=int(self.headers.get('Content-Length','0')) | |
| if not 0<length<16384:raise ValueError('Invalid body') | |
| action(urlparse(self.path).path.removeprefix('/api/'),json.loads(self.rfile.read(length)));self.send({'ok':True}) | |
| except (ValueError,TypeError) as e:self.send({'error':str(e)},status=400) | |
| def log_message(self,fmt,*args): | |
| if '/api/state' not in (args[0] if args else ''):super().log_message(fmt,*args) | |
| if __name__=='__main__': | |
| import argparse | |
| ap=argparse.ArgumentParser();ap.add_argument('--port',type=int,default=8997);args=ap.parse_args() | |
| print(f'Listening on http://127.0.0.1:{args.port}',flush=True) | |
| ThreadingHTTPServer(('127.0.0.1',args.port),Handler).serve_forever() | |