Self-Forcing / pipeline /predictor_v4_dmd_training.py
Cccccz's picture
Add files using upload-large-folder tool
d1d122e verified
Raw
History Blame Contribute Delete
16.6 kB
"""Random-exit DMD rollout for Predictor-v4 and optional Full-Generator tuning."""
from __future__ import annotations
from contextlib import nullcontext
from typing import Any
import torch
import torch.distributed as dist
from predictor_training.rollout_cache import (
assert_clean_history_extent,
build_predictor_workspace,
reset_main_caches,
reset_predictor_workspace,
)
CHUNK_FRAMES = 3
NUM_CHUNKS = 7
NUM_STEPS = 4
TOKENS_PER_FRAME = 1560
TOKENS_PER_CHUNK = CHUNK_FRAMES * TOKENS_PER_FRAME
def _unwrap(module: torch.nn.Module) -> torch.nn.Module:
while hasattr(module, "module"):
module = module.module
return module
class PredictorV4DMDTrainingPipeline:
"""Generate one detached random-exit trajectory for DMD.
``predictor_only`` samples exits P1/P2/P3. ``joint`` additionally samples
the first Full step. One exit is shared by all chunks and distributed
ranks. Only the selected call for each chunk retains a graph.
"""
def __init__(
self,
*,
denoising_step_list: torch.Tensor,
scheduler: Any,
generator: torch.nn.Module,
predictor: torch.nn.Module,
training_mode: str,
context_noise: int = 0,
forced_exit_step: int | None = None,
) -> None:
training_mode = str(training_mode).lower()
if training_mode not in {"predictor_only", "joint"}:
raise ValueError(f"Unknown Predictor DMD mode {training_mode!r}")
if len(denoising_step_list) != NUM_STEPS:
raise ValueError("Predictor-v4 DMD requires exactly four denoising steps")
self.denoising_step_list = denoising_step_list
self.scheduler = scheduler
self.generator = generator
self.predictor = predictor
self.training_mode = training_mode
self.context_noise = int(context_noise)
self.forced_exit_step = (
None if forced_exit_step is None else int(forced_exit_step)
)
if self.forced_exit_step is not None:
minimum = 1 if training_mode == "predictor_only" else 0
if not minimum <= self.forced_exit_step < NUM_STEPS:
raise ValueError(
f"forced_exit_step must be in [{minimum}, {NUM_STEPS - 1}]"
)
self.kv_cache1: list[dict[str, torch.Tensor]] | None = None
self.crossattn_cache: list[dict[str, torch.Tensor | bool]] | None = None
self.last_exit_step: int | None = None
@property
def predictor_module(self) -> torch.nn.Module:
return _unwrap(self.predictor)
@property
def generator_wrapper(self) -> torch.nn.Module:
return _unwrap(self.generator)
def _sample_exit_step(self, device: torch.device) -> int:
if self.forced_exit_step is not None:
value = torch.tensor(
[self.forced_exit_step], device=device, dtype=torch.long
)
elif not dist.is_initialized() or dist.get_rank() == 0:
low = 1 if self.training_mode == "predictor_only" else 0
value = torch.randint(low, NUM_STEPS, (1,), device=device)
else:
value = torch.empty(1, device=device, dtype=torch.long)
if dist.is_initialized():
dist.broadcast(value, src=0)
return int(value.item())
def _initialize_caches(
self,
*,
batch_size: int,
dtype: torch.dtype,
device: torch.device,
) -> None:
capacity = NUM_CHUNKS * TOKENS_PER_CHUNK
self.kv_cache1 = [
{
"k": torch.zeros(
batch_size, capacity, 12, 128, dtype=dtype, device=device
),
"v": torch.zeros(
batch_size, capacity, 12, 128, dtype=dtype, device=device
),
"global_end_index": torch.zeros(
1, dtype=torch.long, device=device
),
"local_end_index": torch.zeros(
1, dtype=torch.long, device=device
),
}
for _ in range(30)
]
self.crossattn_cache = [
{
"k": torch.zeros(
batch_size, 512, 12, 128, dtype=dtype, device=device
),
"v": torch.zeros(
batch_size, 512, 12, 128, dtype=dtype, device=device
),
"is_init": False,
}
for _ in range(30)
]
def _prepare_caches(self, noise: torch.Tensor) -> None:
if (
self.kv_cache1 is None
or self.kv_cache1[0]["k"].shape[0] != noise.shape[0]
or self.kv_cache1[0]["k"].device != noise.device
or self.kv_cache1[0]["k"].dtype != noise.dtype
):
self._initialize_caches(
batch_size=noise.shape[0],
dtype=noise.dtype,
device=noise.device,
)
else:
reset_main_caches(self.kv_cache1, self.crossattn_cache)
def _timestep(
self,
step_id: int,
*,
batch_size: int,
device: torch.device,
) -> torch.Tensor:
return torch.full(
(batch_size, CHUNK_FRAMES),
float(self.denoising_step_list[int(step_id)].item()),
dtype=torch.float32,
device=device,
)
def _renoise(
self,
clean: torch.Tensor,
*,
next_timestep: torch.Tensor,
) -> torch.Tensor:
return self.scheduler.add_noise(
clean.detach().flatten(0, 1),
torch.randn_like(clean.detach().flatten(0, 1)),
next_timestep.flatten(0, 1),
).unflatten(0, clean.shape[:2]).detach()
def _full_step(
self,
*,
latent: torch.Tensor,
conditional_dict: dict[str, torch.Tensor],
timestep: torch.Tensor,
current_start: int,
capture_hidden: bool,
retain_graph: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
captured: list[torch.Tensor] = []
handle = None
if capture_hidden:
def capture_head_input(module, args):
if not args or not torch.is_tensor(args[0]):
raise RuntimeError("Wan head hook did not receive hidden states")
captured.append(args[0])
handle = self.generator_wrapper.model.head.register_forward_pre_hook(
capture_head_input
)
grad_context = nullcontext() if retain_graph else torch.no_grad()
try:
with grad_context, torch.autocast(
device_type="cuda", dtype=torch.bfloat16
):
flow, clean = self.generator(
noisy_image_or_video=latent,
conditional_dict=conditional_dict,
timestep=timestep,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start,
)
finally:
if handle is not None:
handle.remove()
if capture_hidden and len(captured) != 1:
raise RuntimeError(
f"Expected one Full hidden capture, got {len(captured)}"
)
return flow, clean, captured[0] if captured else None
def _predictor_step(
self,
*,
latent: torch.Tensor,
timestep: torch.Tensor,
anchor_hidden: torch.Tensor,
previous_chunk_hidden: torch.Tensor,
predictor_cache: dict[int, dict[str, torch.Tensor]],
current_start: int,
retain_graph: bool,
) -> dict[str, torch.Tensor]:
grad_context = nullcontext() if retain_graph else torch.no_grad()
with grad_context, torch.autocast(
device_type="cuda", dtype=torch.bfloat16
):
return self.predictor(
target_latent=latent,
target_timestep=timestep,
anchor_hidden=anchor_hidden,
previous_chunk_hidden=previous_chunk_hidden,
kv_cache=predictor_cache,
crossattn_cache=self.crossattn_cache,
current_start=current_start,
)
@torch.no_grad()
def _commit_clean_history(
self,
*,
clean: torch.Tensor,
conditional_dict: dict[str, torch.Tensor],
current_start: int,
expected_end: int,
) -> None:
timestep = torch.full(
(clean.shape[0], clean.shape[1]),
float(self.context_noise),
dtype=torch.float32,
device=clean.device,
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
self.generator(
noisy_image_or_video=clean.detach(),
conditional_dict=conditional_dict,
timestep=timestep,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=current_start,
)
assert_clean_history_extent(
self.kv_cache1, expected_tokens=expected_end
)
def inference_with_trajectory(
self,
noise: torch.Tensor,
initial_latent: torch.Tensor | None = None,
return_sim_step: bool = False,
**conditional_dict: torch.Tensor,
):
if initial_latent is not None:
raise NotImplementedError("Predictor-v4 DMD currently supports T2V only")
if noise.shape[1] != NUM_CHUNKS * CHUNK_FRAMES:
raise ValueError(
f"Expected 21 latent frames, got shape {tuple(noise.shape)}"
)
self._prepare_caches(noise)
exit_step = self._sample_exit_step(noise.device)
self.last_exit_step = exit_step
outputs: list[torch.Tensor] = []
previous_hidden: dict[int, torch.Tensor] = {}
for chunk_id in range(NUM_CHUNKS):
current_start = chunk_id * TOKENS_PER_CHUNK
current_end = current_start + TOKENS_PER_CHUNK
latent = noise[
:, chunk_id * CHUNK_FRAMES : (chunk_id + 1) * CHUNK_FRAMES
].detach()
if exit_step == 0:
timestep = self._timestep(
0, batch_size=noise.shape[0], device=noise.device
)
_, clean, _ = self._full_step(
latent=latent,
conditional_dict=conditional_dict,
timestep=timestep,
current_start=current_start,
capture_hidden=False,
retain_graph=self.training_mode == "joint",
)
outputs.append(clean)
self._commit_clean_history(
clean=clean,
conditional_dict=conditional_dict,
current_start=current_start,
expected_end=current_end,
)
continue
if chunk_id == 0:
current_hidden: dict[int, torch.Tensor] = {}
clean = None
for step_id in range(exit_step + 1):
timestep = self._timestep(
step_id,
batch_size=noise.shape[0],
device=noise.device,
)
_, clean, hidden = self._full_step(
latent=latent,
conditional_dict=conditional_dict,
timestep=timestep,
current_start=current_start,
capture_hidden=step_id > 0,
retain_graph=(
self.training_mode == "joint"
and step_id == exit_step
),
)
if step_id > 0:
if hidden is None:
raise RuntimeError(
f"Chunk 0 Full step {step_id} has no hidden"
)
current_hidden[step_id] = hidden.detach()
if step_id < exit_step:
latent = self._renoise(
clean,
next_timestep=self._timestep(
step_id + 1,
batch_size=noise.shape[0],
device=noise.device,
),
)
if clean is None:
raise RuntimeError("Chunk 0 produced no clean latent")
outputs.append(clean)
self._commit_clean_history(
clean=clean,
conditional_dict=conditional_dict,
current_start=current_start,
expected_end=current_end,
)
previous_hidden = current_hidden
continue
predictor_cache = build_predictor_workspace(
self.kv_cache1,
source_block_ids=tuple(self.predictor_module.source_block_ids),
history_tokens=current_start,
current_tokens=TOKENS_PER_CHUNK,
)
timestep0 = self._timestep(
0, batch_size=noise.shape[0], device=noise.device
)
_, clean0, anchor = self._full_step(
latent=latent,
conditional_dict=conditional_dict,
timestep=timestep0,
current_start=current_start,
capture_hidden=True,
retain_graph=False,
)
if anchor is None:
raise RuntimeError(f"Chunk {chunk_id} has no Full step-0 anchor")
latent = self._renoise(
clean0,
next_timestep=self._timestep(
1, batch_size=noise.shape[0], device=noise.device
),
)
current_hidden = {}
pred_clean = None
for step_id in range(1, exit_step + 1):
if step_id not in previous_hidden:
raise RuntimeError(
f"Previous chunk lacks Predictor hidden step {step_id}"
)
reset_predictor_workspace(
predictor_cache, history_tokens=current_start
)
timestep = self._timestep(
step_id,
batch_size=noise.shape[0],
device=noise.device,
)
output = self._predictor_step(
latent=latent,
timestep=timestep,
anchor_hidden=anchor.detach(),
previous_chunk_hidden=previous_hidden[step_id],
predictor_cache=predictor_cache,
current_start=current_start,
retain_graph=step_id == exit_step,
)
flow = output["pred_flow"]
pred_clean = self.generator_wrapper._convert_flow_pred_to_x0(
flow_pred=flow.flatten(0, 1),
xt=latent.flatten(0, 1),
timestep=timestep.flatten(0, 1),
).unflatten(0, flow.shape[:2])
anchor = output["pred_hidden"].detach()
current_hidden[step_id] = anchor
if step_id < exit_step:
latent = self._renoise(
pred_clean,
next_timestep=self._timestep(
step_id + 1,
batch_size=noise.shape[0],
device=noise.device,
),
)
if pred_clean is None:
raise RuntimeError(f"Chunk {chunk_id} produced no Predictor output")
outputs.append(pred_clean)
self._commit_clean_history(
clean=pred_clean,
conditional_dict=conditional_dict,
current_start=current_start,
expected_end=current_end,
)
previous_hidden = current_hidden
result = torch.cat(outputs, dim=1)
if return_sim_step:
return result, 0, 0, exit_step + 1
return result, 0, 0
__all__ = ["PredictorV4DMDTrainingPipeline"]