diff options
Diffstat (limited to 'ep_run/fig_curves.py')
| -rw-r--r-- | ep_run/fig_curves.py | 94 |
1 files changed, 94 insertions, 0 deletions
diff --git a/ep_run/fig_curves.py b/ep_run/fig_curves.py new file mode 100644 index 0000000..9185702 --- /dev/null +++ b/ep_run/fig_curves.py @@ -0,0 +1,94 @@ +"""Validation curves at 135M: EP with the readout fix, EP without it, and the backprop twins. + +The point of the figure is that the first and third are the same curve, and the second stops +descending partway through. Run after fw135m_rlin completes; it reads the training logs directly. +""" +import re +import statistics as st +from pathlib import Path + +import matplotlib +matplotlib.use('Agg') +import matplotlib.pyplot as plt + +RUNS = Path('/home/yurenh2/ept/ep_run/runs') +TOTAL = 440000 +TOK_PER_STEP = 24 * 256 # effective batch x context + + +def curve(tag): + pts = [] + for line in (RUNS / f'{tag}.log').read_bytes().decode(errors='replace').splitlines(): + m = re.match(rf'step (\d+)/{TOTAL} \| train [\d.]+ val ([\d.]+)', line) + if m: + pts.append((int(m.group(1)), float(m.group(2)))) + pts.sort() + return pts + + +def smooth(pts, win=25): + xs = [p[0] for p in pts] + ys = [p[1] for p in pts] + out = [] + for i in range(len(ys)): + lo, hi = max(0, i - win // 2), min(len(ys), i + win // 2 + 1) + out.append(st.median(ys[lo:hi])) + return xs, out + + +def main(): + # EP is drawn thick and underneath; the backprop seeds go on top as thin lines, so the reader + # sees them tracking rather than being hidden by whichever curve was plotted last. + arms = [ + ('fw135m_rlin', 'Equilibrium Propagation', '#1f6feb', 3.0, '-', 2), + ('fw135m_bsign', 'Equilibrium Propagation,\nwithout the readout fix', '#d1442f', 2.6, '-', 2), + ('fw135m_bp', 'Backprop, seed 1', '#2b3540', 1.0, '-', 4), + ('fw135m_bp_s2', 'Backprop, seed 2', '#7d8794', 1.0, '-', 4), + ] + fig, ax = plt.subplots(figsize=(7.6, 4.9)) + finals = {} + for tag, label, color, lw, ls, zo in arms: + pts = curve(tag) + if not pts: + print(f'skip {tag}: no data') + continue + xs, ys = smooth(pts) + toks = [x * TOK_PER_STEP / 1e9 for x in xs] + ax.plot(toks, ys, color=color, lw=lw, ls=ls, label=label, zorder=zo) + tail = [v for s, v in pts if s >= 0.9 * TOTAL] + finals[tag] = st.mean(tail) if tail else float('nan') + print(f'{tag:14s} last step {pts[-1][0]:>7} tail mean {finals[tag]:.4f}') + + ax.set_xlabel('training tokens (billions)') + ax.set_ylabel('validation cross-entropy') + ax.set_ylim(3.15, 4.35) + ax.set_xlim(0, TOTAL * TOK_PER_STEP / 1e9) + ax.grid(alpha=0.18, lw=0.7) + for side in ('top', 'right'): + ax.spines[side].set_visible(False) + h, l = ax.get_legend_handles_labels() + order = [2, 3, 0, 1] # backprop seeds first in the legend, matching how one reads the plot + leg = ax.legend([h[i] for i in order], [l[i] for i in order], + frameon=False, fontsize=9.5, loc='upper right', handlelength=1.6) + for t in leg.get_texts(): + t.set_va('center') + + if 'fw135m_bsign' in finals and 'fw135m_rlin' in finals: + ax.annotate('single-precision rounding in the contrast readout;\nthe run stops improving', + xy=(1.9, finals['fw135m_bsign'] + 0.005), xytext=(1.35, 3.86), + fontsize=9, color='#d1442f', + arrowprops=dict(arrowstyle='-', color='#d1442f', lw=0.9, alpha=0.8)) + ax.annotate(f"EP {finals['fw135m_rlin']:.3f}\nBP {st.mean([finals['fw135m_bp'], finals['fw135m_bp_s2']]):.3f}", + xy=(2.62, 3.245), fontsize=9.5, color='#2b3540', ha='right') + + ax.set_title('135M-parameter transformer language model, FineWeb-Edu', + fontsize=10.5, color='#333', pad=10, loc='left') + fig.tight_layout() + for ext in ('png', 'pdf'): + p = Path(f'/home/yurenh2/ept/assets/figs/fig_135m_curves.{ext}') + fig.savefig(p, dpi=300, bbox_inches='tight') + print('wrote', p) + + +if __name__ == '__main__': + main() |
