Cccccz's picture
Add files using upload-large-folder tool
510ab6b verified
Raw
History Blame Contribute Delete
5.45 kB
"""Tensor and manifest schema for Self-Forcing Predictor v4 data."""
from __future__ import annotations
from collections.abc import Mapping
import torch
SCHEMA_VERSION = "self_forcing_predictor_v4_bf16_v1"
CANDIDATE_BLOCK_IDS = (0, 1, 28, 29)
NUM_STEPS = 4
NUM_CHUNKS = 7
CHUNK_FRAMES = 3
LATENT_CHANNELS = 16
LATENT_HEIGHT = 60
LATENT_WIDTH = 104
HIDDEN_SIZE = 1536
NUM_HEADS = 12
HEAD_DIM = 128
TEXT_TOKENS = 512
TOKENS_PER_FRAME = 1560
TOKENS_PER_CHUNK = CHUNK_FRAMES * TOKENS_PER_FRAME
SUPERVISION_PAIRS = ((0, 1), (1, 2), (2, 3))
def _require_finite_bf16(name: str, tensor: torch.Tensor) -> None:
if tensor.dtype != torch.bfloat16:
raise ValueError(f"{name} must be BF16, got {tensor.dtype}")
if not torch.isfinite(tensor).all():
raise ValueError(f"{name} contains non-finite values")
def validate_step_tensors(tensors: Mapping[str, torch.Tensor]) -> None:
"""Validate one chunk containing all four Full-DiT teacher steps."""
required = {
f"step_{step}_{field}"
for step in range(NUM_STEPS)
for field in ("noisy_latent", "timestep", "final_hidden", "flow")
}
missing = required.difference(tensors)
extra = set(tensors).difference(required)
if missing or extra:
raise ValueError(
f"step tensor keys mismatch; missing={sorted(missing)}, extra={sorted(extra)}"
)
latent_shape = (
1,
CHUNK_FRAMES,
LATENT_CHANNELS,
LATENT_HEIGHT,
LATENT_WIDTH,
)
hidden_shape = (1, TOKENS_PER_CHUNK, HIDDEN_SIZE)
timestep_shape = (1, CHUNK_FRAMES)
for step in range(NUM_STEPS):
for field in ("noisy_latent", "flow"):
name = f"step_{step}_{field}"
value = tensors[name]
if tuple(value.shape) != latent_shape:
raise ValueError(f"{name} has shape {tuple(value.shape)}, expected {latent_shape}")
_require_finite_bf16(name, value)
hidden_name = f"step_{step}_final_hidden"
hidden = tensors[hidden_name]
if tuple(hidden.shape) != hidden_shape:
raise ValueError(
f"{hidden_name} has shape {tuple(hidden.shape)}, expected {hidden_shape}"
)
_require_finite_bf16(hidden_name, hidden)
timestep_name = f"step_{step}_timestep"
timestep = tensors[timestep_name]
if tuple(timestep.shape) != timestep_shape or timestep.dtype != torch.int64:
raise ValueError(
f"{timestep_name} must be INT64 {timestep_shape}, "
f"got {timestep.dtype} {tuple(timestep.shape)}"
)
if not bool(torch.all(timestep == timestep[:, :1])):
raise ValueError(f"{timestep_name} must be constant within the temporal chunk")
def validate_clean_prefeature(
block_id: int,
tensors: Mapping[str, torch.Tensor],
) -> None:
"""Validate the input to ``block.self_attn.k`` from one clean pass."""
if int(block_id) not in CANDIDATE_BLOCK_IDS:
raise ValueError(f"unsupported clean prefeature block: {block_id}")
required = {"self_attn_input", "start_frame", "num_frames"}
missing = required.difference(tensors)
extra = set(tensors).difference(required)
if missing or extra:
raise ValueError(
f"clean prefeature keys mismatch; missing={sorted(missing)}, extra={sorted(extra)}"
)
feature = tensors["self_attn_input"]
expected = (1, TOKENS_PER_CHUNK, HIDDEN_SIZE)
if tuple(feature.shape) != expected:
raise ValueError(
f"block {block_id} self_attn_input has shape {tuple(feature.shape)}, "
f"expected {expected}"
)
_require_finite_bf16("self_attn_input", feature)
for name in ("start_frame", "num_frames"):
value = tensors[name]
if value.dtype != torch.int64 or tuple(value.shape) != (1,):
raise ValueError(f"{name} must be one INT64 value")
if int(tensors["start_frame"].item()) < 0:
raise ValueError("start_frame must be non-negative")
if int(tensors["num_frames"].item()) != CHUNK_FRAMES:
raise ValueError(f"num_frames must equal {CHUNK_FRAMES}")
def validate_case_tensors(
tensors: Mapping[str, torch.Tensor],
block_ids: tuple[int, ...] = CANDIDATE_BLOCK_IDS,
) -> None:
"""Validate the case-level selected-layer text cross-attention KV."""
required = {
f"block_{block_id:02d}_cross_{field}"
for block_id in block_ids
for field in ("k", "v")
}
missing = required.difference(tensors)
extra = set(tensors).difference(required)
if missing or extra:
raise ValueError(
f"case tensor keys mismatch; missing={sorted(missing)}, extra={sorted(extra)}"
)
expected = (1, TEXT_TOKENS, NUM_HEADS, HEAD_DIM)
token_counts = set()
for name in sorted(required):
value = tensors[name]
if value.ndim != 4 or (
value.shape[0] != expected[0]
or value.shape[2] != expected[2]
or value.shape[3] != expected[3]
):
raise ValueError(
f"{name} has shape {tuple(value.shape)}, expected [1, text, 12, 128]"
)
token_counts.add(int(value.shape[1]))
_require_finite_bf16(name, value)
if len(token_counts) != 1 or next(iter(token_counts)) != TEXT_TOKENS:
raise ValueError(f"all text KV tensors must have exactly {TEXT_TOKENS} tokens")