Efficient-Large-Model/Fast-dDrive
4167
1"""Generation utilities for Fast-dDrive.2 3This module provides the three inference paths exposed by the canonical paper4release:5 6* ``mdm_sample_deep_scaffold`` — Section Diffusion (SD): iterative MDM7 denoising over a pre-filled JSON scaffold, no AR verification.8* ``scaffold_speculative_sample`` — Scaffold Spec (SS): scaffold-aware9 self-speculative decoding (MDM draft + AR verify per block).10* ``scaffold_spec_with_ss_multi_traj`` — SS with shared-prefix multi-trajectory11 rollouts (the test-time inference-scaling path).12 13All three are attached as bound methods on14:class:`Fast_dDriveForConditionalGeneration` when this module is15imported (see ``modeling.py`` for the import hook).16"""17 18import os19import re20import sys21import math22import torch23import types24import numpy as np25from transformers.cache_utils import DynamicCache26 27 28def _crop_cache(past_key_values, max_length: int):29 """Crop a DynamicCache to max_length tokens, compatible with Qwen cache layout."""30 new_past_key_values = []31 for layer_num in range(len(past_key_values)):32 layer_past_key_values = ()33 for kv_idx in range(len(past_key_values[layer_num])):34 layer_past_key_values += (past_key_values[layer_num][kv_idx][:, :, :max_length, :],)35 new_past_key_values.append(layer_past_key_values)36 return DynamicCache(new_past_key_values)37 38 39def _sample_from_logits(logits, temperature=0.0):40 """Sample token ids from logits with optional temperature scaling.41 42 When temperature <= 0, falls back to argmax (greedy).43 """44 if temperature <= 0:45 return logits.argmax(dim=-1)46 scaled = logits / temperature47 probs = torch.softmax(scaled, dim=-1)48 original_shape = probs.shape[:-1]49 flat_probs = probs.reshape(-1, probs.shape[-1])50 sampled = torch.multinomial(flat_probs, num_samples=1).squeeze(-1)51 return sampled.reshape(original_shape)52 53 54# ---------------------------------------------------------------------------55# mdm_sample_deep_scaffold — Section Diffusion (SD)56# ---------------------------------------------------------------------------57 58def mdm_sample_deep_scaffold(59 self,60 input_ids,61 tokenizer,62 max_tokens=512,63 pixel_values=None,64 image_grid_thw=None,65 mask_id=151665,66 null_id=151666,67 threshold=0.9,68 stop_token=151645,69 explanation_block_size=32,70 explanation_max_blocks=6,71 block_size=32,72 return_stats=False,73 use_kv_cache=True,74 temperature=0.0,75):76 """77 Deep scaffold MDM generation with train-consistent hybrid block causal mask.78 79 Pre-fills the entire JSON scaffold (including sub-keys for critical_objects,80 future_meta_behavior, trajectory) with MASK tokens at value positions only.81 Then denoises each section's value tokens via iterative unmasking.82 83 The attention mask matches training: prompt tokens use causal attention,84 response tokens use block-causal attention where each section's denoise85 steps form separate blocks. Block i can see all prompt tokens and blocks86 0..i, but NOT blocks i+1..N (which still contain MASK tokens).87 88 For explanation (variable length), NULL tokens in the output signal that89 the section content is complete — trailing NULLs are stripped.90 91 KV-cache path (``use_kv_cache=True``, default):92 Prompt K/V is computed once with vision embedding scatter, then each93 response block, once fully denoised, gets its K/V appended to the94 cache. Subsequent blocks' iterative unmasking only forwards their95 own ~block_size tokens against the cache (plus prior committed96 blocks), avoiding O(seqlen^2) recomputation of the prompt + prior97 blocks every iteration. Correctness is preserved because98 block-causal attention means block k only attends to prompt +99 blocks 0..k, which is exactly what the cache provides.100 """101 import math102 import os as _os103 from .section_utils import (104 build_deep_json_scaffold,105 SECTION_KEYS,106 NULL_TOKEN_ID,107 )108 109 # Env override for A/B testing the KV cache path without editing code.110 _kv_env = _os.environ.get("MDM_DS_USE_KV_CACHE")111 if _kv_env is not None:112 use_kv_cache = _kv_env not in ("0", "false", "False", "")113 114 scaffold_tokens, section_ranges, scaffold_mask_list = build_deep_json_scaffold(115 tokenizer,116 mask_id=mask_id,117 null_id=null_id,118 explanation_block_size=explanation_block_size,119 explanation_max_blocks=explanation_max_blocks,120 )121 122 tokens_per_step = []123 original_input_length = input_ids.shape[1]124 125 # Phase 1: Build sequence with scaffold appended126 scaffold_tensor = torch.tensor(scaffold_tokens, device=self.device, dtype=torch.long).unsqueeze(0)127 x_t = torch.cat([input_ids, scaffold_tensor], dim=1)128 seqlen = x_t.shape[1]129 130 # Track scaffold (frozen) vs value (to denoise) positions in scaffold region131 scaffold_frozen = torch.tensor(scaffold_mask_list, device=self.device, dtype=torch.bool)132 133 # ── Build response_block_idx matching training's compute_section_block_idx_deep_static ──134 response_block_idx = torch.full((seqlen,), -1, device=self.device, dtype=torch.long)135 current_block = 0136 assigned = set()137 138 for section_name in SECTION_KEYS:139 if section_name not in section_ranges:140 continue141 sec_start, sec_end = section_ranges[section_name]142 143 # Find value positions (non-scaffold) in this section144 value_positions = []145 for i in range(sec_start, sec_end):146 if not scaffold_mask_list[i]: # 0 = value token147 value_positions.append(original_input_length + i)148 149 if not value_positions:150 current_block += 1151 continue152 153 # Block assignment MUST match training's154 # compute_section_block_idx_deep_static: n_blocks = ceil(value/block_size)155 # for every section. Previously non-explanation sections were forced to156 # a single block; that broke attention alignment for trajectory157 # (70 value tokens → training 3 blocks vs inference 1 block), causing158 # trajectory over-extrapolation. CO (12) and FMB (6) still resolve159 # to 1 block since their value counts are < block_size.160 tokens_per_step_sec = block_size161 n_steps = max(1, math.ceil(len(value_positions) / tokens_per_step_sec))162 163 # Assign block indices to value tokens164 for vi, abs_pos in enumerate(value_positions):165 block_in_section = min(vi // tokens_per_step_sec, n_steps - 1)166 response_block_idx[abs_pos] = current_block + block_in_section167 assigned.add(abs_pos)168 169 # Assign scaffold tokens to nearest value token's block170 for i in range(sec_start, sec_end):171 abs_pos = original_input_length + i172 if scaffold_mask_list[i] and abs_pos not in assigned:173 best_block = -1174 for delta in range(1, sec_end - sec_start + 10):175 for cand in [abs_pos + delta, abs_pos - delta]:176 if cand in assigned:177 best_block = response_block_idx[cand].item()178 break179 if best_block >= 0:180 break181 if best_block >= 0:182 response_block_idx[abs_pos] = best_block183 assigned.add(abs_pos)184 185 current_block += n_steps186 187 # Assign any remaining unassigned scaffold tokens (e.g. top-level separators)188 for i in range(len(scaffold_tokens)):189 abs_pos = original_input_length + i190 if abs_pos not in assigned:191 # Find nearest assigned position192 best_block = -1193 for delta in range(1, seqlen):194 for cand in [abs_pos + delta, abs_pos - delta]:195 if 0 <= cand < seqlen and cand in assigned:196 best_block = response_block_idx[cand].item()197 break198 if best_block >= 0:199 break200 if best_block >= 0:201 response_block_idx[abs_pos] = best_block202 assigned.add(abs_pos)203 204 # ── Build hybrid block causal mask (computed once, reused for all forward passes) ──205 attention_mask = self.model.eval_hybrid_mask(seqlen, response_block_idx).to(self.device)206 207 # Section-MoE-LoRA: set section_ids before language model forward208 set_section_ids = lambda *a, **kw: None # noqa: E731 (Section-MoE-LoRA disabled in release)209 # Map block indices to section IDs (0=CO, 1=Exp, 2=FMB, 3=Traj, 4=Other/Prompt)210 _sec_ids = torch.full((seqlen,), 4, device=self.device, dtype=torch.long)211 for section_name, (sec_start, sec_end) in section_ranges.items():212 abs_start = original_input_length + sec_start213 abs_end = original_input_length + sec_end214 if section_name == "critical_objects":215 _sec_ids[abs_start:abs_end] = 0216 elif section_name == "explanation":217 _sec_ids[abs_start:abs_end] = 1218 elif section_name == "future_meta_behavior":219 _sec_ids[abs_start:abs_end] = 2220 elif section_name == "trajectory":221 _sec_ids[abs_start:abs_end] = 3222 223 # Add batch dimension224 _sec_ids_batch = _sec_ids.unsqueeze(0)225 set_section_ids(_sec_ids_batch)226 227 # ── Precompute vision embeddings and position_ids once ──228 # BUG FIX: Previously pixel_values was only passed on the first forward229 # (step==0) but with use_cache=False every forward is independent, so all230 # subsequent forwards lost vision information entirely.231 _embed_fn = self.model.get_input_embeddings()232 _cached_image_embeds = None233 _cached_image_mask = None234 235 if pixel_values is not None:236 _cached_image_embeds = self.model.get_image_features(pixel_values, image_grid_thw)237 _cached_image_embeds = torch.cat(_cached_image_embeds, dim=0).to(238 self.device, _embed_fn.weight.dtype239 )240 _tmp_embeds = _embed_fn(x_t)241 _cached_image_mask, _ = self.model.get_placeholder_mask(242 x_t, inputs_embeds=_tmp_embeds, image_features=_cached_image_embeds243 )244 245 # Compute position_ids once with correct image_grid_thw (3D RoPE)246 _position_ids, _rope_deltas = self.model.get_rope_index(247 x_t, image_grid_thw, None248 )249 self.model.rope_deltas = _rope_deltas250 251 # ── Compute contiguous block ranges in the response region ──252 # Each block's absolute [start, end) range in x_t is the maximal253 # contiguous span of positions sharing the same response_block_idx.254 # Blocks are ordered by block_idx and cover the entire response.255 _block_ranges = [] # list of (block_idx, abs_start, abs_end)256 _cur_bi = None257 _cur_start = None258 for _p in range(seqlen):259 _bi = int(response_block_idx[_p].item())260 if _bi < 0:261 if _cur_bi is not None:262 _block_ranges.append((_cur_bi, _cur_start, _p))263 _cur_bi, _cur_start = None, None264 continue265 if _cur_bi is None:266 _cur_bi, _cur_start = _bi, _p267 elif _bi != _cur_bi:268 _block_ranges.append((_cur_bi, _cur_start, _p))269 _cur_bi, _cur_start = _bi, _p270 if _cur_bi is not None:271 _block_ranges.append((_cur_bi, _cur_start, seqlen))272 273 # Map block_idx -> section_name for downstream logic (section-specific274 # behaviors like explanation NULL handling can still be scoped).275 _block_idx_to_section = {}276 for _sname, (_sstart, _send) in section_ranges.items():277 _sabs_start = original_input_length + _sstart278 _sabs_end = original_input_length + _send279 for _bi, _bs, _be in _block_ranges:280 # Assign section by whether the block's range overlaps the section281 if _bs < _sabs_end and _be > _sabs_start:282 _block_idx_to_section.setdefault(_bi, _sname)283 284 # ── Phase 2: Denoise block-by-block with optional KV cache ──285 # Without cache (fallback): each forward replays the entire sequence.286 # With cache: prompt K/V computed once; each block's finalized K/V is287 # appended after denoising, so later blocks only forward their own288 # ~block_size tokens against the cache.289 step = 0290 291 past_kv = None292 prev_last_logit = None # logit at the position just before the next block293 294 if use_kv_cache:295 # Phase 0: prompt prefill. Includes vision scatter; cache becomes296 # the reusable foundation for every scaffold block.297 prompt_tokens = x_t[:, :original_input_length]298 prompt_embeds = _embed_fn(prompt_tokens)299 if _cached_image_embeds is not None:300 prompt_image_mask = _cached_image_mask[:, :original_input_length]301 prompt_embeds = prompt_embeds.masked_scatter(302 prompt_image_mask, _cached_image_embeds303 )304 prompt_position_ids = _position_ids[..., :original_input_length]305 306 # Causal over prompt (matches training's prompt-side attention).307 # When attention_mask=None, the model's eval_mask auto-builds causal308 # because use_block_causal_mask=True and update_kv_cache=True.309 prompt_out = self.forward(310 inputs_embeds=prompt_embeds,311 position_ids=prompt_position_ids,312 attention_mask=None,313 past_key_values=None,314 use_cache=True,315 update_kv_cache=True,316 )317 past_kv = prompt_out.past_key_values318 # Logit at position (original_input_length - 1); used to predict319 # the first token of the first response block via causal shift.320 prev_last_logit = prompt_out.logits[:, -1:, :]321 322 # ── Iterate blocks in order ──323 for _block_idx, block_abs_start, block_abs_end in _block_ranges:324 B = block_abs_end - block_abs_start325 section_name = _block_idx_to_section.get(_block_idx, None)326 327 # Count MASK tokens in this block328 block_slice = x_t[0, block_abs_start:block_abs_end]329 n_masks_in_block = int((block_slice == mask_id).sum().item())330 331 # ── Iterative unmasking within this block (if any MASKs) ──332 if n_masks_in_block > 0:333 max_iter = n_masks_in_block + 5 # safety limit334 for _ in range(max_iter):335 current_block_masks = (x_t[:, block_abs_start:block_abs_end] == mask_id)336 if current_block_masks.sum() == 0:337 break338 339 if use_kv_cache:340 # Feed only this block; past_kv covers prompt + prior blocks.341 block_tokens = x_t[:, block_abs_start:block_abs_end]342 block_embeds = _embed_fn(block_tokens)343 block_position_ids = _position_ids[..., block_abs_start:block_abs_end]344 L_cached = past_kv.get_seq_length() if past_kv is not None else 0345 # Block-causal + bidirectional-within-block ⇒ this346 # block's queries attend to all cached KV plus all347 # fresh block KV ⇒ all-True mask of shape [B, L+B].348 block_attn = torch.ones(349 B, L_cached + B, device=self.device, dtype=torch.bool350 )351 output = self.forward(352 inputs_embeds=block_embeds,353 attention_mask=block_attn,354 position_ids=block_position_ids,355 past_key_values=past_kv,356 use_cache=True,357 update_kv_cache=False, # read-only during iteration358 )359 logits = output.logits # [1, B, V]360 # Shift: pred for abs_pos uses logit at abs_pos-1.361 # logit at block_abs_start-1 is prev_last_logit; the362 # rest come from this forward's earlier positions.363 sec_logits = torch.cat([prev_last_logit, logits[:, :-1, :]], dim=1)364 else:365 # Full-sequence forward (fallback path, same as before)366 _cur_embeds = _embed_fn(x_t)367 if _cached_image_embeds is not None:368 _cur_embeds = _cur_embeds.masked_scatter(369 _cached_image_mask, _cached_image_embeds370 )371 output = self.forward(372 input_ids=x_t,373 inputs_embeds=_cur_embeds,374 attention_mask=attention_mask,375 position_ids=_position_ids,376 use_cache=False,377 )378 logits = output.logits379 sec_logits = logits[:, block_abs_start:block_abs_end, :]380 sec_logits = torch.cat(381 [logits[:, block_abs_start - 1:block_abs_start, :],382 sec_logits[:, :-1, :]], dim=1383 )384 385 if temperature > 0:386 # Temperature sampling for diverse generation (e.g. GRPO rollouts)387 sampling_probs = torch.softmax(sec_logits / temperature, dim=-1)388 x_1 = torch.multinomial(389 sampling_probs.view(-1, sampling_probs.shape[-1]), num_samples=1390 ).view(sampling_probs.shape[:-1])391 else:392 # Greedy (default, backward compatible)393 x_1 = sec_logits.argmax(dim=-1)394 probs = torch.softmax(sec_logits, dim=-1)395 x1_p = torch.gather(probs, dim=-1, index=x_1.unsqueeze(-1)).squeeze(-1)396 397 # Only consider currently-masked positions in this block398 x1_p = torch.where(current_block_masks, x1_p, -torch.inf)399 unmask_idx = (x1_p > threshold)400 401 if unmask_idx.sum() > 0:402 x_t[:, block_abs_start:block_abs_end][unmask_idx] = x_1[unmask_idx]403 tokens_per_step.append(int(unmask_idx.sum()))404 else:405 # Fallback: unmask highest-confidence token406 pos = x1_p.argmax()407 row = 0408 col = pos.item()409 x_t[:, block_abs_start:block_abs_end][row, col] = x_1[row, col]410 tokens_per_step.append(1)411 412 step += 1413 if step > max_tokens:414 break415 416 # ── Commit this block's K/V to the cache ──417 # Run one final forward at block's fully-denoised state with418 # update_kv_cache=True so future blocks can attend to it via cache.419 # prev_last_logit is refreshed to the logit at the last position420 # of this block for the NEXT block's first-position prediction.421 if use_kv_cache:422 block_tokens = x_t[:, block_abs_start:block_abs_end]423 block_embeds = _embed_fn(block_tokens)424 block_position_ids = _position_ids[..., block_abs_start:block_abs_end]425 L_cached = past_kv.get_seq_length() if past_kv is not None else 0426 block_attn = torch.ones(427 B, L_cached + B, device=self.device, dtype=torch.bool428 )429 commit_out = self.forward(430 inputs_embeds=block_embeds,431 attention_mask=block_attn,432 position_ids=block_position_ids,433 past_key_values=past_kv,434 use_cache=True,435 update_kv_cache=True,436 )437 past_kv = commit_out.past_key_values438 prev_last_logit = commit_out.logits[:, -1:, :]439 440 # NOTE: a previous null_ratio>0.3 early-stopping heuristic was441 # removed. It computed the ratio globally across the whole442 # explanation and, when tripped, force-filled every remaining443 # MASK with NULL — including MASKs in middle positions that444 # should have held real text — which cut short explanations445 # mid-sentence. Training always produces 192 value tokens446 # (real text + <|NULL|> padding at the tail) and the model447 # learned to emit NULL cleanly at the tail, so the final448 # NULL-strip below is sufficient. Cost: every sample now449 # denoises all 6 explanation blocks.450 451 # Post-process: strip NULL tokens from the output452 gen_tokens = x_t[0, original_input_length:].tolist()453 cleaned = [t for t in gen_tokens if t != null_id and t != mask_id]454 x_t = torch.cat([455 input_ids,456 torch.tensor([cleaned], device=self.device, dtype=torch.long)457 ], dim=1)458 459 gen_length = x_t.shape[1] - original_input_length460 461 if return_stats:462 stats = {463 "tokens_per_step": tokens_per_step,464 "total_steps": step,465 "gen_length": gen_length,466 "null_tokens_stripped": len(gen_tokens) - len(cleaned),467 "block_size": block_size,468 }469 return x_t, stats470 return x_t471 472@torch.no_grad()473 474 475# ---------------------------------------------------------------------------476# scaffold_speculative_sample — Scaffold Spec (SS)477# ---------------------------------------------------------------------------478 479def scaffold_speculative_sample(480 self,481 input_ids,482 tokenizer,483 block_size=32,484 max_tokens=1024,485 pixel_values=None,486 image_grid_thw=None,487 mask_id=151665,488 null_id=151666,489 threshold=0.9,490 stop_token=151645,491 explanation_block_size=32,492 explanation_max_blocks=6,493 return_stats=False,494 draft_temperature=0.0,495 verify_temperature=0.0,496):497 """498 Scaffold-aware self-speculative decoding.499 500 Minimal modification of standard self-spec501 (speculative_block_causal_sample_cache): scaffold (structural JSON)502 tokens are pre-filled in the draft block instead of MASK and503 auto-accepted during causal verification.504 505 Key design: uses *exactly the same* attention patterns as standard506 self-spec (block-diff for draft, **causal** for verify via507 auto eval_mask). Only the draft block content differs — scaffold508 positions carry known tokens instead of MASK, giving the draft509 better context while scaffold tokens are "free" during acceptance.510 """511 from .section_utils import (512 build_deep_json_scaffold,513 NULL_TOKEN_ID,514 )515 516 scaffold_tokens, section_ranges, scaffold_mask_list = build_deep_json_scaffold(517 tokenizer,518 mask_id=mask_id,519 null_id=null_id,520 explanation_block_size=explanation_block_size,521 explanation_max_blocks=explanation_max_blocks,522 )523 524 scaffold_len = len(scaffold_tokens)525 original_input_length = input_ids.shape[1]526 tokens_per_step = []527 self.model.bd_size = block_size528 529 _ss_profile = bool(os.environ.get("SS_PROFILE"))530 _ss_traj_start = section_ranges.get("trajectory", (None, None))[0]531 if _ss_profile:532 import time as _time533 torch.cuda.synchronize()534 _ss_t = {"start": _time.perf_counter()}535 _ss_marked_traj_start = False536 _ss_n_fwd_prefix = 0537 _ss_n_fwd_traj = 0538 539 # Pre-convert to tensors for vectorized operations in the loop540 scaffold_tok_t = torch.tensor(541 scaffold_tokens, device=self.device, dtype=torch.long542 )543 scaffold_is_fixed = torch.tensor(544 scaffold_mask_list, device=self.device, dtype=torch.bool545 )546 547 # ── Phase 1: Prefill prompt (identical to standard self-spec) ──548 output = self.forward(549 input_ids=input_ids,550 pixel_values=pixel_values,551 image_grid_thw=image_grid_thw,552 use_cache=True,553 update_kv_cache=True,554 )555 logits, past_key_values = output.logits, output.past_key_values556 if _ss_profile:557 torch.cuda.synchronize()558 _ss_t["after_prefill"] = _time.perf_counter()559 560 # First token — use scaffold token (always '{')561 next_token = torch.tensor(562 [[scaffold_tokens[0]]], device=self.device, dtype=torch.long563 )564 input_ids = torch.cat([input_ids, next_token], dim=1)565 tokens_per_step.append(1)566 scaffold_cursor = 1567 step = 1568 569 # ── Phase 2: Self-speculative decoding loop ──570 # Follows the exact same structure as571 # speculative_block_causal_sample_cache, with scaffold-aware draft.572 while scaffold_cursor < scaffold_len:573 if _ss_profile and (not _ss_marked_traj_start) and (574 _ss_traj_start is not None and scaffold_cursor >= _ss_traj_start575 ):576 torch.cuda.synchronize()577 _ss_t["enter_traj"] = _time.perf_counter()578 _ss_marked_traj_start = True579 prompt_length = input_ids.shape[1]580 n_draft = min(block_size - 1, scaffold_len - scaffold_cursor)581 582 # Build draft block: [seed, scaffold_or_MASK × n_draft]583 sc_end = scaffold_cursor + n_draft584 is_fixed = scaffold_is_fixed[scaffold_cursor:sc_end]585 draft_tensor = torch.where(586 is_fixed,587 scaffold_tok_t[scaffold_cursor:sc_end],588 mask_id,589 ).unsqueeze(0)590 x_t = torch.cat([input_ids[:, -1:], draft_tensor], dim=1)591 mask_idx = (x_t == mask_id)592 593 # ── Draft (block-diff bidirectional via auto eval_mask) ──594 logits = self.forward(595 input_ids=x_t,596 use_cache=True,597 past_key_values=past_key_values,598 update_kv_cache=False,599 eval_bd_size=block_size,600 ).logits601 tokens_per_step.append(0)602 step += 1603 604 # Shift logits (same as standard self-spec)605 logits = torch.cat([logits[:, :1, :], logits[:, :-1, :]], dim=1)606 if draft_temperature > 0:607 # Temperature sampling for draft diversity608 scaled = logits / draft_temperature609 draft_probs = torch.softmax(scaled, dim=-1)610 x_1 = torch.multinomial(611 draft_probs.view(-1, draft_probs.shape[-1]), num_samples=1612 ).view(draft_probs.shape[:-1])613 # Confidence uses unscaled probs for thresholding614 probs = torch.softmax(logits, dim=-1)615 x1_p = torch.gather(616 probs, dim=-1, index=x_1.unsqueeze(-1)617 ).squeeze(-1)618 else:619 x_1 = logits.argmax(dim=-1)620 probs = torch.softmax(logits, dim=-1)621 x1_p = torch.gather(622 probs, dim=-1, index=x_1.unsqueeze(-1)623 ).squeeze(-1)624 625 # Only fill MASK positions; scaffold positions keep their tokens626 x1_p = torch.where(mask_idx, x1_p, -torch.inf)627 unmask_idx = (x1_p > 0) # threshold=0 for draft filling628 629 if unmask_idx.sum() > 0:630 x_t[unmask_idx] = x_1[unmask_idx]631 else:632 # Fallback: fill most confident MASK633 mask_only_p = x1_p.clone()634 mask_only_p[~mask_idx] = -torch.inf635 if mask_only_p.max() > -torch.inf:636 best = mask_only_p.argmax()637 x_t.view(-1)[best] = x_1.view(-1)[best]638 639 # ── Verify (causal via auto eval_mask, commit to cache) ──640 output = self.forward(641 input_ids=x_t,642 use_cache=True,643 past_key_values=past_key_values,644 update_kv_cache=True,645 eval_bd_size=block_size,646 )647 past_key_values = output.past_key_values648 if verify_temperature > 0:649 verify_logits = output.logits / verify_temperature650 verify_probs = torch.softmax(verify_logits, dim=-1)651 ar_block_token = torch.multinomial(652 verify_probs.view(-1, verify_probs.shape[-1]), num_samples=1653 ).view(verify_probs.shape[:-1])654 else:655 ar_block_token = output.logits.argmax(dim=-1)656 657 # ── AR acceptance (scaffold positions auto-pass) ──658 ar_matches = (ar_block_token[0, :n_draft] == x_t[0, 1:n_draft + 1])659 accepted_token_num = 0660 for i in range(n_draft):661 if is_fixed[i] or ar_matches[i]:662 accepted_token_num += 1663 else:664 break665 accepted_token_num += 1 # bonus token666 667 tokens_per_step.append(accepted_token_num)668 669 # Force scaffold tokens at scaffold positions, AR predictions elsewhere670 accepted_ids = ar_block_token[:, :accepted_token_num].clone()671 acc_end = min(scaffold_cursor + accepted_token_num, scaffold_len)672 acc_fixed = scaffold_is_fixed[scaffold_cursor:acc_end]673 accepted_ids[0, :len(acc_fixed)][acc_fixed] = \674 scaffold_tok_t[scaffold_cursor:acc_end][acc_fixed]675 676 input_ids = torch.cat([input_ids, accepted_ids], dim=1)677 scaffold_cursor += accepted_token_num678 679 past_key_values = _crop_cache(past_key_values, input_ids.shape[1] - 1)680 681 step += 1682 683 # Stop conditions684 if input_ids.shape[1] - original_input_length > max_tokens:685 break686 if stop_token in input_ids[:, prompt_length:]:687 stop_token_idx = (688 input_ids[:, prompt_length:] == stop_token689 ).nonzero()[0][1]690 if (691 input_ids[:, prompt_length:prompt_length + stop_token_idx]692 == mask_id693 ).sum() == 0:694 break695 696 if _ss_profile:697 torch.cuda.synchronize()698 _ss_t["end"] = _time.perf_counter()699 _t_total = _ss_t["end"] - _ss_t["start"]700 _t_pre = _ss_t["after_prefill"] - _ss_t["start"]701 _t_traj_in = _ss_t.get("enter_traj")702 if _t_traj_in is not None:703 _t_prefix = _t_traj_in - _ss_t["after_prefill"]704 _t_traj = _ss_t["end"] - _t_traj_in705 else:706 _t_prefix = _ss_t["end"] - _ss_t["after_prefill"]707 _t_traj = 0.0708 print(709 f"[ss profile] total={_t_total*1000:.0f}ms "710 f"prefill={_t_pre*1000:.0f}ms "711 f"prefix-decode={_t_prefix*1000:.0f}ms "712 f"traj-decode={_t_traj*1000:.0f}ms",713 flush=True,714 )715 716 # ── Phase 3: Post-process — truncate at stop, strip NULL ──717 if stop_token in input_ids[:, original_input_length:]:718 stop_token_idx = (719 input_ids[:, original_input_length:] == stop_token720 ).nonzero()[0][1]721 input_ids = input_ids[722 :, :stop_token_idx + original_input_length + 1723 ]724 725 gen_tokens = input_ids[0, original_input_length:].tolist()726 cleaned = [t for t in gen_tokens if t != null_id and t != mask_id]727 output_ids = torch.cat(728 [729 input_ids[:, :original_input_length],730 torch.tensor(731 [cleaned], device=self.device, dtype=torch.long732 ),733 ],734 dim=1,735 )736 737 gen_length = output_ids.shape[1] - original_input_length738 739 if return_stats:740 stats = {741 "tokens_per_step": tokens_per_step,742 "total_steps": step,743 "gen_length": gen_length,744 "null_tokens_stripped": len(gen_tokens) - len(cleaned),745 "block_size": block_size,746 "method": "scaffold_speculative_v5",747 }748 return output_ids, stats749 return output_ids750 751@torch.no_grad()752 753 754# ---------------------------------------------------------------------------755# scaffold_spec_with_ss_multi_traj — SS multi-rollout inference scaling756# ---------------------------------------------------------------------------757 758def scaffold_spec_with_ss_multi_traj(759 self,760 input_ids,761 tokenizer,762 block_size=32,763 max_tokens=1024,764 pixel_values=None,765 image_grid_thw=None,766 mask_id=151665,767 null_id=151666,768 threshold=0.9,769 stop_token=151645,770 explanation_block_size=32,771 explanation_max_blocks=6,772 return_stats=False,773 num_traj_rollouts=4,774 traj_verify_temperature=0.5,775 traj_draft_temperature=0.0,776 merge_weights=None,777 batch_parallel=False,778):779 """Scaffold Spec with shared prefix + N SS rollouts on the trajectory section.780 781 Decoding pipeline:782 0) Prompt prefill [shared]783 1) Scaffold Spec for sections 1-3 (CoT) at verify_temp = 0 [shared, deterministic]784 2) Fork KV cache N times [O(N) memory]785 3) For each fork: continue Scaffold Spec on the trajectory786 section with verify_temperature = traj_verify_temperature787 (each rollout draws different samples in the AR-verify step788 because torch.multinomial is invoked with a global RNG).789 4) Parse all N trajectories and return their weighted mean.790 791 Cost: roughly 1 full SS pass (sections 1-3 are ~88%% of decoded tokens792 on our schema) + N x trajectory-only SS passes. For N = 4 this is793 ~1.5x the cost of a single SS, vs ~4x for naive sequential rerolling.794 795 If batch_parallel = True, the N trajectory rollouts are executed in a796 batched (batch_size = N) manner: one shared model.forward per797 speculative draft / verify step over an N-replicated trajectory798 suffix, which removes the per-rollout serial overhead at the cost of799 replicating the per-layer KV cache N-fold along the batch dimension.800 801 Returns: (output_ids, stats) if return_stats else output_ids.802 """803 from .section_utils import (804 build_deep_json_scaffold,805 SECTION_KEYS,806 )807 808 scaffold_tokens, section_ranges, scaffold_mask_list = build_deep_json_scaffold(809 tokenizer,810 mask_id=mask_id,811 null_id=null_id,812 explanation_block_size=explanation_block_size,813 explanation_max_blocks=explanation_max_blocks,814 )815 816 scaffold_len = len(scaffold_tokens)817 original_input_length = input_ids.shape[1]818 tokens_per_step = []819 self.model.bd_size = block_size820 821 scaffold_tok_t = torch.tensor(scaffold_tokens, device=self.device, dtype=torch.long)822 scaffold_is_fixed = torch.tensor(scaffold_mask_list, device=self.device, dtype=torch.bool)823 traj_start_in_scaffold = section_ranges["trajectory"][0]824 825 _profile = bool(os.environ.get("SS_MT_PROFILE"))826 if _profile:827 import time as _time828 torch.cuda.synchronize()829 _t_phase = {"start": _time.perf_counter()}830 _phase_clone_total = 0.0831 _phase_rollout_each = []832 833 # ── Phase 0: Prefill prompt ──834 output = self.forward(835 input_ids=input_ids, pixel_values=pixel_values,836 image_grid_thw=image_grid_thw,837 use_cache=True, update_kv_cache=True,838 )839 logits, past_key_values = output.logits, output.past_key_values840 if _profile:841 torch.cuda.synchronize()842 _t_phase["after_prefill"] = _time.perf_counter()843 844 next_token = torch.tensor(845 [[scaffold_tokens[0]]], device=self.device, dtype=torch.long,846 )847 input_ids = torch.cat([input_ids, next_token], dim=1)848 tokens_per_step.append(1)849 scaffold_cursor = 1850 step = 1851 852 # ── Phase 1: Scaffold Spec for non-trajectory sections (shared, vt=0) ──853 while scaffold_cursor < scaffold_len and scaffold_cursor < traj_start_in_scaffold:854 remaining_before_traj = traj_start_in_scaffold - scaffold_cursor855 n_draft = min(block_size - 1, remaining_before_traj)856 if n_draft <= 0:857 break858 859 sc_end = scaffold_cursor + n_draft860 is_fixed = scaffold_is_fixed[scaffold_cursor:sc_end]861 draft_tensor = torch.where(862 is_fixed, scaffold_tok_t[scaffold_cursor:sc_end], mask_id,863 ).unsqueeze(0)864 x_t = torch.cat([input_ids[:, -1:], draft_tensor], dim=1)865 mask_idx = (x_t == mask_id)866 867 # Draft (block-bidirectional)868 logits = self.forward(869 input_ids=x_t, use_cache=True,870 past_key_values=past_key_values,871 update_kv_cache=False, eval_bd_size=block_size,872 ).logits873 tokens_per_step.append(0)874 step += 1875 876 logits = torch.cat([logits[:, :1, :], logits[:, :-1, :]], dim=1)877 x_1 = logits.argmax(dim=-1)878 probs = torch.softmax(logits, dim=-1)879 x1_p = torch.gather(probs, dim=-1, index=x_1.unsqueeze(-1)).squeeze(-1)880 x1_p = torch.where(mask_idx, x1_p, -torch.inf)881 unmask_idx = (x1_p > 0)882 if unmask_idx.sum() > 0:883 x_t[unmask_idx] = x_1[unmask_idx]884 else:885 mask_only_p = x1_p.clone()886 mask_only_p[~mask_idx] = -torch.inf887 if mask_only_p.max() > -torch.inf:888 best = mask_only_p.argmax()889 x_t.view(-1)[best] = x_1.view(-1)[best]890 891 # Verify (causal, greedy)892 output = self.forward(893 input_ids=x_t, use_cache=True,894 past_key_values=past_key_values,895 update_kv_cache=True, eval_bd_size=block_size,896 )897 past_key_values = output.past_key_values898 ar_block_token = output.logits.argmax(dim=-1)899 900 ar_matches = (ar_block_token[0, :n_draft] == x_t[0, 1:n_draft + 1])901 accepted_token_num = 0902 for i in range(n_draft):903 if is_fixed[i] or ar_matches[i]:904 accepted_token_num += 1905 else:906 break907 accepted_token_num += 1908 909 max_accept = traj_start_in_scaffold - scaffold_cursor910 if accepted_token_num > max_accept:911 accepted_token_num = max_accept912 913 tokens_per_step.append(accepted_token_num)914 accepted_ids = ar_block_token[:, :accepted_token_num].clone()915 acc_end = min(scaffold_cursor + accepted_token_num, scaffold_len)916 acc_fixed = scaffold_is_fixed[scaffold_cursor:acc_end]917 accepted_ids[0, :len(acc_fixed)][acc_fixed] = \918 scaffold_tok_t[scaffold_cursor:acc_end][acc_fixed]919 920 input_ids = torch.cat([input_ids, accepted_ids], dim=1)921 scaffold_cursor += accepted_token_num922 past_key_values = _crop_cache(past_key_values, input_ids.shape[1] - 1)923 step += 1924 925 if input_ids.shape[1] - original_input_length > max_tokens:926 break927 928 if _profile:929 torch.cuda.synchronize()930 _t_phase["after_phase1"] = _time.perf_counter()931 932 # ── Phase 2: Fork KV cache N times (one per trajectory rollout) ──933 prefix_input_ids = input_ids.clone()934 prefix_len = prefix_input_ids.shape[1]935 936 def _clone_cache(kv):937 if _profile:938 torch.cuda.synchronize()939 _t0 = _time.perf_counter()940 cloned = []941 for layer_num in range(len(kv)):942 cloned.append(tuple(t.clone() for t in kv[layer_num]))943 ret = DynamicCache(cloned)944 if _profile:945 torch.cuda.synchronize()946 nonlocal _phase_clone_total947 _phase_clone_total += _time.perf_counter() - _t0948 return ret949 950 # ── Phase 3: N SS rollouts on trajectory section, each with vt > 0 ──951 # All rollouts start from the same prefix; randomness comes from952 # the multinomial calls in draft / verify (RNG is process-global).953 N = max(1, int(num_traj_rollouts))954 955 def _run_one_traj_rollout(start_kv, start_input_ids):956 """Continue Scaffold Spec from start_kv / start_input_ids over the957 trajectory section, applying traj_*_temperature. Returns the958 final ss_input_ids (with trajectory tokens appended) and the959 extracted trajectory value tokens."""960 local_kv = start_kv961 local_input = start_input_ids962 local_cursor = scaffold_cursor963 964 while local_cursor < scaffold_len:965 n_draft = min(block_size - 1, scaffold_len - local_cursor)966 sc_end = local_cursor + n_draft967 is_fixed = scaffold_is_fixed[local_cursor:sc_end]968 draft_tensor = torch.where(969 is_fixed, scaffold_tok_t[local_cursor:sc_end], mask_id,970 ).unsqueeze(0)971 x_t = torch.cat([local_input[:, -1:], draft_tensor], dim=1)972 mask_idx = (x_t == mask_id)973 974 # Draft (block-bidirectional, optionally temp-sampled)975 draft_logits = self.forward(976 input_ids=x_t, use_cache=True, past_key_values=local_kv,977 update_kv_cache=False, eval_bd_size=block_size,978 ).logits979 draft_logits = torch.cat(980 [draft_logits[:, :1, :], draft_logits[:, :-1, :]], dim=1,981 )982 if traj_draft_temperature > 0:983 scaled = draft_logits / traj_draft_temperature984 draft_probs = torch.softmax(scaled, dim=-1)985 x_1 = torch.multinomial(986 draft_probs.view(-1, draft_probs.shape[-1]),987 num_samples=1,988 ).view(draft_probs.shape[:-1])989 else:990 x_1 = draft_logits.argmax(dim=-1)991 probs = torch.softmax(draft_logits, dim=-1)992 x1_p = torch.gather(993 probs, dim=-1, index=x_1.unsqueeze(-1),994 ).squeeze(-1)995 x1_p = torch.where(mask_idx, x1_p, -torch.inf)996 unmask_idx = (x1_p > 0)997 if unmask_idx.sum() > 0:998 x_t[unmask_idx] = x_1[unmask_idx]999 else:1000 mask_only_p = x1_p.clone()1001 mask_only_p[~mask_idx] = -torch.inf1002 if mask_only_p.max() > -torch.inf:1003 x_t.view(-1)[mask_only_p.argmax()] = \1004 x_1.view(-1)[mask_only_p.argmax()]1005 1006 # Verify (causal, optionally temp-sampled)1007 v_out = self.forward(1008 input_ids=x_t, use_cache=True, past_key_values=local_kv,1009 update_kv_cache=True, eval_bd_size=block_size,1010 )1011 local_kv = v_out.past_key_values1012 if traj_verify_temperature > 0:1013 v_logits = v_out.logits / traj_verify_temperature1014 v_probs = torch.softmax(v_logits, dim=-1)1015 ar_block_token = torch.multinomial(1016 v_probs.view(-1, v_probs.shape[-1]),1017 num_samples=1,1018 ).view(v_probs.shape[:-1])1019 else:1020 ar_block_token = v_out.logits.argmax(dim=-1)1021 1022 ar_matches = (ar_block_token[0, :n_draft] == x_t[0, 1:n_draft + 1])1023 accepted_token_num = 01024 for i in range(n_draft):1025 if is_fixed[i] or ar_matches[i]:1026 accepted_token_num += 11027 else:1028 break1029 accepted_token_num += 11030 1031 accepted_ids = ar_block_token[:, :accepted_token_num].clone()1032 acc_end = min(local_cursor + accepted_token_num, scaffold_len)1033 acc_fixed = scaffold_is_fixed[local_cursor:acc_end]1034 accepted_ids[0, :len(acc_fixed)][acc_fixed] = \1035 scaffold_tok_t[local_cursor:acc_end][acc_fixed]1036 1037 local_input = torch.cat([local_input, accepted_ids], dim=1)1038 local_cursor += accepted_token_num1039 local_kv = _crop_cache(local_kv, local_input.shape[1] - 1)1040 1041 if local_input.shape[1] - original_input_length > max_tokens:1042 break1043 if stop_token in local_input[:, prefix_len:]:1044 st_idx = (local_input[:, prefix_len:] == stop_token).nonzero()1045 if st_idx.numel() > 0:1046 cand_st = st_idx[0][1].item()1047 if (local_input[:, prefix_len:prefix_len + cand_st] == mask_id).sum() == 0:1048 break1049 1050 traj_values = [1051 t for i, t in enumerate(local_input[0, original_input_length:].tolist())1052 if i >= traj_start_in_scaffold and i < scaffold_len1053 and not scaffold_mask_list[i] and t != null_id and t != mask_id1054 ]1055 return local_input, traj_values1056 1057 # Sequential N rollouts (Option A; batch_parallel=False).1058 rollout_inputs = []1059 rollout_traj_values = []1060 for _i in range(N):1061 if _profile:1062 torch.cuda.synchronize()1063 _t_r0 = _time.perf_counter()1064 cand_kv = _clone_cache(past_key_values)1065 cand_input = prefix_input_ids.clone()1066 cand_input, traj_vals = _run_one_traj_rollout(cand_kv, cand_input)1067 rollout_inputs.append(cand_input)1068 rollout_traj_values.append(traj_vals)1069 step += 11070 if _profile:1071 torch.cuda.synchronize()1072 _phase_rollout_each.append(_time.perf_counter() - _t_r0)1073 1074 if _profile:1075 torch.cuda.synchronize()1076 _t_phase["after_rollouts"] = _time.perf_counter()1077 _t_total = _t_phase["after_rollouts"] - _t_phase["start"]1078 _t_pre = _t_phase["after_prefill"] - _t_phase["start"]1079 _t_p1 = _t_phase["after_phase1"] - _t_phase["after_prefill"]1080 _t_rolls = _t_phase["after_rollouts"] - _t_phase["after_phase1"]1081 print(1082 f"[ss_mt profile] total={_t_total*1000:.0f}ms "1083 f"prefill(P0)={_t_pre*1000:.0f}ms "1084 f"prefix-decode(P1)={_t_p1*1000:.0f}ms "1085 f"rollouts(P2+P3)={_t_rolls*1000:.0f}ms "1086 f"of which kv-clone={_phase_clone_total*1000:.0f}ms "1087 f"per-rollout={[f'{r*1000:.0f}' for r in _phase_rollout_each]}ms",1088 flush=True,1089 )1090 1091 # ── Phase 4: Parse all rollouts, weighted-merge waypoints ──1092 def _decode_trajectory(traj_tokens):1093 text = tokenizer.decode(traj_tokens, skip_special_tokens=False)1094 text = text.replace("<|NULL|>", "").strip()1095 coords = re.findall(r"[+-]?\d+\.?\d*", text)1096 wps = []1097 for i in range(0, len(coords) - 1, 2):1098 wps.append([float(coords[i]), float(coords[i + 1])])1099 return wps1100 1101 rollout_waypoints = [_decode_trajectory(v) for v in rollout_traj_values]1102 1103 if merge_weights is None or len(merge_weights) != N:1104 ws = [1.0 / N] * N1105 else:1106 total = sum(merge_weights)1107 ws = [w / total for w in merge_weights]1108 1109 if rollout_waypoints and all(len(w) > 0 for w in rollout_waypoints):1110 n_wp = min(len(w) for w in rollout_waypoints)1111 merged_waypoints = []1112 for i in range(n_wp):1113 mx = sum(ws[c] * rollout_waypoints[c][i][0] for c in range(N))1114 my = sum(ws[c] * rollout_waypoints[c][i][1] for c in range(N))1115 merged_waypoints.append([mx, my])1116 else:1117 merged_waypoints = next(1118 (w for w in rollout_waypoints if w), [],1119 )1120 1121 # Output text: take rollout 0's full text but replace its trajectory1122 # with the merged waypoints.1123 base_input = rollout_inputs[0]1124 if stop_token in base_input[:, original_input_length:]:1125 st_idx = (base_input[:, original_input_length:] == stop_token).nonzero()[0][1]1126 base_input = base_input[:, :st_idx + original_input_length + 1]1127 base_raw_tokens = base_input[0, original_input_length:].tolist()1128 base_cleaned = [t for t in base_raw_tokens if t != null_id and t != mask_id]1129 base_null_stripped = len(base_raw_tokens) - len(base_cleaned)1130 base_text = tokenizer.decode(base_cleaned, skip_special_tokens=False)1131 1132 traj_parts = [1133 f"[{x:+07.2f},{y:+06.2f}]" for x, y in merged_waypoints1134 ]1135 merged_traj_str = "[" + ", ".join(traj_parts) + "]"1136 replaced_text = re.sub(1137 r'("trajectory"\s*:\s*")(\[\[.*?\]\])',1138 r"\g<1>" + merged_traj_str, base_text,1139 )1140 1141 merged_tokens = tokenizer.encode(replaced_text, add_special_tokens=False)1142 output_ids = torch.cat([1143 input_ids[:, :original_input_length],1144 torch.tensor([merged_tokens], device=self.device, dtype=torch.long),1145 ], dim=1)1146 1147 gen_length = output_ids.shape[1] - original_input_length1148 1149 if return_stats:1150 stats = {1151 "tokens_per_step": tokens_per_step,1152 "total_steps": step,1153 "gen_length": gen_length,1154 "null_tokens_stripped": base_null_stripped,1155 "block_size": block_size,1156 "method": "scaffold_spec_with_ss_multi_traj",1157 "num_traj_rollouts": N,1158 "traj_verify_temperature": traj_verify_temperature,1159 "rollout_waypoints": rollout_waypoints,1160 "merged_waypoints": merged_waypoints,1161 "merge_weights": ws,1162 }1163 return output_ids, stats1164 return output_ids1165 1166@torch.no_grad()1167 1168 1169# ---------------------------------------------------------------------------1170# Bind decoding methods onto the model class.1171#1172# ``modeling.py`` imports this module at the bottom of the file, after the1173# ``Fast_dDriveForConditionalGeneration`` class has been defined. We1174# attach the three decoding paths as ordinary methods so callers can invoke1175# them as ``model.mdm_sample_deep_scaffold(...)`` etc. without any extra1176# registration step.1177# ---------------------------------------------------------------------------1178 1179def attach_generation_methods(cls):1180 """Attach the three release decoding paths as methods of ``cls``."""1181 cls.mdm_sample_deep_scaffold = mdm_sample_deep_scaffold1182 cls.scaffold_speculative_sample = scaffold_speculative_sample1183 cls.scaffold_spec_with_ss_multi_traj = scaffold_spec_with_ss_multi_traj1184 return cls1185 1186 1187__all__ = [1188 "mdm_sample_deep_scaffold",1189 "scaffold_speculative_sample",1190 "scaffold_spec_with_ss_multi_traj",1191 "attach_generation_methods",1192]1193 