Diffusers
Safetensors
File size: 4,168 Bytes
5f0e4a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
"""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}")