Skip to content

Commit e2acfcc

Browse files
perf(restored): ADR 0015 item #1 - memory-efficient restoration prefill
Unblock long-context multi-tenant for the restored-S5 path: - forward() gains logits_to_keep (default 0 = unchanged; spec-decode still gets all-position logits). Bench passes 1 (only last logits needed). - capture_verifier_own_kv skips the [B,T,vocab] logits entirely (logits_to_keep=1; it only needs K/V via hooks). - restored K/V cast to verifier compute dtype right after f_theta (was held full-length in fp32; cast-at-use was already bf16, so numerically identical). - bench --attn-impl (default eager for exact repro; sdpa for long context routes the patched forward to SDPA, killing the O(T^2) eager score materialisation). 41 v04 unit tests still pass (logits_to_keep default preserves behavior). Co-authored-by: FluffyAIcode <FluffyAIcode@users.noreply.github.com>
1 parent f030a74 commit e2acfcc

2 files changed

Lines changed: 35 additions & 6 deletions

File tree

inference_engine/v04/cross_model_dlm_verifier.py

Lines changed: 27 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -310,9 +310,16 @@ def forward(
310310
eager_attention_forward: Callable,
311311
all_attention_functions: Optional[Any] = None,
312312
capture_kv: Optional[list] = None,
313+
logits_to_keep: int = 0,
313314
):
314315
"""Run a verifier forward with f_θ-mediated K/V Restoration.
315316
317+
``logits_to_keep`` is forwarded to the verifier model: 0 (default)
318+
computes logits for all positions (needed by spec-decode block
319+
verification); set 1 to compute only the last position's logits
320+
(decode/prefill that only needs the next token) — this avoids the
321+
``[B, T, vocab]`` materialisation that dominates long-context memory.
322+
316323
Steps:
317324
1. Compute evicted positions from sink+window per ADR §11.7.
318325
2. Drafter forward + f_θ projection → verifier K/V at every
@@ -335,11 +342,23 @@ def forward(
335342
# needed — run the verifier directly. This is the trivial case
336343
# for short prompts, e.g. T=8 with sink=4 + window=64.
337344
if not evicted_positions:
338-
return self.verifier_model(input_ids=input_ids, use_cache=False)
345+
return self.verifier_model(
346+
input_ids=input_ids, use_cache=False,
347+
logits_to_keep=logits_to_keep,
348+
)
339349

340350
# f_θ projection → per-layer lists (layers can have different
341351
# KV-head counts on Gemma 4).
342352
verifier_k_layers, verifier_v_layers = self.project_drafter_kv(input_ids)
353+
# Hold the restored K/V at the verifier's compute dtype (they are
354+
# cast to it at injection anyway): keeps the full-length [B,T,kv,D]
355+
# tensors from doubling long-context memory when f_θ runs in fp32.
356+
try:
357+
vdtype = next(self.verifier_model.parameters()).dtype
358+
verifier_k_layers = [k.to(vdtype) for k in verifier_k_layers]
359+
verifier_v_layers = [v.to(vdtype) for v in verifier_v_layers]
360+
except StopIteration:
361+
pass
343362

344363
# S5: for exact layers, replace the f_θ-restored K/V with the
345364
# verifier's OWN true K/V (the recall-critical full-attention
@@ -375,7 +394,10 @@ def forward(
375394
all_attention_functions=all_attention_functions,
376395
capture_kv=capture_kv,
377396
)
378-
return self.verifier_model(input_ids=input_ids, use_cache=False)
397+
return self.verifier_model(
398+
input_ids=input_ids, use_cache=False,
399+
logits_to_keep=logits_to_keep,
400+
)
379401
finally:
380402
for layer_idx, layer in enumerate(layers):
381403
layer.self_attn.forward = originals[layer_idx]
@@ -548,7 +570,9 @@ def _vh(_m, _inp, out, idx=i):
548570
else:
549571
v_shared.append(i)
550572
try:
551-
verifier_model(input_ids=input_ids, use_cache=False)
573+
# Only K/V are needed (captured via k_proj/v_proj hooks); skip the
574+
# [B, T, vocab] logits materialisation entirely.
575+
verifier_model(input_ids=input_ids, use_cache=False, logits_to_keep=1)
552576
finally:
553577
for h in handles:
554578
h.remove()

scripts/research/k3_cuda_multitenant_parallel_bench.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -36,7 +36,7 @@
3636
def _ar_batched(model, ids_bt, gen_tokens, device, eos_ids):
3737
"""Batched AR decode. ids_bt: [N, T]. Returns (per_row_tokens, decode_s)."""
3838
N = ids_bt.size(0)
39-
out = model(input_ids=ids_bt, use_cache=True)
39+
out = model(input_ids=ids_bt, use_cache=True, logits_to_keep=1)
4040
cache = out.past_key_values
4141
nxt = out.logits[:, -1, :].argmax(-1) # [N]
4242
gen = [[int(nxt[i].item())] for i in range(N)]
@@ -62,7 +62,7 @@ def _restored_prefill_batched(restored, ids_bt, helpers):
6262
from transformers.cache_utils import DynamicCache
6363
n_layers = len(_decoder_layers(restored.verifier_model))
6464
capture: list = [None] * n_layers
65-
out = restored.forward(ids_bt, capture_kv=capture, **helpers)
65+
out = restored.forward(ids_bt, capture_kv=capture, logits_to_keep=1, **helpers)
6666
logits = out.logits if hasattr(out, "logits") else out
6767
if any(c is None for c in capture):
6868
raise RuntimeError("restored prefill did not capture all layers "
@@ -111,6 +111,10 @@ def main() -> int:
111111
ap.add_argument("--sink", type=int, default=4)
112112
ap.add_argument("--window", type=int, default=64)
113113
ap.add_argument("--seed", type=int, default=0)
114+
ap.add_argument("--attn-impl", default="eager",
115+
help="verifier attn_implementation: 'eager' (default, exact "
116+
"repro) or 'sdpa'/'flash_attention_2' (memory-efficient "
117+
"prefill — required for long context, ADR 0015 item #1).")
114118
ap.add_argument("--output", default=None)
115119
args = ap.parse_args()
116120

@@ -134,7 +138,7 @@ def main() -> int:
134138
print(f"[mt] loading verifier {args.verifier_id}", file=sys.stderr, flush=True)
135139
tok = AutoTokenizer.from_pretrained(args.verifier_id)
136140
verifier = AutoModelForCausalLM.from_pretrained(
137-
args.verifier_id, dtype=dtype, attn_implementation="eager",
141+
args.verifier_id, dtype=dtype, attn_implementation=args.attn_impl,
138142
).to(device).eval()
139143
for p in verifier.parameters():
140144
p.requires_grad_(False)
@@ -233,6 +237,7 @@ def recall(tokens, ans):
233237
"verifier_id": args.verifier_id, "drafter_id": args.drafter_id,
234238
"haystack_lines": args.haystack_lines, "modal_prompt_len": modal_len,
235239
"gen_tokens": args.gen_tokens, "sink": args.sink, "window": args.window,
240+
"attn_impl": args.attn_impl,
236241
"batch_sizes": batch_sizes, "exact_layers": exact_layers,
237242
"note": ("per-session binding via batched decode (each row = a "
238243
"session with its own KV-cache row); recall-preserving S5 "

0 commit comments

Comments
 (0)