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()