"""Schema and validation for exact BF16 Predictor-v2 trajectories.""" from __future__ import annotations from typing import Mapping import torch from .schema import ( CHUNK_LATENT_FRAMES, HIDDEN_SIZE, LATENT_CHANNELS, LATENT_HEIGHT, LATENT_WIDTH, NUM_STEPS, STEP_FIELDS, TOKENS_PER_CHUNK, ) LEGACY_SCHEMA_VERSION_V2 = "predictor_v2_exact_ar_context_kv_bf16" SCHEMA_VERSION_V2 = "predictor_v2_padded_masked_ar_context_kv_bf16" CONTEXT_BLOCK_IDS = (0, 1, 52, 53) ATTENTION_HEADS = 16 HEAD_DIM = HIDDEN_SIZE // ATTENTION_HEADS TOKENS_PER_FRAME = LATENT_HEIGHT * LATENT_WIDTH DEFAULT_TEXT_KV_TOKENS = 903 DEFAULT_CONTEXT_KV_FRAMES = 20 def block_key(block_id: int, field: str) -> str: if block_id not in CONTEXT_BLOCK_IDS: raise ValueError(f"Unsupported context block: {block_id}") return f"block_{block_id:02d}_{field}" def expected_step_keys() -> set[str]: keys = { "action_labels", "target_viewmats", "target_Ks", "rope_temporal_size", "start_rope_start_idx", } for step in range(NUM_STEPS): keys.update(f"step_{step}_{field}" for field in STEP_FIELDS) return keys def _validate_bf16(name: str, tensor: torch.Tensor) -> None: if tensor.is_floating_point() and tensor.dtype != torch.bfloat16: raise ValueError(f"{name} must be BF16, got {tensor.dtype}") if tensor.is_floating_point() and not torch.isfinite(tensor).all(): raise ValueError(f"Non-finite values in {name}") def _validate_mask(name: str, mask: torch.Tensor, expected_shape: tuple[int, int]) -> None: if tuple(mask.shape) != expected_shape: raise ValueError(f"Unexpected {name} shape: {tuple(mask.shape)} != {expected_shape}") if mask.dtype != torch.bool: raise ValueError(f"{name} must be bool, got {mask.dtype}") if mask.shape[1] and not torch.equal( mask, torch.arange(mask.shape[1], device=mask.device)[None] < mask.sum(dim=1)[:, None] ): raise ValueError(f"{name} must contain a contiguous valid prefix") def validate_case_tensors_v2(tensors: Mapping[str, torch.Tensor]) -> int: required = {"image_condition_latent", "text_valid_mask"} for block_id in CONTEXT_BLOCK_IDS: required.update( {block_key(block_id, "k_txt"), block_key(block_id, "v_txt")} ) missing = required.difference(tensors) if missing: raise ValueError(f"Missing v2 case tensor keys: {sorted(missing)}") image_condition = tensors["image_condition_latent"] if tuple(image_condition.shape) != (1, LATENT_CHANNELS, 1, LATENT_HEIGHT, LATENT_WIDTH): raise ValueError(f"Unexpected image condition shape: {tuple(image_condition.shape)}") _validate_bf16("image_condition_latent", image_condition) text_valid_mask = tensors["text_valid_mask"] token_count = None for block_id in CONTEXT_BLOCK_IDS: k_txt = tensors[block_key(block_id, "k_txt")] v_txt = tensors[block_key(block_id, "v_txt")] if k_txt.shape != v_txt.shape: raise ValueError(f"Block {block_id} text K/V shape mismatch") if ( k_txt.ndim != 4 or k_txt.shape[0] != 1 or k_txt.shape[1] != ATTENTION_HEADS or k_txt.shape[3] != HEAD_DIM ): raise ValueError(f"Unexpected block {block_id} text KV shape: {tuple(k_txt.shape)}") if token_count is None: token_count = int(k_txt.shape[2]) elif int(k_txt.shape[2]) != token_count: raise ValueError("Text token count differs between context blocks") _validate_bf16(block_key(block_id, "k_txt"), k_txt) _validate_bf16(block_key(block_id, "v_txt"), v_txt) assert token_count is not None _validate_mask("text_valid_mask", text_valid_mask, (1, token_count)) valid_tokens = int(text_valid_mask.sum()) if valid_tokens <= 0: raise ValueError("text_valid_mask must contain at least one valid token") return valid_tokens def validate_step_tensors_v2(tensors: Mapping[str, torch.Tensor]) -> None: missing = expected_step_keys().difference(tensors) if missing: raise ValueError(f"Missing v2 step tensor keys: {sorted(missing)}") if tuple(tensors["action_labels"].shape) != (1, CHUNK_LATENT_FRAMES): raise ValueError(f"Unexpected action_labels shape: {tuple(tensors['action_labels'].shape)}") if tuple(tensors["target_viewmats"].shape) != (1, CHUNK_LATENT_FRAMES, 4, 4): raise ValueError(f"Unexpected target_viewmats shape: {tuple(tensors['target_viewmats'].shape)}") if tuple(tensors["target_Ks"].shape) != (1, CHUNK_LATENT_FRAMES, 3, 3): raise ValueError(f"Unexpected target_Ks shape: {tuple(tensors['target_Ks'].shape)}") for step in range(NUM_STEPS): noisy = tensors[f"step_{step}_noisy_sample"] hidden = tensors[f"step_{step}_final_hidden"] condition = tensors[f"step_{step}_frame_condition"] velocity = tensors[f"step_{step}_velocity"] timestep = tensors[f"step_{step}_timestep"] if tuple(noisy.shape) != (1, LATENT_CHANNELS, CHUNK_LATENT_FRAMES, LATENT_HEIGHT, LATENT_WIDTH): raise ValueError(f"Step {step} noisy shape: {tuple(noisy.shape)}") if tuple(hidden.shape) != (1, TOKENS_PER_CHUNK, HIDDEN_SIZE): raise ValueError(f"Step {step} hidden shape: {tuple(hidden.shape)}") if tuple(condition.shape) != (1, CHUNK_LATENT_FRAMES, HIDDEN_SIZE): raise ValueError(f"Step {step} condition shape: {tuple(condition.shape)}") if tuple(velocity.shape) != tuple(noisy.shape): raise ValueError(f"Step {step} velocity shape: {tuple(velocity.shape)}") if timestep.numel() != 1 or timestep.dtype != torch.float32: raise ValueError(f"Step {step} timestep must be one FP32 scalar") for name, tensor in ( ("noisy_sample", noisy), ("final_hidden", hidden), ("frame_condition", condition), ("velocity", velocity), ): _validate_bf16(f"step_{step}_{name}", tensor) _validate_bf16("target_viewmats", tensors["target_viewmats"]) _validate_bf16("target_Ks", tensors["target_Ks"]) def validate_vision_context_v2( block_id: int, tensors: Mapping[str, torch.Tensor], ) -> int: required = {"k_vision", "v_vision", "context_valid_mask"} missing = required.difference(tensors) if missing: raise ValueError(f"Block {block_id} missing vision KV keys: {sorted(missing)}") k_vision = tensors["k_vision"] v_vision = tensors["v_vision"] if k_vision.shape != v_vision.shape: raise ValueError(f"Block {block_id} vision K/V shape mismatch") if ( k_vision.ndim != 4 or k_vision.shape[0] != 2 or k_vision.shape[1] != ATTENTION_HEADS or k_vision.shape[3] != HEAD_DIM ): raise ValueError(f"Unexpected block {block_id} vision KV shape: {tuple(k_vision.shape)}") if k_vision.shape[2] % TOKENS_PER_FRAME: raise ValueError( f"Block {block_id} context tokens {k_vision.shape[2]} are not frame-aligned" ) _validate_bf16(f"block_{block_id}_k_vision", k_vision) _validate_bf16(f"block_{block_id}_v_vision", v_vision) context_valid_mask = tensors["context_valid_mask"] _validate_mask( f"block_{block_id}_context_valid_mask", context_valid_mask, (1, int(k_vision.shape[2])), ) valid_tokens = int(context_valid_mask.sum()) if valid_tokens % TOKENS_PER_FRAME: raise ValueError( f"Block {block_id} valid context tokens {valid_tokens} are not frame-aligned" ) return valid_tokens // TOKENS_PER_FRAME