"""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] ), }