Download code/alpamayo_tt/preprocess.py from changh95/Alpamayo2-Super-p300x2: direct link, hf CLI and curl.
- Browser
- Download file 7.22 kB
-
https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/preprocess.py
- Command line
-
hf download hf://changh95/Alpamayo2-Super-p300x2/code/alpamayo_tt/preprocess.py
-
curl -L -o preprocess.py https://huggingface.co/changh95/Alpamayo2-Super-p300x2/resolve/main/code/alpamayo_tt/preprocess.py
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, | |
| } | |