Qevi-2B / qevi /pack.py
zandervdm's picture
Bundle the qevi inference engine in-repo
16e5bd8 verified
Raw History Blame Contribute Delete
7.22 kB
"""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]