changh95's picture
Add files using upload-large-folder tool
0190e6b verified
Raw History Blame Contribute Delete
7.22 kB
"""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"<i{i}>" for i in range(raw["history_vocab_size"])])
tok.add_tokens([f"<i{i + raw['history_vocab_size']}>" 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,
}