"""Brainmu: reproducible DAT reconstruction and grayscale evaluation.""" import argparse,csv,hashlib,json,math,platform,sys,time,traceback from pathlib import Path ROOT=Path(__file__).resolve().parents[1] sys.path.insert(0,str(ROOT/'code')) from project_config import load_config,DEFAULT_CONFIG def sha256(path): h=hashlib.sha256() with path.open('rb') as f: for block in iter(lambda:f.read(1024*1024),b''):h.update(block) return h.hexdigest() def validate_pairs(spike_dir,gt_dir,expected_count=0): from PIL import Image spec=load_config()['input'] files=sorted(spike_dir.glob('*.dat')) if not files:raise ValueError('No DAT files found') if expected_count and len(files)!=expected_count:raise ValueError(f'Expected {expected_count} DAT files, found {len(files)}') rows=[] for f in files: gt=gt_dir/(f.stem+'.png') if f.stat().st_size!=spec['packed_bytes']:raise ValueError('DAT size differs from config.json: '+f.name) if not gt.is_file():raise ValueError('Missing GT: '+gt.name) with Image.open(gt) as im: if im.size!=(spec['width'],spec['height']):raise ValueError('GT must be 400 x 250: '+gt.name) im.verify() rows.append(dict(id=f.stem,spike=str(f),gt=str(gt),spike_sha256=sha256(f),gt_sha256=sha256(gt))) return files,rows def main(): project=load_config();inference=project['generation']['inference'] p=argparse.ArgumentParser(description=__doc__) for name,default in [('model-path','checkpoints/Brainmu'),('adapter','checkpoints/lora/adapter.safetensors'),('frontend','../model.safetensors'),('spike-dir','data/test/spike'),('gt-dir','data/test/gt')]: p.add_argument('--'+name,type=Path,default=ROOT/default) p.add_argument('--output-dir',type=Path,default=ROOT/'outputs'/('test_'+time.strftime('%Y%m%d_%H%M%S'))) p.add_argument('--expected-count',type=int,default=0,help='Require exactly this many pairs; 0 = any count') p.add_argument('--limit',type=int,default=0,help='Smoke test only: sorted prefix; incompatible with expected-count') p.add_argument('--prompt',default=inference['prompt']) p.add_argument('--check-only',action='store_true',help='Validate all inputs without loading models or using GPU') a=p.parse_args() for name in ['model_path','adapter','frontend','spike_dir','gt_dir','output_dir']:setattr(a,name,getattr(a,name).resolve()) if a.limit<0 or a.expected_count<0:p.error('Counts must be nonnegative') if a.limit and a.expected_count:p.error('For subset smoke tests set --expected-count 0; full evaluation must not use --limit') if not a.prompt.strip():p.error('Prompt must not be empty') if a.output_dir.exists():p.error('output-dir must be new') required=[a.model_path/n for n in ['ema.safetensors','ae.safetensors','llm_config.json','vit_config.json','tokenizer_config.json','tokenizer.json']]+[a.adapter,a.adapter.with_name('adapter_config.json'),a.frontend] for f in required: if not f.is_file():p.error('Missing model file: '+str(f)) try:files,manifest=validate_pairs(a.spike_dir,a.gt_dir,a.expected_count) except (ValueError,OSError) as exc:p.error(str(exc)) if a.limit:files=files[:a.limit];manifest=manifest[:a.limit] a.output_dir.mkdir(parents=True) for sub in ['condition','prediction']:(a.output_dir/sub).mkdir() (a.output_dir/'dataset_manifest.json').write_text(json.dumps(manifest,indent=2)) import importlib.metadata as metadata versions={} for name in ['torch','torchvision','transformers','accelerate','flash-attn','numpy','Pillow']: try:versions[name]=metadata.version(name) except metadata.PackageNotFoundError:versions[name]=None config=dict(arguments={k:str(v) if isinstance(v,Path) else v for k,v in vars(a).items()},python=platform.python_version(),packages=versions,project_config=project,config_sha256=sha256(DEFAULT_CONFIG),metric='Pillow RGB->L, [0,1], per-image PSNR then mean; SSIM Gaussian 11x11 sigma=1.5, no border crop; GT resized to prediction',weights=[dict(path=str(f),bytes=f.stat().st_size,mtime_ns=f.stat().st_mtime_ns) for f in required],weight_fingerprint='size/mtime only, not content hashes',source_sha256={str(f.relative_to(ROOT)):sha256(f) for base in ['code','ui','vendor/Brainmu'] for f in sorted((ROOT/base).rglob('*.py'))}) (a.output_dir/'run_config.json').write_text(json.dumps(config,indent=2)) rows=[];begin=time.time() def persist(status,error=None): n=len(rows);complete=status=='complete' and n==len(files) report=dict(status=status,completed=n,requested=len(files),complete=complete,full_evaluation=complete and not a.limit,expected_count=a.expected_count,mean_psnr=sum(x['psnr_db'] for x in rows)/n if n else None,mean_ssim=sum(x['ssim'] for x in rows)/n if n else None,elapsed_seconds=time.time()-begin,error=error,rows=rows) tmp=a.output_dir/'metrics.json.tmp';tmp.write_text(json.dumps(report,indent=2,allow_nan=False));tmp.replace(a.output_dir/'metrics.json') if rows: tmp=a.output_dir/'metrics.csv.tmp' with tmp.open('w',newline='') as h: w=csv.DictWriter(h,fieldnames=list(rows[0]));w.writeheader();w.writerows(rows) tmp.replace(a.output_dir/'metrics.csv') return report if a.check_only: persist('validated');print(f'VALIDATED {len(files)} pairs; no inference performed. Report: {a.output_dir}');return persist('loading') try: sys.path.insert(0,str(ROOT/'ui')) from engine import Engine e=Engine();e.load(a.model_path,a.adapter,a.frontend) for i,f in enumerate(files): row=e.predict(f,a.gt_dir/(f.stem+'.png'),i,a.output_dir,a.prompt) if not all(math.isfinite(row[k]) for k in ['psnr_db','ssim']):raise ValueError('Nonfinite metric: '+f.name) rows.append(row);persist('running') print(f"{len(rows)}/{len(files)} PSNR={row['psnr_db']:.4f} SSIM={row['ssim']:.6f}",flush=True) report=persist('complete') print(f"COMPLETE {len(rows)}/{len(files)} | mean PSNR {report['mean_psnr']:.6f} dB | mean SSIM {report['mean_ssim']:.6f} | {a.output_dir}",flush=True) except BaseException as exc: persist('interrupted' if isinstance(exc,KeyboardInterrupt) else 'failed',str(exc));raise if __name__=='__main__':main()