CoolFace
Modelpublic

Cccccz/HY

sourceHugging Faceupdated 11d agoView on Hugging Face
0likes
v2_schema.py188 linesDownload Raw Back to predictor_data
1"""Schema and validation for exact BF16 Predictor-v2 trajectories."""2 3from __future__ import annotations4 5from typing import Mapping6 7import torch8 9from .schema import (10    CHUNK_LATENT_FRAMES,11    HIDDEN_SIZE,12    LATENT_CHANNELS,13    LATENT_HEIGHT,14    LATENT_WIDTH,15    NUM_STEPS,16    STEP_FIELDS,17    TOKENS_PER_CHUNK,18)19 20 21LEGACY_SCHEMA_VERSION_V2 = "predictor_v2_exact_ar_context_kv_bf16"22SCHEMA_VERSION_V2 = "predictor_v2_padded_masked_ar_context_kv_bf16"23CONTEXT_BLOCK_IDS = (0, 1, 52, 53)24ATTENTION_HEADS = 1625HEAD_DIM = HIDDEN_SIZE // ATTENTION_HEADS26TOKENS_PER_FRAME = LATENT_HEIGHT * LATENT_WIDTH27DEFAULT_TEXT_KV_TOKENS = 90328DEFAULT_CONTEXT_KV_FRAMES = 2029 30 31def block_key(block_id: int, field: str) -> str:32    if block_id not in CONTEXT_BLOCK_IDS:33        raise ValueError(f"Unsupported context block: {block_id}")34    return f"block_{block_id:02d}_{field}"35 36 37def expected_step_keys() -> set[str]:38    keys = {39        "action_labels",40        "target_viewmats",41        "target_Ks",42        "rope_temporal_size",43        "start_rope_start_idx",44    }45    for step in range(NUM_STEPS):46        keys.update(f"step_{step}_{field}" for field in STEP_FIELDS)47    return keys48 49 50def _validate_bf16(name: str, tensor: torch.Tensor) -> None:51    if tensor.is_floating_point() and tensor.dtype != torch.bfloat16:52        raise ValueError(f"{name} must be BF16, got {tensor.dtype}")53    if tensor.is_floating_point() and not torch.isfinite(tensor).all():54        raise ValueError(f"Non-finite values in {name}")55 56 57def _validate_mask(name: str, mask: torch.Tensor, expected_shape: tuple[int, int]) -> None:58    if tuple(mask.shape) != expected_shape:59        raise ValueError(f"Unexpected {name} shape: {tuple(mask.shape)} != {expected_shape}")60    if mask.dtype != torch.bool:61        raise ValueError(f"{name} must be bool, got {mask.dtype}")62    if mask.shape[1] and not torch.equal(63        mask, torch.arange(mask.shape[1], device=mask.device)[None] < mask.sum(dim=1)[:, None]64    ):65        raise ValueError(f"{name} must contain a contiguous valid prefix")66 67 68def validate_case_tensors_v2(tensors: Mapping[str, torch.Tensor]) -> int:69    required = {"image_condition_latent", "text_valid_mask"}70    for block_id in CONTEXT_BLOCK_IDS:71        required.update(72            {block_key(block_id, "k_txt"), block_key(block_id, "v_txt")}73        )74    missing = required.difference(tensors)75    if missing:76        raise ValueError(f"Missing v2 case tensor keys: {sorted(missing)}")77 78    image_condition = tensors["image_condition_latent"]79    if tuple(image_condition.shape) != (1, LATENT_CHANNELS, 1, LATENT_HEIGHT, LATENT_WIDTH):80        raise ValueError(f"Unexpected image condition shape: {tuple(image_condition.shape)}")81    _validate_bf16("image_condition_latent", image_condition)82 83    text_valid_mask = tensors["text_valid_mask"]84    token_count = None85    for block_id in CONTEXT_BLOCK_IDS:86        k_txt = tensors[block_key(block_id, "k_txt")]87        v_txt = tensors[block_key(block_id, "v_txt")]88        if k_txt.shape != v_txt.shape:89            raise ValueError(f"Block {block_id} text K/V shape mismatch")90        if (91            k_txt.ndim != 492            or k_txt.shape[0] != 193            or k_txt.shape[1] != ATTENTION_HEADS94            or k_txt.shape[3] != HEAD_DIM95        ):96            raise ValueError(f"Unexpected block {block_id} text KV shape: {tuple(k_txt.shape)}")97        if token_count is None:98            token_count = int(k_txt.shape[2])99        elif int(k_txt.shape[2]) != token_count:100            raise ValueError("Text token count differs between context blocks")101        _validate_bf16(block_key(block_id, "k_txt"), k_txt)102        _validate_bf16(block_key(block_id, "v_txt"), v_txt)103    assert token_count is not None104    _validate_mask("text_valid_mask", text_valid_mask, (1, token_count))105    valid_tokens = int(text_valid_mask.sum())106    if valid_tokens <= 0:107        raise ValueError("text_valid_mask must contain at least one valid token")108    return valid_tokens109 110 111def validate_step_tensors_v2(tensors: Mapping[str, torch.Tensor]) -> None:112    missing = expected_step_keys().difference(tensors)113    if missing:114        raise ValueError(f"Missing v2 step tensor keys: {sorted(missing)}")115 116    if tuple(tensors["action_labels"].shape) != (1, CHUNK_LATENT_FRAMES):117        raise ValueError(f"Unexpected action_labels shape: {tuple(tensors['action_labels'].shape)}")118    if tuple(tensors["target_viewmats"].shape) != (1, CHUNK_LATENT_FRAMES, 4, 4):119        raise ValueError(f"Unexpected target_viewmats shape: {tuple(tensors['target_viewmats'].shape)}")120    if tuple(tensors["target_Ks"].shape) != (1, CHUNK_LATENT_FRAMES, 3, 3):121        raise ValueError(f"Unexpected target_Ks shape: {tuple(tensors['target_Ks'].shape)}")122 123    for step in range(NUM_STEPS):124        noisy = tensors[f"step_{step}_noisy_sample"]125        hidden = tensors[f"step_{step}_final_hidden"]126        condition = tensors[f"step_{step}_frame_condition"]127        velocity = tensors[f"step_{step}_velocity"]128        timestep = tensors[f"step_{step}_timestep"]129        if tuple(noisy.shape) != (1, LATENT_CHANNELS, CHUNK_LATENT_FRAMES, LATENT_HEIGHT, LATENT_WIDTH):130            raise ValueError(f"Step {step} noisy shape: {tuple(noisy.shape)}")131        if tuple(hidden.shape) != (1, TOKENS_PER_CHUNK, HIDDEN_SIZE):132            raise ValueError(f"Step {step} hidden shape: {tuple(hidden.shape)}")133        if tuple(condition.shape) != (1, CHUNK_LATENT_FRAMES, HIDDEN_SIZE):134            raise ValueError(f"Step {step} condition shape: {tuple(condition.shape)}")135        if tuple(velocity.shape) != tuple(noisy.shape):136            raise ValueError(f"Step {step} velocity shape: {tuple(velocity.shape)}")137        if timestep.numel() != 1 or timestep.dtype != torch.float32:138            raise ValueError(f"Step {step} timestep must be one FP32 scalar")139        for name, tensor in (140            ("noisy_sample", noisy),141            ("final_hidden", hidden),142            ("frame_condition", condition),143            ("velocity", velocity),144        ):145            _validate_bf16(f"step_{step}_{name}", tensor)146 147    _validate_bf16("target_viewmats", tensors["target_viewmats"])148    _validate_bf16("target_Ks", tensors["target_Ks"])149 150 151def validate_vision_context_v2(152    block_id: int,153    tensors: Mapping[str, torch.Tensor],154) -> int:155    required = {"k_vision", "v_vision", "context_valid_mask"}156    missing = required.difference(tensors)157    if missing:158        raise ValueError(f"Block {block_id} missing vision KV keys: {sorted(missing)}")159    k_vision = tensors["k_vision"]160    v_vision = tensors["v_vision"]161    if k_vision.shape != v_vision.shape:162        raise ValueError(f"Block {block_id} vision K/V shape mismatch")163    if (164        k_vision.ndim != 4165        or k_vision.shape[0] != 2166        or k_vision.shape[1] != ATTENTION_HEADS167        or k_vision.shape[3] != HEAD_DIM168    ):169        raise ValueError(f"Unexpected block {block_id} vision KV shape: {tuple(k_vision.shape)}")170    if k_vision.shape[2] % TOKENS_PER_FRAME:171        raise ValueError(172            f"Block {block_id} context tokens {k_vision.shape[2]} are not frame-aligned"173        )174    _validate_bf16(f"block_{block_id}_k_vision", k_vision)175    _validate_bf16(f"block_{block_id}_v_vision", v_vision)176    context_valid_mask = tensors["context_valid_mask"]177    _validate_mask(178        f"block_{block_id}_context_valid_mask",179        context_valid_mask,180        (1, int(k_vision.shape[2])),181    )182    valid_tokens = int(context_valid_mask.sum())183    if valid_tokens % TOKENS_PER_FRAME:184        raise ValueError(185            f"Block {block_id} valid context tokens {valid_tokens} are not frame-aligned"186        )187    return valid_tokens // TOKENS_PER_FRAME188