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