CoolFace
Modelpublic

Cccccz/HY

sourceHugging Faceupdated 11d agoView on Hugging Face
0likes
train_predictor_v2.py375 linesDownload Raw Back to root
1#!/usr/bin/env python32"""Single-GPU training for AR-context HY-WorldPlay Predictor-v2."""3 4from __future__ import annotations5 6import argparse7import json8import math9import os10import random11import time12from pathlib import Path13 14import swanlab15import torch16import torch.nn.functional as F17from torch.optim import AdamW18from torch.optim.lr_scheduler import LambdaLR19from torch.utils.data import DataLoader20 21from hyvideo.commons.infer_state import initialize_infer_state22from models import HYWorldPlayPredictorV2, HYWorldPlayPredictorV3, PredictorV2Config23from predictor_training.checkpoint import (24    save_predictor_weights,25    save_training_state,26    trainable_state_dict,27)28from predictor_training.dataset_v2 import (29    PredictorV2PairDataset,30    move_v2_pair_to_device,31    predictor_v2_collate,32)33from predictor_training.dataset_v3 import (34    PredictorV3PairDataset,35    move_v3_pair_to_device,36    predictor_v3_collate,37)38 39 40def parse_source_blocks(value: str) -> tuple[int, int]:41    result = tuple(int(part.strip()) for part in value.split(",") if part.strip())42    if result not in {(0, 53), (1, 52)}:43        raise argparse.ArgumentTypeError("source_blocks must be 0,53 or 1,52")44    return result45 46 47def parse_args() -> argparse.Namespace:48    parser = argparse.ArgumentParser(description=__doc__)49    parser.add_argument("--manifest", default="datasets/predictor_v2/train_eval_manifest.jsonl")50    parser.add_argument("--predictor_version", choices=("v2", "v3"), default="v2")51    parser.add_argument("--source_blocks", type=parse_source_blocks, required=True)52    parser.add_argument(53        "--teacher_checkpoint",54        default="models-ms/HY-WorldPlay/ar_distilled_action_model/diffusion_pytorch_model.safetensors",55    )56    parser.add_argument("--output_dir", required=True)57    parser.add_argument("--max_steps", type=int, default=2000)58    parser.add_argument("--batch_size", type=int, default=1)59    parser.add_argument("--gradient_accumulation_steps", type=int, default=4)60    parser.add_argument("--warmup_blocks_steps", type=int, default=100)61    parser.add_argument("--scheduler_warmup_steps", type=int, default=100)62    parser.add_argument("--fusion_lr", type=float, default=1e-4)63    parser.add_argument("--blocks_lr", type=float, default=1e-5)64    parser.add_argument("--weight_decay", type=float, default=0.01)65    parser.add_argument("--hidden_weight", type=float, default=0.1)66    parser.add_argument("--velocity_weight", type=float, default=1.0)67    parser.add_argument("--grad_clip", type=float, default=1.0)68    parser.add_argument("--log_every", type=int, default=10)69    parser.add_argument("--save_every", type=int, default=500)70    parser.add_argument("--keep_snapshots", type=int, default=2)71    parser.add_argument("--seed", type=int, default=0)72    parser.add_argument("--num_workers", type=int, default=0)73    parser.add_argument("--max_records", type=int, default=None)74    parser.add_argument("--gradient_checkpointing", action=argparse.BooleanOptionalAction, default=True)75    parser.add_argument("--save_optimizer_state", action=argparse.BooleanOptionalAction, default=True)76    parser.add_argument("--resume", default=None)77    parser.add_argument("--swanlab_project", default="HY-WorldPlay-Predictor-v2")78    parser.add_argument("--swanlab_workspace", default=None)79    parser.add_argument("--swanlab_experiment", default=None)80    parser.add_argument("--swanlab_mode", choices=("cloud", "offline", "local", "disabled"), default="cloud")81    return parser.parse_args()82 83 84def initialize_runtime_state() -> None:85    class Args:86        sage_blocks_range = "0-53"87        use_sageattn = False88        enable_torch_compile = False89        use_fp8_gemm = False90        quant_type = "fp8-per-block"91        include_patterns = "double_blocks"92        use_vae_parallel = False93 94    initialize_infer_state(Args())95 96 97def make_scheduler(optimizer: torch.optim.Optimizer, warmup: int, total: int) -> LambdaLR:98    def factor(step: int) -> float:99        if step < warmup:100            return max(1e-8, float(step + 1) / max(1, warmup))101        progress = min(1.0, float(step - warmup) / max(1, total - warmup))102        return 0.5 * (1.0 + math.cos(math.pi * progress))103 104    return LambdaLR(optimizer, factor)105 106 107def set_ar_blocks_trainable(model: HYWorldPlayPredictorV2, enabled: bool) -> None:108    for block in model.double_blocks:109        for name, parameter in block.named_parameters():110            parameter.requires_grad_(enabled and not name.startswith("txt_"))111 112 113def prune_snapshots(output_dir: Path, keep: int) -> None:114    snapshots = sorted(output_dir.glob("predictor_step_*.safetensors"))115    for path in snapshots[:-keep] if keep > 0 else snapshots:116        path.unlink()117 118 119def init_swanlab(args: argparse.Namespace, config: dict):120    if args.swanlab_mode == "cloud":121        api_key = os.environ.get("SWANLAB_API_KEY")122        swanlab.login(api_key=api_key, save=False) if api_key else swanlab.login()123    fusion_lr = f"{args.fusion_lr:.0e}".replace("e-0", "e-").replace("e+0", "e+")124    blocks_lr = f"{args.blocks_lr:.0e}".replace("e-0", "e-").replace("e+0", "e+")125    experiment = args.swanlab_experiment or (126        f"predictor-{args.predictor_version}-blocks"127        + "-".join(str(value) for value in args.source_blocks)128        + f"-flr{fusion_lr}-blr{blocks_lr}-schedsteps{args.max_steps}"129        + f"-swarmup{args.scheduler_warmup_steps}"130        + f"-bwarmup{args.warmup_blocks_steps}"131    )132    return swanlab.init(133        project=args.swanlab_project,134        workspace=args.swanlab_workspace,135        experiment_name=experiment,136        description=(137            f"HY-WorldPlay {args.predictor_version} AR-context Predictor, "138            "fresh exact 10-case targets"139        ),140        tags=[141            f"predictor-{args.predictor_version}",142            "ar-context-kv",143            f"blocks-{args.source_blocks[0]}-{args.source_blocks[1]}",144        ],145        config=config,146        logdir=str(Path(args.output_dir).resolve() / "swanlab"),147        mode=args.swanlab_mode,148    )149 150 151def main() -> None:152    args = parse_args()153    if not torch.cuda.is_available():154        raise RuntimeError("CUDA is required")155    random.seed(args.seed)156    torch.manual_seed(args.seed)157    torch.cuda.manual_seed_all(args.seed)158    initialize_runtime_state()159    device = torch.device("cuda:0")160    dtype = torch.bfloat16161    output_dir = Path(args.output_dir).resolve()162    output_dir.mkdir(parents=True, exist_ok=True)163    log_path = output_dir / "train_log.jsonl"164 165    dataset_class = (166        PredictorV3PairDataset if args.predictor_version == "v3" else PredictorV2PairDataset167    )168    collate_fn = (169        predictor_v3_collate if args.predictor_version == "v3" else predictor_v2_collate170    )171    move_pair = (172        move_v3_pair_to_device173        if args.predictor_version == "v3"174        else move_v2_pair_to_device175    )176    dataset = dataset_class(177        args.manifest,178        source_block_ids=args.source_blocks,179        max_records=args.max_records,180    )181    generator = torch.Generator().manual_seed(args.seed)182    loader = DataLoader(183        dataset,184        batch_size=args.batch_size,185        shuffle=True,186        num_workers=args.num_workers,187        pin_memory=False,188        collate_fn=collate_fn,189        generator=generator,190    )191 192    model_class = (193        HYWorldPlayPredictorV3 if args.predictor_version == "v3" else HYWorldPlayPredictorV2194    )195    model = model_class(PredictorV2Config(source_block_ids=args.source_blocks))196    model.load_teacher_initialization(args.teacher_checkpoint)197    model.to(device=device, dtype=dtype)198    model.enable_gradient_checkpointing(args.gradient_checkpointing)199    model.train()200 201    fusion_parameters = list(model.img_fusion.parameters()) + list(model.residual_out.parameters())202    block_parameters = [203        parameter for parameter in model.double_blocks.parameters() if parameter.requires_grad204    ]205    intended_trainable_parameters = sum(206        parameter.numel() for parameter in fusion_parameters + block_parameters207    )208    intended_parameter_breakdown = model.trainable_parameter_breakdown()209    optimizer = AdamW(210        [211            {"params": fusion_parameters, "lr": args.fusion_lr},212            {"params": block_parameters, "lr": args.blocks_lr},213        ],214        betas=(0.9, 0.95),215        weight_decay=args.weight_decay,216    )217    scheduler = make_scheduler(optimizer, args.scheduler_warmup_steps, args.max_steps)218    global_step = 0219    micro_step = 0220    set_ar_blocks_trainable(model, args.warmup_blocks_steps <= 0)221 222    if args.resume:223        state = torch.load(args.resume, map_location="cpu", weights_only=False)224        if tuple(state["predictor_config"]["source_block_ids"]) != args.source_blocks:225            raise ValueError("Resume checkpoint source blocks do not match")226        model.load_state_dict(state["model"], strict=False)227        optimizer.load_state_dict(state["optimizer"])228        scheduler.load_state_dict(state["scheduler"])229        global_step = int(state["global_step"])230        micro_step = int(state.get("micro_step", global_step * args.gradient_accumulation_steps))231        if "torch_rng_state" in state:232            torch.set_rng_state(state["torch_rng_state"])233        if "cuda_rng_state" in state:234            torch.cuda.set_rng_state_all(state["cuda_rng_state"])235        set_ar_blocks_trainable(model, global_step >= args.warmup_blocks_steps)236 237    run_config = {238        **vars(args),239        "source_blocks": list(args.source_blocks),240        "manifest": str(Path(args.manifest).resolve()),241        "teacher_checkpoint": str(Path(args.teacher_checkpoint).resolve()),242        "dataset_pairs": len(dataset),243        "effective_batch_size": args.batch_size * args.gradient_accumulation_steps,244        "trainable_parameters": intended_trainable_parameters,245        "parameter_breakdown": intended_parameter_breakdown,246        "dtype": "bfloat16",247        "predictor_version": args.predictor_version,248        "gpu": torch.cuda.get_device_name(device),249    }250    run = init_swanlab(args, run_config)251    print(252        json.dumps(253            {254                "dataset_pairs": len(dataset),255                "source_blocks": args.source_blocks,256                "trainable_parameters": intended_trainable_parameters,257                "breakdown": intended_parameter_breakdown,258                "start_step": global_step,259                "swanlab_run_id": getattr(run, "id", None),260            },261            indent=2,262        ),263        flush=True,264    )265 266    optimizer.zero_grad(set_to_none=True)267    running: list[dict[str, float | str]] = []268    epoch = 0269    try:270        while global_step < args.max_steps:271            epoch += 1272            for item in loader:273                if global_step >= args.max_steps:274                    break275                if global_step == args.warmup_blocks_steps:276                    set_ar_blocks_trainable(model, True)277 278                inputs, target_hidden, target_velocity = move_pair(279                    item, device=device, dtype=dtype280                )281                started = time.perf_counter()282                with torch.autocast(device_type="cuda", dtype=dtype):283                    output = model(**inputs)284                    hidden_loss = F.mse_loss(285                        output["pred_hidden"].float(), target_hidden.float()286                    )287                    velocity_loss = F.mse_loss(288                        output["pred_velocity"].float(), target_velocity.float()289                    )290                    loss = args.hidden_weight * hidden_loss + args.velocity_weight * velocity_loss291                    scaled_loss = loss / args.gradient_accumulation_steps292                scaled_loss.backward()293                micro_step += 1294                running.append(295                    {296                        "pair_id": ",".join(item["pair_id"]),297                        "loss": float(loss.detach()),298                        "hidden_mse": float(hidden_loss.detach()),299                        "velocity_mse": float(velocity_loss.detach()),300                        "micro_time_s": time.perf_counter() - started,301                        "context_frames": sum(item["context_frames"])302                        / len(item["context_frames"]),303                    }304                )305                del (306                    output,307                    loss,308                    scaled_loss,309                    hidden_loss,310                    velocity_loss,311                    inputs,312                    target_hidden,313                    target_velocity,314                    item,315                )316 317                if micro_step % args.gradient_accumulation_steps:318                    continue319                grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)320                optimizer.step()321                scheduler.step()322                optimizer.zero_grad(set_to_none=True)323                global_step += 1324 325                if global_step % args.log_every == 0 or global_step == 1:326                    torch.cuda.synchronize()327                    count = len(running)328                    record = {329                        "global_step": global_step,330                        "epoch": epoch,331                        "loss": sum(float(x["loss"]) for x in running) / count,332                        "hidden_mse": sum(float(x["hidden_mse"]) for x in running) / count,333                        "velocity_mse": sum(float(x["velocity_mse"]) for x in running) / count,334                        "grad_norm": float(grad_norm),335                        "lr_fusion": optimizer.param_groups[0]["lr"],336                        "lr_blocks": optimizer.param_groups[1]["lr"],337                        "peak_memory_gib": torch.cuda.max_memory_allocated() / 2**30,338                        "avg_micro_time_s": sum(float(x["micro_time_s"]) for x in running) / count,339                        "avg_context_frames": sum(float(x["context_frames"]) for x in running) / count,340                    }341                    with log_path.open("a", encoding="utf-8") as handle:342                        handle.write(json.dumps(record, sort_keys=True) + "\n")343                    swanlab.log(record, step=global_step)344                    print(json.dumps(record, sort_keys=True), flush=True)345                    running.clear()346 347                should_save = global_step % args.save_every == 0 or global_step == args.max_steps348                if should_save:349                    weights_path = output_dir / f"predictor_step_{global_step:05d}.safetensors"350                    save_predictor_weights(model, weights_path)351                    if args.save_optimizer_state:352                        latest = {353                            "model": trainable_state_dict(model),354                            "optimizer": optimizer.state_dict(),355                            "scheduler": scheduler.state_dict(),356                            "global_step": global_step,357                            "micro_step": micro_step,358                            "torch_rng_state": torch.get_rng_state(),359                            "cuda_rng_state": torch.cuda.get_rng_state_all(),360                            "args": vars(args),361                            "predictor_config": model.config_dict,362                            "weights_path": str(weights_path),363                        }364                        save_training_state(output_dir / "training_latest.pt", latest)365                    prune_snapshots(output_dir, args.keep_snapshots)366    except BaseException as error:367        swanlab.finish(error=error)368        raise369    else:370        swanlab.finish()371 372 373if __name__ == "__main__":374    main()375