Download scripts/build_predictor_chunk0_context.py from Cccccz/Causal-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 7.57 kB
-
https://huggingface.co/Cccccz/Causal-Forcing/resolve/main/scripts/build_predictor_chunk0_context.py
- Command line
-
hf download hf://Cccccz/Causal-Forcing/scripts/build_predictor_chunk0_context.py
-
curl -L -o build_predictor_chunk0_context.py https://huggingface.co/Cccccz/Causal-Forcing/resolve/main/scripts/build_predictor_chunk0_context.py
7.57 kB
| #!/usr/bin/env python3 | |
| """Add the minimal chunk-0 context required by Stage-1 Predictor training.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| import shutil | |
| import time | |
| from pathlib import Path | |
| import build_predictor_offline_data as base | |
| torch = base.torch | |
| OmegaConf = base.OmegaConf | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--gpu", default=base.PHYSICAL_GPU) | |
| parser.add_argument("--dataset_root", type=Path, required=True) | |
| parser.add_argument( | |
| "--config_path", | |
| type=Path, | |
| default=Path("configs/causal_forcing_dmd_chunkwise.yaml"), | |
| ) | |
| parser.add_argument( | |
| "--checkpoint_path", | |
| type=Path, | |
| default=Path("checkpoints/chunkwise/causal_forcing.pt"), | |
| ) | |
| parser.add_argument("--prompt_ids", type=int, nargs="*", default=None) | |
| parser.add_argument("--generation_seed", type=int, default=0) | |
| parser.add_argument("--max_new_prompts", type=int, default=None) | |
| parser.add_argument("--min_free_gib", type=float, default=500.0) | |
| return parser.parse_args() | |
| def generate_chunk0( | |
| pipeline, | |
| recorder: base.TrajectoryRecorder, | |
| prompt: str, | |
| generation_seed: int, | |
| device: torch.device, | |
| ): | |
| base.set_seed(generation_seed) | |
| base.reset_caches(pipeline, 1, torch.bfloat16, device) | |
| recorder.clean_prefeatures = {layer: [] for layer in recorder.layers} | |
| conditional = pipeline.text_encoder(text_prompts=[prompt]) | |
| # Draw the full 21-latent noise tensor so the RNG state and chunk-0 slice | |
| # exactly match the original offline rollout. | |
| full_noise = torch.randn( | |
| 1, 21, base.LATENT_CHANNELS, base.LATENT_HEIGHT, base.LATENT_WIDTH, | |
| dtype=torch.bfloat16, device=device, | |
| ) | |
| noisy_input = full_noise[:, : pipeline.num_frame_per_block] | |
| timesteps = pipeline.denoising_step_list.to(device=device) | |
| trajectory = {} | |
| torch.cuda.reset_peak_memory_stats() | |
| torch.cuda.synchronize() | |
| started = time.perf_counter() | |
| timestep = None | |
| denoised_pred = None | |
| for step, current_timestep in enumerate(timesteps): | |
| timestep = torch.ones( | |
| [1, pipeline.num_frame_per_block], device=device, dtype=torch.int64 | |
| ) * current_timestep | |
| recorder.start_denoising_step() | |
| _, denoised_pred = pipeline.generator( | |
| noisy_image_or_video=noisy_input, | |
| conditional_dict=conditional, | |
| timestep=timestep, | |
| kv_cache=pipeline.kv_cache1, | |
| crossattn_cache=pipeline.crossattn_cache, | |
| current_start=0, | |
| ) | |
| trajectory[f"chunk_00_step_{step:02d}_final_hidden"] = ( | |
| recorder.finish_denoising_step() | |
| ) | |
| if step < len(timesteps) - 1: | |
| next_timestep = timesteps[step + 1] | |
| flat = denoised_pred.flatten(0, 1) | |
| noisy_input = pipeline.scheduler.add_noise( | |
| flat, | |
| torch.randn_like(flat), | |
| next_timestep * torch.ones( | |
| [pipeline.num_frame_per_block], | |
| device=device, | |
| dtype=torch.long, | |
| ), | |
| ).unflatten(0, denoised_pred.shape[:2]) | |
| if denoised_pred is None or timestep is None: | |
| raise RuntimeError("Chunk-0 denoising produced no output") | |
| recorder.start_clean_pass() | |
| pipeline.generator( | |
| noisy_image_or_video=denoised_pred, | |
| conditional_dict=conditional, | |
| timestep=torch.ones_like(timestep) * pipeline.args.context_noise, | |
| kv_cache=pipeline.kv_cache1, | |
| crossattn_cache=pipeline.crossattn_cache, | |
| current_start=0, | |
| ) | |
| recorder.finish_clean_pass() | |
| torch.cuda.synchronize() | |
| return ( | |
| trajectory, | |
| recorder.clean_prefeatures, | |
| time.perf_counter() - started, | |
| torch.cuda.max_memory_allocated() / 1024**3, | |
| ) | |
| def save_context(prompt_dir, trajectory, prefeatures, elapsed_s, peak_gib): | |
| destination = prompt_dir / "chunk0_context" | |
| partial = prompt_dir / "chunk0_context.partial" | |
| if partial.exists(): | |
| shutil.rmtree(partial) | |
| partial.mkdir(parents=True) | |
| metadata = {"dataset_version": "2", "kind": "chunk0_context"} | |
| base.atomic_save_safetensors( | |
| trajectory, partial / "trajectory.safetensors", metadata | |
| ) | |
| for layer in range(30): | |
| values = prefeatures[layer] | |
| if len(values) != 1: | |
| raise RuntimeError(f"Layer {layer} captured {len(values)} values") | |
| base.atomic_save_safetensors( | |
| {"chunk_00": values[0]}, | |
| partial / "clean_prefeatures" / f"block_{layer:02d}.safetensors", | |
| {**metadata, "block_id": str(layer)}, | |
| ) | |
| base.atomic_write_json( | |
| partial / "metadata.json", | |
| { | |
| "kind": "context_only", | |
| "chunk": 0, | |
| "is_training_target": False, | |
| "hidden_steps": [0, 1, 2, 3], | |
| "layers": list(range(30)), | |
| "elapsed_s": elapsed_s, | |
| "peak_gpu_gib": peak_gib, | |
| }, | |
| ) | |
| (partial / "_SUCCESS").write_text("ok\n", encoding="utf-8") | |
| if destination.exists(): | |
| shutil.rmtree(destination) | |
| os.replace(partial, destination) | |
| return sum(p.stat().st_size for p in destination.rglob("*") if p.is_file()) | |
| def main() -> None: | |
| args = parse_args() | |
| root = args.dataset_root.resolve() | |
| config_path = base.resolve_path(args.config_path) | |
| checkpoint_path = base.resolve_path(args.checkpoint_path) | |
| selected = json.loads((root / "prompt_selection.json").read_text())["prompts"] | |
| requested = ( | |
| set(range(len(selected))) | |
| if args.prompt_ids is None | |
| else set(args.prompt_ids) | |
| ) | |
| pending = [ | |
| (index, item) | |
| for index, item in enumerate(selected) | |
| if index in requested | |
| and not (root / f"prompt_{index:04d}" / "chunk0_context" / "_SUCCESS").exists() | |
| ] | |
| if not pending: | |
| print("[context] all requested prompt contexts exist", flush=True) | |
| return | |
| config = OmegaConf.merge( | |
| OmegaConf.load(base.REPO_ROOT / "configs/default_config.yaml"), | |
| OmegaConf.load(config_path), | |
| ) | |
| device = torch.device("cuda") | |
| pipeline = base.build_pipeline(config, checkpoint_path, device) | |
| recorder = base.TrajectoryRecorder(pipeline.generator.model, list(range(30))) | |
| generated = 0 | |
| try: | |
| for index, selection in pending: | |
| if args.max_new_prompts is not None and generated >= args.max_new_prompts: | |
| break | |
| free_gib = shutil.disk_usage(root).free / 1024**3 | |
| if free_gib < args.min_free_gib: | |
| raise RuntimeError(f"Only {free_gib:.1f} GiB free") | |
| trajectory, prefeatures, elapsed_s, peak_gib = generate_chunk0( | |
| pipeline, recorder, selection["prompt"], args.generation_seed, device | |
| ) | |
| size = save_context( | |
| root / f"prompt_{index:04d}", trajectory, prefeatures, | |
| elapsed_s, peak_gib, | |
| ) | |
| generated += 1 | |
| print( | |
| f"[context] saved prompt_{index:04d}: {size / 1024**3:.3f} GiB, " | |
| f"{elapsed_s:.1f}s, peak={peak_gib:.1f} GiB", | |
| flush=True, | |
| ) | |
| del trajectory, prefeatures | |
| torch.cuda.empty_cache() | |
| finally: | |
| recorder.close() | |
| print(f"[context] generated {generated} prompt contexts", flush=True) | |
| if __name__ == "__main__": | |
| main() | |