Download code/reporting.py from FidelityWM/planning-baselines: direct link, hf CLI and curl.
- Browser
- Download file 5.23 kB
-
https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/reporting.py
- Command line
-
hf download hf://FidelityWM/planning-baselines/code/reporting.py
-
curl -L -o reporting.py https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/reporting.py
5.23 kB
| """Report only measured 200-sample scores; aggregate across all five seeds.""" | |
| import csv, hashlib, json, statistics, sys | |
| from pathlib import Path | |
| ROOT=Path(__file__).resolve().parents[1] | |
| sys.path.insert(0,str(ROOT/'runtime')) | |
| METHODS=['dinowm','pldm','fast-lewm','gcbc','gciql','gcivl'] | |
| TASKS=['tworoom','cube','pusht'] | |
| OFFSETS=[25,50,75,100] | |
| SEEDS=[42,43,44,45,46] | |
| def result_path(method,task,offset,seed): | |
| suffix='' if seed==42 else f'_seed{seed}' | |
| return ROOT/'results'/f'{method}_{task}_{offset}{suffix}.json' | |
| def atomic_json(path,data): | |
| tmp=path.with_suffix('.tmp');tmp.write_text(json.dumps(data,indent=2));tmp.replace(path) | |
| def aggregate(rows): | |
| values=[r['success_percent'] for r in rows if r['status']=='complete'] | |
| complete=len(values)==len(SEEDS) | |
| return len(values),statistics.mean(values) if complete else None,statistics.stdev(values) if complete else None | |
| def save_sheet(): | |
| rows=[];summary=[] | |
| plan_path=ROOT/'results/training_plan.json' | |
| planned={(p['model'],p['dataset']) for p in json.loads(plan_path.read_text()) if p['train']} if plan_path.exists() else set() | |
| for method in METHODS: | |
| for task in TASKS: | |
| for offset in OFFSETS: | |
| group=[] | |
| for seed in SEEDS: | |
| p=result_path(method,task,offset,seed) | |
| d=json.loads(p.read_text()) if p.exists() else {} | |
| count=len(d.get('successes',[])) | |
| assert count<=200 | |
| status='complete' if count==200 else ('running' if count else 'queued') | |
| if not count and not (method=='fast-lewm' or (method=='dinowm' and task=='pusht')): | |
| status='unavailable: compatible checkpoint not found' | |
| if not count and (method,task) in planned: | |
| training_path=ROOT/'results'/f'training_{method}_{task}.json' | |
| training=json.loads(training_path.read_text()) if training_path.exists() else {} | |
| status='training' if training.get('status')=='training' else ('queued evaluation' if training.get('status')=='trained' else 'queued training') | |
| r=dict(model=method,dataset=task,offset=offset,seed=seed,samples_completed=count,samples_required=200,successes=sum(d['successes']) if count else None,success_percent=100*sum(d['successes'])/count if count==200 else None,status=status,execution_budget=50,checkpoint_release=d.get('checkpoint_release','naiverer/fast-leworldmodel' if method=='fast-lewm' and count else None)) | |
| r.update(training_epochs=d.get('training',{}).get('epochs_completed'),training_seed=d.get('training',{}).get('training_seed'),model_variant=d.get('training',{}).get('variant')) | |
| rows.append(r);group.append(r) | |
| n,mean,std=aggregate(group) | |
| summary.append(dict(model=method,dataset=task,offset=offset,**{f'seed_{r["seed"]}':r['success_percent'] for r in group},seeds_completed=n,mean_percent=mean,std_percent=std,mean_plus_std=f'{mean:.2f} ± {std:.2f}' if mean is not None else None)) | |
| for name,data in [('planning_success',rows),('planning_summary',summary)]: | |
| atomic_json(ROOT/'results'/f'{name}.json',data) | |
| p=ROOT/'results'/f'{name}.csv';tmp=p.with_suffix('.tmp') | |
| with tmp.open('w',newline='') as f: | |
| w=csv.DictWriter(f,fieldnames=list(data[0]));w.writeheader();w.writerows(data) | |
| tmp.replace(p) | |
| import openpyxl | |
| wb=openpyxl.Workbook();wb.remove(wb.active) | |
| for name,data in [('Seeds and mean std',summary),('Planning success',rows)]: | |
| ws=wb.create_sheet(name);ws.append(list(data[0])) | |
| for r in data:ws.append(['null' if v is None else v for v in r.values()]) | |
| ws.freeze_panes='D2';ws.auto_filter.ref=ws.dimensions | |
| for col in ws.columns:ws.column_dimensions[col[0].column_letter].width=min(60,max(len(str(c.value)) for c in col)+2) | |
| ws=wb.create_sheet('Protocol');ws.append(['Setting','Value']) | |
| for row in [('Seeds','42, 43, 44, 45, 46'),('Samples per seed / offset',200),('Goal offsets','25, 50, 75, 100'),('Mean and std','Across all five completed seed success percentages; sample std, ddof=1'),('Sampling','Different pairs per seed; identical pairs across models'),('CEM seed','Run seed + batch starting index'),('Execution budget',50),('Planning window',25),('CEM candidates',300),('CEM iterations',30),('CEM elites',30),('Environment batch',2),('Checkpoint deletion','After all five seeds and four offsets for that model/task finish'),('Meaning of null','Unavailable or incomplete; aggregate requires all five seeds'),('Full protocol','See PROTOCOL.md')]:ws.append(row) | |
| ws.column_dimensions['A'].width=30;ws.column_dimensions['B'].width=95 | |
| workbook=ROOT/'results/planning_success.xlsx';tmp=workbook.with_suffix('.tmp.xlsx');wb.save(tmp);tmp.replace(workbook) | |
| atomic_json(ROOT/'results/validation.json',dict(completed_runs=sum(r['status']=='complete' for r in rows),completed_samples=sum(r['samples_completed'] for r in rows if r['status']=='complete'),in_progress_samples=sum(r['samples_completed'] for r in rows if r['status']!='complete'),xlsx_sha256=hashlib.sha256(workbook.read_bytes()).hexdigest(),downloaded_checkpoints_remaining=len([p for p in (ROOT/'checkpoints').rglob('*') if p.suffix in ['.pth','.pt','.ckpt']]),scores=[r for r in rows if r['status']=='complete'])) | |
| if __name__=='__main__':save_sheet() | |