Team Ai
Modelpublic

Efficient-Large-Model/Fast-dDrive

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
4likes167downloads
section_utils.py804 linesDownload Raw Back to root
1"""2Section-aware block scheduling for JSON structured output.3 4Inspired by the S3 (Self-adaptive Schema Scaffolding) paper (arXiv:2507.04504),5this module provides utilities to:61. Parse tokenized JSON output into sections (critical_objects, explanation, etc.)72. Assign section-aware block indices for variable block sizes per section83. Build JSON scaffolds for inference (pre-fill structural tokens)9 10The DVLM-AD output schema has 4 sections:11  - critical_objects: ~88 tokens (12 yes/no fields, nearly constant)12  - explanation: ~114 tokens (variable, 72-172)13  - future_meta_behavior: ~40 tokens (nearly constant)14  - trajectory: ~80 tokens (nearly constant)15"""16 17import torch18from typing import Dict, List, Optional, Tuple19import math20 21 22# Ordered list of section keys as they appear in the JSON output23SECTION_KEYS = [24    "critical_objects",25    "explanation",26    "future_meta_behavior",27    "trajectory",28]29 30# Default token budgets per section (based on training data analysis)31DEFAULT_TOKEN_BUDGETS = {32    "critical_objects": 88,33    "explanation": 128,34    "future_meta_behavior": 40,35    "trajectory": 80,36}37 38# Default steps per section39DEFAULT_SECTION_STEPS = {40    "critical_objects": 1,41    "explanation": 3,42    "future_meta_behavior": 1,43    "trajectory": 1,44}45 46 47 48 49def _v1_removed_parse_json_sections(*args, **kwargs):50    raise NotImplementedError("DS v1 parse_json_sections has been removed. Use deep scaffold v2.")51 52 53def _v1_removed_compute_section_block_idx(*args, **kwargs):54    raise NotImplementedError("DS v1 compute_section_block_idx has been removed. Use compute_section_block_idx_deep_static.")55 56 57def _v1_removed_build_json_scaffold(*args, **kwargs):58    raise NotImplementedError("DS v1 build_json_scaffold has been removed. Use build_deep_json_scaffold.")59 60 61def _v1_removed_compute_section_block_sizes(*args, **kwargs):62    raise NotImplementedError("DS v1 compute_section_block_sizes has been removed.")63 64 65def build_static_scaffold_sequences(tokenizer) -> Dict[str, List[int]]:66    """Pre-compute token sequences for top-level JSON boundary matching.67 68    Used internally by :func:`build_deep_scaffold_sequences`.69    """70    return {71        "prefix": tokenizer.encode('{"critical_objects":', add_special_tokens=False),72        "between_co_exp": tokenizer.encode(' "explanation":', add_special_tokens=False),73        "between_exp_fmb": tokenizer.encode(' "future_meta_behavior":', add_special_tokens=False),74        "between_fmb_traj": tokenizer.encode(' "trajectory":', add_special_tokens=False),75    }76 77 78def _v1_removed_compute_section_block_idx_static(*args, **kwargs):79    raise NotImplementedError("DS v1 compute_section_block_idx_static has been removed. Use compute_section_block_idx_deep_static.")80 81 82# Backward-compatible aliases so stale imports produce clear errors83parse_json_sections = _v1_removed_parse_json_sections84compute_section_block_idx = _v1_removed_compute_section_block_idx85build_json_scaffold = _v1_removed_build_json_scaffold86compute_section_block_sizes = _v1_removed_compute_section_block_sizes87compute_section_block_idx_static = _v1_removed_compute_section_block_idx_static88 89 90# ═══════════════════════════════════════════════════════════════91# Deep scaffold v2: constants and utilities92# ═══════════════════════════════════════════════════════════════93 94NULL_TOKEN_ID = 15166695 96# critical_objects: 12 sub-keys, each value is exactly 1 token (yes=9693 / no=2152)97CRITICAL_OBJECTS_SUBKEYS = [98    "nearby_vehicle", "pedestrian", "cyclist", "construction",99    "traffic_element", "weather_condition", "road_hazard",100    "emergency_vehicle", "animal", "special_vehicle",101    "conflicting_vehicle", "door_opening_vehicle",102]103 104# future_meta_behavior: each sub-key value is exactly 3 tokens105# (e.g., "keep speed" → [4867, 4732, 151667] or "go straight" → [2849, 7833, 151667])106FMB_VALUE_BUDGET = 3107 108 109def build_deep_json_scaffold(110    tokenizer,111    section_token_budgets: Optional[Dict[str, int]] = None,112    mask_id: Optional[int] = None,113    null_id: Optional[int] = None,114    explanation_block_size: int = 32,115    explanation_max_blocks: int = 6,116) -> Tuple[List[int], Dict[str, Tuple[int, int]], List[int]]:117    """Build a deep JSON scaffold for inference (v2).118 119    Constructs a template response by building a Python dict and120    processing it through the **exact same pipeline** as the training121    dataloader (``multi_modal_dataset.py``):122 123    1. Build a realistic dict with placeholder values.124    2. Pad explanation with ``<|NULL|>`` to ``exp_budget`` tokens.125    3. Pad FMB values with ``<|NULL|>`` to 3 tokens each.126    4. Normalize trajectory to ``+XXX.XX`` format with spaces.127    5. Serialize with ``json.dumps(obj, ensure_ascii=False)``.128    6. Tokenize the whole string as one piece.129    7. Run ``compute_section_block_idx_deep_static`` to get scaffold/value.130    8. Replace value positions with MASK tokens.131 132    This guarantees identical BPE tokenization as training data.133 134    Returns135    -------136    scaffold_tokens : list[int]137        Token IDs with MASK at value positions.138    section_ranges : dict139        Section name -> (start, end) within scaffold_tokens.140    scaffold_mask : list[int]141        0 = value (to denoise), 1 = scaffold (frozen).142    """143    import torch as _torch144    import json as _json145    import re as _re146 147    if mask_id is None:148        mask_tok = tokenizer.encode("|<MASK>|", add_special_tokens=False)149        mask_id = mask_tok[0] if len(mask_tok) == 1 else 151665150    if null_id is None:151        null_id = NULL_TOKEN_ID152 153    exp_budget = explanation_block_size * explanation_max_blocks  # default 192154 155    # ── Step 1: Build a Python dict matching training data structure ──156    # Placeholder explanation text (will be replaced with MASK anyway).157    filler_explanation = (158        "The ego vehicle is driving forward on the road. "159        "There are nearby vehicles ahead that may affect the path. "160        "No pedestrians or cyclists are detected in the immediate area. "161        "The road conditions appear normal with no hazards present. "162        "Speed adjustment may be needed based on the traffic ahead. "163        "No lateral maneuvering is required at this time."164    )165 166    def _build_template(n_exp_nulls: int) -> str:167        """Build template via json.dumps — identical to dataloader output."""168        null_pad = "<|NULL|>" * n_exp_nulls169 170        data_obj = {171            "critical_objects": {172                "nearby_vehicle": "no", "pedestrian": "no", "cyclist": "no",173                "construction": "no", "traffic_element": "no",174                "weather_condition": "no", "road_hazard": "no",175                "emergency_vehicle": "no", "animal": "no",176                "special_vehicle": "no", "conflicting_vehicle": "no",177                "door_opening_vehicle": "no",178            },179            "explanation": filler_explanation + null_pad,180            "future_meta_behavior": {181                "longitudinal": "come to stop",182                "lateral": "go straight<|NULL|>",183            },184            # Raw trajectory — will be normalized below185            "trajectory": "[[+14.70,-00.04], [+29.55,-00.21], [+44.51,-00.56], [+59.50,-01.06], [+74.39,-01.69]]",186        }187 188        # Apply exact same trajectory normalization as dataloader (lines 851-863)189        traj = data_obj["trajectory"]190        def _fmt_coord(m):191            raw = m.group(0)192            sign = raw[0]193            num = float(raw[1:])194            return f"{sign}{num:06.2f}"195        traj = _re.sub(r'[+-]\d+\.\d+', _fmt_coord, traj)196        traj = _re.sub(r',([+-])', r', \1', traj)197        traj = _re.sub(r'\[([+-])', r'[ \1', traj)198        data_obj["trajectory"] = traj199 200        # Serialize with json.dumps — identical to dataloader line 865201        return _json.dumps(data_obj, ensure_ascii=False)202 203    # ── Step 2: Iteratively adjust NULL count for exp_budget ──204    deep_seqs = build_deep_scaffold_sequences(tokenizer)205    top_seqs = deep_seqs["top"]206 207    def _count_exp_value_tokens(tok_list):208        """Count explanation VALUE tokens (between boundary patterns)."""209        co_exp_pat = top_seqs["between_co_exp"]210        exp_fmb_pat = top_seqs["between_exp_fmb"]211        co_exp_pos = _find_subseq(tok_list, co_exp_pat, 0)212        if co_exp_pos < 0:213            return None214        exp_start = co_exp_pos + len(co_exp_pat)215        exp_fmb_pos = _find_subseq(tok_list, exp_fmb_pat, exp_start)216        if exp_fmb_pos < 0:217            return None218        # exp_start..exp_fmb_pos includes opening/closing quotes (scaffold)219        # value tokens = total - 2 (quotes)220        return (exp_fmb_pos - exp_start) - 2221 222    # Measure base explanation tokens (no NULLs)223    toks_0 = tokenizer.encode(_build_template(0), add_special_tokens=False)224    base_exp = _count_exp_value_tokens(toks_0)225    if base_exp is not None:226        needed_nulls = max(0, exp_budget - base_exp)227    else:228        needed_nulls = exp_budget // 2  # fallback229 230    # Build and measure, adjust once231    template = _build_template(needed_nulls)232    template_tokens = tokenizer.encode(template, add_special_tokens=False)233    actual_exp = _count_exp_value_tokens(template_tokens)234    if actual_exp is not None and actual_exp != exp_budget:235        needed_nulls = max(0, needed_nulls + (exp_budget - actual_exp))236        template = _build_template(needed_nulls)237        template_tokens = tokenizer.encode(template, add_special_tokens=False)238 239    # ── Step 3: Run training scaffold detection ──240    prompt_len = 10241    all_tokens = [1] * prompt_len + template_tokens242    labels_list = [-100] * prompt_len + template_tokens243 244    labels = _torch.tensor([labels_list])245    token_ids = _torch.tensor([all_tokens])246 247    _, _, _, scaffold_mask_tensor, _ = compute_section_block_idx_deep_static(248        labels, token_ids, deep_seqs, fallback_block_size=32,249    )250 251    # ── Step 4: Extract scaffold/value and replace value with MASK ──252    scaffold_tokens = list(template_tokens)253    scaffold_mask_list: List[int] = []254    for i in range(len(template_tokens)):255        abs_pos = prompt_len + i256        is_scaffold = scaffold_mask_tensor[abs_pos].item()257        scaffold_mask_list.append(1 if is_scaffold else 0)258 259    for i in range(len(scaffold_tokens)):260        if scaffold_mask_list[i] == 0:261            scaffold_tokens[i] = mask_id262 263    # ── Step 5: Compute section ranges ──264    section_ranges: Dict[str, Tuple[int, int]] = {}265    boundary_order = [266        ("prefix", "critical_objects"),267        ("between_co_exp", "explanation"),268        ("between_exp_fmb", "future_meta_behavior"),269        ("between_fmb_traj", "trajectory"),270    ]271 272    search_from = 0273    prev_section_name = None274    prev_value_start = None275 276    for boundary_key, section_name in boundary_order:277        pattern = top_seqs.get(boundary_key)278        if pattern is None:279            continue280        pos = _find_subseq(template_tokens, pattern, search_from)281        if pos < 0:282            continue283        if prev_section_name is not None and prev_value_start is not None:284            section_ranges[prev_section_name] = (prev_value_start, pos)285        value_start = pos + len(pattern)286        prev_section_name = section_name287        prev_value_start = value_start288        search_from = value_start289 290    if prev_section_name is not None and prev_value_start is not None:291        section_ranges[prev_section_name] = (prev_value_start, len(template_tokens))292 293    return scaffold_tokens, section_ranges, scaffold_mask_list294 295 296def _find_subseq(seq: List[int], pattern: List[int], start: int = 0) -> int:297    """Find first occurrence of *pattern* in *seq* starting at *start*. Returns -1 if not found."""298    n = len(pattern)299    for i in range(start, len(seq) - n + 1):300        if seq[i : i + n] == pattern:301            return i302    return -1303 304 305def build_deep_scaffold_sequences(tokenizer) -> Dict[str, object]:306    """307    Pre-compute token sequences for deep scaffold matching.308 309    Returns a dict with:310      - Top-level boundary patterns (same as build_static_scaffold_sequences)311      - Sub-key patterns for critical_objects, future_meta_behavior, trajectory312    """313    seqs: Dict[str, object] = {}314 315    # ── Top-level boundaries (reuse existing) ──316    seqs["top"] = build_static_scaffold_sequences(tokenizer)317 318    # ── critical_objects sub-key patterns ──319    # In context, CO value starts with ' {"nearby_vehicle": "yes", ...'320    # Token 5212 = ' {"' merges space+brace+quote in context321    # First entry: ' {"key": "'322    # Subsequent: '", "key": "'  (token 497='","' merges quote+comma)323    co_patterns = []324    for i, key in enumerate(CRITICAL_OBJECTS_SUBKEYS):325        if i == 0:326            pattern = tokenizer.encode(' {"' + key + '": "', add_special_tokens=False)327        else:328            pattern = tokenizer.encode('", "' + key + '": "', add_special_tokens=False)329        co_patterns.append({"key": key, "pattern": pattern, "index": i})330    seqs["co_subkeys"] = co_patterns331    seqs["co_closing"] = tokenizer.encode('"}', add_special_tokens=False)332    # json.dumps produces "}," which may merge into a single token333    seqs["co_closing_comma"] = tokenizer.encode('"},', add_special_tokens=False)334 335    # ── future_meta_behavior sub-key patterns ──336    # After dataloader processing (mdm markers removed, NULLs cleaned):337    #   ' {"longitudinal": "keep speed", "lateral": "go straight"}'338    # Scaffold = everything except the value content between quotes.339    seqs["fmb_prefix"] = tokenizer.encode(' {"longitudinal": "', add_special_tokens=False)340    seqs["fmb_closing"] = tokenizer.encode('"}', add_special_tokens=False)341    seqs["fmb_closing_comma"] = tokenizer.encode('"},', add_special_tokens=False)342    # Between longitudinal value and lateral value: '", "lateral": "'343    seqs["fmb_between"] = tokenizer.encode('", "lateral": "', add_special_tokens=False)344 345    # ── trajectory structure patterns ──346    # After dataloader processing (no mdm markers), traj is:347    #   ' "[[+14.70,-00.04], [+29.55,-00.21], ...]"'348    seqs["traj_open"] = tokenizer.encode(' "[[', add_special_tokens=False)349    # After dataloader inserts spaces (e.g. [+14.70,-00.04] → [ +14.70, -00.04]),350    # tokens split cleanly: '],'(1125), ' ['(508), ','(11) are all independent.351    seqs["traj_wp_sep"] = tokenizer.encode('],', add_special_tokens=False)   # [1125]352    seqs["traj_wp_open"] = tokenizer.encode(' [', add_special_tokens=False)  # [508]353    seqs["traj_coord_comma"] = tokenizer.encode(',', add_special_tokens=False)  # [11]354    seqs["traj_close"] = tokenizer.encode(']]"}', add_special_tokens=False)355    seqs["traj_close_split"] = tokenizer.encode(']]"', add_special_tokens=False)356    seqs["traj_close_split2"] = tokenizer.encode(']]', add_special_tokens=False)357    # Trajectory-only output support, e.g. {"trajectory": "..."}.358    seqs["traj_only_boundaries"] = [359        tokenizer.encode('{"trajectory":', add_special_tokens=False),360        tokenizer.encode(' {"trajectory":', add_special_tokens=False),361        tokenizer.encode('"trajectory":', add_special_tokens=False),362        tokenizer.encode(' "trajectory":', add_special_tokens=False),363    ]364 365    return seqs366 367 368def _mark_scaffold_range(scaffold_positions: List[int], start: int, length: int):369    """Add positions [start, start+length) to scaffold_positions."""370    for i in range(length):371        scaffold_positions.append(start + i)372 373 374def compute_section_block_idx_deep_static(375    labels: torch.Tensor,376    token_ids: torch.Tensor,377    deep_scaffold_sequences: Dict[str, object],378    fallback_block_size: int = 32,379) -> Tuple[torch.Tensor, torch.Tensor, int, torch.Tensor]:380    """381    Deep-scaffold v2 block index computation.382 383    Freezes sub-keys within sections:384      - critical_objects: only yes/no values are denoised385      - future_meta_behavior: only value tokens are denoised386      - trajectory: only coordinate digits are denoised387      - explanation: all content is denoised388 389    Block count per section is computed dynamically:390    ``n_blocks = ceil(num_value_tokens / fallback_block_size)``.391 392    Args:393        labels:                [B, seq_len]394        token_ids:             [B, seq_len]395        deep_scaffold_sequences: output of ``build_deep_scaffold_sequences``396        fallback_block_size:   block size (bd_size), default 32397 398    Returns:399        response_block_idx, turn_idx, n_blocks, scaffold_mask400    """401    labels_single = labels[0]402    token_list = token_ids[0].tolist()403    seq_len = labels_single.shape[0]404    device = labels.device405 406    response_mask = (labels_single != -100)407    response_block_idx = torch.full((seq_len,), -1, device=device, dtype=torch.int64)408    turn_idx = torch.zeros((seq_len,), device=device, dtype=torch.int64)409    scaffold_mask = torch.zeros((seq_len,), device=device, dtype=torch.bool)410 411    response_positions = response_mask.nonzero(as_tuple=True)[0]412    if len(response_positions) == 0:413        return response_block_idx, turn_idx, 0, scaffold_mask414 415    resp_start = response_positions[0].item()416    resp_end = response_positions[-1].item() + 1417    effective_resp_end = resp_end418    resp_tokens = token_list[resp_start:resp_end]419 420    top_seqs = deep_scaffold_sequences["top"]421 422    # ── Step 1: Find top-level section boundaries (same as static version) ──423    boundary_order = [424        ("prefix",           "critical_objects"),425        ("between_co_exp",   "explanation"),426        ("between_exp_fmb",  "future_meta_behavior"),427        ("between_fmb_traj", "trajectory"),428    ]429 430    sections: Dict[str, Tuple[int, int]] = {}431    scaffold_positions: List[int] = []432    # Top-level boundary scaffold tokens should belong to the *following*433    # section's first block (e.g. `"explanation":` -> explanation block 0).434    boundary_scaffold_to_section: Dict[str, List[int]] = {}435 436    search_from = 0437    prev_section_name: Optional[str] = None438    prev_value_start: Optional[int] = None439 440    for boundary_key, section_name in boundary_order:441        pattern = top_seqs.get(boundary_key)442        if pattern is None:443            continue444        pos = _find_subseq(resp_tokens, pattern, search_from)445        if pos < 0:446            continue447 448        if prev_section_name is not None and prev_value_start is not None:449            sections[prev_section_name] = (prev_value_start, pos)450 451        _mark_scaffold_range(scaffold_positions, pos, len(pattern))452        boundary_scaffold_to_section.setdefault(section_name, []).extend(453            list(range(pos, pos + len(pattern)))454        )455 456        value_start = pos + len(pattern)457        prev_section_name = section_name458        prev_value_start = value_start459        search_from = value_start460 461    if prev_section_name is not None and prev_value_start is not None:462        sections[prev_section_name] = (prev_value_start, len(resp_tokens))463 464    # New dataset compatibility: response may contain only trajectory.465    # If the 4-section boundaries are not found, try direct trajectory key match.466    if "trajectory" not in sections:467        traj_only_patterns = deep_scaffold_sequences.get("traj_only_boundaries", [])468        # Reuse legacy boundary pattern as additional fallback (contains469        # `"trajectory":` in old-format responses).470        between_fmb_traj = top_seqs.get("between_fmb_traj")471        if between_fmb_traj:472            traj_only_patterns = list(traj_only_patterns) + [between_fmb_traj]473 474        traj_pos = -1475        traj_pat: Optional[List[int]] = None476        for pat in traj_only_patterns:477            if not pat:478                continue479            pos = _find_subseq(resp_tokens, pat, 0)480            if pos >= 0:481                traj_pos = pos482                traj_pat = pat483                break484 485        if traj_pos >= 0 and traj_pat is not None:486            _mark_scaffold_range(scaffold_positions, traj_pos, len(traj_pat))487            boundary_scaffold_to_section.setdefault("trajectory", []).extend(488                list(range(traj_pos, traj_pos + len(traj_pat)))489            )490            sections["trajectory"] = (traj_pos + len(traj_pat), len(resp_tokens))491    # print(f"sections: {sections}")492    # ── Step 2: Deep scaffold within critical_objects ──493    if "critical_objects" in sections:494        co_start, co_end = sections["critical_objects"]495        co_tokens = resp_tokens[co_start:co_end]496 497        co_search = 0498        for entry in deep_scaffold_sequences["co_subkeys"]:499            pattern = entry["pattern"]500            pos = _find_subseq(co_tokens, pattern, co_search)501            if pos < 0:502                continue503            _mark_scaffold_range(scaffold_positions, co_start + pos, len(pattern))504            # The single value token is right after the pattern — skip it505            co_search = pos + len(pattern) + 1506 507        # Mark closing '"}' or "}," as scaffold508        co_close = deep_scaffold_sequences["co_closing"]509        close_pos = _find_subseq(co_tokens, co_close,510                                  max(0, len(co_tokens) - len(co_close) - 2))511        if close_pos >= 0:512            _mark_scaffold_range(scaffold_positions, co_start + close_pos, len(co_close))513        else:514            # json.dumps may produce "}," as a single token515            co_close_comma = deep_scaffold_sequences.get("co_closing_comma")516            if co_close_comma:517                close_pos = _find_subseq(co_tokens, co_close_comma,518                                          max(0, len(co_tokens) - len(co_close_comma) - 2))519                if close_pos >= 0:520                    _mark_scaffold_range(scaffold_positions, co_start + close_pos, len(co_close_comma))521 522    # ── Step 2b: Explanation opening/closing quotes as scaffold ──523    # Explanation content is all VALUE, but the surrounding quotes must be524    # SCAFFOLD so that VALUE tokens are exactly block-aligned (multiple of bd_size).525    if "explanation" in sections:526        exp_start, exp_end = sections["explanation"]527        if exp_start < exp_end:528            # Opening quote: first token of explanation section (e.g. ' "')529            scaffold_positions.append(exp_start)530            # Closing quote+comma: last token (e.g. '",')531            scaffold_positions.append(exp_start + (exp_end - exp_start) - 1)532 533    # ── Step 3: Deep scaffold within future_meta_behavior ──534    # After dataloader processing, FMB has no <|mdm_start|>/<|mdm_end|> markers.535    # Format: ' {"longitudinal": "keep speed", "lateral": "go straight"}'536    # Strategy: use fmb_prefix to find start, fmb_between to split long/lat values,537    # and fmb_closing to find end. Everything except value content is scaffold.538    if "future_meta_behavior" in sections:539        fmb_start, fmb_end = sections["future_meta_behavior"]540        fmb_tokens = resp_tokens[fmb_start:fmb_end]541 542        fmb_scaffold_positions = set()543 544        # 1. Mark fmb_prefix as scaffold: ' {"longitudinal": "'545        fmb_prefix = deep_scaffold_sequences["fmb_prefix"]546        prefix_pos = _find_subseq(fmb_tokens, fmb_prefix, 0)547        if prefix_pos >= 0:548            for i in range(prefix_pos, prefix_pos + len(fmb_prefix)):549                fmb_scaffold_positions.add(i)550 551            long_value_start = prefix_pos + len(fmb_prefix)552 553            # 2. Mark fmb_between as scaffold: '", "lateral": "'554            fmb_between = deep_scaffold_sequences.get("fmb_between")555            if fmb_between:556                between_pos = _find_subseq(fmb_tokens, fmb_between, long_value_start)557                if between_pos >= 0:558                    for i in range(between_pos, between_pos + len(fmb_between)):559                        fmb_scaffold_positions.add(i)560 561                    lat_value_start = between_pos + len(fmb_between)562 563                    # 3. Mark closing '"}'  or "}," as scaffold564                    fmb_close = deep_scaffold_sequences["fmb_closing"]565                    close_pos = _find_subseq(fmb_tokens, fmb_close,566                                              max(0, len(fmb_tokens) - len(fmb_close) - 2))567                    if close_pos < 0:568                        fmb_close_comma = deep_scaffold_sequences.get("fmb_closing_comma")569                        if fmb_close_comma:570                            close_pos = _find_subseq(fmb_tokens, fmb_close_comma,571                                                      max(0, len(fmb_tokens) - len(fmb_close_comma) - 2))572                            if close_pos >= 0:573                                fmb_close = fmb_close_comma574                    if close_pos >= 0:575                        for i in range(close_pos, close_pos + len(fmb_close)):576                            fmb_scaffold_positions.add(i)577 578        for i in fmb_scaffold_positions:579            scaffold_positions.append(fmb_start + i)580 581    # ── Step 4: Deep scaffold within trajectory ──582    # After dataloader processing (no mdm markers), trajectory is:583    #   ' "[[+14.70,-00.04], [+29.55,-00.21], ...]"'584    if "trajectory" in sections:585        traj_start, traj_end = sections["trajectory"]586        traj_tokens = resp_tokens[traj_start:traj_end]587 588        # Opening "[[589        traj_open = deep_scaffold_sequences["traj_open"]590        open_pos = _find_subseq(traj_tokens, traj_open, 0)591        if open_pos >= 0:592            _mark_scaffold_range(scaffold_positions, traj_start + open_pos, len(traj_open))593 594        # Waypoint separators ], (4 of them between 5 waypoints)595        traj_wp_sep = deep_scaffold_sequences["traj_wp_sep"]596        sep_search = 0597        for _ in range(4):598            sep_pos = _find_subseq(traj_tokens, traj_wp_sep, sep_search)599            if sep_pos < 0:600                break601            _mark_scaffold_range(scaffold_positions, traj_start + sep_pos, len(traj_wp_sep))602            sep_search = sep_pos + len(traj_wp_sep)603 604        # Intermediate waypoint opening ' [' (4 of them, between 5 waypoints)605        traj_wp_open = deep_scaffold_sequences.get("traj_wp_open")606        if traj_wp_open:607            wo_search = 0608            for _ in range(4):609                wo_pos = _find_subseq(traj_tokens, traj_wp_open, wo_search)610                if wo_pos < 0:611                    break612                _mark_scaffold_range(scaffold_positions, traj_start + wo_pos, len(traj_wp_open))613                wo_search = wo_pos + len(traj_wp_open)614 615        # Coordinate comma ',' between x and y within each waypoint (5 of them)616        traj_coord_comma = deep_scaffold_sequences.get("traj_coord_comma")617        if traj_coord_comma:618            cc_search = 0619            for _ in range(5):620                cc_pos = _find_subseq(traj_tokens, traj_coord_comma, cc_search)621                if cc_pos < 0:622                    break623                _mark_scaffold_range(scaffold_positions, traj_start + cc_pos, len(traj_coord_comma))624                cc_search = cc_pos + len(traj_coord_comma)625 626        # Closing ]]" or just ]]627        traj_close = deep_scaffold_sequences["traj_close"]628        close_pos = _find_subseq(traj_tokens, traj_close,629                                  max(0, len(traj_tokens) - len(traj_close) - 6))630        if close_pos < 0:631            for split_key in ["traj_close_split", "traj_close_split2"]:632                tcs = deep_scaffold_sequences.get(split_key)633                if tcs:634                    close_pos = _find_subseq(traj_tokens, tcs,635                                              max(0, len(traj_tokens) - len(tcs) - 6))636                    if close_pos >= 0:637                        traj_close = tcs638                        break639        if close_pos >= 0:640            _mark_scaffold_range(scaffold_positions, traj_start + close_pos, len(traj_close))641            # Align training with inference scaffold: exclude trailing tokens642            # after the JSON closing of trajectory (e.g. "<|im_end|>\n") from643            # section/block scheduling.644            effective_resp_end = min(645                effective_resp_end,646                resp_start + traj_start + close_pos + len(traj_close),647            )648 649        # Opening quote " (first token of traj value)650        if len(traj_tokens) > 0:651            scaffold_positions.append(traj_start)652 653    # ── Mark scaffold mask (absolute positions) ──654    scaffold_positions_set = set(scaffold_positions)655    for sp in scaffold_positions_set:656        abs_pos = resp_start + sp657        if abs_pos < seq_len:658            scaffold_mask[abs_pos] = True659 660    # ── Assign block indices per section ──661    current_block = 0662    assigned = set()663    block_to_section = {}  # block_idx -> section_name (for SASD compatibility)664    section_first_block: Dict[str, int] = {}665 666    for section_name in SECTION_KEYS:667        if section_name not in sections:668            continue669 670        rel_start, rel_end = sections[section_name]671        abs_start = resp_start + rel_start672        abs_end = resp_start + rel_end673        abs_start = max(abs_start, resp_start)674        abs_end = min(abs_end, effective_resp_end)675 676        num_tokens = abs_end - abs_start677        if num_tokens <= 0:678            continue679 680        # Count only non-scaffold tokens for block sizing681        value_positions = [p for p in range(abs_start, abs_end)682                           if response_mask[p] and (p - resp_start) not in scaffold_positions_set]683        num_value_tokens = len(value_positions)684 685        if num_value_tokens <= 0:686            section_first_block[section_name] = current_block687            block_to_section[current_block] = section_name688            current_block += 1689            continue690 691        # Use fixed block size (bd_size) and compute number of blocks dynamically692        tokens_per_step = fallback_block_size693        n_steps = max(1, math.ceil(num_value_tokens / tokens_per_step))694 695        for b in range(n_steps):696            block_to_section[current_block + b] = section_name697 698        section_first_block[section_name] = current_block699        for vi, pos in enumerate(value_positions):700            block_in_section = min(vi // tokens_per_step, n_steps - 1)701            response_block_idx[pos] = current_block + block_in_section702            assigned.add(pos)703 704        current_block += n_steps705 706    # Assign scaffold tokens within each section to the nearest value token707    # in the SAME section. This keeps section-closing tokens such as `"},`708    # with their section instead of drifting to the next section.709    for section_name in SECTION_KEYS:710        if section_name not in sections:711            continue712        rel_start, rel_end = sections[section_name]713        abs_start = max(resp_start + rel_start, resp_start)714        abs_end = min(resp_start + rel_end, resp_end)715        if abs_end <= abs_start:716            continue717 718        for abs_pos in range(abs_start, abs_end):719            rel_pos = abs_pos - resp_start720            if (721                abs_pos >= seq_len722                or not response_mask[abs_pos]723                or abs_pos in assigned724                or rel_pos not in scaffold_positions_set725            ):726                continue727 728            best_block = -1729            max_delta = max(1, abs_end - abs_start)730            for delta in range(1, max_delta + 1):731                # Prefer left first so closing punctuation tends to stay with732                # the preceding content in the same section.733                for cand in [abs_pos - delta, abs_pos + delta]:734                    if abs_start <= cand < abs_end and cand in assigned:735                        best_block = response_block_idx[cand].item()736                        break737                if best_block >= 0:738                    break739 740            if best_block < 0:741                best_block = section_first_block.get(section_name, -1)742 743            if best_block >= 0:744                response_block_idx[abs_pos] = best_block745                assigned.add(abs_pos)746 747    # Top-level boundary tokens are explicitly attached to the following748    # section's first block, instead of nearest-neighbor assignment.749    for section_name, rel_positions in boundary_scaffold_to_section.items():750        first_block = section_first_block.get(section_name)751        if first_block is None:752            continue753        for rel_pos in rel_positions:754            abs_pos = resp_start + rel_pos755            if abs_pos >= seq_len or not response_mask[abs_pos]:756                continue757            response_block_idx[abs_pos] = first_block758            assigned.add(abs_pos)759 760    # Scaffold tokens → block index of nearest assigned neighbour761    for sp in scaffold_positions_set:762        abs_pos = resp_start + sp763        if (764            abs_pos >= seq_len765            or abs_pos >= effective_resp_end766            or not response_mask[abs_pos]767            or abs_pos in assigned768        ):769            continue770        best_block = -1771        for delta in range(1, seq_len):772            for cand in [abs_pos + delta, abs_pos - delta]:773                if 0 <= cand < seq_len and cand in assigned:774                    best_block = response_block_idx[cand].item()775                    break776            if best_block >= 0:777                break778        if best_block >= 0:779            response_block_idx[abs_pos] = best_block780            assigned.add(abs_pos)781 782    # Fallback for unassigned response tokens783    for pos in range(resp_start, effective_resp_end):784        if response_mask[pos] and pos not in assigned:785            offset = pos - resp_start786            response_block_idx[pos] = current_block + offset // fallback_block_size787            assigned.add(pos)788 789    fallback_positions = [p for p in range(resp_start, effective_resp_end)790                          if response_mask[p] and response_block_idx[p].item() >= current_block]791    if fallback_positions:792        current_block = max(response_block_idx[p].item() for p in fallback_positions) + 1793 794    n_blocks = current_block795 796    # Turn index797    for i in range(1, seq_len):798        if response_block_idx[i] != response_block_idx[i - 1]:799            turn_idx[i] = turn_idx[i - 1] + 1800        else:801            turn_idx[i] = turn_idx[i - 1]802 803    return response_block_idx, turn_idx, n_blocks, scaffold_mask, block_to_section804