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:
# pip install -U transformers accelerate # Load model directly from transformers import SpikeConvFrontend model = SpikeConvFrontend.from_pretrained("BAAI/Brainmu-SpikeCamera", device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 9,493 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 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 | #!/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()
|