Files
esh-pfi-infrastructure/scripts/training-probes/step1_profile.py
T
vh 7b5fd91d3c docs(gemma4-erp-tune): root-cause the 8.6% MFU — attention on Ampere kernels, 29.9% padding
Run-01 was killed at step 19 by operator instruction to root-cause before
spending a ~13.9h window. Two independent methods now agree on where the step
time went, and neither was the hypothesis the consult panel converged on.

Scaling fit (3 points, 2 params, residuals <3ms over an 8x range):
  A = 6.87e-4 s/token, B = 8.85e-8 s/token^2
  quadratic share 20.9% @ w=2048 -> 67.8% @ w=16384
  No fixed term was needed, which refutes launch-bound outright.

Profiler kernel table (device rows only):
  attention   22,835.8 ms  65.2%   fmha_cutlass*_sm80
  dense GEMM   2,774.0 ms   7.9%
  other        5,739.0 ms  16.4%

The attention kernels are sm80 — Ampere-generation CUTLASS running on an
sm_120 Blackwell card, with the forward on the gmem fallback tier. That is the
mechanism behind 100% SM utilisation at 27 of 304 available TFLOPS.

Correctness cleared separately: the sliding mask asserts at max 1024
allowed/row, so the 25 windowed layers were genuinely windowed. The same probe
found that right-padding is what pins the 5 global layers to an explicit 4D
mask and off the is_causal fast path — measured at 9.4% slower for 24% less
loss work at fixed width.

The largest available win is not the attention kernel. The corpus is 29.9%
padding, and bucket-to-pair + shuffle-to-mix takes it to 0.0% for >=35.5% wall
clock, no new dependency, unchanged peak memory. Bucket size turned out not to
be a diversity knob — roots per accumulation window are flat across a 256x
range, so the global micro-batch shuffle does that work alone and the bucket
should be tight.

Adds docs/pfi/training-throughput-playbook.md as the durable model-agnostic
home (sibling to the quantization playbook), the four probes under
scripts/training-probes/ with raw output kept for re-derivation, and a §6 to
the sizing doc carrying the Gemma-4-specific numbers and round-2 restart
parameters.

Measured negatives recorded so they are not re-chased: grouped_mm (0.9%
slower, and MoE is only 7.9% of the step), CUDA graphs / torch.compile over
the expert loop (no fixed cost to amortise), liger fused CE (~1-3% lever),
FA4 on sm_120.

Round-1 state preserved: 609MB encode cache, order manifest, truncation
report, resume script. No checkpoints — it died at step 19 and the first was
due at 100, so the lora_B inert-adapter gate never ran and moves to the
restart.
2026-08-24 22:10:51 -07:00

207 lines
7.7 KiB
Python

"""Steps 1/2/4 - profiler kernel table, sequence scaling, isolated CE timing.
Loads the real model exactly as erp_sft_harness.runtime does (same
from_pretrained args, same PEFT config, same gradient checkpointing, same
chunked-CE compute_loss) and measures:
A. sequence scaling 2x2048 / 2x8192 / 2x16384 fwd+bwd
linear-dominated -> time falls ~8x from 16384 to 2048
launch-bound -> time barely falls
quadratic-dominated -> time falls ~64x
B. isolated CE wall clock (CUDA events around the chunked-CE block)
C. torch.profiler kernel table, sorted by self CUDA time
D. expert-GEMM launch counts (settles grouped_mm without kernel-name
archaeology: 128 sequential launches per layer = no-op, 1 = grouped)
Runs on GPU0, which is reserved and idle. Nothing else touches it.
"""
import json
import sys
import time
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, get_peft_model
sys.path.insert(0, "/tank/erp-tune/eitri-smithy")
from erp_sft_harness.core import IGNORE_INDEX, discover_target_modules
MODEL = "/tank/aimodels/gemma4-26b-a4b-it-heretic-bf16"
CHUNK = 1024
MB = 2
print("=" * 72)
print("loading model")
print("=" * 72, flush=True)
t0 = time.time()
model = AutoModelForCausalLM.from_pretrained(
MODEL, dtype=torch.bfloat16, device_map={"": 0}, attn_implementation="sdpa",
)
targets = discover_target_modules(model)
model = get_peft_model(model, LoraConfig(
r=64, lora_alpha=128, lora_dropout=0.05, target_modules=targets,
bias="none", task_type="CAUSAL_LM",
))
model.enable_input_require_grads()
model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
model.train()
print("loaded in %.1fs targets=%d" % (time.time() - t0, len(targets)), flush=True)
base = model.base_model.model if hasattr(model, "base_model") else model
body = base.model
lm_head = base.get_output_embeddings()
softcap = getattr(model.config.get_text_config(), "final_logit_softcapping", None)
print("final_logit_softcapping = %s" % softcap)
print("attn_implementation = %s" % model.config.get_text_config()._attn_implementation)
print(flush=True)
ce_ms = {"fwd": 0.0}
def compute_loss(input_ids, attention_mask, labels, time_ce=False):
"""Byte-for-byte the harness's compute_loss, with optional CE timing."""
hidden = body(input_ids=input_ids, attention_mask=attention_mask,
use_cache=False).last_hidden_state
flat_hidden = hidden[:, :-1, :].reshape(-1, hidden.size(-1))
flat_labels = labels[:, 1:].reshape(-1)
keep = flat_labels != IGNORE_INDEX
kept_hidden = flat_hidden[keep]
kept_labels = flat_labels[keep]
kept = int(kept_labels.numel())
def chunk_loss(chunk_hidden, chunk_labels):
logits = lm_head(chunk_hidden).float()
if softcap is not None:
logits = torch.tanh(logits / softcap) * softcap
return torch.nn.functional.cross_entropy(logits, chunk_labels, reduction="sum")
if time_ce:
s, e = torch.cuda.Event(True), torch.cuda.Event(True)
torch.cuda.synchronize()
s.record()
total = torch.zeros((), device=kept_hidden.device, dtype=torch.float32)
for start in range(0, kept, CHUNK):
total = total + torch.utils.checkpoint.checkpoint(
chunk_loss, kept_hidden[start:start + CHUNK],
kept_labels[start:start + CHUNK], use_reentrant=False,
)
if time_ce:
e.record()
torch.cuda.synchronize()
ce_ms["fwd"] = s.elapsed_time(e)
return total / kept, kept
def make_batch(n, pad_frac=0.0):
"""Synthetic batch. pad_frac trims the SECOND row and right-pads it,
mimicking collate_mixed on a heterogeneous pair."""
ids = torch.randint(100, 200000, (MB, n), device="cuda")
am = torch.ones(MB, n, dtype=torch.long, device="cuda")
labels = ids.clone()
if pad_frac > 0:
keep = int(n * (1 - pad_frac))
am[1, keep:] = 0
labels[1, keep:] = IGNORE_INDEX
# ~40% of real tokens carry loss (measured mean 2188/2752 is higher, but
# rp-dialogue assistant-only masking pulls the mix down); use the measured
# global ratio 57.7M ctx -> 45.9M targets = 0.795
m = torch.rand(labels.shape, device="cuda") > 0.795
labels[m] = IGNORE_INDEX
return ids, am, labels
def timed(n, pad_frac=0.0, reps=2, time_ce=False, label=""):
ids, am, labels = make_batch(n, pad_frac)
for _ in range(1): # warmup
loss, kept = compute_loss(ids, am, labels)
loss.backward()
model.zero_grad(set_to_none=True)
torch.cuda.synchronize()
best = None
for _ in range(reps):
torch.cuda.reset_peak_memory_stats()
t = time.perf_counter()
loss, kept = compute_loss(ids, am, labels, time_ce=time_ce)
loss.backward()
torch.cuda.synchronize()
dt = time.perf_counter() - t
best = dt if best is None else min(best, dt)
model.zero_grad(set_to_none=True)
peak = torch.cuda.max_memory_allocated() / 2**30
print(" %-34s %7.3f s kept=%-6d peak=%5.1f GiB%s" % (
label or ("2x%d pad=%.0f%%" % (n, pad_frac * 100)),
best, kept, peak,
(" CE=%.0f ms (%.1f%%)" % (ce_ms["fwd"], 100 * ce_ms["fwd"] / 1000 / best)) if time_ce else ""))
return best
print("=" * 72)
print("A. SEQUENCE SCALING (no padding - isolates n)")
print("=" * 72, flush=True)
t2048 = timed(2048, 0.0, label="2 x 2,048")
t8192 = timed(8192, 0.0, label="2 x 8,192")
t16384 = timed(16384, 0.0, label="2 x 16,384")
print()
print(" 16384 -> 2048 ratio %.2fx (linear ~8x, launch-bound ~1x, quadratic ~64x)"
% (t16384 / t2048))
print(" 16384 -> 8192 ratio %.2fx (linear ~2x, quadratic ~4x)"
% (t16384 / t8192))
print(flush=True)
print("=" * 72)
print("B. PADDING PENALTY (same real tokens, with vs without pad)")
print("=" * 72, flush=True)
timed(16384, 0.0, label="2 x 16,384 no padding")
timed(16384, 0.5, label="2 x 16,384 50% pad on row 1")
print(flush=True)
print("=" * 72)
print("C. ISOLATED CE WALL CLOCK")
print("=" * 72, flush=True)
timed(16384, 0.0, reps=2, time_ce=True, label="2 x 16,384 (CE timed)")
timed(4096, 0.0, reps=2, time_ce=True, label="2 x 4,096 (CE timed)")
print(flush=True)
print("=" * 72)
print("D. KERNEL TABLE - one fwd+bwd at 2 x 16,384")
print("=" * 72, flush=True)
ids, am, labels = make_batch(16384, 0.0)
loss, _ = compute_loss(ids, am, labels)
loss.backward()
model.zero_grad(set_to_none=True)
torch.cuda.synchronize()
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA],
record_shapes=False, with_stack=False,
) as prof:
loss, _ = compute_loss(ids, am, labels)
loss.backward()
torch.cuda.synchronize()
model.zero_grad(set_to_none=True)
print(prof.key_averages().table(sort_by="self_cuda_time_total", row_limit=45))
print()
print("=" * 72)
print("E. LAUNCH COUNTS (grouped_mm: 128/layer sequential = no-op, 1 = grouped)")
print("=" * 72)
rows = []
for ev in prof.key_averages():
if ev.self_device_time_total <= 0:
continue
rows.append((ev.key, ev.count, ev.self_device_time_total / 1000.0))
rows.sort(key=lambda r: -r[2])
print(" %-58s %8s %10s" % ("kernel", "count", "self ms"))
for k, c, ms in rows[:30]:
print(" %-58s %8d %10.1f" % (k[:58], c, ms))
total_ms = sum(r[2] for r in rows)
print()
print(" total self CUDA time %.1f ms" % total_ms)
gemm = sum(ms for k, c, ms in rows if any(t in k.lower() for t in
("gemm", "cutlass", "sm90", "sm100", "sm120", "nvjet", "ampere", "tensor")))
print(" GEMM-ish kernels %.1f ms (%.1f%%)" % (gemm, 100 * gemm / total_ms))
print(" non-GEMM %.1f ms (%.1f%%)" % (total_ms - gemm, 100 * (total_ms - gemm) / total_ms))