"""Dataset schema constants and tensor validation.""" from __future__ import annotations from typing import Mapping import torch SCHEMA_VERSION = "predictor_v1_txt_features_no_context_kv" NUM_STEPS = 4 LATENT_CHANNELS = 32 MODEL_INPUT_CHANNELS = 65 CHUNK_LATENT_FRAMES = 4 LATENT_HEIGHT = 30 LATENT_WIDTH = 52 HIDDEN_SIZE = 2048 TOKENS_PER_CHUNK = CHUNK_LATENT_FRAMES * LATENT_HEIGHT * LATENT_WIDTH STEP_FIELDS = ( "timestep", "noisy_sample", "frame_condition", "final_hidden", "velocity", ) def expected_chunk_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_case_tensors(tensors: Mapping[str, torch.Tensor]) -> None: required = { "image_condition_latent", "current_txt", "cached_txt", "vec_txt", } missing = required.difference(tensors) if missing: raise ValueError(f"Missing case tensor keys: {sorted(missing)}") image_condition = tensors["image_condition_latent"] if tuple(image_condition.shape) != (1, 32, 1, 30, 52): raise ValueError( "image_condition_latent must be [1,32,1,30,52], got " f"{tuple(image_condition.shape)}" ) current_txt = tensors["current_txt"] cached_txt = tensors["cached_txt"] if current_txt.ndim != 3 or current_txt.shape[0] != 1 or current_txt.shape[-1] != HIDDEN_SIZE: raise ValueError(f"current_txt must be [1,S,2048], got {tuple(current_txt.shape)}") if tuple(cached_txt.shape) != tuple(current_txt.shape): raise ValueError( f"cached_txt {tuple(cached_txt.shape)} != current_txt {tuple(current_txt.shape)}" ) if tuple(tensors["vec_txt"].shape) != (1, HIDDEN_SIZE): raise ValueError(f"vec_txt must be [1,2048], got {tuple(tensors['vec_txt'].shape)}") for name, tensor in tensors.items(): if not torch.isfinite(tensor).all(): raise ValueError(f"Non-finite values in case tensor {name}") def validate_chunk_tensors(tensors: Mapping[str, torch.Tensor]) -> None: missing = expected_chunk_keys().difference(tensors) if missing: raise ValueError(f"Missing chunk 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, 32, 4, 30, 52): 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) != (1, 32, 4, 30, 52): raise ValueError(f"step {step} velocity shape: {tuple(velocity.shape)}") if timestep.numel() != 1: raise ValueError(f"step {step} timestep must be scalar, got {tuple(timestep.shape)}") for name, tensor in tensors.items(): if tensor.is_floating_point() and not torch.isfinite(tensor).all(): raise ValueError(f"Non-finite values in chunk tensor {name}")