"""Serving-side preprocessing: camera frames + ego history -> model inputs (no dependence on the CPU reference env). Reproduces `alpamayo2_super.helper.prepare_model_inputs` + `fuse_traj_tokens` + Qwen3-VL `get_rope_index` with the transformers version of the tt-metal environment (5.x); verified bit-exact against samples prepared with 4.57.1 (scripts/check_preprocess_tf5.py). """ from __future__ import annotations import os from pathlib import Path import torch from .config import ModelConfig, default_snapshot def get_rope_index(input_ids: torch.Tensor, image_grid_thw: torch.Tensor, image_token_id: int, vision_start_id: int, merge: int = 2) -> tuple[torch.Tensor, torch.Tensor]: """Qwen3-VL mrope position ids for a batch-1 image prompt: [3, 1, L] and rope_deltas [1, 1]. Text tokens advance all three axes together; each image occupies a (t, h/merge, w/merge) grid whose three axes start at the text position that follows the preceding text; the next text token starts at max(image) + 1. """ ids = input_ids[0].tolist() L = len(ids) pos = torch.zeros(3, L, dtype=torch.long) st = 0 # next token index to assign cur = 0 # next position value img = 0 grids = image_grid_thw.tolist() i = 0 while i < L: if ids[i] == image_token_id: # text before this image run is already assigned; assign the image grid t, h, w = grids[img] img += 1 gh, gw = h // merge, w // merge n = t * gh * gw t_idx = torch.arange(t).view(t, 1, 1).expand(t, gh, gw).flatten() h_idx = torch.arange(gh).view(1, gh, 1).expand(t, gh, gw).flatten() w_idx = torch.arange(gw).view(1, 1, gw).expand(t, gh, gw).flatten() pos[0, i:i + n] = cur + t_idx pos[1, i:i + n] = cur + h_idx pos[2, i:i + n] = cur + w_idx cur = int(pos[:, i:i + n].max()) + 1 i += n else: pos[:, i] = cur cur += 1 i += 1 assert img == len(grids), f"used {img} image grids of {len(grids)}" rope_deltas = torch.tensor([[int(pos.max()) + 1 - L]], dtype=torch.long) return pos.unsqueeze(1), rope_deltas class Preprocessor: """Builds the exact prompt of the reference pipeline for the 'trajectory' task.""" def __init__(self, cfg: ModelConfig | None = None, snapshot: str | Path | None = None): import transformers from transformers import AutoProcessor import hydra.utils as hyu from alpamayo2_super.config import build_alpamayo2_super_tokenizer self.cfg = cfg or ModelConfig.load(snapshot) snap = str(self.cfg.snapshot) self.transformers_version = transformers.__version__ raw = self.cfg.raw kw = {"min_pixels": raw["min_pixels"], "max_pixels": raw["max_pixels"]} try: self.processor = AutoProcessor.from_pretrained(snap, fix_mistral_regex=True, **kw) except TypeError: self.processor = AutoProcessor.from_pretrained(snap, **kw) try: self.tokenizer = build_alpamayo2_super_tokenizer(snap, raw["history_vocab_size"], raw["future_vocab_size"]) except TypeError: # older/newer AutoTokenizer signature without fix_mistral_regex from transformers import AutoTokenizer from alpamayo2_super.models.utils import SPECIAL_TOKENS tok = AutoTokenizer.from_pretrained(snap) tok.add_tokens([f"" for i in range(raw["history_vocab_size"])]) tok.add_tokens([f"" for i in range(raw["future_vocab_size"])]) tok.add_tokens(list(SPECIAL_TOKENS.values()), special_tokens=True) self.tokenizer = tok self.processor.tokenizer = self.tokenizer self.hist_tok = hyu.instantiate(raw["hist_traj_tokenizer_cfg"], load_weights=False) self.fut_tok = hyu.instantiate(raw["future_traj_tokenizer_cfg"], load_weights=False) self.traj_ids = dict(raw["traj_ids"]) self.include_camera_ids = raw.get("include_camera_ids", True) self.include_frame_nums = raw.get("frame_label", "frame_num") == "frame_num" self.tokens_per_history_traj = raw["tokens_per_history_traj"] self.tokens_per_future_traj = raw["tokens_per_future_traj"] def messages(self, image_frames: torch.Tensor, camera_ids: list[int], no_cot: bool): from alpamayo2_super.chat_template.conversation import build_conversation data = {"image_frames": image_frames, "camera_indices": torch.tensor(camera_ids)} msgs = build_conversation( data=data, num_tokens_per_history_traj=self.tokens_per_history_traj, num_tokens_per_future_traj=self.tokens_per_future_traj, components_order=["image", "traj_history", "prompt"], components_prompt=["traj_future"] if no_cot else ["cot", "traj_future"], generation_mode=True, include_camera_ids=self.include_camera_ids, camera_ids=torch.tensor(camera_ids), include_frame_nums=self.include_frame_nums) if msgs[-1]["role"] == "assistant" and not msgs[-1]["content"]: msgs = msgs[:-1] return msgs def __call__(self, image_frames: torch.Tensor, camera_ids: list[int], ego_history_xyz: torch.Tensor, ego_history_rot: torch.Tensor, no_cot: bool = False) -> dict: """image_frames [cams, frames, 3, H, W] uint8 (camera order = camera_ids ascending); ego history [1,1,16,3] / [1,1,16,3,3] (or [16,3] / [16,3,3]). Returns the dict consumed by Alpamayo2Pipeline.run.""" from alpamayo2_super.models.utils import fuse_traj_tokens if ego_history_xyz.dim() == 2: ego_history_xyz = ego_history_xyz[None, None] if ego_history_rot.dim() == 3: ego_history_rot = ego_history_rot[None, None] msgs = self.messages(image_frames, camera_ids, no_cot) text = self.processor.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True, add_vision_id=False) images = image_frames.flatten(0, 1) images = (images.float() / 255.0) if images.dtype == torch.uint8 else images.float() tok = dict(self.processor(text=text, images=images, videos=None, padding=False, return_tensors="pt", do_rescale=False)) traj = {"ego_history_xyz": ego_history_xyz.float(), "ego_history_rot": ego_history_rot.float()} input_ids = fuse_traj_tokens(self.hist_tok, self.fut_tok, tok["input_ids"], traj, self.traj_ids) position_ids, rope_deltas = get_rope_index(input_ids, tok["image_grid_thw"], self.cfg.traj.image_token, self.cfg.traj.vision_start, self.cfg.vision.merge) return { "input_ids": input_ids, "attention_mask": tok["attention_mask"], "pixel_values": tok["pixel_values"], "image_grid_thw": tok["image_grid_thw"], "position_ids": position_ids, "rope_deltas": rope_deltas, "ego_history_xyz": traj["ego_history_xyz"], "ego_history_rot": traj["ego_history_rot"], "camera_indices": torch.tensor(camera_ids), "no_cot": no_cot, }