Diffusers
Safetensors
File size: 8,331 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
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
"""Read-only hooks that capture exact Full-DiT teacher trajectories."""

from __future__ import annotations

from dataclasses import dataclass
from typing import Any, Callable

import torch

from .schema import LATENT_HEIGHT, LATENT_WIDTH, NUM_STEPS


def _clone_cpu(tensor: torch.Tensor) -> torch.Tensor:
    return tensor.detach().to(device="cpu").contiguous()


@dataclass
class _ActiveStep:
    chunk_id: int
    step_id: int
    tensors: dict[str, torch.Tensor]
    shared: dict[str, torch.Tensor]


class PredictorTeacherCapture:
    """Capture denoising inputs, final hidden/condition, velocity, and dense txt features.

    The hook assumes a single positive AR stream (few-step guidance=1) and a fixed
    number of denoising steps per chunk. History-prefill calls are excluded through
    ``cache_vision``.
    """

    def __init__(
        self,
        transformer: torch.nn.Module,
        *,
        on_chunk: Callable[[int, dict[str, torch.Tensor]], None],
        num_steps: int = NUM_STEPS,
    ) -> None:
        self.transformer = transformer
        self.on_chunk = on_chunk
        self.num_steps = num_steps
        self.call_index = 0
        self.active: _ActiveStep | None = None
        self.chunk_steps: list[_ActiveStep] = []
        self.current_txt: torch.Tensor | None = None
        self.cached_txt: torch.Tensor | None = None
        self.vec_txt: torch.Tensor | None = None
        self.image_condition_latent: torch.Tensor | None = None
        self._handles: list[Any] = []
        self._original_get_text_and_mask = None

    def __enter__(self) -> "PredictorTeacherCapture":
        self._original_get_text_and_mask = self.transformer.get_text_and_mask

        def wrapped_get_text_and_mask(*args, **kwargs):
            txt, text_mask, vec_txt = self._original_get_text_and_mask(*args, **kwargs)
            if self.current_txt is None:
                if txt.shape[0] != 1:
                    raise ValueError("Predictor capture currently requires text batch size 1")
                valid = text_mask[0].bool().to(txt.device)
                self.current_txt = _clone_cpu(txt[:, valid])
                self.vec_txt = _clone_cpu(vec_txt)
            return txt, text_mask, vec_txt

        self.transformer.get_text_and_mask = wrapped_get_text_and_mask
        self._handles.append(
            self.transformer.register_forward_pre_hook(self._transformer_pre, with_kwargs=True)
        )
        self._handles.append(
            self.transformer.register_forward_hook(self._transformer_post, with_kwargs=True)
        )
        self._handles.append(
            self.transformer.final_layer.register_forward_pre_hook(self._final_pre, with_kwargs=True)
        )
        self._handles.append(
            self.transformer.double_blocks[-1].register_forward_hook(
                self._last_block_post, with_kwargs=True
            )
        )
        return self

    def __exit__(self, exc_type, exc, traceback) -> bool:
        for handle in self._handles:
            handle.remove()
        self._handles.clear()
        if self._original_get_text_and_mask is not None:
            self.transformer.get_text_and_mask = self._original_get_text_and_mask
        self.active = None
        return False

    def _last_block_post(self, module, args, kwargs, output) -> None:
        if kwargs.get("ar_txt_inference", False):
            txt = output[0] if isinstance(output, tuple) else output
            self.cached_txt = _clone_cpu(txt)

    def _transformer_pre(self, module, args, kwargs) -> None:
        is_denoise = (
            kwargs.get("ar_vision_inference", False)
            and not kwargs.get("cache_vision", False)
        )
        if not is_denoise:
            return
        if self.active is not None:
            raise RuntimeError("Nested denoising capture is not supported")

        chunk_id, step_id = divmod(self.call_index, self.num_steps)
        model_input = kwargs["hidden_states"]
        if model_input.shape[1] != 65:
            raise ValueError(f"Teacher denoising input must have 65 channels, got {model_input.shape}")
        if self.image_condition_latent is None:
            self.image_condition_latent = _clone_cpu(model_input[:, 32:64, 0:1])
            mask = model_input[:, 64:65]
            if not torch.all(mask[:, :, 0] == 1) or not torch.all(mask[:, :, 1:] == 0):
                raise ValueError("Unexpected I2V condition mask in first chunk")

        timestep = kwargs["timestep"].reshape(-1)[0:1]
        shared = {
            "action_labels": _clone_cpu(kwargs["action"].reshape(1, -1).round().long()),
            "target_viewmats": _clone_cpu(kwargs["viewmats"]),
            "target_Ks": _clone_cpu(kwargs["Ks"]),
            "rope_temporal_size": torch.tensor([int(kwargs["rope_temporal_size"])], dtype=torch.int64),
            "start_rope_start_idx": torch.tensor(
                [int(kwargs["start_rope_start_idx"])], dtype=torch.int64
            ),
        }
        self.active = _ActiveStep(
            chunk_id=chunk_id,
            step_id=step_id,
            tensors={
                "timestep": _clone_cpu(timestep.float()),
                "noisy_sample": _clone_cpu(model_input[:, :32]),
            },
            shared=shared,
        )

    def _final_pre(self, module, args, kwargs) -> None:
        if self.active is None:
            return
        hidden, condition = args[0], args[1]
        batch, tokens, hidden_size = hidden.shape
        spatial_tokens = LATENT_HEIGHT * LATENT_WIDTH
        if tokens % spatial_tokens:
            raise ValueError(f"Final hidden token count {tokens} is not divisible by {spatial_tokens}")
        frames = tokens // spatial_tokens
        compact = condition.reshape(batch, frames, spatial_tokens, hidden_size)[:, :, 0]
        expanded = compact[:, :, None].expand(batch, frames, spatial_tokens, hidden_size)
        if not torch.equal(expanded.reshape(batch, tokens, hidden_size), condition.reshape(batch, tokens, hidden_size)):
            raise ValueError("Final-layer condition varies inside a latent frame")
        self.active.tensors["frame_condition"] = _clone_cpu(compact)
        self.active.tensors["final_hidden"] = _clone_cpu(hidden)

    def _transformer_post(self, module, args, kwargs, output) -> None:
        if self.active is None:
            return
        velocity = output[0] if isinstance(output, tuple) else output
        self.active.tensors["velocity"] = _clone_cpu(velocity)
        required = {"timestep", "noisy_sample", "frame_condition", "final_hidden", "velocity"}
        missing = required.difference(self.active.tensors)
        if missing:
            raise RuntimeError(f"Incomplete teacher step capture: {sorted(missing)}")
        self.chunk_steps.append(self.active)
        completed_step = self.active.step_id
        self.active = None
        self.call_index += 1
        if completed_step == self.num_steps - 1:
            self._flush_chunk()

    def _flush_chunk(self) -> None:
        if len(self.chunk_steps) != self.num_steps:
            raise RuntimeError(f"Expected {self.num_steps} captured steps, got {len(self.chunk_steps)}")
        chunk_id = self.chunk_steps[0].chunk_id
        if any(step.chunk_id != chunk_id for step in self.chunk_steps):
            raise RuntimeError("Captured steps cross chunk boundary")
        tensors = dict(self.chunk_steps[0].shared)
        for step in self.chunk_steps:
            for name, tensor in step.tensors.items():
                tensors[f"step_{step.step_id}_{name}"] = tensor
        self.on_chunk(chunk_id, tensors)
        self.chunk_steps.clear()

    def case_tensors(self) -> dict[str, torch.Tensor]:
        missing = [
            name
            for name, value in (
                ("image_condition_latent", self.image_condition_latent),
                ("current_txt", self.current_txt),
                ("cached_txt", self.cached_txt),
                ("vec_txt", self.vec_txt),
            )
            if value is None
        ]
        if missing:
            raise RuntimeError(f"Missing case captures: {missing}")
        return {
            "image_condition_latent": self.image_condition_latent,
            "current_txt": self.current_txt,
            "cached_txt": self.cached_txt,
            "vec_txt": self.vec_txt,
        }