Team Ai
Modelpublic

Efficient-Large-Model/Fast-dDrive

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
4likes167downloads
generation_utils.py1193 linesDownload Raw Back to root
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