Self-Forcing / trainer /predictor_v4.py
Cccccz's picture
Add files using upload-large-folder tool
510ab6b verified
Raw
History Blame Contribute Delete
26.4 kB
"""Eight-GPU DDP trainer for the offline Self-Forcing Predictor-v4."""
from __future__ import annotations
import gc
import json
import math
import os
import random
import time
from pathlib import Path
from typing import Any
import torch
import torch.distributed as dist
import torch.nn.functional as F
from omegaconf import OmegaConf
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.optim import AdamW
from torch.optim.lr_scheduler import LambdaLR
from torch.utils.data import DataLoader
from predictor_training import (
DistributedContextBucketBatchSampler,
PredictorV4PairDataset,
move_batch_to_device,
predictor_v4_collate,
)
from predictor_training.cache import build_cross_attention_cache
from predictor_training.checkpoint import (
atomic_torch_save,
capture_rng_state,
restore_rng_state,
save_predictor_weights,
trainable_state_dict,
)
def _cosine_with_warmup(step: int, warmup: int, total: int) -> float:
if step < warmup:
return max(1e-8, float(step + 1) / max(1, warmup))
progress = min(1.0, float(step - warmup) / max(1, total - warmup))
return 0.5 * (1.0 + math.cos(math.pi * progress))
def _make_scheduler(
optimizer: torch.optim.Optimizer,
*,
scheduler_warmup_steps: int,
block_warmup_steps: int,
max_steps: int,
) -> LambdaLR:
return LambdaLR(
optimizer,
[
lambda step: _cosine_with_warmup(
step, scheduler_warmup_steps, max_steps
),
lambda step: (
0.0
if step < block_warmup_steps
else _cosine_with_warmup(
step - block_warmup_steps,
scheduler_warmup_steps,
max(1, max_steps - block_warmup_steps),
)
),
],
)
def _distributed_mean(values: torch.Tensor) -> torch.Tensor:
dist.all_reduce(values, op=dist.ReduceOp.SUM)
return values / dist.get_world_size()
def _atomic_json(path: Path, payload: dict[str, Any]) -> None:
temporary = path.with_suffix(path.suffix + f".tmp.{os.getpid()}")
temporary.write_text(
json.dumps(payload, indent=2, ensure_ascii=False, sort_keys=True, default=str)
+ "\n",
encoding="utf-8",
)
os.replace(temporary, path)
def _prune_snapshots(output_dir: Path, keep: int) -> None:
snapshots = sorted(output_dir.glob("predictor_step_*.safetensors"))
for path in snapshots[:-keep] if keep > 0 else snapshots:
path.unlink()
def _load_teacher_and_predictor(
cfg,
) -> torch.nn.Module:
"""Load DMD EMA on CPU and retain only Predictor/frozen projections."""
from model.predictor_v4 import SelfForcingPredictorV4
from utils.wan_wrapper import WanDiffusionWrapper
wrapper = WanDiffusionWrapper(
model_name=str(cfg.teacher_model_name),
timestep_shift=float(cfg.timestep_shift),
is_causal=True,
local_attn_size=int(cfg.local_attn_size),
sink_size=int(cfg.sink_size),
)
checkpoint = torch.load(
str(cfg.teacher_checkpoint),
map_location="cpu",
weights_only=False,
)
checkpoint_key = str(cfg.teacher_checkpoint_key)
if checkpoint_key not in checkpoint:
raise KeyError(
f"{cfg.teacher_checkpoint} has no {checkpoint_key!r}; "
f"available keys are {sorted(checkpoint)}"
)
result = wrapper.load_state_dict(
checkpoint[checkpoint_key],
strict=bool(cfg.strict_teacher_load),
)
if not bool(cfg.strict_teacher_load):
if result.unexpected_keys:
raise RuntimeError(
f"Unexpected teacher weights: {result.unexpected_keys[:20]}"
)
teacher_model = wrapper.model
source_blocks = tuple(int(value) for value in cfg.source_block_ids)
predictor = SelfForcingPredictorV4.from_teacher(
teacher_model,
source_block_ids=source_blocks,
)
del checkpoint, wrapper, teacher_model
gc.collect()
return predictor
def _configure_precision(
module: torch.nn.Module,
*,
device: torch.device,
activation_dtype: torch.dtype,
fp32_trainable_params: bool,
) -> None:
module.to(device=device, dtype=activation_dtype)
if fp32_trainable_params:
with torch.no_grad():
for parameter in module.parameters():
if parameter.requires_grad:
parameter.data = parameter.data.float()
def _optimizer_groups(
model: torch.nn.Module,
) -> tuple[list[torch.nn.Parameter], list[torch.nn.Parameter], dict[str, str]]:
fusion_parameters = []
block_parameters = []
parameter_group_names: dict[str, str] = {}
for name, parameter in model.named_parameters():
if not parameter.requires_grad:
continue
lowered = name.lower()
is_block = (
"predictor_blocks" in lowered
or "source_blocks" in lowered
or "double_blocks" in lowered
or lowered.startswith("blocks.")
)
if is_block:
block_parameters.append(parameter)
parameter_group_names[name] = "blocks"
else:
fusion_parameters.append(parameter)
parameter_group_names[name] = "fusion_residual"
if not fusion_parameters or not block_parameters:
raise RuntimeError(
"Could not form both Predictor optimizer groups. "
f"fusion={len(fusion_parameters)}, blocks={len(block_parameters)}. "
"Block modules must include 'predictor_blocks', 'source_blocks', "
"'double_blocks', or begin with 'blocks'."
)
return fusion_parameters, block_parameters, parameter_group_names
def _set_gradient_checkpointing(model: torch.nn.Module, enabled: bool) -> None:
if hasattr(model, "enable_gradient_checkpointing"):
try:
model.enable_gradient_checkpointing(enabled)
except TypeError:
if enabled:
model.enable_gradient_checkpointing()
elif hasattr(model, "disable_gradient_checkpointing"):
model.disable_gradient_checkpointing()
class Trainer:
"""Trainer dispatch target used by ``train.py``."""
def __init__(self, config) -> None:
if not torch.cuda.is_available():
raise RuntimeError("Predictor-v4 training requires CUDA")
required_env = {"RANK", "WORLD_SIZE", "LOCAL_RANK"}
if not required_env.issubset(os.environ):
raise RuntimeError(
"Launch Predictor-v4 training with torchrun; missing "
f"{sorted(required_env.difference(os.environ))}"
)
self.root_config = config
self.cfg = config.predictor_v4
self.rank = int(os.environ["RANK"])
self.world_size = int(os.environ["WORLD_SIZE"])
self.local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(self.local_rank)
self.device = torch.device("cuda", self.local_rank)
dist.init_process_group(backend="nccl")
self.is_main_process = self.rank == 0
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.set_float32_matmul_precision("high")
base_seed = int(self.root_config.seed)
seed = base_seed + self.rank
random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
self.activation_dtype = torch.bfloat16
self.source_blocks = tuple(int(value) for value in self.cfg.source_block_ids)
self.output_dir = Path(str(self.cfg.output_dir)).resolve()
if self.is_main_process:
self.output_dir.mkdir(parents=True, exist_ok=True)
dist.barrier()
self.log_path = self.output_dir / "train_log.jsonl"
self.dataset = PredictorV4PairDataset(
self.cfg.manifest,
source_block_ids=self.source_blocks,
max_records=(
None
if self.cfg.max_records is None
else int(self.cfg.max_records)
),
)
self.sampler = DistributedContextBucketBatchSampler(
self.dataset,
batch_size=int(self.cfg.batch_size),
rank=self.rank,
world_size=self.world_size,
seed=base_seed,
drop_last=True,
)
self.loader = DataLoader(
self.dataset,
batch_sampler=self.sampler,
num_workers=int(self.cfg.num_workers),
pin_memory=bool(self.cfg.pin_memory),
persistent_workers=(
int(self.cfg.num_workers) > 0
and bool(self.cfg.persistent_workers)
),
collate_fn=predictor_v4_collate,
)
if len(self.loader) == 0:
raise ValueError(
"Predictor-v4 loader has zero batches; lower batch_size or "
"provide more records for every context bucket"
)
self.model = _load_teacher_and_predictor(self.cfg)
_configure_precision(
self.model,
device=self.device,
activation_dtype=self.activation_dtype,
fp32_trainable_params=bool(self.cfg.fp32_trainable_params),
)
_set_gradient_checkpointing(
self.model, bool(self.cfg.gradient_checkpointing)
)
self.model.train()
fusion_parameters, block_parameters, group_names = _optimizer_groups(
self.model
)
self.optimizer = AdamW(
[
{
"params": fusion_parameters,
"lr": float(self.cfg.fusion_lr),
"name": "fusion_residual",
},
{
"params": block_parameters,
"lr": float(self.cfg.blocks_lr),
"name": "blocks",
},
],
betas=(float(self.cfg.beta1), float(self.cfg.beta2)),
weight_decay=float(self.cfg.weight_decay),
)
self.scheduler = _make_scheduler(
self.optimizer,
scheduler_warmup_steps=int(self.cfg.scheduler_warmup_steps),
block_warmup_steps=int(self.cfg.block_warmup_steps),
max_steps=int(self.cfg.max_steps),
)
self.global_step = 0
self.micro_step = 0
self.epoch = 0
self.batch_in_epoch = 0
self._resume(group_names)
self.ddp = DDP(
self.model,
device_ids=[self.local_rank],
output_device=self.local_rank,
broadcast_buffers=False,
gradient_as_bucket_view=True,
find_unused_parameters=bool(self.cfg.find_unused_parameters),
)
self.run_config = self._run_config(group_names)
self.wandb_run = self._initialize_wandb()
self.swanlab_run = self._initialize_swanlab()
if self.is_main_process:
_atomic_json(self.output_dir / "train_config.json", self.run_config)
print(json.dumps(self.run_config, indent=2, default=str), flush=True)
def _resume(self, group_names: dict[str, str]) -> None:
resume = self.cfg.resume
if resume is None or str(resume).lower() in {"", "null", "none"}:
return
resume_path = Path(str(resume)).resolve()
state = torch.load(resume_path, map_location="cpu", weights_only=False)
if tuple(state["source_block_ids"]) != self.source_blocks:
raise ValueError("Resume source_block_ids do not match current config")
if state.get("parameter_groups") != group_names:
raise ValueError("Resume optimizer parameter grouping has changed")
result = self.model.load_state_dict(state["model"], strict=False)
trainable = {
name
for name, parameter in self.model.named_parameters()
if parameter.requires_grad
}
missing_trainable = sorted(trainable.intersection(result.missing_keys))
if result.unexpected_keys or missing_trainable:
raise RuntimeError(
"Resume model mismatch: "
f"unexpected={result.unexpected_keys}, "
f"missing_trainable={missing_trainable}"
)
self.optimizer.load_state_dict(state["optimizer"])
self.scheduler.load_state_dict(state["scheduler"])
self.global_step = int(state["global_step"])
self.micro_step = int(state["micro_step"])
self.epoch = int(state["epoch"])
self.batch_in_epoch = int(state["batch_in_epoch"])
rng_path = resume_path.parent / f"rng_rank_{self.rank:02d}.pt"
if not rng_path.is_file():
raise FileNotFoundError(
f"Full resume requires per-rank RNG checkpoint {rng_path}"
)
restore_rng_state(
torch.load(rng_path, map_location="cpu", weights_only=False)
)
def _run_config(self, parameter_groups: dict[str, str]) -> dict[str, Any]:
trainable = sum(
parameter.numel()
for parameter in self.model.parameters()
if parameter.requires_grad
)
return {
**OmegaConf.to_container(self.cfg, resolve=True),
"manifest": str(Path(str(self.cfg.manifest)).resolve()),
"teacher_checkpoint": str(
Path(str(self.cfg.teacher_checkpoint)).resolve()
),
"output_dir": str(self.output_dir),
"source_block_ids": list(self.source_blocks),
"prompt_count": int(self.cfg.prompt_count),
"supervision_pairs": [list(pair) for pair in self.dataset.PAIRS],
"manifest_records": len(self.dataset.records),
"dataset_pairs": len(self.dataset),
"steps_per_epoch_per_rank": len(self.loader),
"world_size": self.world_size,
"global_batch_size": (
int(self.cfg.batch_size)
* self.world_size
* int(self.cfg.gradient_accumulation_steps)
),
"trainable_parameters": trainable,
"parameter_groups": parameter_groups,
"activation_dtype": "bfloat16",
"trainable_parameter_dtype": (
"float32"
if bool(self.cfg.fp32_trainable_params)
else "bfloat16"
),
}
def _initialize_wandb(self):
if (
not self.is_main_process
or bool(self.root_config.disable_wandb)
or not bool(self.cfg.use_wandb)
):
return None
import wandb
return wandb.init(
project=str(self.cfg.wandb_project),
entity=(
None
if self.cfg.wandb_entity is None
else str(self.cfg.wandb_entity)
),
name=str(self.cfg.wandb_name),
dir=str(self.output_dir),
config=self.run_config,
)
def _initialize_swanlab(self):
if (
not self.is_main_process
or not bool(getattr(self.cfg, "use_swanlab", False))
):
return None
import swanlab
mode = str(getattr(self.cfg, "swanlab_mode", "cloud"))
if mode == "cloud":
api_key = os.environ.get("SWANLAB_API_KEY")
if api_key:
swanlab.login(api_key=api_key, save=False)
else:
# Reuse the host's existing ~/.swanlab/.netrc credential.
swanlab.login()
workspace = getattr(self.cfg, "swanlab_workspace", None)
run = swanlab.init(
project=str(self.cfg.swanlab_project),
workspace=None if workspace is None else str(workspace),
experiment_name=str(self.cfg.swanlab_experiment),
description=str(self.cfg.swanlab_description),
tags=list(self.cfg.swanlab_tags),
config=self.run_config,
logdir=str(self.output_dir / "swanlab"),
mode=mode,
)
run_id = getattr(run, "id", None)
self.run_config["swanlab_run_id"] = run_id
print(json.dumps({"swanlab_run_id": run_id}), flush=True)
return run
def _save(self) -> None:
dist.barrier()
atomic_torch_save(
capture_rng_state(),
self.output_dir / f"rng_rank_{self.rank:02d}.pt",
)
if self.is_main_process:
weights_path = (
self.output_dir
/ f"predictor_step_{self.global_step:05d}.safetensors"
)
save_predictor_weights(
self.model,
weights_path,
metadata={
"source_block_ids": list(self.source_blocks),
"global_step": self.global_step,
"teacher_checkpoint": str(self.cfg.teacher_checkpoint),
"teacher_checkpoint_key": str(
self.cfg.teacher_checkpoint_key
),
"predictor_config": self.model.config_dict,
"schema_version": str(self.cfg.schema_version),
},
)
atomic_torch_save(
{
"model": trainable_state_dict(self.model),
"optimizer": self.optimizer.state_dict(),
"scheduler": self.scheduler.state_dict(),
"global_step": self.global_step,
"micro_step": self.micro_step,
"epoch": self.epoch,
"batch_in_epoch": self.batch_in_epoch,
"source_block_ids": self.source_blocks,
"parameter_groups": self.run_config["parameter_groups"],
"config": self.run_config,
"weights_path": str(weights_path),
},
self.output_dir / "training_latest.pt",
)
_prune_snapshots(self.output_dir, int(self.cfg.keep_snapshots))
dist.barrier()
def _forward_loss(
self,
batch: dict[str, Any],
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
context_frames_values = set(int(value) for value in batch["context_frames"])
if len(context_frames_values) != 1:
raise ValueError("A Predictor batch contains mixed Context lengths")
context_frames = context_frames_values.pop()
current_tokens = int(batch["anchor_hidden"].shape[1])
tokens_per_frame = int(self.model.predictor_config.tokens_per_frame)
current_start = context_frames * tokens_per_frame
with torch.no_grad():
kv_cache = self.model.build_history_kv_cache(
batch["clean_prefeature"],
current_start=current_start,
current_tokens=current_tokens,
start_frames=0,
)
crossattn_cache = build_cross_attention_cache(
batch["text_kv"], self.source_blocks
)
output = self.ddp(
target_latent=batch["target_latent"],
target_timestep=batch["target_timestep"],
anchor_hidden=batch["anchor_hidden"],
previous_chunk_hidden=batch["previous_chunk_hidden"],
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=current_start,
)
if not isinstance(output, dict) or not {
"pred_hidden",
"pred_flow",
}.issubset(output):
raise TypeError(
"SelfForcingPredictorV4.forward must return a dict containing "
"pred_hidden and pred_flow"
)
hidden_loss = F.mse_loss(
output["pred_hidden"].float(), batch["target_hidden"].float()
)
flow_loss = F.mse_loss(
output["pred_flow"].float(), batch["target_flow"].float()
)
loss = (
float(self.cfg.hidden_loss_weight) * hidden_loss
+ float(self.cfg.flow_loss_weight) * flow_loss
)
return loss, hidden_loss, flow_loss
def train(self) -> None:
accumulation = int(self.cfg.gradient_accumulation_steps)
max_steps = int(self.cfg.max_steps)
self.optimizer.zero_grad(set_to_none=True)
running = torch.zeros(4, device=self.device, dtype=torch.float64)
running_count = 0
try:
while self.global_step < max_steps:
self.sampler.set_epoch(self.epoch)
for batch_index, cpu_batch in enumerate(self.loader):
if batch_index < self.batch_in_epoch:
continue
self.batch_in_epoch = batch_index + 1
started = time.perf_counter()
batch = move_batch_to_device(
cpu_batch,
device=self.device,
dtype=self.activation_dtype,
)
sync_gradients = (self.micro_step + 1) % accumulation == 0
sync_context = (
torch.enable_grad()
if sync_gradients
else self.ddp.no_sync()
)
with sync_context:
with torch.autocast(
device_type="cuda", dtype=self.activation_dtype
):
loss, hidden_loss, flow_loss = self._forward_loss(batch)
scaled_loss = loss / accumulation
scaled_loss.backward()
self.micro_step += 1
running += torch.tensor(
[
float(loss.detach()),
float(hidden_loss.detach()),
float(flow_loss.detach()),
time.perf_counter() - started,
],
device=self.device,
dtype=torch.float64,
)
running_count += 1
del (
cpu_batch,
batch,
loss,
hidden_loss,
flow_loss,
scaled_loss,
)
if not sync_gradients:
continue
grad_norm = torch.nn.utils.clip_grad_norm_(
self.model.parameters(), float(self.cfg.grad_clip)
)
self.optimizer.step()
self.scheduler.step()
self.optimizer.zero_grad(set_to_none=True)
self.global_step += 1
if (
self.global_step == 1
or self.global_step % int(self.cfg.log_every) == 0
):
torch.cuda.synchronize(self.device)
averaged = _distributed_mean(running.clone())
averaged /= running_count
max_grad = torch.tensor(
float(grad_norm),
device=self.device,
dtype=torch.float64,
)
dist.all_reduce(max_grad, op=dist.ReduceOp.MAX)
if self.is_main_process:
record = {
"global_step": self.global_step,
"epoch": self.epoch,
"loss": float(averaged[0]),
"hidden_mse": float(averaged[1]),
"flow_mse": float(averaged[2]),
"avg_micro_time_s": float(averaged[3]),
"grad_norm_max": float(max_grad),
"lr_fusion": self.optimizer.param_groups[0]["lr"],
"lr_blocks": self.optimizer.param_groups[1]["lr"],
"peak_memory_gib": (
torch.cuda.max_memory_allocated(self.device)
/ 2**30
),
}
with self.log_path.open("a", encoding="utf-8") as handle:
handle.write(
json.dumps(record, sort_keys=True) + "\n"
)
print(json.dumps(record, sort_keys=True), flush=True)
if self.wandb_run is not None:
self.wandb_run.log(record, step=self.global_step)
if self.swanlab_run is not None:
import swanlab
swanlab.log(record, step=self.global_step)
running.zero_()
running_count = 0
should_save = (
not bool(self.root_config.no_save)
and (
self.global_step % int(self.cfg.save_every) == 0
or self.global_step == max_steps
)
)
if should_save:
self._save()
if self.global_step >= max_steps:
break
if self.global_step < max_steps:
self.epoch += 1
self.batch_in_epoch = 0
except BaseException as error:
if self.swanlab_run is not None:
try:
import swanlab
swanlab.finish(error=str(error))
except Exception:
pass
raise
else:
if self.swanlab_run is not None:
import swanlab
swanlab.finish()
finally:
if self.wandb_run is not None:
self.wandb_run.finish()
dist.destroy_process_group()