"""Packed multi-question sequence construction for the Qevi architecture spike. Builds, from an image token count and a list of typed questions, the packed sequence layout specified in plan/02-ARCHITECTURE.md section 3: one image prefix followed by one context plus readout segment group per question, under a block-diagonal attention mask so that questions cannot see each other and options within one question cannot see each other. Model-agnostic. Produces a 4D additive attention bias and a position id tensor that any transformers causal LM accepts as ordinary keyword arguments. SCOPE LIMIT, current as of 19 September 2026: position_ids here is a single scalar stream, correct for plain 1D-RoPE decoders. Qwen3-VL-2B-Instruct, the settled primary base model, uses 3D multimodal RoPE (three streams: temporal, height, width; see research/18-qwen3-vl-2b-verification.md section 6). For text segments this scalar stream transfers directly, since Qwen3-VL gives all three axes the identical value for text tokens; a caller targeting Qwen3-VL broadcasts this module's scalar position_ids across all three axes for text positions. For the image segment itself, this module's arange(image_len) placeholder does NOT match Qwen3-VL's real 2D grid convention and must be replaced with the model's own get_vision_position_ids output before a real forward pass on that model. Not yet done here. """ from dataclasses import dataclass, field from typing import Literal import torch NEG_INF = torch.finfo(torch.float32).min @dataclass class Segment: name: str start: int end: int # exclusive kind: Literal["image", "context", "readout"] question_idx: int | None = None readout_idx: int | None = None @property def length(self) -> int: return self.end - self.start @property def sentinel_pos(self) -> int: """Absolute index of the last token in this segment, where the typed head reads.""" return self.end - 1 @dataclass class Question: qtype: Literal["noul", "choice", "score"] context_len: int readout_lens: list[int] = field(default_factory=lambda: [1]) @dataclass class PackedBatch: seq_len: int segments: list[Segment] attn_bias: torch.Tensor # (1, 1, seq_len, seq_len) additive, 0 = allowed, NEG_INF = blocked position_ids: torch.Tensor # (1, seq_len) sentinel_positions: dict # (question_idx, readout_idx) -> absolute position def _build_segments(image_len: int, questions: list[Question]) -> list[Segment]: segments = [Segment("V", 0, image_len, kind="image")] cursor = image_len for qi, q in enumerate(questions): c_start = cursor c_end = c_start + q.context_len segments.append(Segment(f"C{qi}", c_start, c_end, kind="context", question_idx=qi)) cursor = c_end for ri, rlen in enumerate(q.readout_lens): r_start = cursor r_end = r_start + rlen segments.append( Segment(f"R{qi}_{ri}", r_start, r_end, kind="readout", question_idx=qi, readout_idx=ri) ) cursor = r_end return segments def _group_key(seg: Segment): """Two positions may attend to each other (subject to causal order) only if they share a group. A readout's group includes its own context segment, which is why a readout sees all of C_i without needing a separate 'full attention to context' rule: C_i is already earlier in sequence order, so ordinary intra-group causal masking gives it full visibility.""" if seg.kind == "image": return ("V",) if seg.kind == "context": return ("Q", seg.question_idx) return ("Q", seg.question_idx, seg.readout_idx) def _segment_at(segments: list[Segment], pos: int) -> Segment: for seg in segments: if seg.start <= pos < seg.end: return seg raise IndexError(pos) def build_packed_batch( image_len: int, questions: list[Question], position_scheme: Literal["reset", "monotonic"] = "reset", text_start_offset: int | None = None, ) -> PackedBatch: """text_start_offset is the position value each question's context segment resets to under the 'reset' scheme. Defaults to image_len, correct for a plain 1D-RoPE decoder where the image occupies image_len sequential positions. It is WRONG for M-RoPE models such as Qwen3-VL, whose own convention resumes text at max(grid_h, grid_w) // spatial_merge_size, not the raw image token count; see research/18-qwen3-vl-2b-verification.md section 6. Callers targeting such a model must compute the correct value themselves (typically via the model's own get_vision_position_ids or equivalent) and pass it here explicitly.""" if text_start_offset is None: text_start_offset = image_len segments = _build_segments(image_len, questions) seq_len = segments[-1].end bias = torch.full((seq_len, seq_len), NEG_INF, dtype=torch.float32) for i in range(seq_len): seg_i = _segment_at(segments, i) group_i = _group_key(seg_i) for seg_j in segments: group_j = _group_key(seg_j) same_group = group_i == group_j is_image = seg_j.kind == "image" # A readout's group is (Q, qi, ri); its context segment's group is (Q, qi). # Give the readout visibility into its own context by matching on question index. same_question_context = ( seg_i.kind == "readout" and seg_j.kind == "context" and seg_i.question_idx == seg_j.question_idx ) if not (same_group or is_image or same_question_context): continue j_start = seg_j.start j_end = min(seg_j.end, i + 1) # causal: never attend beyond position i if j_end > j_start: bias[i, j_start:j_end] = 0.0 attn_bias = bias.unsqueeze(0).unsqueeze(0) # (1, 1, seq_len, seq_len) position_ids = torch.zeros(seq_len, dtype=torch.long) if position_scheme == "monotonic": position_ids[:] = torch.arange(seq_len) else: position_ids[: image_len] = torch.arange(image_len) for qi, q in enumerate(questions): c_seg = next(s for s in segments if s.kind == "context" and s.question_idx == qi) local_len = c_seg.length position_ids[c_seg.start : c_seg.end] = text_start_offset + torch.arange(local_len) cursor = text_start_offset + local_len for ri, seg in enumerate(s for s in segments if s.kind == "readout" and s.question_idx == qi): rlen = seg.length position_ids[seg.start : seg.end] = cursor + torch.arange(rlen) cursor += rlen sentinel_positions = { (seg.question_idx, seg.readout_idx): seg.sentinel_pos for seg in segments if seg.kind == "readout" } return PackedBatch( seq_len=seq_len, segments=segments, attn_bias=attn_bias, position_ids=position_ids.unsqueeze(0), sentinel_positions=sentinel_positions, ) def permute_questions(questions: list[Question], order: list[int]) -> list[Question]: return [questions[i] for i in order]