File size: 1,206 Bytes
f048438 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 | 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()
|