Files
esh-pfi-infrastructure/scripts/training-probes/bench_lora_step.py
T
vh 7db6c44bcd feat(r49-prep): author-voice LoRA regime prep on gx10 — carriers staged, throughput measured, adapters secured
Prep for the BabyBronte / brokkr-smithy R49 author-voice adapter regime, plus
the operator's "keep the adapter" ruling made durable.

Measured on pfi-gx10 (GB10, sm_121), n=10 per arm after 3 warmup steps, seq
4096, LoRA r=32 on q/k/v/o + MLP, bf16, sdpa, grad-checkpointing on:

  Qwen3-0.6B-Base    dense    0.616 B   1.707 s/step   2,399 tok/s
  Qwen3-1.7B-Base    dense    1.755 B   2.895 s/step   1,415 tok/s
  Qwen3.5-0.8B-Base  hybrid   0.765 B   7.581 s/step     540 tok/s

The dense 1.755 B carrier trains 2.6x faster than the hybrid 0.765 B one on 2.3x
the parameters (~6x per parameter), with more LoRA modules adapted (196 vs 96).
Spreads of 0.6-2.6% put instrument noise an order of magnitude below the effect.
Cause: Qwen3.5 is 18 linear-attention (SSM) layers to 6 attention, and no fused
linear-attention kernel is installed on the box. Grad checkpointing is not the
culprit (19%, and saves 2.6x memory). Batching is not the lever for either
family -- both sit at this box's roofline at batch 1.

Projected per voice on a Brontë-scale corpus: dense 0.6B 2.7 h, dense 1.7B
4.6 h, hybrid 0.8B 12 h. The hybrid would take longer than the 7 h 26B-A4B tune
the regime exists to replace, so the carrier family is now an open decision with
a recommendation for the dense Qwen3 line -- the design doc's original pin.

Two further Qwen3.5 findings, both measured rather than read off the config: the
Base checkpoints ship a vision tower (153/297 model.visual.* Linear tensors that
target_modules="all-linear" would train on text) and an MTP head, both dropped
for free by loading through AutoModelForCausalLM -- which renames modules
relative to the vLLM serving path, so adapter binding needs the
sampled-target-changed check on the serving side; and cross-document packing is
unsafe because SSM state ignores the attention mask, breaking the per-copy
name-consistency invariant the design doc calls sacred. Neither exists on dense.

Adapter disposition, per the operator's ruling: all five gx10-resident ERP
adapters (run-03c/04/05/06/07) mirrored to ana-ml2:/tank/erp-tune/run-<N>/adapter
matching the layout runs 01-03 already used, byte-totals identical both sides and
sha256 matching on every adapter_model.safetensors. /tank/* is deliberately
excluded from ana-ml2's restic sources, so the profile gains one documented
carve-out for /tank/erp-tune/run-*/adapter, verified by resticprofile --dry-run
to expand to exactly those eight paths.

Nothing is training and nothing is queued.
2026-09-09 22:41:47 -07:00

94 lines
3.6 KiB
Python

"""Throughput floor for an R49 author-voice LoRA step on pfi-gx10 (GB10, sm_121).
Measures the cost of ONE forward+backward+optimizer microbatch on synthetic
tokens, so a full-corpus wall-clock can be projected before any corpus exists.
Deliberately synthetic: random token ids exercise the same kernels at the same
shapes as real text, and this is a THROUGHPUT harness only -- it says nothing
about loss, quality, or voice transfer. The harness is part of the number, so
every knob is printed with the result.
python bench_lora_step.py <model-dir> --seq 4096 --targets attn_mlp|all_linear_text
"""
import argparse, json, statistics, time, os
import torch
from transformers import AutoModelForCausalLM, AutoConfig
from peft import LoraConfig, get_peft_model
ATTN_MLP = ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]
PLUS_SSM = ATTN_MLP + ["in_proj_qkv", "in_proj_a", "in_proj_b", "in_proj_z", "out_proj"]
ap = argparse.ArgumentParser()
ap.add_argument("model")
ap.add_argument("--seq", type=int, default=4096)
ap.add_argument("--batch", type=int, default=1)
ap.add_argument("--rank", type=int, default=32)
ap.add_argument("--targets", choices=["attn_mlp", "plus_ssm"], default="attn_mlp")
ap.add_argument("--warmup", type=int, default=3)
ap.add_argument("--steps", type=int, default=10)
ap.add_argument("--no-grad-ckpt", action="store_true")
ap.add_argument("--attn", default="sdpa")
a = ap.parse_args()
torch.manual_seed(0)
cfg = AutoConfig.from_pretrained(a.model)
vocab = getattr(getattr(cfg, "text_config", cfg), "vocab_size")
model = AutoModelForCausalLM.from_pretrained(
a.model, dtype=torch.bfloat16, attn_implementation=a.attn,
).to("cuda")
targets = ATTN_MLP if a.targets == "attn_mlp" else PLUS_SSM
peft_cfg = LoraConfig(
r=a.rank, lora_alpha=2 * a.rank, lora_dropout=0.0, bias="none",
task_type="CAUSAL_LM", target_modules=targets,
)
model = get_peft_model(model, peft_cfg)
if not a.no_grad_ckpt:
model.gradient_checkpointing_enable()
model.enable_input_require_grads()
model.train()
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
total = sum(p.numel() for p in model.parameters())
n_adapted = sum(1 for n, _ in model.named_modules() if n.endswith("lora_A.default"))
opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=1e-4)
ids = torch.randint(0, vocab, (a.batch, a.seq), device="cuda")
def step():
opt.zero_grad(set_to_none=True)
out = model(input_ids=ids, labels=ids)
out.loss.backward()
opt.step()
return float(out.loss)
for _ in range(a.warmup):
step()
torch.cuda.synchronize()
lat = []
for _ in range(a.steps):
t0 = time.perf_counter()
step()
torch.cuda.synchronize()
lat.append(time.perf_counter() - t0)
tok = a.batch * a.seq
res = dict(
model=os.path.basename(a.model.rstrip("/")),
total_params_B=round(total / 1e9, 3),
lora_rank=a.rank, targets=a.targets, adapted_modules=n_adapted,
trainable_params_M=round(trainable / 1e6, 2),
trainable_pct=round(100 * trainable / total, 3),
batch=a.batch, seq=a.seq, tokens_per_microbatch=tok,
grad_checkpointing=not a.no_grad_ckpt, attn_impl=a.attn,
dtype="bfloat16", device=torch.cuda.get_device_name(0),
torch=torch.__version__, warmup=a.warmup, n=a.steps,
s_per_step_median=round(statistics.median(lat), 4),
s_per_step_min=round(min(lat), 4), s_per_step_max=round(max(lat), 4),
s_per_step_spread_pct=round(100 * (max(lat) - min(lat)) / statistics.median(lat), 1),
tok_per_s_median=round(tok / statistics.median(lat), 1),
peak_mem_GiB=round(torch.cuda.max_memory_allocated() / 2**30, 2),
)
print(json.dumps(res))