Download predictor_training/offline_data.py from Cccccz/Causal-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 15 kB
-
https://huggingface.co/Cccccz/Causal-Forcing/resolve/main/predictor_training/offline_data.py
- Command line
-
hf download hf://Cccccz/Causal-Forcing/predictor_training/offline_data.py
-
curl -L -o offline_data.py https://huggingface.co/Cccccz/Causal-Forcing/resolve/main/predictor_training/offline_data.py
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 | |
| 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] | |
| 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, | |
| ) | |
| 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() | |
| 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] | |
| ), | |
| } | |