Diffusers
Safetensors
HY / predictor_data /prefeature_schema.py
Cccccz's picture
Upload batch 64: 500 files (0.40 GiB)
5f0e4a2 verified
Raw History Blame Contribute Delete
3.65 kB
"""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