sunbaby's picture
Upload 69 files
4719196
Raw History Blame Contribute Delete
6.02 kB
"""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()