#!/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