Buckets:
| """Post-process HF Job 2 outputs (mem_sweep.json + GPT arm CSVs/JSONs) into | |
| verdict tables and Plotly figures for the logbook. | |
| Usage: python analyze_job2.py <job2_dir> (a local download of bucket /job2) | |
| """ | |
| import csv | |
| import json | |
| import math | |
| import os | |
| import sys | |
| import plotly.graph_objects as go | |
| from plotly.subplots import make_subplots | |
| # expected bytes/param (analytic; scale overhead = 2/32 per quantized moment) | |
| EXPECTED = { | |
| "ref": {"params": 4, "state": 8, "grads": 4}, | |
| "flash": {"params": 2, "state": 3.125, "grads": 2}, | |
| "flash_gr": {"params": 2, "state": 3.125, "grads": 0}, | |
| "split_only": {"params": 2, "state": None, "grads": 2}, # states: library | |
| # stores param-dtype | |
| "quant_only": {"params": 4, "state": 2.125, "grads": 4}, | |
| } | |
| COLORS = {"ref": "#2a78d6", "flash": "#008300", "flash_gr": "#1baf7a", | |
| "split_only": "#e87ba4", "quant_only": "#eda100", "linear": "#e34948"} | |
| INK, MUTED, GRID, SURF = "#0b0b0b", "#898781", "#e1e0d9", "#fcfcfb" | |
| LLAMA_N = 8_030_261_248 | |
| GIB = 2 ** 30 | |
| def analyze_mem(d, out): | |
| recs = [r for r in json.load(open(f"{d}/mem_sweep.json")) if "n_params" in r] | |
| fails = [r for r in json.load(open(f"{d}/mem_sweep.json")) if "n_params" not in r] | |
| rows, verdicts = [], {} | |
| gpt = [r for r in recs if r["model"].startswith("gpt")] | |
| for r in recs: | |
| exp = EXPECTED.get(r["variant"], {}) | |
| rows.append(dict(model=r["model"], variant=r["variant"], n=r["n_params"], | |
| bpp_params=round(r["bpp_params"], 3), | |
| bpp_state=round(r["bpp_state"], 3), | |
| bpp_grads=round(r["bpp_grads"], 3), | |
| bpp_total=round(r["bpp_params_state_grads"], 3), | |
| exp_params=exp.get("params"), exp_state=exp.get("state"), | |
| exp_grads=exp.get("grads"), | |
| peak_bpp=round(r["peak_steady"] / r["n_params"], 3))) | |
| # slope check: bytes/param stability across sizes for ref & flash | |
| for variant in ["ref", "flash", "flash_gr", "split_only", "quant_only"]: | |
| vs = sorted([r for r in gpt if r["variant"] == variant], key=lambda r: r["n_params"]) | |
| if len(vs) >= 2: | |
| tot = [r["bpp_params_state_grads"] for r in vs] | |
| verdicts[variant] = dict( | |
| bpp_by_size={f'{r["n_params"]/1e6:.0f}M': round(r["bpp_params_state_grads"], 3) for r in vs}, | |
| bpp_at_largest=round(tot[-1], 3), | |
| extrapolated_8B_params_state_grads_gib=round(tot[-1] * LLAMA_N / GIB, 1)) | |
| # Table 4 comparison from largest measured slope | |
| t4 = {} | |
| for variant, paper_params, paper_optim in [("ref", 29.9, 59.8), ("flash", 15.0, 23.4)]: | |
| vs = sorted([r for r in gpt if r["variant"] == variant], key=lambda r: r["n_params"]) | |
| if vs: | |
| r = vs[-1] | |
| t4[variant] = dict( | |
| pred_params_gib=round(r["bpp_params"] * LLAMA_N / GIB, 1), paper_params=paper_params, | |
| pred_optim_gib=round(r["bpp_state"] * LLAMA_N / GIB, 1), paper_optim=paper_optim) | |
| summary = dict(rows=rows, verdicts=verdicts, table4_extrapolation=t4, | |
| failed_probes=fails) | |
| json.dump(summary, open(f"{out}/mem_analysis.json", "w"), indent=2) | |
| # figure: bytes/param by model size per variant (params+state+grads) | |
| fig = go.Figure() | |
| for variant in ["ref", "flash", "flash_gr", "split_only", "quant_only"]: | |
| vs = sorted([r for r in gpt if r["variant"] == variant], key=lambda r: r["n_params"]) | |
| if not vs: | |
| continue | |
| fig.add_trace(go.Scatter( | |
| x=[r["n_params"] / 1e6 for r in vs], | |
| y=[r["bpp_params_state_grads"] for r in vs], | |
| mode="lines+markers", name=variant, | |
| line=dict(color=COLORS[variant], width=2), marker=dict(size=8))) | |
| for y, lab in [(16, "paper: AdamW 16 B/param"), (7.125, "paper: FlashAdamW 7.125"), | |
| (5.125, "paper: +grad release 5.125")]: | |
| fig.add_hline(y=y, line_dash="dot", line_color=MUTED, | |
| annotation_text=lab, annotation_font_color=MUTED) | |
| fig.update_xaxes(type="log", title="model parameters (millions)", gridcolor=GRID, color=MUTED) | |
| fig.update_yaxes(title="measured bytes / parameter (params+optimizer+grads)", | |
| gridcolor=GRID, color=MUTED, rangemode="tozero") | |
| fig.update_layout(title="Measured CUDA bytes/param vs model size (A10G, one variant per process)", | |
| font=dict(family="system-ui, sans-serif", color=INK), | |
| paper_bgcolor=SURF, plot_bgcolor=SURF, height=480, | |
| legend=dict(orientation="h", y=1.08, x=0)) | |
| fig.write_html(f"{out}/mem_scaling.html", include_plotlyjs="cdn") | |
| return summary | |
| def analyze_gpt(d, out): | |
| arms = {} | |
| for name in ["ref_s0", "flash_s0", "ref_s1", "linear_s0"]: | |
| try: | |
| arms[name] = dict( | |
| summary=json.load(open(f"{d}/{name}.json")), | |
| curve=[(int(r["step"]), float(r["loss"])) | |
| for r in csv.DictReader(open(f"{d}/{name}.csv"))]) | |
| except FileNotFoundError: | |
| arms[name] = None | |
| ok = {k: v for k, v in arms.items() if v} | |
| def ema(xs, a=0.02): | |
| out_, m = [], None | |
| for x in xs: | |
| m = x if m is None else (1 - a) * m + a * x | |
| out_.append(m) | |
| return out_ | |
| metrics = {} | |
| if "ref_s0" in ok and "flash_s0" in ok: | |
| n = min(len(ok["ref_s0"]["curve"]), len(ok["flash_s0"]["curve"])) | |
| r0 = ema([x[1] for x in ok["ref_s0"]["curve"][:n]]) | |
| f0 = ema([x[1] for x in ok["flash_s0"]["curve"][:n]]) | |
| metrics["flash_vs_ref_mean_abs_ema_delta_2nd_half"] = ( | |
| sum(abs(a - b) for a, b in zip(r0[n // 2:], f0[n // 2:])) / (n - n // 2)) | |
| metrics["final_val_loss"] = {k: ok[k]["summary"]["val_loss"] for k in ok} | |
| if "ref_s1" in ok: | |
| n1 = min(n, len(ok["ref_s1"]["curve"])) | |
| s1 = ema([x[1] for x in ok["ref_s1"]["curve"][:n1]]) | |
| metrics["seed_noise_mean_abs_ema_delta_2nd_half"] = ( | |
| sum(abs(a - b) for a, b in zip(r0[n1 // 2:n1], s1[n1 // 2:])) / (n1 - n1 // 2)) | |
| if "linear_s0" in ok: | |
| metrics["linear_diverged"] = bool(ok["linear_s0"]["summary"].get("diverged")) | |
| metrics["linear_final_loss"] = ok["linear_s0"]["summary"].get("final_loss") | |
| json.dump(metrics, open(f"{out}/gpt_analysis.json", "w"), indent=2) | |
| fig = make_subplots(rows=1, cols=2, column_widths=[0.62, 0.38], horizontal_spacing=0.08, | |
| subplot_titles=("Training loss (EMA-smoothed)", "Raw loss, divergence arm")) | |
| names = {"ref_s0": ("AdamW (ref, seed0)", "#2a78d6"), "flash_s0": ("FlashAdamW (seed0)", "#008300"), | |
| "ref_s1": ("AdamW (ref, seed1)", "#e87ba4"), "linear_s0": ("Linear-quant FlashAdamW", "#e34948")} | |
| for k, (label, color) in names.items(): | |
| if not arms.get(k): | |
| continue | |
| xs = [p[0] for p in arms[k]["curve"]] | |
| ys = ema([p[1] for p in arms[k]["curve"]]) if k != "linear_s0" else [p[1] for p in arms[k]["curve"]] | |
| fig.add_trace(go.Scatter(x=xs, y=ys, mode="lines", name=label, | |
| line=dict(color=color, width=2)), | |
| row=1, col=(2 if k == "linear_s0" else 1)) | |
| if k == "linear_s0" and arms.get("ref_s0"): | |
| fig.add_trace(go.Scatter(x=[p[0] for p in arms["ref_s0"]["curve"]], | |
| y=[p[1] for p in arms["ref_s0"]["curve"]], | |
| mode="lines", name="AdamW raw (ref)", showlegend=False, | |
| line=dict(color="#2a78d6", width=1)), row=1, col=2) | |
| fig.update_xaxes(title="step", gridcolor=GRID, color=MUTED) | |
| fig.update_yaxes(title="train loss", gridcolor=GRID, color=MUTED) | |
| fig.update_layout(title="GPT-2 124M on FineWeb10B (A10G): FlashAdamW vs AdamW, matched init+data", | |
| font=dict(family="system-ui, sans-serif", color=INK), | |
| paper_bgcolor=SURF, plot_bgcolor=SURF, height=460, | |
| legend=dict(orientation="h", y=1.12, x=0)) | |
| fig.write_html(f"{out}/gpt_trajectories.html", include_plotlyjs="cdn") | |
| return metrics | |
| if __name__ == "__main__": | |
| d = sys.argv[1] | |
| out = "repro_flashoptim/outputs/job2" | |
| os.makedirs(out, exist_ok=True) | |
| m1 = analyze_mem(d, out) | |
| m2 = analyze_gpt(d, out) | |
| print(json.dumps({"mem": {k: v for k, v in m1.items() if k != "rows"}, | |
| "gpt": m2}, indent=2, default=str)) | |
Xet Storage Details
- Size:
- 8.66 kB
- Xet hash:
- af5630b6150437ffa0f7dfbfce2b863a4c3466339367493bf44ab3563b543849
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.