"""Schema for BF16 joint-window Context pre-KV features.""" from __future__ import annotations from typing import Mapping import torch from .schema import HIDDEN_SIZE from .v2_schema import ( CONTEXT_BLOCK_IDS, TOKENS_PER_FRAME, validate_case_tensors_v2, validate_step_tensors_v2, ) SCHEMA_VERSION_PREFEATURE = "predictor_context_prefeature_bf16_v1" SUPERVISION_PAIRS = ((0, 1), (1, 2), (2, 3)) DEFAULT_ACTIONS = ( ("w_s", "w-15,s-16"), ("a_d", "a-15,d-16"), ("up_down", "up-15,down-16"), ("left_right", "left-15,right-16"), ) def validate_case_tensors(tensors: Mapping[str, torch.Tensor]) -> int: return validate_case_tensors_v2(tensors) def validate_step_tensors(tensors: Mapping[str, torch.Tensor]) -> None: validate_step_tensors_v2(tensors) def validate_context_prefeature( block_id: int, tensors: Mapping[str, torch.Tensor], ) -> int: if block_id not in CONTEXT_BLOCK_IDS: raise ValueError(f"Unsupported context block: {block_id}") required = { "img_modulated", "context_valid_mask", "selected_frame_indices", "context_viewmats", "context_Ks", "rope_temporal_size", "start_rope_start_idx", } missing = required.difference(tensors) if missing: raise ValueError(f"Block {block_id} missing prefeature keys: {sorted(missing)}") feature = tensors["img_modulated"] if feature.ndim != 3 or feature.shape[0] != 1 or feature.shape[2] != HIDDEN_SIZE: raise ValueError(f"Unexpected block {block_id} feature shape: {tuple(feature.shape)}") if feature.dtype != torch.bfloat16: raise ValueError(f"Block {block_id} img_modulated must be BF16") if not torch.isfinite(feature).all(): raise ValueError(f"Block {block_id} img_modulated contains non-finite values") tokens = int(feature.shape[1]) if tokens % TOKENS_PER_FRAME: raise ValueError(f"Block {block_id} token count {tokens} is not frame-aligned") frames = tokens // TOKENS_PER_FRAME mask = tensors["context_valid_mask"] if tuple(mask.shape) != (1, tokens) or mask.dtype != torch.bool or not bool(mask.all()): raise ValueError("On-disk context_valid_mask must be all-True and unpadded") indices = tensors["selected_frame_indices"] if tuple(indices.shape) != (frames,) or indices.dtype != torch.int64: raise ValueError("selected_frame_indices must be int64 [context_frames]") if frames and (int(indices.min()) < 0 or not bool(torch.all(indices[1:] > indices[:-1]))): raise ValueError("selected_frame_indices must be non-negative and strictly increasing") if tuple(tensors["context_viewmats"].shape) != (1, frames, 4, 4): raise ValueError("Unexpected context_viewmats shape") if tuple(tensors["context_Ks"].shape) != (1, frames, 3, 3): raise ValueError("Unexpected context_Ks shape") for name in ("context_viewmats", "context_Ks"): value = tensors[name] if value.dtype != torch.bfloat16 or not torch.isfinite(value).all(): raise ValueError(f"{name} must contain finite BF16 values") for name in ("rope_temporal_size", "start_rope_start_idx"): value = tensors[name] if value.dtype != torch.int64 or value.numel() != 1: raise ValueError(f"{name} must be one int64 scalar") if int(tensors["rope_temporal_size"].item()) != frames: raise ValueError("rope_temporal_size must equal the compact context frame count") if int(tensors["start_rope_start_idx"].item()) != 0: raise ValueError("Context prefill must start RoPE at zero") return frames