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.
This commit is contained in:
@@ -62,6 +62,23 @@ repeats the playbook, you are re-litigating — record the delta in the playbook
|
||||
instead. When a playbook claim turns out wrong, don't just fix it: add a dated
|
||||
row to its superseded-claims table so old docs stop misleading people.
|
||||
|
||||
## Training throughput
|
||||
|
||||
Same contract as quantization, different subject: **`docs/pfi/training-throughput-playbook.md`
|
||||
is the durable home** for why a training run is slow — the 10-minute scaling
|
||||
triage that names the regime before you profile, the padding/masking landmines,
|
||||
the profiler traps, and its own superseded-claims table. Read it before
|
||||
hypothesising about kernels.
|
||||
|
||||
The instruments are committed at [`scripts/training-probes/`](scripts/training-probes/)
|
||||
with raw output kept alongside, so the claims can be re-derived rather than
|
||||
taken on faith.
|
||||
|
||||
⚠ **Measure before you argue.** The playbook exists because a four-model
|
||||
frontier panel produced four self-retractions in ninety minutes on this
|
||||
question, and every one of them was a derivation while every survivor was a
|
||||
measurement.
|
||||
|
||||
## Purpose
|
||||
|
||||
- Inventory of servers and their state
|
||||
|
||||
@@ -321,3 +321,217 @@ Also fold in:
|
||||
`vllm/vllm-openai:v0.27.1` already on disk. **If it still no-ops, the harness
|
||||
must emit merged weights** — and Eitri needs that requirement while he is
|
||||
early, not after the run.
|
||||
|
||||
---
|
||||
|
||||
## 6. Round-1 aborted; throughput root-caused (measured 2026-08-24 22:00 PDT)
|
||||
|
||||
Run-01 launched, reached step 19 of 1,312 at ~35–46 s/it, and was **killed by
|
||||
operator instruction** — not a crash, not an OOM. ETA was ~13.9 h at 8.6% MFU
|
||||
and the operator elected to root-cause before spending the window.
|
||||
|
||||
Nothing was destroyed: the 609 MB encode cache, `order-manifest.jsonl`,
|
||||
`truncation-report.json` and `resume-run-01.sh` are all preserved at
|
||||
`/tank/erp-tune/run-01/`. **There are no checkpoints** — the first was due at
|
||||
step 100, so brokkr's `lora_B` inert-adapter gate never ran. That question is
|
||||
open and moves to the restart.
|
||||
|
||||
Model-agnostic lessons from this investigation are in
|
||||
[`training-throughput-playbook.md`](training-throughput-playbook.md); the
|
||||
probes are at [`scripts/training-probes/`](../../scripts/training-probes/).
|
||||
What follows is Gemma-4-specific.
|
||||
|
||||
### 6.1 Where the step time goes
|
||||
|
||||
Real checkpoint, GPU0, `attn_implementation="sdpa"`, PEFT + gradient
|
||||
checkpointing + the chunked CE, fwd+bwd, best-of-2 after warmup:
|
||||
|
||||
| shape | time | peak |
|
||||
|---|---:|---:|
|
||||
| 2 × 2,048 | 1.776 s | 53.2 GiB |
|
||||
| 2 × 8,192 | 11.570 s | 62.3 GiB |
|
||||
| 2 × 16,384 | **35.017 s** | 76.6 GiB |
|
||||
|
||||
Fitting `t(w) = A·w + B·w²` over all three (per-sequence `w`, batch 2):
|
||||
|
||||
A = 6.8715e-04 s/token B = 8.8509e-08 s/token²
|
||||
|
||||
| w | predicted | measured | linear | quadratic | quad share |
|
||||
|---:|---:|---:|---:|---:|---:|
|
||||
| 2,048 | 1.779 | 1.776 | 1.407 | 0.371 | 20.9% |
|
||||
| 8,192 | 11.569 | 11.570 | 5.629 | 5.940 | 51.3% |
|
||||
| 16,384 | 35.017 | 35.017 | 11.258 | 23.759 | **67.8%** |
|
||||
|
||||
**Two terms, three points, residuals under 3 ms across an 8× range.** No fixed
|
||||
per-batch term was needed, which refutes the launch-bound hypothesis outright —
|
||||
~3,840 expert-GEMM launches per forward are not the cost.
|
||||
|
||||
Independently, the profiler kernel table (device rows only — see playbook §3.4):
|
||||
|
||||
| device kernel | ms | of step |
|
||||
|---|---:|---:|
|
||||
| `fmha_cutlassB_bf16_aligned_128x64_k65536_sm80` (attn BWD) | 16,144.6 | 46.1% |
|
||||
| `fmha_cutlassF_bf16_aligned_32x128_gmem_sm80` (attn FWD) | 6,691.2 | 19.1% |
|
||||
| `cutlass_80_tensorop_bf16_s16816gemm` ×3 (dense GEMM) | 2,774.0 | 7.9% |
|
||||
| elementwise / vectorized / unrolled ×9 | 4,787.4 | 13.7% |
|
||||
| gather / Memcpy DtoD / dropout | 951.6 | 2.7% |
|
||||
| **attention total** | **22,835.8** | **65.2%** |
|
||||
|
||||
**Scaling fit says 67.8% quadratic; kernel table says 65.2% attention. Two
|
||||
independent methods, 2.6 points apart.**
|
||||
|
||||
### 6.2 ⚠ The attention kernels are Ampere, on a Blackwell card
|
||||
|
||||
`fmha_cutlass*_sm80` on sm_120. There is no Blackwell-tuned attention kernel in
|
||||
this path at all, and the forward is additionally on `gmem` — the
|
||||
global-memory fallback tier of the memory-efficient backend, selected when the
|
||||
working set will not fit in shared memory.
|
||||
|
||||
This is the mechanism behind the 100%-SM / 27-TFLOPS / 304-TFLOPS-capable
|
||||
reading: the chip is saturated running a kernel generation behind on the
|
||||
dominant cost centre.
|
||||
|
||||
The candidate fix is a purpose-built kernel for this architecture's mixed
|
||||
256/512 head-dim split — `zzhhjjj/gemma-triton-flash-attn`
|
||||
(`register_triton_attention()`, then `_attn_implementation = "triton_gqa"`),
|
||||
reported 9.23× over SDPA at N=16K D=256 SWA and 2.94× fwd+bwd at D=512.
|
||||
`flex_attention` + `BlockMask` is the no-new-dependency alternative.
|
||||
|
||||
⚠ **Prefer a UNIFORM backend over a per-layer split.** vLLM special-cased this
|
||||
exact mixed-head-dim architecture and measured mixed backends **8% slower** than
|
||||
uniform. And `attn_implementation` is all-or-nothing at `from_pretrained` /
|
||||
`set_attn_implementation` — per-layer routing requires a custom function
|
||||
registered on `ALL_ATTENTION_FUNCTIONS` branching on `module.head_dim` /
|
||||
`sliding_window`.
|
||||
|
||||
⚠ **FA2 is not available for this model**: it caps head_dim at 256 and the 5
|
||||
global layers are at 512. FA3 is Hopper-only. Do not bet on FA4 on sm_120.
|
||||
|
||||
### 6.3 Masking is CORRECT — and padding is what costs
|
||||
|
||||
Band structure asserted directly against the real config at n=16,384:
|
||||
|
||||
sliding_attention max 1,024 allowed/row, saturates at row 1,023 PASS
|
||||
|
||||
Constraints were **not** silently dropped; the 25 sliding layers were genuinely
|
||||
windowed. Run-01 was training the model we intended.
|
||||
|
||||
The same probe found the mechanism nobody had measured:
|
||||
|
||||
| 2D mask supplied | `full_attention` mask returned |
|
||||
|---|---|
|
||||
| `None` | **`None`** → `is_causal` fast path AVAILABLE |
|
||||
| all-ones (no padding) | **`None`** → `is_causal` fast path AVAILABLE |
|
||||
| right-padded (what `collate_mixed` emits) | 4D `16384²` → **fast path LOST** |
|
||||
|
||||
**Padding is what pins the 5 global layers to an explicit mask.** The 25
|
||||
sliding layers get a 4D tensor either way — `sdpa_attention_forward` sets
|
||||
`is_causal=True` only when `attention_mask is None`, and a 1024 window cannot
|
||||
be expressed as `is_causal`.
|
||||
|
||||
Isolated, same width, only the mask differing:
|
||||
|
||||
2 × 16,384, no padding 35.244 s 26,048 loss targets
|
||||
2 × 16,384, 50% pad on row 1 38.567 s 19,640 loss targets
|
||||
|
||||
**9.4% slower for 24% less work.**
|
||||
|
||||
### 6.4 The corpus is 29.9% padding — and bucketing is the biggest win available
|
||||
|
||||
Measured off the preserved encode cache in true `SequentialSampler` order:
|
||||
|
||||
records 20,982 (3,583 rp-dialogue / 12,003 prose-chunk / 5,396 actual-play)
|
||||
seq len min/mean/max 142 / 2,752 / 16,384
|
||||
micro-batches (mb=2) 10,491
|
||||
real tokens 57,733,156
|
||||
padded tokens 82,337,318
|
||||
PADDING WASTE 29.9%
|
||||
mb width p50/p90/p99 2,092 / 10,634 / 16,341
|
||||
micro-batches at 16,384 3 of 10,491 (0.0%)
|
||||
|
||||
⚠ Note the last line against §6.1: **the 2 × 16,384 benchmark shape occurs in
|
||||
three micro-batches out of 10,491.** Weighted over the real distribution the
|
||||
quadratic share is ~51%, not 67.8%.
|
||||
|
||||
**Bucket-to-pair, shuffle-to-mix** (brokkr's design, validated on measured
|
||||
lengths — form micro-batches within length buckets, then shuffle the resulting
|
||||
*micro-batches* globally):
|
||||
|
||||
| bucket | waste | predicted step | zero-pad mb | roots/accum window |
|
||||
|---:|---:|---:|---:|---:|
|
||||
| current | 29.9% | 44.3 s → 16.13 h | 0.1% | 3.68 |
|
||||
| **2** | **0.0%** | **28.6 s → 10.40 h** | **78.3%** | 3.56 |
|
||||
| 8 | 0.0% | 28.6 s → 10.41 h | 65.3% | 3.54 |
|
||||
| 32 | 0.1% | 28.6 s → 10.43 h | 41.9% | 3.55 |
|
||||
| 128 | 0.6% | 28.8 s → 10.51 h | 14.7% | 3.55 |
|
||||
| 512 | 2.4% | 29.7 s → 10.82 h | 4.1% | 3.61 |
|
||||
|
||||
**≥35.5% wall clock, no kernel work, no new dependency, peak memory unchanged.**
|
||||
|
||||
Two findings that changed the design:
|
||||
|
||||
- **Bucket size is not a diversity knob.** Roots per accumulation window are
|
||||
flat at 3.54–3.61 across a 256× range. The global micro-batch shuffle does
|
||||
all the mixing. Use the tightest bucket.
|
||||
- **35.5% is a floor.** Zero-pad micro-batches go 0.1% → 78.3%, which puts the
|
||||
5 global layers back on `is_causal` for most of the run (§6.3). The cost
|
||||
model does not capture that. Direction certain, magnitude not yet measured at
|
||||
representative shapes.
|
||||
|
||||
⚠ **Source-homogeneity is a real hazard here** — length correlates hard with
|
||||
root (kvasir short, chunked RP windows long), so length-homogeneous batches are
|
||||
root-homogeneous batches. The global micro-batch shuffle is what prevents an
|
||||
accumulation window drawing its whole gradient from one source. It is
|
||||
load-bearing, not decoration.
|
||||
|
||||
### 6.5 The chunked CE is fine — do not swap it
|
||||
|
||||
2 × 16,384 CE forward 374 ms of 35.329 s = 1.1%
|
||||
2 × 4,096 CE forward 93 ms of 4.387 s = 2.1%
|
||||
|
||||
⚠ **Forward only** — the `torch.utils.checkpoint` recompute runs inside
|
||||
`.backward()`, outside the timing window. Even at 3× it is ~3%.
|
||||
|
||||
`liger-kernel` fused linear CE is a ~1–3% lever on this shape. §2's finding
|
||||
stands unchanged: chunking is what makes seq 16384 *reachable*, and it is not
|
||||
what makes it slow.
|
||||
|
||||
### 6.6 MoE is ~8% — stop optimising it
|
||||
|
||||
Dense GEMM is 7.9% of the step, confirming the earlier decomposition bound of
|
||||
~10% from the kernel side.
|
||||
|
||||
On `grouped_mm`: **the trace does not adjudicate it.** Run-01 was relaunched on
|
||||
`eager`, so the profile shows the *default* path — 25,463 `aten::mm` dispatches
|
||||
in one fwd+bwd, far more than the ~90 a grouped path would produce, so the
|
||||
default is per-expert sequential. Whether the flag changes that when set is a
|
||||
different measurement and was not run. At 7.9% it is not worth running.
|
||||
|
||||
### 6.7 Restart parameters for round 2
|
||||
|
||||
**Do not relaunch without the sampler change.** It is the only lever that wins
|
||||
under every branch of the diagnosis.
|
||||
|
||||
1. **Implement bucket-to-pair + shuffle-to-mix** in the harness, tightest
|
||||
bucket, global micro-batch shuffle. Expected ~16.1 h → ~10.4 h or better.
|
||||
2. **Re-assert the mask band structure** after the sampler change —
|
||||
`scripts/training-probes/step0_mask.py`, 30 s, no GPU. The sampler touches
|
||||
batch composition, which is what drives mask construction.
|
||||
3. **Resume with `/tank/erp-tune/resume-run-01.sh`, NEVER the original launch
|
||||
command** — it begins `rm -rf /tank/erp-tune/run-01` and would destroy the
|
||||
609 MB encode cache (2.5 min to reuse, ~4.3 h to rebuild). ⚠ A sampler change
|
||||
alters record *order*, not encoding, so the cache stays valid — but bump
|
||||
`encode_version` if anything upstream of `input_ids` changes.
|
||||
4. **Run the `lora_B` inert-adapter gate at step 100.** It never ran in round 1.
|
||||
Norm every `lora_B` tensor in the checkpoint: all-non-zero = real, all-zero =
|
||||
INERT (kill the run), partial = module-selection problem. This is the one
|
||||
failure that stays invisible until brokkr's acceptance gate reports
|
||||
base-identical numbers.
|
||||
5. **The corpus override is ONE RUN ONLY** (`operator-2026-08-25-rnd-run`). A
|
||||
second run needs a second operator grant.
|
||||
6. **Attention backend is round 2's second lever**, gated on an A/B on the
|
||||
replica — not on argument. It can run while the tuned job trains.
|
||||
|
||||
⚠ GPU0 is currently **reserved and idle** by operator instruction; `sec` /
|
||||
mog-sec remains down. The window is still open, so
|
||||
`playbooks/ana-ml2-training-window-close.yaml` has NOT been run.
|
||||
|
||||
@@ -0,0 +1,359 @@
|
||||
# Training throughput playbook — how to find where the step time went
|
||||
|
||||
_Sibling to [`model-quantization-playbook.md`](model-quantization-playbook.md).
|
||||
That one is for making a model small; this one is for making a training run
|
||||
fast. Same contract: **model-agnostic lessons live here, model-specific ones
|
||||
stay in the per-model artifact and link up.**_
|
||||
|
||||
First written 2026-08-24 out of the Gemma-4 26B-A4B ERP/RP tune, which ran at
|
||||
8.6% MFU and cost a four-model frontier panel and most of a night to explain.
|
||||
The worked example in §7 is that run. The lessons above it are not about
|
||||
Gemma-4.
|
||||
|
||||
> **Read this before hypothesising about kernels.** The single most expensive
|
||||
> failure in that investigation was not a wrong hypothesis. It was *four
|
||||
> people, including four frontier models, reasoning confidently from
|
||||
> arithmetic instead of spending ten minutes on a measurement that settled
|
||||
> it.* Two of the panel's conclusions were retracted by their own authors
|
||||
> within the hour. Every retraction was a derivation; every survivor was a
|
||||
> measurement.
|
||||
|
||||
---
|
||||
|
||||
## 1. The 10-minute triage — do this FIRST, always
|
||||
|
||||
Before you profile, before you read a modelling file, before you ask anyone:
|
||||
**measure the step's scaling curve.** Three sequence lengths, fixed batch,
|
||||
fwd+bwd, best-of-2 after a warmup.
|
||||
|
||||
t(w) = A·w + B·w² w = per-sequence length
|
||||
|
||||
Fit two parameters to three points. The residuals tell you which regime you
|
||||
are in, and the regime tells you which lever exists:
|
||||
|
||||
| observed `t(4w)/t(w)` | regime | the lever |
|
||||
|---|---|---|
|
||||
| ~4× | **linear** — per-token work dominates | fewer tokens; fused elementwise |
|
||||
| ~16× | **quadratic** — attention dominates | attention backend / kernel |
|
||||
| ~1× | **launch-bound** — fixed per-batch cost | CUDA graphs, `torch.compile`, bigger batch |
|
||||
|
||||
**If the two-term fit closes with residuals under ~1%, launch-bound is
|
||||
refuted.** You did not need a constant term, so there is not a meaningful one.
|
||||
This is the cheapest possible refutation of the most seductive wrong answer,
|
||||
and it costs one extra data point.
|
||||
|
||||
### ⚠ 1.1 ⭐⭐ Three points minimum. A two-point fit with three plausible terms is UNDETERMINED
|
||||
|
||||
This is the lesson that cost the most. A two-point fit over {quadratic,
|
||||
linear, fixed} has infinitely many solutions, and which one you land on is
|
||||
decided by whichever per-step number you happened to quote. In the worked
|
||||
example a peer produced **two confident, opposite conclusions from the same
|
||||
method inside an hour** — "attention is ~5 s of 35" and then "attention is
|
||||
21–33 s of 35" — because the inputs drifted between attempts.
|
||||
|
||||
Three points, two parameters, and check the residuals. If they do not close,
|
||||
you have a third term and you need a fourth point.
|
||||
|
||||
### ⚠ 1.2 ⭐⭐ Benchmark the shape you RUN, not the worst case you can construct
|
||||
|
||||
The quadratic share is **strongly shape-dependent** — in the worked example it
|
||||
ran 20.9% at w=2,048, 51.3% at w=8,192, 67.8% at w=16,384. A synthetic
|
||||
`max_seq_len` benchmark therefore measures the shape where attention looks
|
||||
worst, and generalising from it overstates the attention prize by ~1.3×.
|
||||
|
||||
Get the real distribution off the encode cache and weight by it:
|
||||
|
||||
E[t] = A·E[w] + B·E[w²]
|
||||
|
||||
**`E[w²]` is not `E[w]²`.** For a bimodal length distribution they can differ
|
||||
by 2× or more, and a quadratic term is dominated by the rare long batches that
|
||||
an `E[w]²` shortcut averages away. In the worked example `E[n²]/E[n]²` was
|
||||
**2.08**.
|
||||
|
||||
Sanity check the weighted prediction against the observed `s/it` before you
|
||||
trust any of it.
|
||||
|
||||
---
|
||||
|
||||
## 2. The reference probe set
|
||||
|
||||
Committed at [`scripts/training-probes/`](../../scripts/training-probes/).
|
||||
Run them in this order; each is minutes and none needs the real checkpoint
|
||||
except the profiler.
|
||||
|
||||
| probe | what it settles | needs GPU? |
|
||||
|---|---|---|
|
||||
| `step0_mask.py` | mask band structure + which layers keep the fast path | no |
|
||||
| `step2_padding.py` | padding waste, length distribution, CE chunk sizing | no |
|
||||
| `step_bucket.py` | bucketing gain, bucket-size sweep, root diversity | no |
|
||||
| `step1_profile.py` | scaling fit, padding penalty, CE wall clock, kernel table | yes |
|
||||
|
||||
`step1_profile.py` loads the real model but reuses the harness's own
|
||||
`discover_target_modules` and `compute_loss`, so it measures the thing that
|
||||
actually runs rather than a re-implementation. **Keep that property when you
|
||||
adapt it** — a probe that reimplements the training step measures the probe.
|
||||
|
||||
---
|
||||
|
||||
## 3. The recurring landmines
|
||||
|
||||
### 3.1 ⭐⭐ Right-padding is a compute tax AND a backend tax
|
||||
|
||||
Everyone knows padding wastes tokens. The second effect is the one that gets
|
||||
missed: **an explicit padding mask can knock fast-path-eligible layers off
|
||||
`is_causal`.**
|
||||
|
||||
`scaled_dot_product_attention` takes `is_causal=True` **or** an `attn_mask`,
|
||||
never both usefully. HF sets `is_causal=True` only when `attention_mask is
|
||||
None`. Right-pad a batch and you hand it a 2D mask, it materialises a 4D
|
||||
tensor, and every layer that could have taken the clean causal route now takes
|
||||
a masked dense one.
|
||||
|
||||
Measured, same width, same `n`, only the mask differing:
|
||||
|
||||
no padding 35.244 s 26,048 loss targets
|
||||
50% pad on one row 38.567 s 19,640 loss targets
|
||||
|
||||
**9.4% slower for 24% less work.** Verify this on your own stack with
|
||||
`step0_mask.py` — it prints whether `create_causal_mask` returns `None` or a
|
||||
tensor for each mask case.
|
||||
|
||||
### 3.2 ⭐⭐ Length-bucket to PAIR, shuffle micro-batches to MIX — and the bucket should be TIGHT
|
||||
|
||||
Naive length-bucketing has a real hazard: length correlates with data source,
|
||||
so length-homogeneous batches are **source-homogeneous batches**, and an
|
||||
accumulation window can end up drawing its entire gradient from one root.
|
||||
|
||||
The fix costs nothing: **form micro-batches within length buckets, then
|
||||
shuffle the resulting micro-batches globally.** Padding efficiency is a
|
||||
property of the pairing alone, so all of the saving survives the shuffle.
|
||||
|
||||
**The non-obvious part — bucket size is not a diversity knob.** Measured
|
||||
across a 256× range of bucket sizes, roots per accumulation window stayed flat
|
||||
at 3.54–3.61 (against 3.68 for a pure shuffle). The *global micro-batch
|
||||
shuffle* does all of the mixing; the bucket contributes nothing to diversity
|
||||
and only costs padding. So use the tightest bucket you can — which in the
|
||||
limit is a full length sort.
|
||||
|
||||
| bucket | padding waste | zero-pad micro-batches | roots/window |
|
||||
|---|---|---|---|
|
||||
| current (shuffle) | 29.9% | 0.1% | 3.68 |
|
||||
| 2 | 0.0% | **78.3%** | 3.56 |
|
||||
| 32 | 0.1% | 41.9% | 3.55 |
|
||||
| 512 | 2.4% | 4.1% | 3.61 |
|
||||
|
||||
Note the `zero-pad micro-batches` column — that is §3.1 compounding. A tight
|
||||
bucket does not merely cut tokens, it puts most batches back on the causal
|
||||
fast path.
|
||||
|
||||
**Peak memory does not rise.** `padded = batch × max(len)`, so one long record
|
||||
forces a full-width batch regardless of its partner. Bucketing pairs long
|
||||
records *with each other*, which roughly halves the number of worst-case
|
||||
batches.
|
||||
|
||||
### 3.3 ⭐⭐ Check the kernel GENERATION, not just the backend name
|
||||
|
||||
The backend name (`EFFICIENT_ATTENTION`, `FLASH_ATTENTION`, …) is not the whole
|
||||
story. Read the actual kernel symbols out of the profiler:
|
||||
|
||||
fmha_cutlassF_bf16_aligned_32x128_gmem_sm80
|
||||
fmha_cutlassB_bf16_aligned_128x64_k65536_sm80
|
||||
^^^^
|
||||
|
||||
`sm80` is **Ampere**. Those were running on an sm_120 Blackwell card, on the
|
||||
dominant cost centre of the step. A backend can be "selected correctly" and
|
||||
still be a generation behind, and nothing in the config surface tells you.
|
||||
|
||||
Also read the variant suffix: `gmem` on the forward kernel is the
|
||||
**global-memory fallback tier** of the memory-efficient path, chosen when the
|
||||
working set will not fit in shared memory. Wrong backend *and* that backend's
|
||||
slow path.
|
||||
|
||||
### 3.4 ⭐ `key_averages()` double-counts — use device-kernel rows only
|
||||
|
||||
`torch.profiler`'s `key_averages()` table lists both the ATen op and the CUDA
|
||||
kernel it launched, each carrying the same `self_device_time_total`. Summing
|
||||
the whole table gives you roughly **2× the real step time**.
|
||||
|
||||
The tell is exact equality between an `aten::` row and a kernel row:
|
||||
|
||||
aten::_efficient_attention_backward 30 16144.6
|
||||
fmha_cutlassB_bf16_aligned_128x64_k65536 30 16144.6
|
||||
|
||||
Filter to device kernels (`void …`, `fmha_…`, `cutlass::…`, `Memcpy…`) and
|
||||
sanity-check the total against the measured wall clock. In the worked example
|
||||
the filtered total came to 89.5% of the step, which is the right shape; the
|
||||
unfiltered total came to 202%.
|
||||
|
||||
### 3.5 ⭐ Time the loss forward AND account for its backward recompute separately
|
||||
|
||||
If the loss head is gradient-checkpointed, a CUDA-event window around the
|
||||
forward loop measures **half the story at best** — the recompute happens inside
|
||||
`.backward()`, outside your window.
|
||||
|
||||
State the caveat explicitly when you report the number. In the worked example
|
||||
the CE forward measured 374 ms of a 35.3 s step (1.1%); even at 3× for
|
||||
recompute-plus-backward it is ~3%, which was enough to kill a proposed
|
||||
dependency swap — but "1.1%" alone would have been an unearned claim.
|
||||
|
||||
### 3.6 ⭐⭐ Assert mask band structure directly; never infer it from performance
|
||||
|
||||
`transformers` can **silently skip mask creation** and pass
|
||||
`attention_mask=None` when a mask function is not registered. If that fires on
|
||||
a sliding-window model, the windowed layers do full causal attention — not a
|
||||
speed bug, **a different model from the one you will serve**.
|
||||
|
||||
There is a tempting alibi: "if constraints were dropped we would be on the
|
||||
fast path and fast; we are slow, therefore correct." It is decent evidence and
|
||||
it is not an assertion. Materialise the mask once and count allowed positions
|
||||
per row:
|
||||
|
||||
sliding_attention max 1,024 allowed/row, saturates at row 1,023 PASS
|
||||
|
||||
Thirty seconds, on CPU, no weights. Do it before every run that changes the
|
||||
masking path, and before believing any optimisation result.
|
||||
|
||||
### 3.7 ⭐ "Bit-identical output from a different backend" — ask *could this have disagreed?*
|
||||
|
||||
A backend flag that produces `max_abs_diff == 0.0` against the reference is
|
||||
either (a) legitimately the same GEMMs behind a different launcher, or (b) a
|
||||
flag that never took. **Argument cannot separate these** — in the worked
|
||||
example three frontier models split 2–1 on it and the majority was not
|
||||
obviously right.
|
||||
|
||||
Do not resolve it by vote. **Count kernel launches.** A per-expert loop leaves
|
||||
`n_experts` dispatches per layer visible; a grouped path leaves one. That is
|
||||
unambiguous and falls out of a trace you are running anyway.
|
||||
|
||||
Related trap: **a trace of the default path does not test the flag.** If the
|
||||
run was relaunched without the flag set, the profile tells you what the default
|
||||
does and nothing about the flag. Say so rather than over-claiming.
|
||||
|
||||
### 3.8 ⭐ MFU is a denominator argument waiting to happen — report the decomposition instead
|
||||
|
||||
MFU invites an unwinnable fight about what counts as a FLOP (active vs dense
|
||||
params for MoE, whether checkpoint recompute counts, whether frozen-base
|
||||
skipped GEMMs count). That fight consumed an hour of the worked example and
|
||||
produced nothing.
|
||||
|
||||
Report these **beside** MFU, not instead of it:
|
||||
|
||||
- tokens/s, and **real (unpadded) tokens/s** separately
|
||||
- achieved hardware FLOPs straight from the profiler
|
||||
- the time decomposition (attention / GEMM / elementwise / other)
|
||||
|
||||
Then the denominator stops mattering.
|
||||
|
||||
**The reading that actually diagnosed it** was not an MFU number at all:
|
||||
|
||||
> 100% SM utilisation at 279–292 W, running 27 TFLOPS, on a card that does
|
||||
> 304 TFLOPS on a dense GEMM at the same power.
|
||||
|
||||
**SM-busy, tensor-core-idle.** The chip is fully occupied doing work that is
|
||||
not matrix multiplication. No FLOP-counting convention changes that, and it
|
||||
points straight at the kernel table.
|
||||
|
||||
### 3.9 Frozen-base LoRA is ~4ND, not ~6ND — and the arithmetic intensity does NOT drop
|
||||
|
||||
A claim that circulated and was wrong: "frozen-base LoRA has structurally lower
|
||||
arithmetic intensity, so a dense-GEMM ceiling is unreachable in principle."
|
||||
|
||||
The correct accounting: forward is 2ND, input-gradient backward through the
|
||||
frozen weights is 2ND, and only the weight-gradient (~2ND) is skipped. So
|
||||
**~4ND against ~6ND — two-thirds of the work, at the same arithmetic intensity
|
||||
per remaining GEMM.** You do fewer GEMMs; the ones you do are exactly as dense.
|
||||
|
||||
Gradient checkpointing is a separate, real ~⅓ recompute tax. Account for it
|
||||
separately rather than folding it into an intensity story.
|
||||
|
||||
---
|
||||
|
||||
## 4. Panel / consult discipline for perf work
|
||||
|
||||
Perf investigations are unusually good at generating confident wrong answers,
|
||||
because the arithmetic is easy and the ground truth is expensive. Specific
|
||||
guards, learned the hard way:
|
||||
|
||||
- **Every arm's claim gets a measurement or an expiry date.** In the worked
|
||||
example the panel produced four self-retractions in ninety minutes. The
|
||||
measurements produced zero.
|
||||
- **Treat cross-arm agreement as weak evidence.** Ask arms to attack a
|
||||
hypothesis rather than extend it; agreement among similarly-primed readers of
|
||||
the same artifact is not independent confirmation.
|
||||
- **A dispute about what a specific dispatcher does is a question of fact.**
|
||||
Do not put it to a panel. Instrument it.
|
||||
- **When an arm says "you missed X," check what they read.** If your settled
|
||||
artifact was not on their reading list, the "miss" is usually
|
||||
restatement-of-a-settled-prior, not a genuine gap.
|
||||
|
||||
---
|
||||
|
||||
## 5. Superseded claims — do not follow these
|
||||
|
||||
| claim | status | replaced by |
|
||||
|---|---|---|
|
||||
| "Explicit mask → `EFFICIENT_ATTENTION`" is over-specific; Blackwell defaults to `CUDNN_ATTENTION` | **WRONG** (2026-08-24) | Measured: sm_120 selects `fmha_cutlass*_sm80`, i.e. `EFFICIENT_ATTENTION`. The original claim was right. |
|
||||
| Attention's quadratic share is ~5 s of a 35 s step | **WRONG** (2026-08-24) | Measured 22.8 s / 65.2% at w=16,384; 67.8% by independent scaling fit |
|
||||
| Frozen-base LoRA has structurally lower arithmetic intensity | **WRONG** (2026-08-24) | ~4ND vs 6ND at unchanged intensity — see §3.9 |
|
||||
| The chunked CE is a 2–5× under-estimated cost centre | **WRONG** (2026-08-24) | Measured 1.1% of step forward, ≲3% with recompute |
|
||||
| `attn_implementation="flash_attention_2"` is the per-layer lever | **NOT A FLAG** (2026-08-24) | All-or-nothing at `from_pretrained`; per-layer needs a custom fn on `ALL_ATTENTION_FUNCTIONS`. FA2 also caps head_dim at 256. |
|
||||
| Bucket size ~256 is needed to preserve source diversity | **UNNECESSARY** (2026-08-24) | Diversity is flat in bucket size; the global micro-batch shuffle does that work — see §3.2 |
|
||||
|
||||
## 6. Measured negatives — don't re-chase
|
||||
|
||||
- **Fused MoE kernel (`grouped_mm`) as the throughput fix.** Measured 0.9%
|
||||
*slower* than the Python loop and bit-identical. Independently, dense GEMM is
|
||||
only 7.9% of the step, so the whole category is capped near 10%.
|
||||
- **CUDA graphs / `torch.compile` over the expert loop.** The two-term scaling
|
||||
fit closed without a constant term, so there is no meaningful fixed per-batch
|
||||
cost to amortise. ~3,840 expert-GEMM launches per forward are not what you
|
||||
are paying for.
|
||||
- **`liger-kernel` fused linear CE.** Real and correct, but a ~1–3% lever on
|
||||
this shape. Not a project.
|
||||
- **FlashAttention-4 on sm_120.** Public reports are sour — one measurement of
|
||||
1.07× over FA2, and an sm_120 patch people could not get working that fell
|
||||
back to torch SDPA. Do not bet a round on it.
|
||||
|
||||
---
|
||||
|
||||
## 7. Worked example — Gemma-4 26B-A4B ERP/RP tune, 2026-08-24
|
||||
|
||||
Model-specific detail lives in
|
||||
[`gemma4-erp-tune-sizing.md`](gemma4-erp-tune-sizing.md) §6. The short version,
|
||||
because the *shape* of the investigation is the transferable part:
|
||||
|
||||
**Symptom.** 8.6% MFU, ~35–46 s/it, 1,312 steps, ~13.9 h ETA.
|
||||
|
||||
**What the panel produced.** Four frontier arms plus an orchestrator, over
|
||||
ninety minutes: a sliding-window hypothesis, a retraction of it, a retraction
|
||||
of the retraction, a correctness scare that resolved itself, two mutually
|
||||
contradictory readings of one dispatcher, and four self-corrections.
|
||||
|
||||
**What settled it, in about twenty minutes of GPU time:**
|
||||
|
||||
scaling fit (3 points, 2 params, residuals <3 ms over an 8× range)
|
||||
A = 6.87e-4 s/token B = 8.85e-8 s/token²
|
||||
quadratic share: 20.9% @ w=2,048 → 67.8% @ w=16,384
|
||||
|
||||
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%
|
||||
|
||||
Two independent methods, 2.6 points apart. **Attention was the answer, on
|
||||
Ampere-generation kernels, with the forward on a global-memory fallback tier.**
|
||||
|
||||
**The largest actionable win was not the attention kernel.** It was a sampler
|
||||
change — bucket-to-pair, shuffle-to-mix — worth 29.9% of tokens and ~35.5% of
|
||||
wall clock, with no new dependency, no kernel work, and unchanged peak memory.
|
||||
It also wins under *every* branch of the diagnosis, which is why it was
|
||||
recommended while the rest was still unresolved.
|
||||
|
||||
**The transferable ordering:**
|
||||
|
||||
1. Assert correctness (mask band structure). Everything downstream assumes it.
|
||||
2. Scaling curve. Names the regime in ten minutes.
|
||||
3. Kernel table. Names the cost centre.
|
||||
4. Data-side levers first (padding, bucketing) — they need no dependency and
|
||||
they multiply into every other cost.
|
||||
5. Kernel/backend levers last, gated on 2 and 3.
|
||||
@@ -0,0 +1,53 @@
|
||||
# Training throughput probes
|
||||
|
||||
Instruments for finding where a training step's time actually went. Written
|
||||
2026-08-24 during the Gemma-4 26B-A4B ERP/RP tune investigation; the lessons
|
||||
they produced live in
|
||||
[`docs/pfi/training-throughput-playbook.md`](../../docs/pfi/training-throughput-playbook.md).
|
||||
|
||||
**These are diagnostic instruments, not production code.** They hard-code paths
|
||||
for that run. Adapt the constants at the top; keep the measurement design.
|
||||
|
||||
## The probes
|
||||
|
||||
| script | settles | GPU | runtime |
|
||||
|---|---|---|---|
|
||||
| `step0_mask.py` | mask band structure; which layers keep the `is_causal` fast path | no | ~30 s |
|
||||
| `step2_padding.py` | padding waste, length distribution, CE chunk sizing | no | ~2 min |
|
||||
| `step_bucket.py` | bucketing gain, bucket-size sweep, source diversity | no | ~3 min |
|
||||
| `step1_profile.py` | scaling fit, padding penalty, CE wall clock, kernel table | **yes** | ~15 min |
|
||||
|
||||
Run in that order. Only the last needs the real checkpoint, and it wants an
|
||||
idle card — it loads ~48 GiB and peaks near 77 GiB at `2 × 16,384`.
|
||||
|
||||
## Design rules worth preserving when you adapt these
|
||||
|
||||
**`step1_profile.py` reuses the harness's own `discover_target_modules` and
|
||||
replicates its `compute_loss` byte-for-byte** rather than re-implementing the
|
||||
step. A probe that reimplements the training step measures the probe. If you
|
||||
port this, keep the import from the real harness.
|
||||
|
||||
**`step0_mask.py` needs no weights and no GPU** — SDPA backend selection and
|
||||
mask construction depend on shapes, dtype and mask presence, not on weight
|
||||
values. That is what makes the correctness assertion cheap enough to run before
|
||||
every job.
|
||||
|
||||
**The scaling test takes three points, not two.** Two points over three
|
||||
plausible terms (quadratic, linear, fixed-per-batch) is underdetermined; see
|
||||
playbook §1.1 for the hour that cost.
|
||||
|
||||
**`step_bucket.py` sweeps bucket size deliberately.** The first version
|
||||
re-sorted within each bucket, which silently collapsed every bucket size to a
|
||||
full global sort and made the sweep a no-op. If you change the pairing logic,
|
||||
check that the sweep still varies something.
|
||||
|
||||
## Raw evidence
|
||||
|
||||
`step1-profile-output-2026-08-24.txt` is the unedited output of the run the
|
||||
playbook's numbers come from — scaling points, padding penalty, CE timing, and
|
||||
the full `key_averages()` kernel table. Kept so the claims can be re-derived
|
||||
rather than taken on faith.
|
||||
|
||||
⚠ That table **double-counts**: `key_averages()` lists both the ATen op and the
|
||||
CUDA kernel it launched, each carrying the same self device time. Sum device
|
||||
kernel rows only. See playbook §3.4.
|
||||
@@ -0,0 +1,75 @@
|
||||
"""Step 0 - assert the sliding mask band structure, and record which path
|
||||
mask creation actually takes under the run config (attn_implementation=sdpa).
|
||||
|
||||
Correctness gate: transformers can SILENTLY skip mask creation and pass
|
||||
attention_mask=None, which would make the 25 sliding layers do full causal
|
||||
attention - a different model from the one vLLM serves. This converts
|
||||
"probably fine because we are slow" into a measurement.
|
||||
|
||||
CPU only. No weights. No GPU.
|
||||
"""
|
||||
import torch
|
||||
from transformers import AutoConfig
|
||||
from transformers.masking_utils import (
|
||||
create_causal_mask, create_sliding_window_causal_mask,
|
||||
)
|
||||
|
||||
MODEL = "/tank/aimodels/gemma4-26b-a4b-it-heretic-bf16"
|
||||
N = 16384
|
||||
W = 1024
|
||||
PAD = " " + " " * 20
|
||||
|
||||
cfg = AutoConfig.from_pretrained(MODEL)
|
||||
text = cfg.get_text_config()
|
||||
text._attn_implementation = "sdpa"
|
||||
print("sliding_window %s" % text.sliding_window)
|
||||
print("layers %d (%d sliding / %d full)" % (
|
||||
len(text.layer_types),
|
||||
text.layer_types.count("sliding_attention"),
|
||||
text.layer_types.count("full_attention")))
|
||||
print("_attn_implementation %s" % text._attn_implementation)
|
||||
print()
|
||||
|
||||
|
||||
def build(attn_2d, label):
|
||||
batch = attn_2d.shape[0] if attn_2d is not None else 1
|
||||
embeds = torch.zeros(batch, N, 8, dtype=torch.bfloat16)
|
||||
pos = torch.arange(N).unsqueeze(0)
|
||||
kw = dict(config=text, inputs_embeds=embeds, attention_mask=attn_2d,
|
||||
past_key_values=None, position_ids=pos)
|
||||
full = create_causal_mask(**kw)
|
||||
slide = create_sliding_window_causal_mask(**kw)
|
||||
print("--- %s ---" % label)
|
||||
for name, m in (("full_attention", full), ("sliding_attention", slide)):
|
||||
if m is None:
|
||||
print(" %-20s None -> flash / is_causal path AVAILABLE" % name)
|
||||
continue
|
||||
print(" %-20s tensor shape=%s dtype=%s" % (name, tuple(m.shape), m.dtype))
|
||||
allowed = m if m.dtype == torch.bool else (m == 0)
|
||||
per_row = allowed[0, 0].sum(-1)
|
||||
print("%sallowed/row min=%d max=%d mean=%.1f" % (
|
||||
PAD, per_row.min().item(), per_row.max().item(),
|
||||
per_row.float().mean().item()))
|
||||
if name == "sliding_attention":
|
||||
ok = per_row.max().item() <= W
|
||||
print("%sBAND <= %d ? %s" % (PAD, W, "PASS" if ok else "FAIL"))
|
||||
sat = (per_row >= W).nonzero()
|
||||
if sat.numel():
|
||||
print("%ssaturates at row %d" % (PAD, sat[0].item()))
|
||||
else:
|
||||
print("%slast row allows %d of %d (%s)" % (
|
||||
PAD, per_row[-1].item(), N,
|
||||
"causal-full OK" if per_row[-1].item() == N else "UNEXPECTED"))
|
||||
print()
|
||||
|
||||
|
||||
# 1. no 2D mask at all - the "constraints silently dropped" scenario
|
||||
build(None, "attention_mask=None (no padding info)")
|
||||
|
||||
# 2. all-ones 2D mask - equal-length batch, no padding
|
||||
build(torch.ones(2, N, dtype=torch.long), "all-ones 2D (no padding)")
|
||||
|
||||
# 3. REAL right-padded batch - what collate_mixed actually produces
|
||||
real = torch.ones(2, N, dtype=torch.long)
|
||||
real[1, 6000:] = 0
|
||||
build(real, "right-padded 2D (what collate_mixed emits)")
|
||||
@@ -0,0 +1,126 @@
|
||||
========================================================================
|
||||
loading model
|
||||
========================================================================
|
||||
|
||||
Loading weights: 0%| | 0/1013 [00:00<?, ?it/s]
|
||||
Loading weights: 0%| | 2/1013 [00:00<01:19, 12.71it/s]
|
||||
Loading weights: 0%| | 4/1013 [00:00<01:10, 14.40it/s]
|
||||
Loading weights: 2%|▏ | 25/1013 [00:00<00:16, 61.02it/s]
|
||||
Loading weights: 5%|▍ | 48/1013 [00:00<00:12, 79.42it/s]
|
||||
Loading weights: 7%|▋ | 70/1013 [00:00<00:10, 87.32it/s]
|
||||
Loading weights: 9%|▉ | 92/1013 [00:01<00:10, 91.14it/s]
|
||||
Loading weights: 11%|█ | 113/1013 [00:01<00:09, 92.58it/s]
|
||||
Loading weights: 13%|█▎ | 135/1013 [00:01<00:09, 95.13it/s]
|
||||
Loading weights: 15%|█▌ | 157/1013 [00:01<00:08, 97.82it/s]
|
||||
Loading weights: 18%|█▊ | 179/1013 [00:02<00:08, 101.21it/s]
|
||||
Loading weights: 20%|█▉ | 201/1013 [00:02<00:08, 100.37it/s]
|
||||
Loading weights: 22%|██▏ | 223/1013 [00:02<00:07, 102.12it/s]
|
||||
Loading weights: 24%|██▍ | 244/1013 [00:02<00:07, 104.64it/s]
|
||||
Loading weights: 26%|██▋ | 266/1013 [00:02<00:07, 105.68it/s]
|
||||
Loading weights: 28%|██▊ | 288/1013 [00:03<00:06, 106.87it/s]
|
||||
Loading weights: 30%|███ | 308/1013 [00:03<00:05, 122.26it/s]
|
||||
Loading weights: 32%|███▏ | 322/1013 [00:03<00:05, 117.76it/s]
|
||||
Loading weights: 33%|███▎ | 335/1013 [00:03<00:06, 99.26it/s]
|
||||
Loading weights: 35%|███▍ | 354/1013 [00:03<00:06, 100.53it/s]
|
||||
Loading weights: 37%|███▋ | 375/1013 [00:03<00:06, 104.62it/s]
|
||||
Loading weights: 39%|███▉ | 397/1013 [00:04<00:05, 103.34it/s]
|
||||
Loading weights: 41%|████▏ | 419/1013 [00:04<00:05, 106.05it/s]
|
||||
Loading weights: 44%|████▎ | 441/1013 [00:04<00:05, 109.33it/s]
|
||||
Loading weights: 46%|████▌ | 463/1013 [00:04<00:05, 107.04it/s]
|
||||
Loading weights: 48%|████▊ | 485/1013 [00:04<00:04, 107.29it/s]
|
||||
Loading weights: 50%|█████ | 507/1013 [00:05<00:04, 105.14it/s]
|
||||
Loading weights: 52%|█████▏ | 528/1013 [00:05<00:04, 103.18it/s]
|
||||
Loading weights: 54%|█████▍ | 550/1013 [00:05<00:04, 102.05it/s]
|
||||
Loading weights: 56%|█████▋ | 572/1013 [00:05<00:04, 101.82it/s]
|
||||
Loading weights: 59%|█████▊ | 594/1013 [00:05<00:03, 106.80it/s]
|
||||
Loading weights: 61%|██████ | 616/1013 [00:06<00:03, 104.33it/s]
|
||||
Loading weights: 63%|██████▎ | 638/1013 [00:06<00:03, 103.57it/s]
|
||||
Loading weights: 77%|███████▋ | 778/1013 [00:06<00:00, 320.91it/s]
|
||||
Loading weights: 91%|█████████ | 920/1013 [00:06<00:00, 532.70it/s]
|
||||
Loading weights: 100%|██████████| 1013/1013 [00:06<00:00, 152.35it/s]
|
||||
loaded in 9.8s targets=205
|
||||
final_logit_softcapping = 30.0
|
||||
attn_implementation = sdpa
|
||||
|
||||
========================================================================
|
||||
A. SEQUENCE SCALING (no padding - isolates n)
|
||||
========================================================================
|
||||
2 x 2,048 1.776 s kept=3227 peak= 53.2 GiB
|
||||
2 x 8,192 11.570 s kept=13050 peak= 62.3 GiB
|
||||
2 x 16,384 35.017 s kept=25989 peak= 76.6 GiB
|
||||
|
||||
16384 -> 2048 ratio 19.71x (linear ~8x, launch-bound ~1x, quadratic ~64x)
|
||||
16384 -> 8192 ratio 3.03x (linear ~2x, quadratic ~4x)
|
||||
|
||||
========================================================================
|
||||
B. PADDING PENALTY (same real tokens, with vs without pad)
|
||||
========================================================================
|
||||
2 x 16,384 no padding 35.244 s kept=26048 peak= 76.6 GiB
|
||||
2 x 16,384 50% pad on row 1 38.567 s kept=19640 peak= 77.8 GiB
|
||||
|
||||
========================================================================
|
||||
C. ISOLATED CE WALL CLOCK
|
||||
========================================================================
|
||||
2 x 16,384 (CE timed) 35.329 s kept=26210 peak= 76.6 GiB CE=374 ms (1.1%)
|
||||
2 x 4,096 (CE timed) 4.387 s kept=6512 peak= 55.3 GiB CE=93 ms (2.1%)
|
||||
|
||||
========================================================================
|
||||
D. KERNEL TABLE - one fwd+bwd at 2 x 16,384
|
||||
========================================================================
|
||||
USDT:2026-08-24 22:03:51 574811:574811 SyncActivityProfilerHandler.cpp:52] profiler_start
|
||||
USDT:2026-08-24 22:04:27 574811:574811 SyncActivityProfilerHandler.cpp:59] profiler_stop
|
||||
------------------------------------------------------- ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------
|
||||
Name Self CPU % Self CPU CPU total % CPU total CPU time avg Self CUDA Self CUDA % CUDA total CUDA time avg # of Calls
|
||||
------------------------------------------------------- ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------
|
||||
aten::_efficient_attention_backward 0.00% 573.671us 0.00% 1.681ms 56.028us 16.145s 45.71% 16.156s 538.546ms 30
|
||||
fmha_cutlassB_bf16_aligned_128x64_k65536_sm80(PyTorc... 0.00% 0.000us 0.00% 0.000us 0.000us 16.145s 45.71% 16.145s 538.152ms 30
|
||||
aten::_efficient_attention_forward 0.00% 795.660us 0.01% 1.933ms 32.218us 6.691s 18.94% 6.691s 111.519ms 60
|
||||
fmha_cutlassF_bf16_aligned_32x128_gmem_sm80(PyTorchM... 0.00% 0.000us 0.00% 0.000us 0.000us 6.691s 18.94% 6.691s 111.519ms 60
|
||||
aten::mm 0.51% 180.757ms 0.77% 270.268ms 10.614us 3.745s 10.60% 3.745s 147.063us 25463
|
||||
aten::mul 0.13% 45.428ms 0.18% 64.624ms 12.129us 2.721s 7.70% 2.721s 510.606us 5328
|
||||
aten::copy_ 0.06% 19.823ms 92.46% 32.574s 6.100ms 1.977s 5.60% 1.977s 370.189us 5340
|
||||
void cutlass::Kernel2<cutlass_80_tensorop_bf16_s1681... 0.00% 0.000us 0.00% 0.000us 0.000us 1.542s 4.37% 1.542s 656.804us 2348
|
||||
void at::native::elementwise_kernel<128, 2, at::nati... 0.00% 0.000us 0.00% 0.000us 0.000us 1.033s 2.92% 1.033s 545.219us 1894
|
||||
void cutlass::Kernel2<cutlass_80_tensorop_bf16_s1681... 0.00% 0.000us 0.00% 0.000us 0.000us 868.642ms 2.46% 868.642ms 583.373us 1489
|
||||
void at::native::vectorized_elementwise_kernel<4, at... 0.00% 0.000us 0.00% 0.000us 0.000us 787.028ms 2.23% 787.028ms 395.095us 1992
|
||||
void at::native::unrolled_elementwise_kernel<at::nat... 0.00% 0.000us 0.00% 0.000us 0.000us 697.963ms 1.98% 697.963ms 304.521us 2292
|
||||
aten::masked_fill_ 0.01% 3.403ms 0.01% 4.927ms 27.373us 576.927ms 1.63% 576.927ms 3.205ms 180
|
||||
void at::native::vectorized_elementwise_kernel<4, at... 0.00% 0.000us 0.00% 0.000us 0.000us 561.815ms 1.59% 561.815ms 413.403us 1359
|
||||
void at::native::vectorized_elementwise_kernel<4, at... 0.00% 0.000us 0.00% 0.000us 0.000us 469.818ms 1.33% 469.818ms 459.255us 1023
|
||||
aten::add_ 0.01% 3.730ms 0.02% 6.640ms 7.209us 448.665ms 1.27% 448.665ms 487.150us 921
|
||||
void cutlass::Kernel2<cutlass_80_tensorop_bf16_s1681... 0.00% 0.000us 0.00% 0.000us 0.000us 363.206ms 1.03% 363.206ms 394.789us 920
|
||||
aten::add 0.03% 9.745ms 0.04% 14.287ms 10.205us 356.226ms 1.01% 356.226ms 254.447us 1400
|
||||
void at::native::elementwise_kernel<128, 2, at::nati... 0.00% 0.000us 0.00% 0.000us 0.000us 352.705ms 1.00% 352.705ms 1.959ms 180
|
||||
aten::index 0.01% 4.542ms 0.09% 32.242ms 132.682us 352.404ms 1.00% 352.430ms 1.450ms 243
|
||||
void at::native::vectorized_gather_kernel<16, long>(... 0.00% 0.000us 0.00% 0.000us 0.000us 351.768ms 1.00% 351.768ms 1.933ms 182
|
||||
Memcpy DtoD (Device -> Device) 0.00% 0.000us 0.00% 0.000us 0.000us 341.588ms 0.97% 341.588ms 634.922us 538
|
||||
aten::pow 0.06% 21.418ms 0.11% 37.019ms 18.659us 332.072ms 0.94% 493.271ms 248.624us 1984
|
||||
void at::native::vectorized_elementwise_kernel<4, at... 0.00% 0.000us 0.00% 0.000us 0.000us 330.315ms 0.94% 330.315ms 499.720us 661
|
||||
void at::native::vectorized_elementwise_kernel<4, at... 0.00% 0.000us 0.00% 0.000us 0.000us 316.108ms 0.89% 316.108ms 383.161us 825
|
||||
aten::sum 0.01% 5.258ms 0.02% 7.700ms 13.461us 286.461ms 0.81% 286.464ms 500.811us 572
|
||||
aten::native_dropout 0.02% 6.453ms 0.03% 11.139ms 27.168us 258.244ms 0.73% 258.244ms 629.865us 410
|
||||
void at::native::(anonymous namespace)::fused_dropou... 0.00% 0.000us 0.00% 0.000us 0.000us 258.244ms 0.73% 258.244ms 629.865us 410
|
||||
void at::native::unrolled_elementwise_kernel<at::nat... 0.00% 0.000us 0.00% 0.000us 0.000us 239.055ms 0.68% 239.055ms 583.061us 410
|
||||
aten::_index_put_impl_ 0.01% 4.170ms 3.22% 1.134s 7.508ms 235.816ms 0.67% 236.930ms 1.569ms 151
|
||||
void at::native::elementwise_kernel<128, 4, at::nati... 0.00% 0.000us 0.00% 0.000us 0.000us 231.434ms 0.66% 231.434ms 385.081us 601
|
||||
void at::native::vectorized_elementwise_kernel<4, at... 0.00% 0.000us 0.00% 0.000us 0.000us 226.451ms 0.64% 226.451ms 692.511us 327
|
||||
void at::native::elementwise_kernel<128, 4, at::nati... 0.00% 0.000us 0.00% 0.000us 0.000us 224.222ms 0.63% 224.222ms 2.491ms 90
|
||||
void cutlass::Kernel2<cutlass_80_simt_sgemm_64x128_8... 0.00% 0.000us 0.00% 0.000us 0.000us 219.030ms 0.62% 219.030ms 534.219us 410
|
||||
void at::native::elementwise_kernel<128, 4, at::nati... 0.00% 0.000us 0.00% 0.000us 0.000us 209.409ms 0.59% 209.409ms 1.745ms 120
|
||||
aten::div 0.01% 4.911ms 0.02% 6.652ms 13.278us 192.925ms 0.55% 192.925ms 385.080us 501
|
||||
void (anonymous namespace)::indexing_backward_kernel... 0.00% 0.000us 0.00% 0.000us 0.000us 183.000ms 0.52% 183.000ms 6.100ms 30
|
||||
void at::native::unrolled_elementwise_kernel<at::nat... 0.00% 0.000us 0.00% 0.000us 0.000us 148.335ms 0.42% 148.335ms 988.899us 150
|
||||
aten::mean 0.02% 6.214ms 0.02% 8.406ms 12.717us 144.904ms 0.41% 144.904ms 219.220us 661
|
||||
void at::native::reduce_kernel<512, 1, at::native::R... 0.00% 0.000us 0.00% 0.000us 0.000us 144.904ms 0.41% 144.904ms 219.220us 661
|
||||
aten::_log_softmax 0.00% 527.616us 0.00% 716.309us 13.775us 139.294ms 0.39% 139.294ms 2.679ms 52
|
||||
void at::native::(anonymous namespace)::cunn_SoftMax... 0.00% 0.000us 0.00% 0.000us 0.000us 139.294ms 0.39% 139.294ms 2.679ms 52
|
||||
void at::native::reduce_kernel<512, 1, at::native::R... 0.00% 0.000us 0.00% 0.000us 0.000us 137.743ms 0.39% 137.743ms 286.368us 481
|
||||
void at::native::reduce_kernel<128, 4, at::native::R... 0.00% 0.000us 0.00% 0.000us 0.000us 130.526ms 0.37% 130.526ms 1.088ms 120
|
||||
aten::native_dropout_backward 0.00% 1.378ms 0.01% 3.156ms 15.394us 118.690ms 0.34% 118.690ms 578.975us 205
|
||||
------------------------------------------------------- ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------ ------------
|
||||
Self CPU time total: 35.229s
|
||||
Self CUDA time total: 35.322s
|
||||
|
||||
|
||||
========================================================================
|
||||
E. LAUNCH COUNTS (grouped_mm: 128/layer sequential = no-op, 1 = grouped)
|
||||
@@ -0,0 +1,206 @@
|
||||
"""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))
|
||||
@@ -0,0 +1,77 @@
|
||||
"""Step 2 — padding ratio. Data-side, no GPU, no model.
|
||||
|
||||
Replicates the exact batching the trainer used: SequentialSampler over the
|
||||
encode-cache order, per_device_batch_size=2, collate_mixed right-padding to
|
||||
the pair max. Reports real vs padded token counts and the loss-target count
|
||||
that sizes the chunked CE.
|
||||
"""
|
||||
import json, sys
|
||||
from collections import Counter
|
||||
|
||||
CACHE = "/tank/erp-tune/run-01/encode-cache/encoded-a4b0796de1260930.jsonl"
|
||||
IGNORE_INDEX = -100
|
||||
MB = 2 # per_device_batch_size
|
||||
ACCUM = 8 # gradient_accumulation_steps
|
||||
|
||||
lens, kept_counts, kinds = [], [], []
|
||||
with open(CACHE) as fh:
|
||||
for line in fh:
|
||||
row = json.loads(line)
|
||||
ids = row["input_ids"]
|
||||
labels = row["labels"]
|
||||
lens.append(len(ids))
|
||||
kept_counts.append(sum(1 for x in labels if x != IGNORE_INDEX))
|
||||
kinds.append(row.get("sample_kind", "?"))
|
||||
|
||||
n = len(lens)
|
||||
print(f"records {n:,}")
|
||||
print(f"sample_kind mix {dict(Counter(kinds))}")
|
||||
print()
|
||||
print(f"seq len min/mean/max {min(lens)} / {sum(lens)/n:.0f} / {max(lens)}")
|
||||
print(f"loss targets min/mean/max {min(kept_counts)} / {sum(kept_counts)/n:.0f} / {max(kept_counts)}")
|
||||
print()
|
||||
|
||||
# --- micro-batch padding, exactly as collate_mixed builds it ---
|
||||
real = padded = 0
|
||||
mb_widths, mb_waste, mb_kept = [], [], []
|
||||
for i in range(0, n - n % MB, MB):
|
||||
group = lens[i:i + MB]
|
||||
width = max(group)
|
||||
r = sum(group)
|
||||
p = width * MB
|
||||
real += r
|
||||
padded += p
|
||||
mb_widths.append(width)
|
||||
mb_waste.append(1 - r / p)
|
||||
mb_kept.append(sum(kept_counts[i:i + MB]))
|
||||
|
||||
nb = len(mb_widths)
|
||||
print(f"micro-batches (mb={MB}) {nb:,}")
|
||||
print(f"real tokens {real:,}")
|
||||
print(f"padded tokens {padded:,}")
|
||||
print(f"PADDING WASTE {100 * (1 - real / padded):.1f}% ({padded - real:,} pad tokens)")
|
||||
print()
|
||||
print(f"mb width min/mean/max {min(mb_widths)} / {sum(mb_widths)/nb:.0f} / {max(mb_widths)}")
|
||||
srt = sorted(mb_widths)
|
||||
for q in (50, 75, 90, 95, 99):
|
||||
print(f" p{q} width {srt[int(nb*q/100)]}")
|
||||
print(f"mb at max_seq_len 16384 {sum(1 for w in mb_widths if w >= 16384):,} ({100*sum(1 for w in mb_widths if w>=16384)/nb:.1f}%)")
|
||||
print()
|
||||
srtw = sorted(mb_waste)
|
||||
print(f"per-mb waste p50/p90/max {100*srtw[nb//2]:.1f}% / {100*srtw[int(nb*0.9)]:.1f}% / {100*max(mb_waste):.1f}%")
|
||||
print()
|
||||
print(f"loss targets per mb min/mean/max {min(mb_kept)} / {sum(mb_kept)/nb:.0f} / {max(mb_kept)}")
|
||||
print(f" -> CE chunks per mb (1024) min/mean/max {min(mb_kept)//1024+1} / {sum(mb_kept)/nb/1024:.1f} / {max(mb_kept)//1024+1}")
|
||||
print()
|
||||
|
||||
# --- what length-bucketing would recover (sort by length, then batch) ---
|
||||
order = sorted(range(n), key=lambda i: lens[i])
|
||||
b_real = b_padded = 0
|
||||
for i in range(0, n - n % MB, MB):
|
||||
group = [lens[j] for j in order[i:i + MB]]
|
||||
b_real += sum(group)
|
||||
b_padded += max(group) * MB
|
||||
print("--- counterfactual: length-bucketed sampler ---")
|
||||
print(f"bucketed padded tokens {b_padded:,}")
|
||||
print(f"bucketed waste {100 * (1 - b_real / b_padded):.1f}%")
|
||||
print(f"TOKEN REDUCTION vs current {100 * (1 - b_padded / padded):.1f}%")
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Measure bucket-to-pair / shuffle-to-mix against the REAL encode cache.
|
||||
|
||||
Brokkr's design, validated on measured record lengths rather than a calibrated
|
||||
length model:
|
||||
|
||||
1. sort records by length
|
||||
2. cut into buckets of BUCKET records
|
||||
3. form micro-batches of 2 WITHIN each bucket (adjacent after sort)
|
||||
4. shuffle the resulting MICRO-BATCHES globally, seeded
|
||||
|
||||
Padding efficiency is a property of the pairing only, so step 4 costs nothing
|
||||
and restores root-mixing inside each accumulation window.
|
||||
|
||||
Also applies the fitted cost model from the replica scaling test to convert
|
||||
token savings into predicted wall clock.
|
||||
"""
|
||||
import json
|
||||
import random
|
||||
from collections import Counter
|
||||
|
||||
CACHE = "/tank/erp-tune/run-01/encode-cache/encoded-a4b0796de1260930.jsonl"
|
||||
IGNORE_INDEX = -100
|
||||
MB = 2
|
||||
ACCUM = 8
|
||||
SEED = 20260824
|
||||
|
||||
# fitted on the replica: t(w) = A*w + B*w^2 for a batch of 2 sequences of len w
|
||||
A = 6.8715e-04
|
||||
B = 8.8509e-08
|
||||
|
||||
rows = []
|
||||
with open(CACHE) as fh:
|
||||
for line in fh:
|
||||
r = json.loads(line)
|
||||
rows.append((len(r["input_ids"]), r.get("dataset_id", "?"),
|
||||
r.get("sample_kind", "?")))
|
||||
n = len(rows)
|
||||
print("records %d" % n)
|
||||
print()
|
||||
|
||||
|
||||
def evaluate(order, label, show_roots=False):
|
||||
real = padded = 0
|
||||
widths = []
|
||||
batches = []
|
||||
for i in range(0, n - n % MB, MB):
|
||||
grp = [rows[j] for j in order[i:i + MB]]
|
||||
w = max(g[0] for g in grp)
|
||||
real += sum(g[0] for g in grp)
|
||||
padded += w * MB
|
||||
widths.append(w)
|
||||
batches.append([g[1] for g in grp])
|
||||
nb = len(widths)
|
||||
Ew = sum(widths) / nb
|
||||
Ew2 = sum(w * w for w in widths) / nb
|
||||
t_mb = A * Ew + B * Ew2
|
||||
srt = sorted(widths)
|
||||
print("--- %s ---" % label)
|
||||
print(" padded tokens %s" % f"{padded:,}")
|
||||
print(" waste %.1f%%" % (100 * (1 - real / padded)))
|
||||
print(" E[w] (per-seq) %.0f" % Ew)
|
||||
print(" E[w^2] %.3e" % Ew2)
|
||||
print(" width p50/p90/p99 %d / %d / %d" % (
|
||||
srt[nb // 2], srt[int(nb * .9)], srt[int(nb * .99)]))
|
||||
print(" predicted micro-batch %.3f s (lin %.3f + quad %.3f, quad %.0f%%)" % (
|
||||
t_mb, A * Ew, B * Ew2, 100 * B * Ew2 / t_mb))
|
||||
print(" predicted step (x%d) %.1f s -> %.2f h over 1312 steps" % (
|
||||
ACCUM, t_mb * ACCUM, t_mb * ACCUM * 1312 / 3600))
|
||||
# unpadded micro-batches take the is_causal fast path on the 5 global layers
|
||||
exact = sum(1 for i in range(0, n - n % MB, MB)
|
||||
if len(set(rows[j][0] for j in order[i:i + MB])) == 1)
|
||||
print(" ZERO-PAD micro-batches %d / %d (%.1f%%) <- global layers on is_causal" % (
|
||||
exact, nb, 100 * exact / nb))
|
||||
if show_roots:
|
||||
# root diversity inside an accumulation window
|
||||
div = []
|
||||
for i in range(0, nb - nb % ACCUM, ACCUM):
|
||||
win = [d for b in batches[i:i + ACCUM] for d in b]
|
||||
div.append(len(set(win)))
|
||||
print(" roots per accum window mean %.2f min %d (of %d roots)" % (
|
||||
sum(div) / len(div), min(div), len({r[1] for r in rows})))
|
||||
print()
|
||||
return padded, t_mb
|
||||
|
||||
|
||||
# --- current: encode-cache order, SequentialSampler ---
|
||||
cur_padded, cur_t = evaluate(list(range(n)), "CURRENT (SequentialSampler)", True)
|
||||
|
||||
# --- bucket-to-pair + shuffle-to-mix ---
|
||||
# BUCKET controls the efficiency-vs-diversity trade: records are globally
|
||||
# sorted, cut into buckets of BUCKET, SHUFFLED WITHIN the bucket (not
|
||||
# re-sorted), then paired adjacently. BUCKET=2 is a perfect global sort
|
||||
# (0% waste, worst root mixing); larger buckets admit more length spread
|
||||
# inside a pair but draw partners from a wider slice of the corpus.
|
||||
for BUCKET in (2, 8, 32, 128, 512):
|
||||
by_len = sorted(range(n), key=lambda i: rows[i][0])
|
||||
rng = random.Random(SEED)
|
||||
micro = []
|
||||
for s in range(0, n, BUCKET):
|
||||
chunk = by_len[s:s + BUCKET]
|
||||
rng.shuffle(chunk) # mix WITHIN the length bucket
|
||||
for k in range(0, len(chunk) - len(chunk) % MB, MB):
|
||||
micro.append(chunk[k:k + MB])
|
||||
rng.shuffle(micro) # shuffle-to-mix across buckets
|
||||
order = [i for b in micro for i in b]
|
||||
placed = set(order)
|
||||
order += [i for i in by_len if i not in placed]
|
||||
p, t = evaluate(order, "BUCKET=%d, shuffle within + global micro-batch shuffle" % BUCKET,
|
||||
True)
|
||||
print(" >>> vs current: %.1f%% fewer padded tokens, %.1f%% less wall clock" % (
|
||||
100 * (1 - p / cur_padded), 100 * (1 - t / cur_t)))
|
||||
print()
|
||||
Reference in New Issue
Block a user