Efficient-Large-Model/Fast-dDrive
4167
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 