ProCreations's picture
download
raw
8.66 kB
"""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.