Cccccz's picture
Upload code and configuration only
ae8ade0 verified
Raw History Blame Contribute Delete
15 kB
"""In-memory access to the prompt-sharded Causal-Forcing predictor dataset."""
from __future__ import annotations
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Iterable
import torch
from safetensors import safe_open
from wan.modules.causal_model import causal_rope_apply
TOKENS_PER_FRAME = 30 * 52
FRAMES_PER_CHUNK = 3
TOKENS_PER_CHUNK = TOKENS_PER_FRAME * FRAMES_PER_CHUNK
@dataclass
class PromptCommon:
hidden: list[list[torch.Tensor]]
noisy: dict[tuple[int, int], torch.Tensor]
flow: dict[tuple[int, int], torch.Tensor]
timestep: dict[tuple[int, int], torch.Tensor]
@dataclass
class LayerPromptCache:
history_k: torch.Tensor
history_v: torch.Tensor
cross_k: torch.Tensor
cross_v: torch.Tensor
class OfflinePredictorStore:
"""Keep common trajectories in RAM; load one block's KV cache at a time."""
def __init__(
self,
root: str | Path,
prompt_ids: Iterable[int],
num_chunks: int = 7,
max_history_chunks: int = 7,
layer_cache_device: torch.device | str | None = None,
) -> None:
self.root = Path(root).resolve()
self.prompt_ids = sorted(set(int(value) for value in prompt_ids))
self.num_chunks = int(num_chunks)
self.max_history_chunks = int(max_history_chunks)
self.layer_cache_device = (
torch.device(layer_cache_device)
if layer_cache_device is not None
else torch.device("cpu")
)
if self.num_chunks < 2:
raise ValueError("num_chunks must be at least 2")
if self.max_history_chunks < 1:
raise ValueError("max_history_chunks must be positive")
self.common: dict[int, PromptCommon] = {}
self.layer_cache: dict[int, LayerPromptCache] = {}
self.layer_caches: dict[int, dict[int, LayerPromptCache]] = {}
self.layer_id: int | None = None
self._load_common()
def _load_common(self) -> None:
started = time.perf_counter()
for offset, prompt_id in enumerate(self.prompt_ids, start=1):
path = (
self.root
/ f"prompt_{prompt_id:04d}"
/ "trajectory.safetensors"
)
context_path = (
self.root
/ f"prompt_{prompt_id:04d}"
/ "chunk0_context"
/ "trajectory.safetensors"
)
hidden: list[list[torch.Tensor]] = []
noisy: dict[tuple[int, int], torch.Tensor] = {}
flow: dict[tuple[int, int], torch.Tensor] = {}
timestep: dict[tuple[int, int], torch.Tensor] = {}
with safe_open(
context_path, framework="pt", device="cpu"
) as context_handle, safe_open(
path, framework="pt", device="cpu"
) as handle:
for chunk in range(self.num_chunks):
chunk_hidden = []
for step in range(4):
prefix = f"chunk_{chunk:02d}_step_{step:02d}"
source = context_handle if chunk == 0 else handle
chunk_hidden.append(
source.get_tensor(
f"{prefix}_final_hidden"
).squeeze(0)
)
if chunk >= 1 and step >= 1:
noisy[(chunk, step)] = handle.get_tensor(
f"{prefix}_noisy_latent"
).squeeze(0)
flow[(chunk, step)] = handle.get_tensor(
f"{prefix}_flow"
).squeeze(0)
timestep[(chunk, step)] = handle.get_tensor(
f"{prefix}_timestep"
).squeeze(0)
hidden.append(chunk_hidden)
self.common[prompt_id] = PromptCommon(
hidden=hidden,
noisy=noisy,
flow=flow,
timestep=timestep,
)
if offset % 10 == 0 or offset == len(self.prompt_ids):
elapsed = time.perf_counter() - started
print(
f"[data] common {offset}/{len(self.prompt_ids)} "
f"({elapsed:.1f}s)",
flush=True,
)
@torch.inference_mode()
def load_layer_cache(
self,
layer_id: int,
teacher_model: torch.nn.Module,
device: torch.device,
) -> None:
"""Project clean prefeatures once with frozen Teacher K/V weights."""
self.layer_caches = {}
self.layer_cache = {}
self.layer_id = int(layer_id)
teacher_block = teacher_model.blocks[layer_id]
heads = teacher_block.num_heads
head_dim = teacher_block.dim // heads
if teacher_model.freqs.device != device:
teacher_model.freqs = teacher_model.freqs.to(device)
started = time.perf_counter()
history_chunks = self.num_chunks - 1
grid_sizes = torch.tensor(
[[history_chunks * FRAMES_PER_CHUNK, 30, 52]], dtype=torch.long
)
for offset, prompt_id in enumerate(self.prompt_ids, start=1):
prompt_dir = self.root / f"prompt_{prompt_id:04d}"
prefeature_path = (
prompt_dir
/ "clean_prefeatures"
/ f"block_{layer_id:02d}.safetensors"
)
context_prefeature_path = (
prompt_dir
/ "chunk0_context"
/ "clean_prefeatures"
/ f"block_{layer_id:02d}.safetensors"
)
with safe_open(
prefeature_path, framework="pt", device="cpu"
) as handle, safe_open(
context_prefeature_path, framework="pt", device="cpu"
) as context_handle:
prefeature = torch.cat(
[
(
context_handle.get_tensor("chunk_00")
if chunk == 0
else handle.get_tensor(f"chunk_{chunk:02d}")
)
for chunk in range(history_chunks)
],
dim=1,
)
prefeature = prefeature.to(
device=device,
dtype=torch.bfloat16,
non_blocking=False,
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
key = teacher_block.self_attn.norm_k(
teacher_block.self_attn.k(prefeature)
).view(1, -1, heads, head_dim)
value = teacher_block.self_attn.v(prefeature).view(
1, -1, heads, head_dim
)
key = causal_rope_apply(
key,
grid_sizes,
teacher_model.freqs,
start_frame=0,
)
key = (
key.reshape(
1, history_chunks, TOKENS_PER_CHUNK, heads, head_dim
)
.squeeze(0)
.to(device=self.layer_cache_device, dtype=torch.bfloat16)
.contiguous()
)
value = (
value.reshape(
1, history_chunks, TOKENS_PER_CHUNK, heads, head_dim
)
.squeeze(0)
.to(device=self.layer_cache_device, dtype=torch.bfloat16)
.contiguous()
)
cross_path = prompt_dir / "cross_attention.safetensors"
with safe_open(
cross_path, framework="pt", device="cpu"
) as handle:
cross_k = handle.get_tensor(
f"block_{layer_id:02d}_k"
).squeeze(0).to(self.layer_cache_device)
cross_v = handle.get_tensor(
f"block_{layer_id:02d}_v"
).squeeze(0).to(self.layer_cache_device)
self.layer_cache[prompt_id] = LayerPromptCache(
history_k=key,
history_v=value,
cross_k=cross_k,
cross_v=cross_v,
)
del prefeature, key, value
if offset % 10 == 0 or offset == len(self.prompt_ids):
elapsed = time.perf_counter() - started
print(
f"[data] block {layer_id:02d} cache "
f"{offset}/{len(self.prompt_ids)} ({elapsed:.1f}s)",
flush=True,
)
torch.cuda.empty_cache()
@torch.inference_mode()
def load_layer_caches(
self,
layer_ids: Iterable[int],
teacher_model: torch.nn.Module,
device: torch.device,
) -> None:
"""Load Teacher-layer caches, retaining overlap with the previous group."""
requested = list(dict.fromkeys(int(value) for value in layer_ids))
if not requested:
raise ValueError("At least one layer cache is required")
loaded: dict[int, dict[int, LayerPromptCache]] = {
layer_id: self.layer_caches[layer_id]
for layer_id in requested
if layer_id in self.layer_caches
}
reused = sorted(loaded)
# Drop layers that are no longer requested before projecting a new one.
# The dictionaries in ``loaded`` keep only the overlapping layers alive.
self.layer_caches = {}
self.layer_cache = {}
self.layer_id = None
if reused:
print(f"[data] reusing layer caches {reused}", flush=True)
for layer_id in requested:
if layer_id in loaded:
continue
self.load_layer_cache(layer_id, teacher_model, device)
loaded[layer_id] = self.layer_cache
self.layer_caches = loaded
def batch_layers(
self,
prompt_ids: list[int],
chunk: int,
target_step: int,
layer_ids: Iterable[int],
) -> dict[str, Any]:
"""Build one sample batch with independent K/V inputs for each block."""
requested = [int(value) for value in layer_ids]
if not requested:
raise ValueError("At least one layer ID is required")
missing = sorted(set(requested) - set(self.layer_caches))
if missing:
raise RuntimeError(f"Layer caches not loaded: {missing}")
output = self._base_batch(prompt_ids, chunk, target_step)
output["layer_ids"] = requested
for position, layer_id in enumerate(requested):
cache = [self.layer_caches[layer_id][prompt_id] for prompt_id in prompt_ids]
output[f"history_k_{position}"] = torch.stack(
[
item.history_k[:chunk].reshape(
chunk * TOKENS_PER_CHUNK,
item.history_k.shape[-2],
item.history_k.shape[-1],
)
for item in cache
]
)
output[f"history_v_{position}"] = torch.stack(
[
item.history_v[:chunk].reshape(
chunk * TOKENS_PER_CHUNK,
item.history_v.shape[-2],
item.history_v.shape[-1],
)
for item in cache
]
)
output[f"cross_k_{position}"] = torch.stack(
[item.cross_k for item in cache]
)
output[f"cross_v_{position}"] = torch.stack(
[item.cross_v for item in cache]
)
return output
def batch(
self,
prompt_ids: list[int],
chunk: int,
target_step: int,
) -> dict[str, Any]:
if self.layer_id is None or not self.layer_cache:
raise RuntimeError("load_layer_cache must be called first")
output = self._base_batch(prompt_ids, chunk, target_step)
cache = [self.layer_cache[prompt_id] for prompt_id in prompt_ids]
history_start = max(0, chunk - self.max_history_chunks)
history_chunks = chunk - history_start
output.update(
{
"history_k": torch.stack(
[
item.history_k[history_start:chunk].reshape(
history_chunks * TOKENS_PER_CHUNK,
item.history_k.shape[-2],
item.history_k.shape[-1],
)
for item in cache
]
),
"history_v": torch.stack(
[
item.history_v[history_start:chunk].reshape(
history_chunks * TOKENS_PER_CHUNK,
item.history_v.shape[-2],
item.history_v.shape[-1],
)
for item in cache
]
),
"cross_k": torch.stack([item.cross_k for item in cache]),
"cross_v": torch.stack([item.cross_v for item in cache]),
}
)
return output
def _base_batch(
self,
prompt_ids: list[int],
chunk: int,
target_step: int,
) -> dict[str, Any]:
if chunk < 1 or chunk >= self.num_chunks:
raise ValueError(
f"Trainable chunk must be 1..{self.num_chunks - 1}, got {chunk}"
)
if target_step < 1 or target_step > 3:
raise ValueError(
f"Target denoising step must be 1..3, got {target_step}"
)
anchor_step = target_step - 1
common = [self.common[prompt_id] for prompt_id in prompt_ids]
return {
"prompt_ids": prompt_ids,
"chunk": chunk,
"anchor_step": anchor_step,
"target_step": target_step,
"noisy_latent": torch.stack(
[item.noisy[(chunk, target_step)] for item in common]
),
"anchor_hidden": torch.stack(
[item.hidden[chunk][anchor_step] for item in common]
),
"previous_hidden": torch.stack(
[item.hidden[chunk - 1][target_step] for item in common]
),
"target_hidden": torch.stack(
[item.hidden[chunk][target_step] for item in common]
),
"target_flow": torch.stack(
[item.flow[(chunk, target_step)] for item in common]
),
"timestep": torch.stack(
[item.timestep[(chunk, target_step)] for item in common]
),
}