Download scripts/plot_results.py from kiruluta/rsil-benchmark: direct link, hf CLI and curl.
- Browser
- Download file 1.21 kB
-
https://huggingface.co/kiruluta/rsil-benchmark/resolve/main/scripts/plot_results.py
- Command line
-
hf download hf://kiruluta/rsil-benchmark/scripts/plot_results.py
-
curl -L -o plot_results.py https://huggingface.co/kiruluta/rsil-benchmark/resolve/main/scripts/plot_results.py
1.21 kB
| import argparse,csv,os | |
| import matplotlib.pyplot as plt | |
| p=argparse.ArgumentParser(); p.add_argument('--run',default='outputs/real_preliminary'); p.add_argument('--out',default='publish/plots'); a=p.parse_args(); os.makedirs(a.out,exist_ok=True) | |
| rows=list(csv.DictReader(open(os.path.join(a.run,'leaderboard.csv')))); names=[r['model'] for r in rows] | |
| for metric,title in [('validation_loss','Held-out validation loss'),('tokens_per_second','Training throughput (tokens/s)')]: | |
| if not rows or metric not in rows[0]: continue | |
| vals=[float(r[metric]) for r in rows]; plt.figure(figsize=(8,4.5)); plt.bar(names,vals); plt.title(title); plt.xticks(rotation=30,ha='right'); plt.tight_layout(); plt.savefig(os.path.join(a.out,metric+'.png'),dpi=160); plt.close() | |
| plt.figure(figsize=(8,4.5)) | |
| for name in names: | |
| fn=os.path.join(a.run,name,'history.csv') | |
| if os.path.exists(fn): | |
| h=list(csv.DictReader(open(fn))); plt.plot([int(x['step']) for x in h],[float(x['train_loss']) for x in h],label=name) | |
| plt.xlabel('step'); plt.ylabel('training loss'); plt.title('Training curves'); plt.legend(fontsize=8); plt.tight_layout(); plt.savefig(os.path.join(a.out,'training_curves.png'),dpi=160); plt.close() | |