Causal-Forcing / scripts /build_predictor_chunk0_context.py
Cccccz's picture
Upload code and configuration only
ae8ade0 verified
Raw History Blame Contribute Delete
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()
@torch.inference_mode()
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()