Cccccz/HY
0
1"""Schema and validation for exact BF16 Predictor-v2 trajectories."""2 3from __future__ import annotations4 5from typing import Mapping6 7import torch8 9from .schema import (10 CHUNK_LATENT_FRAMES,11 HIDDEN_SIZE,12 LATENT_CHANNELS,13 LATENT_HEIGHT,14 LATENT_WIDTH,15 NUM_STEPS,16 STEP_FIELDS,17 TOKENS_PER_CHUNK,18)19 20 21LEGACY_SCHEMA_VERSION_V2 = "predictor_v2_exact_ar_context_kv_bf16"22SCHEMA_VERSION_V2 = "predictor_v2_padded_masked_ar_context_kv_bf16"23CONTEXT_BLOCK_IDS = (0, 1, 52, 53)24ATTENTION_HEADS = 1625HEAD_DIM = HIDDEN_SIZE // ATTENTION_HEADS26TOKENS_PER_FRAME = LATENT_HEIGHT * LATENT_WIDTH27DEFAULT_TEXT_KV_TOKENS = 90328DEFAULT_CONTEXT_KV_FRAMES = 2029 30 31def block_key(block_id: int, field: str) -> str:32 if block_id not in CONTEXT_BLOCK_IDS:33 raise ValueError(f"Unsupported context block: {block_id}")34 return f"block_{block_id:02d}_{field}"35 36 37def expected_step_keys() -> set[str]:38 keys = {39 "action_labels",40 "target_viewmats",41 "target_Ks",42 "rope_temporal_size",43 "start_rope_start_idx",44 }45 for step in range(NUM_STEPS):46 keys.update(f"step_{step}_{field}" for field in STEP_FIELDS)47 return keys48 49 50def _validate_bf16(name: str, tensor: torch.Tensor) -> None:51 if tensor.is_floating_point() and tensor.dtype != torch.bfloat16:52 raise ValueError(f"{name} must be BF16, got {tensor.dtype}")53 if tensor.is_floating_point() and not torch.isfinite(tensor).all():54 raise ValueError(f"Non-finite values in {name}")55 56 57def _validate_mask(name: str, mask: torch.Tensor, expected_shape: tuple[int, int]) -> None:58 if tuple(mask.shape) != expected_shape:59 raise ValueError(f"Unexpected {name} shape: {tuple(mask.shape)} != {expected_shape}")60 if mask.dtype != torch.bool:61 raise ValueError(f"{name} must be bool, got {mask.dtype}")62 if mask.shape[1] and not torch.equal(63 mask, torch.arange(mask.shape[1], device=mask.device)[None] < mask.sum(dim=1)[:, None]64 ):65 raise ValueError(f"{name} must contain a contiguous valid prefix")66 67 68def validate_case_tensors_v2(tensors: Mapping[str, torch.Tensor]) -> int:69 required = {"image_condition_latent", "text_valid_mask"}70 for block_id in CONTEXT_BLOCK_IDS:71 required.update(72 {block_key(block_id, "k_txt"), block_key(block_id, "v_txt")}73 )74 missing = required.difference(tensors)75 if missing:76 raise ValueError(f"Missing v2 case tensor keys: {sorted(missing)}")77 78 image_condition = tensors["image_condition_latent"]79 if tuple(image_condition.shape) != (1, LATENT_CHANNELS, 1, LATENT_HEIGHT, LATENT_WIDTH):80 raise ValueError(f"Unexpected image condition shape: {tuple(image_condition.shape)}")81 _validate_bf16("image_condition_latent", image_condition)82 83 text_valid_mask = tensors["text_valid_mask"]84 token_count = None85 for block_id in CONTEXT_BLOCK_IDS:86 k_txt = tensors[block_key(block_id, "k_txt")]87 v_txt = tensors[block_key(block_id, "v_txt")]88 if k_txt.shape != v_txt.shape:89 raise ValueError(f"Block {block_id} text K/V shape mismatch")90 if (91 k_txt.ndim != 492 or k_txt.shape[0] != 193 or k_txt.shape[1] != ATTENTION_HEADS94 or k_txt.shape[3] != HEAD_DIM95 ):96 raise ValueError(f"Unexpected block {block_id} text KV shape: {tuple(k_txt.shape)}")97 if token_count is None:98 token_count = int(k_txt.shape[2])99 elif int(k_txt.shape[2]) != token_count:100 raise ValueError("Text token count differs between context blocks")101 _validate_bf16(block_key(block_id, "k_txt"), k_txt)102 _validate_bf16(block_key(block_id, "v_txt"), v_txt)103 assert token_count is not None104 _validate_mask("text_valid_mask", text_valid_mask, (1, token_count))105 valid_tokens = int(text_valid_mask.sum())106 if valid_tokens <= 0:107 raise ValueError("text_valid_mask must contain at least one valid token")108 return valid_tokens109 110 111def validate_step_tensors_v2(tensors: Mapping[str, torch.Tensor]) -> None:112 missing = expected_step_keys().difference(tensors)113 if missing:114 raise ValueError(f"Missing v2 step tensor keys: {sorted(missing)}")115 116 if tuple(tensors["action_labels"].shape) != (1, CHUNK_LATENT_FRAMES):117 raise ValueError(f"Unexpected action_labels shape: {tuple(tensors['action_labels'].shape)}")118 if tuple(tensors["target_viewmats"].shape) != (1, CHUNK_LATENT_FRAMES, 4, 4):119 raise ValueError(f"Unexpected target_viewmats shape: {tuple(tensors['target_viewmats'].shape)}")120 if tuple(tensors["target_Ks"].shape) != (1, CHUNK_LATENT_FRAMES, 3, 3):121 raise ValueError(f"Unexpected target_Ks shape: {tuple(tensors['target_Ks'].shape)}")122 123 for step in range(NUM_STEPS):124 noisy = tensors[f"step_{step}_noisy_sample"]125 hidden = tensors[f"step_{step}_final_hidden"]126 condition = tensors[f"step_{step}_frame_condition"]127 velocity = tensors[f"step_{step}_velocity"]128 timestep = tensors[f"step_{step}_timestep"]129 if tuple(noisy.shape) != (1, LATENT_CHANNELS, CHUNK_LATENT_FRAMES, LATENT_HEIGHT, LATENT_WIDTH):130 raise ValueError(f"Step {step} noisy shape: {tuple(noisy.shape)}")131 if tuple(hidden.shape) != (1, TOKENS_PER_CHUNK, HIDDEN_SIZE):132 raise ValueError(f"Step {step} hidden shape: {tuple(hidden.shape)}")133 if tuple(condition.shape) != (1, CHUNK_LATENT_FRAMES, HIDDEN_SIZE):134 raise ValueError(f"Step {step} condition shape: {tuple(condition.shape)}")135 if tuple(velocity.shape) != tuple(noisy.shape):136 raise ValueError(f"Step {step} velocity shape: {tuple(velocity.shape)}")137 if timestep.numel() != 1 or timestep.dtype != torch.float32:138 raise ValueError(f"Step {step} timestep must be one FP32 scalar")139 for name, tensor in (140 ("noisy_sample", noisy),141 ("final_hidden", hidden),142 ("frame_condition", condition),143 ("velocity", velocity),144 ):145 _validate_bf16(f"step_{step}_{name}", tensor)146 147 _validate_bf16("target_viewmats", tensors["target_viewmats"])148 _validate_bf16("target_Ks", tensors["target_Ks"])149 150 151def validate_vision_context_v2(152 block_id: int,153 tensors: Mapping[str, torch.Tensor],154) -> int:155 required = {"k_vision", "v_vision", "context_valid_mask"}156 missing = required.difference(tensors)157 if missing:158 raise ValueError(f"Block {block_id} missing vision KV keys: {sorted(missing)}")159 k_vision = tensors["k_vision"]160 v_vision = tensors["v_vision"]161 if k_vision.shape != v_vision.shape:162 raise ValueError(f"Block {block_id} vision K/V shape mismatch")163 if (164 k_vision.ndim != 4165 or k_vision.shape[0] != 2166 or k_vision.shape[1] != ATTENTION_HEADS167 or k_vision.shape[3] != HEAD_DIM168 ):169 raise ValueError(f"Unexpected block {block_id} vision KV shape: {tuple(k_vision.shape)}")170 if k_vision.shape[2] % TOKENS_PER_FRAME:171 raise ValueError(172 f"Block {block_id} context tokens {k_vision.shape[2]} are not frame-aligned"173 )174 _validate_bf16(f"block_{block_id}_k_vision", k_vision)175 _validate_bf16(f"block_{block_id}_v_vision", v_vision)176 context_valid_mask = tensors["context_valid_mask"]177 _validate_mask(178 f"block_{block_id}_context_valid_mask",179 context_valid_mask,180 (1, int(k_vision.shape[2])),181 )182 valid_tokens = int(context_valid_mask.sum())183 if valid_tokens % TOKENS_PER_FRAME:184 raise ValueError(185 f"Block {block_id} valid context tokens {valid_tokens} are not frame-aligned"186 )187 return valid_tokens // TOKENS_PER_FRAME188 