File size: 6,016 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
"""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()