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()