"""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 --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))