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