--- /tmp/base_eval.bak 2026-07-03 13:59:36.997249952 -0400 +++ deepspec/eval/base_evaluator.py 2026-07-03 20:44:32.070666055 -0400 @@ -317,6 +317,9 @@ propose: Callable[..., DraftProposal], update: Callable[[Any, VerificationResult], None], post_verify: Callable[[DraftProposal, VerificationResult], None] | None = None, + prefill_keep_hidden_layers: list[int] | None = None, + stream_callback=None, + prefill_mm: dict | None = None, ) -> SimpleNamespace: """Speculative-decoding loop. @@ -343,14 +346,49 @@ from deepspec.eval.windowed_cache import build_target_cache past_key_values_target = build_target_cache(target_model, pad=max(64, int(max_proposal_tokens)+8)) - output = target_model( - input_ids=input_ids, - position_ids=position_ids[:, :num_input_tokens], - past_key_values=past_key_values_target, - use_cache=True, - output_hidden_states=True, - logits_to_keep=1, - ) + _chunk = int(os.environ.get("DSPARK_PREFILL_CHUNK", "4096")) + if num_input_tokens > _chunk and prefill_keep_hidden_layers is not None and not prefill_mm: + # chunked prefill: bound activation memory (never materialize all layers x all positions); + # keep only the draft target layers hidden states, concatenated across chunks. + _keep = sorted(set(int(l) for l in prefill_keep_hidden_layers)) + _acc = {li: [] for li in _keep} + _last_logits = None + for _i in range(0, num_input_tokens, _chunk): + _j = min(_i + _chunk, num_input_tokens) + _islast = _j == num_input_tokens + _o = target_model( + input_ids=input_ids[:, _i:_j], + position_ids=position_ids[:, _i:_j], + past_key_values=past_key_values_target, + use_cache=True, + output_hidden_states=True, + logits_to_keep=1 if _islast else 0, + ) + for li in _keep: + _acc[li].append(_o.hidden_states[li]) + if _islast: + _last_logits = _o.logits + del _o + _hw = int(os.environ.get("DSPARK_DRAFT_CTX_WINDOW", "0")) + if _hw: + for li in _keep: + _tot = sum(t.shape[1] for t in _acc[li]) + while len(_acc[li]) > 1 and _tot - _acc[li][0].shape[1] >= _hw: + _tot -= _acc[li].pop(0).shape[1] + _hs = [None] * (max(_keep) + 1) + for li in _keep: + _hs[li] = torch.cat(_acc[li], dim=1) + output = SimpleNamespace(logits=_last_logits, hidden_states=tuple(_hs)) + else: + output = target_model( + input_ids=input_ids, + position_ids=position_ids[:, :num_input_tokens], + past_key_values=past_key_values_target, + use_cache=True, + output_hidden_states=True, + logits_to_keep=1, + **(prefill_mm or {}), + ) output_ids[:, :num_input_tokens] = input_ids output_ids[:, num_input_tokens : num_input_tokens + 1] = sample_from_probs( @@ -384,6 +422,8 @@ ) while start < max_length: + if stream_callback is not None: + stream_callback(output_ids[:, num_input_tokens : start + 1]) proposal = propose( context=context, output_ids=output_ids, @@ -429,6 +469,8 @@ if has_stop_token(new_token_ids, stop_token_ids): break + if stream_callback is not None: + stream_callback(output_ids[:, num_input_tokens : start + 1]) output_ids = output_ids[:, : min(start + 1, max_length)] output_ids = trim_output_ids(output_ids, num_input_tokens, stop_token_ids) return SimpleNamespace(