Compare commits

...
41 changed files with 5253 additions and 51 deletions
+1
View File
@@ -31,6 +31,7 @@ env
**/build/
**.pyc
**.txt
!examples/train/requirements-diffusion-nft.txt
*.log
weights/
logs/
+4
View File
@@ -80,6 +80,10 @@ training:
run_name: my_run
```
## RL Examples
- DiffusionNFT Wan video RL: see `examples/train/diffusion_nft_wan_video.md`.
## Directory Layout
```
@@ -0,0 +1,113 @@
# DiffusionNFT video RL: Wan 2.1 T2V 1.3B on text-only prompts with VideoAlign rewards.
#
# This follows the modular RL layout from diffusion_nft_pick_clip.yaml and only
# swaps the reward suite plus video-sized latent/media settings.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
old:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
reference:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod
reward_backend: genrl
reward_fn:
rewards:
videoalign_vq: 1.0
videoalign_mq: 1.0
videoalign_ta: 1.0
sampling:
num_steps: 50
scheduler: flow_unipc
trajectory: ode
flow_shift: 8.0
guidance_scale: 6.0
validation:
every_steps: 10
num_steps: 50
num_prompts: 16
batch_size: 4
log_samples: true
max_samples: 4
fps: 16
seed: 42
data_path:
sample_train_batch_size: 1
train_batch_size: 1
num_batches_per_epoch: 24
num_video_per_prompt: 8
num_inner_epochs: 1
timestep_fraction: 0.99
beta: 0.1
kl_beta: 0.0001
decay_type: 1
adv_mode: all
adv_clip_max: 5
max_grad_norm: 1.0
ema:
enabled: true
decay: 0.9
update_after_step: 0
validation: true
terminal_progress: true
training:
distributed:
num_gpus: 1
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 1
data:
data_path: data/train_text_only_diffusion_nft_preprocessed
preprocessed_data_type: text_only
dataloader_num_workers: 0
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 20
num_height: 448
num_width: 832
num_frames: 77
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.0001
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 100000
gradient_accumulation_steps: 24
checkpoint:
output_dir: outputs/wan2.1_diffusion_nft_videoalign
training_state_checkpointing_steps: 30
checkpoints_total_limit: 3
tracker:
project_name: diffusion_nft_wan
run_name: wan2.1_diffusion_nft_videoalign
model:
enable_gradient_checkpointing_type: full
pipeline:
flow_shift: 8
+108
View File
@@ -0,0 +1,108 @@
# DiffusionNFT Wan Video RL
This guide reproduces the DiffusionNFT Wan video RL run without any private
launcher. Modal scripts used by individual developers should only call these
tracked commands.
## 1. Prepare Prompts, Rewards, Parquet, and Run Config
Install the reward-stack pins after installing FastVideo:
```bash
uv pip install -e .
uv pip install -r examples/train/requirements-diffusion-nft.txt
```
```bash
python examples/train/prepare_diffusion_nft_assets.py \
--config examples/train/configs/rl/wan/diffusion_nft_videoalign.yaml \
--data-root data/diffusion_nft \
--cache-root .cache/diffusion_nft \
--output-dir outputs/wan2.1_diffusion_nft_videoalign \
--dataset world-r1-enhanced \
--reward videoalign \
--max-prompts quarter \
--num-frames 77 \
--num-gpus 4 \
--hsdp-replicate-dim 1 \
--hsdp-shard-dim 4 \
--max-train-steps 100 \
--gradient-accumulation-steps 24 \
--sample-num-steps 50 \
--sample-flow-shift 8 \
--sample-guidance-scale 6 \
--preprocess-batch-size 128 \
--check-rewards \
--json
```
The script:
- resolves `num_latent_t` from Wan's frame rule,
- downloads World-R1 prompts or reads a DiffusionNFT `dataset/<name>/train.txt`,
- clones `DiffusionNFT` only when a selected dataset or fallback reward needs it,
- downloads the `KwaiVGI/VideoReward` snapshot when VideoAlign rewards are selected,
- preprocesses prompts into FastVideo text-only parquet,
- writes `outputs/diffusion_nft_run_configs/diffusion_nft_wan_run.yaml`.
- optionally loads the selected reward suite on a dummy video with
`--check-rewards`.
Set `VIDEOALIGN_CHECKPOINT_PATH` or `DIFFUSION_NFT_ROOT` through the matching
CLI flags if those assets already live somewhere else.
## 2. Launch Training
```bash
NUM_GPUS=4 bash examples/train/run.sh \
outputs/diffusion_nft_run_configs/diffusion_nft_wan_run.yaml
```
W&B logging records the resolved training config and scalar/video validation
artifacts from the tracker.
The video config intentionally uses the Wan/MAY-30 rollout sampler profile:
`flow_unipc`, 50 denoising steps, flow shift 8, and CFG 6. Validation uses the
same 50-step sampler so W&B qualitative videos are comparable to rollout
quality.
Validation reward metrics are logged under `validation/reward/*` every
`method.validation.every_steps`. Qualitative W&B videos are only a sanity check:
`--log-sample-max-videos 0` disables them, while a positive value caps how many
validation samples are logged per eval.
For short 17-frame H100 runs, keep gradient accumulation matched to the number
of rollout train batches. For example, if you use `--num-batches-per-epoch 4`,
`--collection-batch-size 4`, `--train-batch-size 4`, and 50 sample steps, use
`--gradient-accumulation-steps 4` for one full optimizer update per outer
DiffusionNFT step. A much larger accumulation value, such as 30, still runs, but
only performs a partial final update for that outer step and W&B will report it
under `nft/partial_optimizer_step_ratio`.
DiffusionNFT relies on repeated samples for the same prompt to estimate
per-prompt advantages. `--num-samples-per-prompt 4` is a cheap smoke setting,
but it is high variance. If the reward curves look noisy and memory allows it,
try `--num-samples-per-prompt 8` or increase `--num-batches-per-epoch` before
judging the training run.
For Wan video reward rollouts, keep `method.sampling.guidance_scale` enabled
unless you are intentionally testing unguided generation. The May 30 video run
used CFG sampling at `6.0`; unguided rollouts can look foggy and static before
the reward model ever sees them.
For offline logging:
```bash
WANDB_MODE=offline NUM_GPUS=4 bash examples/train/run.sh \
outputs/diffusion_nft_run_configs/diffusion_nft_wan_run.yaml
```
## Reward Presets
- `--reward videoalign`: `videoalign_vq`, `videoalign_mq`, `videoalign_ta`
- `--reward videoalign_hpsv3`: VideoAlign plus reward-only `hpsv3_general`
- `--reward multi_reward`: DiffusionNFT image reward preset
- `--reward <name>`: one explicit reward name
Use `examples/train/configs/rl/wan/diffusion_nft_videoalign.yaml` as the
checked-in base config and treat the generated run config as the exact
experiment instance.
@@ -0,0 +1,517 @@
# SPDX-License-Identifier: Apache-2.0
"""Prepare reproducible assets for DiffusionNFT Wan RL training.
This script is intentionally environment-agnostic. It contains the prompt,
preprocessing, reward-checkpoint, and run-config preparation needed to reproduce
the DiffusionNFT video run without relying on private Modal launchers.
"""
from __future__ import annotations
import argparse
import gc
import json
import os
from pathlib import Path
import subprocess
import sys
from typing import Any
import yaml
MODEL_ID = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DIFFUSION_NFT_REPO = "https://github.com/NVlabs/DiffusionNFT.git"
VIDEO_REWARD_REPO = "KwaiVGI/VideoReward"
DEFAULT_CONFIG = "examples/train/configs/rl/wan/diffusion_nft_videoalign.yaml"
DEFAULT_OUTPUT_DIR = "outputs/wan2.1_diffusion_nft_videoalign"
IMAGE_MULTI_REWARD_NAMES = ("pickscore", "hpsv2", "clipscore")
VIDEO_MULTI_REWARD_NAMES = (
"videoalign_vq",
"videoalign_mq",
"videoalign_ta",
)
VIDEO_BALANCED_REWARD_WEIGHTS = {
"videoalign_ta": 1.0,
"videoalign_mq": 1.0,
"videoalign_vq": 0.75,
"hpsv3_general": 0.25,
}
GENRL_REWARD_NAMES = frozenset({
"hpsv3_general",
"hpsv3_percentile",
"videoalign_vq",
"videoalign_mq",
"videoalign_ta",
})
def derive_wan_num_latent_t(frame_count: int) -> int:
if frame_count <= 0:
raise ValueError("--num-frames must be positive")
if (frame_count - 1) % 4 != 0:
raise ValueError("Wan frame counts must satisfy num_frames = (num_latent_t - 1) * 4 + 1. "
f"Got num_frames={frame_count}; try 1, 5, 9, 13, 17, ...")
return ((frame_count - 1) // 4) + 1
def resolve_reward_map(reward: str) -> tuple[dict[str, float], str]:
reward = reward.strip().lower()
if reward in {"videoalign", "video_reward", "video_multi_reward"}:
reward_map = {name: 1.0 for name in VIDEO_MULTI_REWARD_NAMES}
elif reward in {"videoalign_hpsv3", "video_reward_hpsv3", "balanced_video", "quality_video"}:
reward_map = dict(VIDEO_BALANCED_REWARD_WEIGHTS)
elif reward in {"multi_reward", "image_multi_reward"}:
reward_map = {name: 1.0 for name in IMAGE_MULTI_REWARD_NAMES}
else:
reward_map = {reward: 1.0}
backend = "genrl" if any(name in GENRL_REWARD_NAMES for name in reward_map) else "diffusion_nft"
return reward_map, backend
def resolve_max_prompts(
max_prompts: str,
*,
total_prompts: int,
max_train_steps: int,
gradient_accumulation_steps: int,
) -> int:
mode = str(max_prompts).strip().lower()
if mode in {"", "0", "all", "full", "none"}:
return 0
estimated_prompt_batches = int(max_train_steps) * int(gradient_accumulation_steps)
if mode in {"tenth", "1/10"}:
return max(1, min(total_prompts, estimated_prompt_batches // 10))
if mode in {"quarter", "1/4"}:
return max(1, min(total_prompts, estimated_prompt_batches // 4))
if mode in {"half", "1/2"}:
return max(1, min(total_prompts, estimated_prompt_batches // 2))
if mode in {"steps", "used"}:
return max(1, min(total_prompts, estimated_prompt_batches))
value = int(mode)
if value < 0:
raise ValueError("--max-prompts must be >= 0, tenth, quarter, half, steps, or full")
return min(total_prompts, value)
def has_parquet(path: Path) -> bool:
return path.exists() and any(path.rglob("*.parquet"))
def verify_text_only_dataset(path: Path, expected_rows: int) -> int:
import pyarrow.parquet as pq
parquet_files = sorted(path.rglob("*.parquet"))
if not parquet_files:
raise RuntimeError(f"Text-only preprocessing produced no parquet under {path}")
row_count = sum(pq.ParquetFile(file_path).metadata.num_rows for file_path in parquet_files)
if row_count < expected_rows:
raise RuntimeError(f"Expected at least {expected_rows} prompt rows in {path}, found {row_count}.")
return int(row_count)
def ensure_diffusion_nft_repo(root: Path) -> None:
if (root / "flow_grpo").is_dir():
return
root.parent.mkdir(parents=True, exist_ok=True)
subprocess.run(
["git", "clone", "--depth", "1", DIFFUSION_NFT_REPO, str(root)],
stdout=sys.stdout,
stderr=sys.stderr,
check=True,
)
def ensure_videoalign_checkpoint(path: Path) -> None:
if has_video_reward_checkpoint(path):
return
try:
from huggingface_hub import snapshot_download
except ImportError as exc:
raise ImportError(f"huggingface_hub is required to download {VIDEO_REWARD_REPO}. "
"Install examples/train/requirements-diffusion-nft.txt and rerun.") from exc
snapshot_download(
repo_id=VIDEO_REWARD_REPO,
repo_type="model",
local_dir=str(path),
local_dir_use_symlinks=False,
)
if not has_video_reward_checkpoint(path):
raise RuntimeError(f"Downloaded {VIDEO_REWARD_REPO}, but no VideoReward checkpoint was found under {path}.")
def has_video_reward_checkpoint(root: Path) -> bool:
model_config = root / "model_config.json"
if not model_config.exists():
return False
if (root / "model.pth").exists():
return True
if ((root / "adapter_model.safetensors").exists() and (root / "non_lora_state_dict.pth").exists()):
return True
for checkpoint in root.glob("checkpoint-*"):
if (checkpoint / "model.pth").exists():
return True
if ((checkpoint / "adapter_model.safetensors").exists()
and (checkpoint / "non_lora_state_dict.pth").exists()):
return True
return False
def load_prompts(
dataset: str,
*,
diffusion_nft_root: Path,
) -> tuple[list[str], str]:
dataset = dataset.strip().lower()
if dataset in {"world-r1", "world-r1-final", "world_r1", "world_r1_final"}:
from datasets import load_dataset
rows = load_dataset("microsoft/World-R1", "final", split="train")
source_label = "microsoft/World-R1 final/train"
elif dataset in {"world-r1-dynamic", "world_r1_dynamic", "world-r1-final-dynamic", "world_r1_final_dynamic"}:
from datasets import load_dataset
rows = load_dataset("microsoft/World-R1", "final", split="dynamic")
source_label = "microsoft/World-R1 final/dynamic"
elif dataset in {"world-r1-enhanced", "world_r1_enhanced"}:
from datasets import load_dataset
rows = load_dataset("microsoft/World-R1", "enhanced", split="train")
source_label = "microsoft/World-R1 enhanced/train"
elif dataset in {"world-r1-enhanced-dynamic", "world_r1_enhanced_dynamic"}:
from datasets import load_dataset
rows = load_dataset("microsoft/World-R1", "enhanced", split="dynamic")
source_label = "microsoft/World-R1 enhanced/dynamic"
else:
source = diffusion_nft_root / "dataset" / dataset / "train.txt"
if not source.is_file():
raise RuntimeError(f"Dataset {dataset!r} is not a built-in World-R1 dataset and does not provide "
f"{source}. Use a World-R1 dataset alias or prepare a DiffusionNFT checkout.")
prompts = [line.strip() for line in source.read_text(encoding="utf-8").splitlines() if line.strip()]
return prompts, str(source)
prompts = [str(row["prompt"]).strip() for row in rows if str(row.get("prompt", "")).strip()]
return prompts, source_label
def prepare_prompt_file(args: argparse.Namespace, *, num_latent_t: int) -> tuple[Path, Path, int, str]:
del num_latent_t
prompts, source_label = load_prompts(args.dataset, diffusion_nft_root=args.diffusion_nft_root)
prompt_limit = resolve_max_prompts(
args.max_prompts,
total_prompts=len(prompts),
max_train_steps=args.max_train_steps,
gradient_accumulation_steps=args.gradient_accumulation_steps,
)
if prompt_limit > 0:
prompts = prompts[:prompt_limit]
if not prompts:
raise RuntimeError(f"No prompts found in {source_label}")
if len(prompts) < args.num_gpus:
raise RuntimeError(f"Need at least {args.num_gpus} prompts for {args.num_gpus} ranks with drop_last=True; "
f"resolved only {len(prompts)} prompt(s). Increase --max-prompts.")
prompt_suffix = "full" if prompt_limit <= 0 else f"first{prompt_limit}"
dataset_root = args.data_root / f"diffusion_nft_{args.dataset}_text_only_f{args.num_frames}_{prompt_suffix}"
parquet_dir = dataset_root / "combined_parquet_dataset"
prompt_file = args.data_root / "prompts" / f"diffusion_nft_{args.dataset}_{prompt_suffix}_train.txt"
prompt_file.parent.mkdir(parents=True, exist_ok=True)
prompt_file.write_text("\n".join(prompts) + "\n", encoding="utf-8")
return prompt_file, parquet_dir, len(prompts), source_label
def run_text_only_preprocess(
args: argparse.Namespace,
*,
prompt_file: Path,
dataset_root: Path,
num_latent_t: int,
) -> None:
dataset_root.mkdir(parents=True, exist_ok=True)
cmd = [
"torchrun",
"--nnodes=1",
"--nproc_per_node",
str(args.preprocess_num_gpus),
"--master_port",
str(args.preprocess_master_port),
"fastvideo/pipelines/preprocess/v1_preprocess.py",
"--model_path",
args.model_id,
"--data_merge_path",
str(prompt_file),
"--preprocess_video_batch_size",
str(args.preprocess_batch_size),
"--seed",
str(args.seed),
"--max_height",
str(args.num_height),
"--max_width",
str(args.num_width),
"--num_frames",
str(args.num_frames),
"--num_latent_t",
str(num_latent_t),
"--dataloader_num_workers",
str(args.dataloader_num_workers),
"--output_dir",
str(dataset_root),
"--samples_per_file",
str(args.samples_per_file),
"--flush_frequency",
str(args.flush_frequency),
"--preprocess_task",
"text_only",
]
subprocess.run(cmd, cwd=args.repo_root, stdout=sys.stdout, stderr=sys.stderr, check=True)
def write_run_config(
args: argparse.Namespace,
*,
parquet_dir: Path,
num_latent_t: int,
reward_map: dict[str, float],
reward_backend: str,
) -> Path:
raw_config = yaml.safe_load(args.config.read_text(encoding="utf-8"))
method_config = raw_config.setdefault("method", {})
training_config = raw_config.setdefault("training", {})
distributed_config = training_config.setdefault("distributed", {})
data_config = training_config.setdefault("data", {})
loop_config = training_config.setdefault("loop", {})
optimizer_config = training_config.setdefault("optimizer", {})
checkpoint_config = training_config.setdefault("checkpoint", {})
tracker_config = training_config.setdefault("tracker", {})
sampling_config = method_config.setdefault("sampling", {})
method_config["reward_backend"] = reward_backend
method_config["reward_fn"] = {"rewards": reward_map}
method_config["num_video_per_prompt"] = int(args.num_samples_per_prompt)
method_config["sample_train_batch_size"] = int(args.collection_batch_size)
method_config["num_inner_epochs"] = int(args.inner_epochs)
method_config["train_batch_size"] = int(args.train_batch_size)
validation_config = method_config.setdefault("validation", {})
validation_config["log_samples"] = bool(args.log_sample_max_videos > 0)
if args.log_sample_max_videos > 0:
validation_config["max_samples"] = int(args.log_sample_max_videos)
else:
validation_config.pop("max_samples", None)
if args.sample_num_steps is not None:
sampling_config["num_steps"] = int(args.sample_num_steps)
if args.sample_flow_shift is not None:
sampling_config["flow_shift"] = float(args.sample_flow_shift)
if args.sample_guidance_scale is not None:
sampling_config["guidance_scale"] = float(args.sample_guidance_scale)
if args.reward in {"multi_reward", "image_multi_reward"}:
method_config["beta"] = 0.1
distributed_config["num_gpus"] = int(args.num_gpus)
distributed_config["tp_size"] = int(args.tp_size)
distributed_config["sp_size"] = int(args.sp_size)
distributed_config["hsdp_replicate_dim"] = int(args.hsdp_replicate_dim)
distributed_config["hsdp_shard_dim"] = int(args.hsdp_shard_dim)
data_config["data_path"] = str(parquet_dir)
data_config["preprocessed_data_type"] = "text_only"
data_config["num_frames"] = int(args.num_frames)
data_config["num_latent_t"] = int(num_latent_t)
data_config["num_height"] = int(args.num_height)
data_config["num_width"] = int(args.num_width)
data_config["dataloader_num_workers"] = int(args.dataloader_num_workers)
loop_config["max_train_steps"] = int(args.max_train_steps)
loop_config["gradient_accumulation_steps"] = int(args.gradient_accumulation_steps)
if args.learning_rate is not None:
optimizer_config["learning_rate"] = float(args.learning_rate)
checkpoint_config["output_dir"] = str(args.output_dir)
tracker_config["project_name"] = args.project_name
tracker_config["run_name"] = args.run_name or args.output_dir.name
args.run_config_dir.mkdir(parents=True, exist_ok=True)
run_config_path = args.run_config_dir / "diffusion_nft_wan_run.yaml"
run_config_path.write_text(yaml.safe_dump(raw_config, sort_keys=False), encoding="utf-8")
return run_config_path
def check_reward_runtime(
reward_map: dict[str, float],
*,
reward_backend: str,
device: str = "auto",
) -> None:
import torch
selected_device = device
if selected_device == "auto":
selected_device = "cuda:0" if torch.cuda.is_available() else "cpu"
torch_device = torch.device(selected_device)
print("=== DiffusionNFT reward preflight ===", flush=True)
print(f"reward_backend={reward_backend} reward_map={reward_map}", flush=True)
from fastvideo.train.methods.rl.rewards import build_multi_reward_scorer
media = torch.zeros(1, 3, 8, 224, 224, device=torch_device)
prompts = ["A small red block moves steadily from left to right."]
scorer = build_multi_reward_scorer(
reward_map,
backend=reward_backend,
device=torch_device,
)
scores = scorer(media, prompts)
score_summary = {name: float(value.detach().float().cpu()[0]) for name, value in scores.items()}
print(f"reward preflight scores: {score_summary}", flush=True)
del scorer, media, scores
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
print("=== DiffusionNFT reward preflight OK ===", flush=True)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--repo-root", type=Path, default=Path.cwd())
parser.add_argument("--config", type=Path, default=Path(DEFAULT_CONFIG))
parser.add_argument("--data-root", type=Path, default=Path("data/diffusion_nft"))
parser.add_argument("--cache-root", type=Path, default=Path(".cache/diffusion_nft"))
parser.add_argument("--output-dir", type=Path, default=Path(DEFAULT_OUTPUT_DIR))
parser.add_argument("--run-config-dir", type=Path, default=Path("outputs/diffusion_nft_run_configs"))
parser.add_argument("--model-id", default=MODEL_ID)
parser.add_argument("--dataset", default="world-r1-enhanced")
parser.add_argument("--reward", default="videoalign")
parser.add_argument("--max-prompts", default="quarter")
parser.add_argument("--num-frames", type=int, default=77)
parser.add_argument("--num-latent-t", type=int, default=0)
parser.add_argument("--num-height", type=int, default=448)
parser.add_argument("--num-width", type=int, default=832)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--preprocess-batch-size", type=int, default=128)
parser.add_argument("--preprocess-num-gpus", type=int, default=1)
parser.add_argument("--preprocess-master-port", type=int, default=29541)
parser.add_argument("--dataloader-num-workers", type=int, default=0)
parser.add_argument("--samples-per-file", type=int, default=1024)
parser.add_argument("--flush-frequency", type=int, default=1024)
parser.add_argument("--num-gpus", type=int, default=4)
parser.add_argument("--sp-size", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=1)
parser.add_argument("--hsdp-replicate-dim", type=int, default=1)
parser.add_argument("--hsdp-shard-dim", type=int, default=4)
parser.add_argument("--max-train-steps", type=int, default=100)
parser.add_argument("--gradient-accumulation-steps", type=int, default=24)
parser.add_argument("--learning-rate", type=float)
parser.add_argument("--num-samples-per-prompt", type=int, default=24)
parser.add_argument("--collection-batch-size", type=int, default=6)
parser.add_argument("--inner-epochs", type=int, default=1)
parser.add_argument("--train-batch-size", type=int, default=6)
parser.add_argument("--sample-num-steps", type=int)
parser.add_argument("--sample-flow-shift", type=float)
parser.add_argument("--sample-guidance-scale", type=float)
parser.add_argument("--log-sample-max-videos", type=int, default=2)
parser.add_argument("--project-name", default="diffusion_nft_wan")
parser.add_argument("--run-name")
parser.add_argument("--diffusion-nft-root", type=Path)
parser.add_argument("--videoalign-checkpoint-path", type=Path)
parser.add_argument("--skip-preprocess", action="store_true")
parser.add_argument("--check-rewards", action="store_true")
parser.add_argument("--reward-device", default="auto", help="Device for --check-rewards: auto, cpu, cuda, cuda:0.")
parser.add_argument("--json", action="store_true", help="Print a machine-readable summary as the final line.")
args = parser.parse_args()
args.repo_root = args.repo_root.resolve()
args.config = (args.repo_root / args.config).resolve() if not args.config.is_absolute() else args.config.resolve()
args.data_root = (args.repo_root / args.data_root).resolve() if not args.data_root.is_absolute() else args.data_root
args.cache_root = ((args.repo_root / args.cache_root).resolve()
if not args.cache_root.is_absolute() else args.cache_root)
args.output_dir = ((args.repo_root / args.output_dir).resolve()
if not args.output_dir.is_absolute() else args.output_dir)
args.run_config_dir = ((args.repo_root / args.run_config_dir).resolve()
if not args.run_config_dir.is_absolute() else args.run_config_dir)
args.diffusion_nft_root = args.diffusion_nft_root or (args.cache_root / "DiffusionNFT")
args.videoalign_checkpoint_path = args.videoalign_checkpoint_path or (args.cache_root / "VideoReward")
return args
def main() -> None:
args = parse_args()
requested_num_latent_t = int(args.num_latent_t)
derived_num_latent_t = derive_wan_num_latent_t(args.num_frames)
if requested_num_latent_t <= 0:
num_latent_t = derived_num_latent_t
elif requested_num_latent_t != derived_num_latent_t:
raise ValueError(f"For Wan num_frames={args.num_frames} implies num_latent_t={derived_num_latent_t}, "
f"but got num_latent_t={requested_num_latent_t}.")
else:
num_latent_t = requested_num_latent_t
reward_map, reward_backend = resolve_reward_map(args.reward)
if args.dataset.strip().lower() not in {
"world-r1",
"world-r1-final",
"world_r1",
"world_r1_final",
"world-r1-dynamic",
"world_r1_dynamic",
"world-r1-final-dynamic",
"world_r1_final_dynamic",
"world-r1-enhanced",
"world_r1_enhanced",
"world-r1-enhanced-dynamic",
"world_r1_enhanced_dynamic",
} or reward_backend == "diffusion_nft":
ensure_diffusion_nft_repo(args.diffusion_nft_root)
os.environ["DIFFUSION_NFT_ROOT"] = str(args.diffusion_nft_root)
if any(name.startswith("videoalign_") for name in reward_map):
ensure_videoalign_checkpoint(args.videoalign_checkpoint_path)
os.environ["VIDEOALIGN_CHECKPOINT_PATH"] = str(args.videoalign_checkpoint_path)
if args.check_rewards:
check_reward_runtime(
reward_map,
reward_backend=reward_backend,
device=args.reward_device,
)
prompt_file, parquet_dir, prompt_count, source_label = prepare_prompt_file(args, num_latent_t=num_latent_t)
if has_parquet(parquet_dir):
row_count = verify_text_only_dataset(parquet_dir, prompt_count)
else:
if args.skip_preprocess:
raise RuntimeError(f"No parquet files found under {parquet_dir} and --skip-preprocess was set.")
run_text_only_preprocess(
args,
prompt_file=prompt_file,
dataset_root=parquet_dir.parent,
num_latent_t=num_latent_t,
)
row_count = verify_text_only_dataset(parquet_dir, prompt_count)
run_config_path = write_run_config(
args,
parquet_dir=parquet_dir,
num_latent_t=num_latent_t,
reward_map=reward_map,
reward_backend=reward_backend,
)
summary: dict[str, Any] = {
"prompt_file": str(prompt_file),
"prompt_source": source_label,
"prompt_count": prompt_count,
"parquet_dir": str(parquet_dir),
"parquet_rows": row_count,
"run_config": str(run_config_path),
"output_dir": str(args.output_dir),
"num_frames": args.num_frames,
"num_latent_t": num_latent_t,
"reward_backend": reward_backend,
"reward_map": reward_map,
}
print("Prepared DiffusionNFT assets:")
for key, value in summary.items():
print(f" {key}: {value}")
if args.json:
print(json.dumps(summary, sort_keys=True))
if __name__ == "__main__":
main()
@@ -0,0 +1,11 @@
# Reward runtime pins for DiffusionNFT video reward training.
#
# The vendored VideoAlign/HPSv3 inference runtimes follow the Qwen2-VL package layout
# used by Transformers 4.x. Newer Transformers 5.x moves Qwen2VL internals and
# can load VideoReward with missing language-model weights.
datasets==3.6.0
peft==0.19.1
qwen-vl-utils==0.0.11
safetensors==0.5.3
timm==1.0.15
transformers==4.57.3
@@ -80,6 +80,7 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
batch_captions,
fastvideo_args,
encoder_index=[0],
padding=True,
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2024 HPSv3 Team
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
@@ -0,0 +1 @@
from .inference import HPSv3RewardInferencer
@@ -0,0 +1,59 @@
# Model Configuration
rm_head_type: "ranknet"
lora_enable: False
vision_lora: False
freeze_vision_tower: False
freeze_llm: False
tune_merger: True
model_name_or_path: "Qwen/Qwen2-VL-7B-Instruct"
num_lora_modules: -1
lora_r: 512
lora_alpha: 1024
lora_namespan_exclude: ['lm_head', 'rm_head', 'embed_tokens']
# Data Configuration
confidence_threshold: 0.95
tied_threshold: null
max_pixels: 200704 # 256 * 28 * 28
min_pixels: 200704
with_instruction: true
train_json_list:
- example_train.json
test_json_list:
- ["Valid Set 1", ["example_set_1_part1.json", "example_set_1_part2.json"]]
- ['Valid Set 2',["example_set_2_part1.json"]]
soft_label: False
output_dir: output_models
use_special_tokens: true
reward_token: "special"
output_dim: 2
loss_type: "uncertainty"
# Training Configuration
disable_flash_attn2: False
per_device_train_batch_size: 2
per_device_eval_batch_size: 8
gradient_accumulation_steps: 4
num_train_epochs: 10
learning_rate: 2.0e-6
special_token_lr: 2.0e-6
warmup_ratio: 0.05
lr_scheduler_type: "constant_with_warmup"
gradient_checkpointing: True
gradient_checkpointing_kwargs: {"use_reentrant": False}
# Evaluation and Logging
eval_strategy: "steps"
logging_epochs: 0.01
eval_epochs: 0.1
save_epochs: 0.1
report_to: tensorboard
# System Configuration
bf16: True
torch_dtype: "bfloat16"
save_only_model: True
save_full_model: True
dataloader_num_workers: 8
@@ -0,0 +1,440 @@
from __future__ import annotations
## This file is modified from https://github.com/kq-chen/qwen-vl-utils/blob/main/src/qwen_vl_utils/vision_process.py
import base64
import logging
import math
import os
import sys
import time
import warnings
from functools import lru_cache
from io import BytesIO
import requests
import torch
import torchvision
from packaging import version
from PIL import Image
from torchvision import io, transforms
from torchvision.transforms import InterpolationMode
logger = logging.getLogger(__name__)
IMAGE_FACTOR = 28
MIN_PIXELS = 4 * 28 * 28
MAX_PIXELS = 16384 * 28 * 28
MAX_RATIO = 200
VIDEO_MIN_PIXELS = 128 * 28 * 28
VIDEO_MAX_PIXELS = 768 * 28 * 28
VIDEO_TOTAL_PIXELS = 24576 * 28 * 28
FRAME_FACTOR = 2
FPS = 2.0
FPS_MIN_FRAMES = 4
FPS_MAX_FRAMES = 768
def round_by_factor(number: int, factor: int) -> int:
"""Returns the closest integer to 'number' that is divisible by 'factor'."""
return round(number / factor) * factor
def ceil_by_factor(number: int, factor: int) -> int:
"""Returns the smallest integer greater than or equal to 'number' that is divisible by 'factor'."""
return math.ceil(number / factor) * factor
def floor_by_factor(number: int, factor: int) -> int:
"""Returns the largest integer less than or equal to 'number' that is divisible by 'factor'."""
return math.floor(number / factor) * factor
def smart_resize(
height: int,
width: int,
factor: int = IMAGE_FACTOR,
min_pixels: int = MIN_PIXELS,
max_pixels: int = MAX_PIXELS,
) -> tuple[int, int]:
"""
Rescales the image so that the following conditions are met:
1. Both dimensions (height and width) are divisible by 'factor'.
2. The total number of pixels is within the range ['min_pixels', 'max_pixels'].
3. The aspect ratio of the image is maintained as closely as possible.
"""
if max(height, width) / min(height, width) > MAX_RATIO:
raise ValueError(
f"absolute aspect ratio must be smaller than {MAX_RATIO}, got {max(height, width) / min(height, width)}"
)
h_bar = max(factor, round_by_factor(height, factor))
w_bar = max(factor, round_by_factor(width, factor))
if h_bar * w_bar > max_pixels:
beta = math.sqrt((height * width) / max_pixels)
h_bar = floor_by_factor(height / beta, factor)
w_bar = floor_by_factor(width / beta, factor)
elif h_bar * w_bar < min_pixels:
beta = math.sqrt(min_pixels / (height * width))
h_bar = ceil_by_factor(height * beta, factor)
w_bar = ceil_by_factor(width * beta, factor)
return h_bar, w_bar
def fetch_image(ele: dict[str, str | Image.Image], size_factor: int = IMAGE_FACTOR) -> Image.Image:
if "image" in ele:
image = ele["image"]
else:
image = ele["image_url"]
image_obj = None
if isinstance(image, Image.Image) or isinstance(image, torch.Tensor):
image_obj = image
elif image.startswith("http://") or image.startswith("https://"):
image_obj = Image.open(requests.get(image, stream=True, timeout=10).raw)
elif image.startswith("file://"):
image_obj = Image.open(image[7:])
elif image.startswith("data:image"):
if "base64," in image:
_, base64_data = image.split("base64,", 1)
data = base64.b64decode(base64_data)
image_obj = Image.open(BytesIO(data))
else:
image_obj = Image.open(image)
if image_obj is None:
raise ValueError(f"Unrecognized image input, support local path, http url, base64 and PIL.Image, got {image}")
if isinstance(image_obj, Image.Image):
image = image_obj.convert("RGB")
## resize
if "resized_height" in ele and "resized_width" in ele:
resized_height, resized_width = smart_resize(
ele["resized_height"],
ele["resized_width"],
factor=size_factor,
)
else:
if isinstance(image, torch.Tensor):
shape = image.shape
if len(shape) == 4:
if shape[1] in [1, 3]: # Likely [B, C, H, W]
height, width = shape[2], shape[3]
image_mode = "NCHW"
elif shape[3] in [1, 3]: # Likely [B, H, W, C]
height, width = shape[1], shape[2]
image_mode = "NHWC"
elif len(shape) == 3:
if shape[0] in [1, 3]: # Likely [C, H, W]
height, width = shape[1], shape[2]
image_mode = "CHW"
elif shape[2] in [1, 3]: # Likely [H, W, C]
height, width = shape[0], shape[1]
image_mode = "HWC"
else:
raise ValueError(f"Cannot determine tensor image format from shape {shape}")
else:
raise ValueError(f"Unsupported tensor image shape: {shape}")
else:
width, height = image.size
min_pixels = ele.get("min_pixels", MIN_PIXELS)
max_pixels = ele.get("max_pixels", MAX_PIXELS)
resized_height, resized_width = smart_resize(
height,
width,
factor=size_factor,
min_pixels=min_pixels,
max_pixels=max_pixels,
)
if isinstance(image, torch.Tensor):
if image_mode == "NCHW":
image = transforms.functional.resize(
image,
[resized_height, resized_width],
interpolation=InterpolationMode.BICUBIC,
antialias=True,
)
elif image_mode == "NHWC":
image = transforms.functional.resize(
image.permute(0, 3, 1, 2),
[resized_height, resized_width],
interpolation=InterpolationMode.BICUBIC,
antialias=True,
)
elif image_mode == "CHW":
image = image.unsqueeze(0) # Add batch dimension
image = transforms.functional.resize(
image,
[resized_height, resized_width],
interpolation=InterpolationMode.BICUBIC,
antialias=True,
)
elif image_mode == "HWC":
image = image.permute(2, 0, 1).unsqueeze(0) # Add batch dimension and change to CHW
image = transforms.functional.resize(
image,
[resized_height, resized_width],
interpolation=InterpolationMode.BICUBIC,
antialias=True,
)
else:
# If the image is a PIL Image, we resize it using PIL.
if image.mode != "RGB":
image = image.convert("RGB")
image = image.resize((resized_width, resized_height), Image.BICUBIC)
return image
def smart_nframes(
ele: dict,
total_frames: int,
video_fps: float,
) -> int:
"""calculate the number of frames for video used for model inputs.
Args:
ele (dict): a dict contains the configuration of video.
support either `fps` or `nframes`:
- nframes: the number of frames to extract for model inputs.
- fps: the fps to extract frames for model inputs.
- min_frames: the minimum number of frames of the video, only used when fps is provided.
- max_frames: the maximum number of frames of the video, only used when fps is provided.
total_frames (int): the original total number of frames of the video.
video_fps (int | float): the original fps of the video.
Raises:
ValueError: nframes should in interval [FRAME_FACTOR, total_frames].
Returns:
int: the number of frames for video used for model inputs.
"""
assert not ("fps" in ele and "nframes" in ele), "Only accept either `fps` or `nframes`"
if "nframes" in ele:
nframes = round_by_factor(ele["nframes"], FRAME_FACTOR)
else:
fps = ele.get("fps", FPS)
min_frames = ceil_by_factor(ele.get("min_frames", FPS_MIN_FRAMES), FRAME_FACTOR)
max_frames = floor_by_factor(ele.get("max_frames", min(FPS_MAX_FRAMES, total_frames)), FRAME_FACTOR)
nframes = total_frames / video_fps * fps
nframes = min(max(nframes, min_frames), max_frames)
nframes = round_by_factor(nframes, FRAME_FACTOR)
nframes = min(nframes, total_frames)
if not (nframes >= FRAME_FACTOR and nframes <= total_frames):
raise ValueError(f"nframes should in interval [{FRAME_FACTOR}, {total_frames}], but got {nframes}.")
return nframes
def _read_video_torchvision(
ele: dict,
) -> torch.Tensor:
"""read video using torchvision.io.read_video
Args:
ele (dict): a dict contains the configuration of video.
support keys:
- video: the path of video. support "file://", "http://", "https://" and local path.
- video_start: the start time of video.
- video_end: the end time of video.
Returns:
torch.Tensor: the video tensor with shape (T, C, H, W).
"""
video_path = ele["video"]
if version.parse(torchvision.__version__) < version.parse("0.19.0"):
if "http://" in video_path or "https://" in video_path:
warnings.warn("torchvision < 0.19.0 does not support http/https video path, please upgrade to 0.19.0.")
if "file://" in video_path:
video_path = video_path[7:]
st = time.time()
video, audio, info = io.read_video(
video_path,
start_pts=ele.get("video_start", 0.0),
end_pts=ele.get("video_end"),
pts_unit="sec",
output_format="TCHW",
)
total_frames, video_fps = video.size(0), info["video_fps"]
# logger.info(f"torchvision: {video_path=}, {total_frames=}, {video_fps=}, time={time.time() - st:.3f}s")
if ele["sample_type"] == "uniform":
nframes = smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
elif ele["sample_type"] == "multi_pts":
frames_each_pts = 6
num_pts = 4
fps = 8
nframes = int(total_frames * fps // video_fps)
frames_idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
start_pt = int(frames_each_pts // 2)
end_pt = int(nframes - frames_each_pts // 2 - 1)
pts = torch.linspace(start_pt, end_pt, num_pts).round().long().tolist()
idx = []
for pt in pts:
idx.extend(frames_idx[pt - frames_each_pts // 2 : pt + frames_each_pts // 2])
video = video[idx]
return video
def is_decord_available() -> bool:
import importlib.util
return importlib.util.find_spec("decord") is not None
def _read_video_decord(
ele: dict,
) -> torch.Tensor:
"""read video using decord.VideoReader
Args:
ele (dict): a dict contains the configuration of video.
support keys:
- video: the path of video. support "file://", "http://", "https://" and local path.
- video_start: the start time of video.
- video_end: the end time of video.
Returns:
torch.Tensor: the video tensor with shape (T, C, H, W).
"""
import decord
video_path = ele["video"]
st = time.time()
vr = decord.VideoReader(video_path)
# TODO: support start_pts and end_pts
if "video_start" in ele or "video_end" in ele:
raise NotImplementedError("not support start_pts and end_pts in decord for now.")
total_frames, video_fps = len(vr), vr.get_avg_fps()
# logger.info(f"decord: {video_path=}, {total_frames=}, {video_fps=}, time={time.time() - st:.3f}s")
if ele["sample_type"] == "uniform":
nframes = smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
# nframes = max(nframes, 8)
# import pdb; pdb.set_trace()
idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
elif ele["sample_type"] == "multi_pts":
frames_each_pts = 6
num_pts = 4
fps = 8
nframes = int(total_frames * fps // video_fps)
frames_idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
start_pt = int(frames_each_pts // 2)
end_pt = int(nframes - frames_each_pts // 2 - 1)
pts = torch.linspace(start_pt, end_pt, num_pts).round().long().tolist()
idx = []
for pt in pts:
idx.extend(frames_idx[pt - frames_each_pts // 2 : pt + frames_each_pts // 2])
video = vr.get_batch(idx).asnumpy()
video = torch.tensor(video).permute(0, 3, 1, 2) # Convert to TCHW format
return video
VIDEO_READER_BACKENDS = {
"decord": _read_video_decord,
"torchvision": _read_video_torchvision,
}
FORCE_QWENVL_VIDEO_READER = os.getenv("FORCE_QWENVL_VIDEO_READER", None)
@lru_cache(maxsize=1)
def get_video_reader_backend() -> str:
if FORCE_QWENVL_VIDEO_READER is not None:
video_reader_backend = FORCE_QWENVL_VIDEO_READER
elif is_decord_available():
video_reader_backend = "decord"
else:
video_reader_backend = "torchvision"
print(f"qwen-vl-utils using {video_reader_backend} to read video.", file=sys.stderr)
return video_reader_backend
def fetch_video(ele: dict, image_factor: int = IMAGE_FACTOR) -> torch.Tensor | list[Image.Image]:
if isinstance(ele["video"], str):
video_reader_backend = get_video_reader_backend()
video = VIDEO_READER_BACKENDS[video_reader_backend](ele)
# import pdb; pdb.set_trace()
nframes, _, height, width = video.shape
min_pixels = ele.get("min_pixels", VIDEO_MIN_PIXELS)
total_pixels = ele.get("total_pixels", VIDEO_TOTAL_PIXELS)
max_pixels = max(
min(VIDEO_MAX_PIXELS, total_pixels / nframes * FRAME_FACTOR),
int(min_pixels * 1.05),
)
max_pixels = ele.get("max_pixels", max_pixels)
if "resized_height" in ele and "resized_width" in ele:
resized_height, resized_width = smart_resize(
ele["resized_height"],
ele["resized_width"],
factor=image_factor,
)
else:
resized_height, resized_width = smart_resize(
height,
width,
factor=image_factor,
min_pixels=min_pixels,
max_pixels=max_pixels,
)
video = transforms.functional.resize(
video,
[resized_height, resized_width],
interpolation=InterpolationMode.BICUBIC,
antialias=True,
).float()
return video
assert isinstance(ele["video"], (list, tuple))
process_info = ele.copy()
process_info.pop("type", None)
process_info.pop("video", None)
images = [
fetch_image({"image": video_element, **process_info}, size_factor=image_factor)
for video_element in ele["video"]
]
nframes = ceil_by_factor(len(images), FRAME_FACTOR)
if len(images) < nframes:
images.extend([images[-1]] * (nframes - len(images)))
return images
def extract_vision_info(conversations: list[dict] | list[list[dict]]) -> list[dict]:
vision_infos = []
if isinstance(conversations[0], dict):
conversations = [conversations]
for conversation in conversations:
for message in conversation:
if isinstance(message["content"], list):
for ele in message["content"]:
if (
"image" in ele
or "image_url" in ele
or "video" in ele
or ele["type"] in ("image", "image_url", "video")
):
vision_infos.append(ele)
return vision_infos
def process_vision_info(
conversations: list[dict] | list[list[dict]],
) -> tuple[list[Image.Image] | None, list[torch.Tensor | list[Image.Image]] | None]:
vision_infos = extract_vision_info(conversations)
## Read images or videos
image_inputs = []
video_inputs = []
for vision_info in vision_infos:
if "image" in vision_info or "image_url" in vision_info:
image_inputs.append(fetch_image(vision_info))
elif "video" in vision_info:
video_inputs.append(fetch_video(vision_info))
else:
raise ValueError("image, image_url or video should in content.")
if len(image_inputs) == 0:
image_inputs = None
if len(video_inputs) == 0:
video_inputs = None
return image_inputs, video_inputs
@@ -0,0 +1,187 @@
import os
from collections.abc import Mapping
from pathlib import Path
import huggingface_hub
import torch
from .dataset.utils import process_vision_info
from .runtime import create_model_and_processor
from .utils.parser import (
DataConfig,
ModelConfig,
PEFTLoraConfig,
TrainingConfig,
parse_args_with_yaml,
)
_MODEL_CONFIG_PATH = Path(__file__).parent / "config/"
INSTRUCTION = """
You are tasked with evaluating a generated image based on Visual Quality and
Text Alignment and give a overall score to estimate the human preference.
Please provide a rating from 0 to 10, with 0 being the worst and 10 being the
best.
Textual prompt - {text_prompt}
"""
prompt_with_special_token = """
Please provide the overall ratings of this image: <|Reward|>
END
"""
prompt_without_special_token = """
Please provide the overall ratings of this image:
"""
class HPSv3RewardInferencer:
def __init__(
self,
config_path=None,
checkpoint_path=None,
device="cuda",
differentiable=False,
):
if differentiable:
raise ValueError("The vendored HPSv3 runtime is inference-only and does not support differentiable mode.")
if config_path is None:
config_path = os.path.join(_MODEL_CONFIG_PATH, "HPSv3_7B.yaml")
if checkpoint_path is None:
checkpoint_path = huggingface_hub.hf_hub_download("MizzenAI/HPSv3", "HPSv3.safetensors", repo_type="model")
(
(
data_config,
training_args,
model_config,
peft_lora_config,
),
config_path,
) = parse_args_with_yaml(
(DataConfig, TrainingConfig, ModelConfig, PEFTLoraConfig),
config_path,
is_train=False,
)
training_args.output_dir = os.path.join(training_args.output_dir, config_path.split("/")[-1].split(".")[0])
model, processor, peft_config = create_model_and_processor(
model_config=model_config,
peft_lora_config=peft_lora_config,
training_args=training_args,
)
self.device = device
self.use_special_tokens = model_config.use_special_tokens
if checkpoint_path.endswith(".safetensors"):
import safetensors.torch
state_dict = safetensors.torch.load_file(checkpoint_path, device="cpu")
else:
state_dict = torch.load(checkpoint_path, map_location="cpu")
if "model" in state_dict:
state_dict = state_dict["model"]
model.load_state_dict(state_dict, strict=True)
model.eval()
self.model = model
self.processor = processor
self.model.to(self.device)
self.data_config = data_config
def _pad_sequence(self, sequences, attention_mask, max_len, padding_side="right"):
"""
Pad the sequences to the maximum length.
"""
assert padding_side in ["right", "left"]
if sequences.shape[1] >= max_len:
return sequences, attention_mask
pad_len = max_len - sequences.shape[1]
padding = (0, pad_len) if padding_side == "right" else (pad_len, 0)
sequences_padded = torch.nn.functional.pad(
sequences, padding, "constant", self.processor.tokenizer.pad_token_id
)
attention_mask_padded = torch.nn.functional.pad(attention_mask, padding, "constant", 0)
return sequences_padded, attention_mask_padded
def _prepare_input(self, data):
"""
Prepare `inputs` before feeding them to the model, converting them to tensors if they are not already and
handling potential state.
"""
if isinstance(data, Mapping):
return type(data)({k: self._prepare_input(v) for k, v in data.items()})
if isinstance(data, (tuple, list)):
return type(data)(self._prepare_input(v) for v in data)
if isinstance(data, torch.Tensor):
kwargs = {"device": self.device}
return data.to(**kwargs)
return data
def _prepare_inputs(self, inputs):
"""
Prepare `inputs` before feeding them to the model, converting them to tensors if they are not already and
handling potential state.
"""
inputs = self._prepare_input(inputs)
if len(inputs) == 0:
raise ValueError("The input batch is empty.")
return inputs
def prepare_batch(self, image_paths, prompts):
max_pixels = 256 * 28 * 28
min_pixels = 256 * 28 * 28
message_list = []
for text, image in zip(prompts, image_paths):
out_message = [
{
"role": "user",
"content": [
{
"type": "image",
"image": image,
"min_pixels": max_pixels,
"max_pixels": max_pixels,
},
{
"type": "text",
"text": (
INSTRUCTION.format(text_prompt=text) + prompt_with_special_token
if self.use_special_tokens
else prompt_without_special_token
),
},
],
}
]
message_list.append(out_message)
image_inputs, _ = process_vision_info(message_list)
batch = self.processor(
text=self.processor.apply_chat_template(message_list, tokenize=False, add_generation_prompt=True),
images=image_inputs,
padding=True,
return_tensors="pt",
videos_kwargs={"do_rescale": True},
)
batch = self._prepare_inputs(batch)
return batch
@torch.inference_mode()
def reward(self, prompts, image_paths):
batch = self.prepare_batch(image_paths, prompts)
rewards = self.model(return_dict=True, **batch)["logits"]
return rewards
@@ -0,0 +1,170 @@
import torch
from torch import nn
from transformers import Qwen2VLForConditionalGeneration
class Qwen2VLRewardModelBT(Qwen2VLForConditionalGeneration):
def __init__(
self,
config,
output_dim=4,
reward_token="last",
special_token_ids=None,
rm_head_type="default",
rm_head_kwargs=None,
**kwargs,
):
del kwargs
super().__init__(config)
self.output_dim = output_dim
hidden_size = getattr(config, "hidden_size", None)
if hidden_size is None and hasattr(config, "text_config"):
hidden_size = getattr(config.text_config, "hidden_size", None)
if hidden_size is None:
raise AttributeError("Qwen2VL reward model config must define hidden_size or text_config.hidden_size")
if rm_head_type == "default":
self.rm_head = nn.Linear(hidden_size, output_dim, bias=False)
elif rm_head_type == "ranknet":
if rm_head_kwargs is None:
self.rm_head = nn.Sequential(
nn.Linear(hidden_size, 1024),
nn.ReLU(),
nn.Dropout(0.05),
nn.Linear(1024, 16),
nn.ReLU(),
nn.Linear(16, output_dim),
)
else:
for layer in range(rm_head_kwargs.get("num_layers", 3)):
if layer == 0:
self.rm_head = nn.Sequential(
nn.Linear(hidden_size, rm_head_kwargs["hidden_size"]),
nn.ReLU(),
nn.Dropout(rm_head_kwargs.get("dropout", 0.1)),
)
elif layer < rm_head_kwargs.get("num_layers", 3) - 1:
self.rm_head.add_module(
f"layer_{layer}",
nn.Sequential(
nn.Linear(
rm_head_kwargs["hidden_size"],
rm_head_kwargs["hidden_size"],
),
nn.ReLU(),
nn.Dropout(rm_head_kwargs.get("dropout", 0.1)),
),
)
else:
self.rm_head.add_module(
"output_layer",
nn.Linear(
rm_head_kwargs["hidden_size"],
output_dim,
bias=rm_head_kwargs.get("bias", False),
),
)
else:
raise ValueError(f"Unsupported rm_head_type: {rm_head_type}")
self.rm_head.to(torch.float32)
self.reward_token = reward_token
self.special_token_ids = special_token_ids
if self.special_token_ids is not None:
self.reward_token = "special"
def forward(
self,
input_ids: torch.LongTensor = None,
attention_mask: torch.Tensor | None = None,
position_ids: torch.LongTensor | None = None,
past_key_values: list[torch.FloatTensor] | None = None,
inputs_embeds: torch.FloatTensor | None = None,
labels: torch.LongTensor | None = None,
use_cache: bool | None = None,
output_attentions: bool | None = None,
output_hidden_states: bool | None = None,
return_dict: bool | None = None,
pixel_values: torch.Tensor | None = None,
pixel_values_videos: torch.FloatTensor | None = None,
image_grid_thw: torch.LongTensor | None = None,
video_grid_thw: torch.LongTensor | None = None,
rope_deltas: torch.LongTensor | None = None,
**kwargs,
):
del kwargs, labels, rope_deltas
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if inputs_embeds is None:
inputs_embeds = self.model.embed_tokens(input_ids)
if pixel_values is not None:
pixel_values = pixel_values.type(self.visual.get_dtype())
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
image_mask = (input_ids == self.config.image_token_id).unsqueeze(-1).expand_as(inputs_embeds)
image_embeds = image_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)
if pixel_values_videos is not None:
pixel_values_videos = pixel_values_videos.type(self.visual.get_dtype())
video_embeds = self.visual(pixel_values_videos, grid_thw=video_grid_thw)
video_mask = (input_ids == self.config.video_token_id).unsqueeze(-1).expand_as(inputs_embeds)
video_embeds = video_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)
if attention_mask is not None:
attention_mask = attention_mask.to(inputs_embeds.device)
outputs = self.model(
input_ids=None,
position_ids=position_ids,
attention_mask=attention_mask,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
hidden_states = outputs[0]
autocast_device = hidden_states.device.type
with torch.autocast(device_type=autocast_device, dtype=torch.float32, enabled=autocast_device != "cpu"):
logits = self.rm_head(hidden_states)
if input_ids is not None:
batch_size = input_ids.shape[0]
else:
batch_size = inputs_embeds.shape[0]
if self.config.pad_token_id is None and batch_size != 1:
raise ValueError("Cannot handle batch sizes > 1 if no padding token is defined.")
if self.config.pad_token_id is None:
sequence_lengths = -1
elif input_ids is not None:
sequence_lengths = torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1
sequence_lengths = sequence_lengths % input_ids.shape[-1]
sequence_lengths = sequence_lengths.to(logits.device)
else:
sequence_lengths = -1
if self.reward_token == "last":
pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths]
elif self.reward_token == "mean":
valid_lengths = torch.clamp(sequence_lengths, min=0, max=logits.size(1) - 1)
pooled_logits = torch.stack([logits[i, :valid_lengths[i]].mean(dim=0) for i in range(batch_size)])
elif self.reward_token == "special":
special_token_mask = torch.zeros_like(input_ids, dtype=torch.bool)
for special_token_id in self.special_token_ids:
special_token_mask = special_token_mask | (input_ids == special_token_id)
pooled_logits = logits[special_token_mask, ...]
pooled_logits = pooled_logits.view(batch_size, 1, -1)
pooled_logits = pooled_logits.view(batch_size, -1)
else:
raise ValueError("Invalid reward_token")
return {"logits": pooled_logits}
+119
View File
@@ -0,0 +1,119 @@
from importlib import util
import torch
from peft import LoraConfig, get_peft_model
from transformers import AutoProcessor
from hpsv3.model.reward_model import Qwen2VLRewardModelBT
def find_target_linear_names(model, num_lora_modules=-1, lora_namespan_exclude=None):
linear_cls = torch.nn.Linear
embedding_cls = torch.nn.Embedding
excluded = lora_namespan_exclude or []
lora_module_names = []
for name, module in model.named_modules():
if any(ex_keyword in name for ex_keyword in excluded):
continue
if isinstance(module, (linear_cls, embedding_cls)):
lora_module_names.append(name)
if num_lora_modules > 0:
lora_module_names = lora_module_names[-num_lora_modules:]
return lora_module_names
def _get_quantization_config(model_config):
if not model_config.load_in_8bit and not model_config.load_in_4bit:
return None
from transformers import BitsAndBytesConfig
if model_config.load_in_8bit:
return BitsAndBytesConfig(load_in_8bit=True)
return BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type=model_config.bnb_4bit_quant_type,
bnb_4bit_use_double_quant=model_config.use_bnb_nested_quant,
)
def create_model_and_processor(
model_config,
peft_lora_config,
training_args,
cache_dir=None,
):
torch_dtype = (
model_config.torch_dtype
if model_config.torch_dtype in ["auto", None]
else getattr(torch, model_config.torch_dtype)
)
quantization_config = _get_quantization_config(model_config)
model_kwargs = dict(
revision=model_config.model_revision,
device_map="auto" if quantization_config is not None else None,
quantization_config=quantization_config,
use_cache=False,
)
processor = AutoProcessor.from_pretrained(
model_config.model_name_or_path, padding_side="right", cache_dir=cache_dir
)
special_token_ids = None
if model_config.use_special_tokens:
special_tokens = ["<|Reward|>"]
processor.tokenizer.add_special_tokens({"additional_special_tokens": special_tokens})
special_token_ids = processor.tokenizer.convert_tokens_to_ids(special_tokens)
has_flash_attn = util.find_spec("flash_attn") is not None
model = Qwen2VLRewardModelBT.from_pretrained(
model_config.model_name_or_path,
output_dim=model_config.output_dim,
reward_token=model_config.reward_token,
special_token_ids=special_token_ids,
torch_dtype=torch_dtype,
attn_implementation=(
"flash_attention_2" if not training_args.disable_flash_attn2 and has_flash_attn else "sdpa"
),
cache_dir=cache_dir,
rm_head_type=model_config.rm_head_type,
rm_head_kwargs=model_config.rm_head_kwargs,
**model_kwargs,
)
if model_config.use_special_tokens:
model.resize_token_embeddings(len(processor.tokenizer))
if training_args.bf16:
model.to(torch.bfloat16)
if training_args.fp16:
model.to(torch.float16)
model.rm_head.to(torch.float32)
if peft_lora_config.lora_enable:
target_modules = find_target_linear_names(
model,
num_lora_modules=peft_lora_config.num_lora_modules,
lora_namespan_exclude=peft_lora_config.lora_namespan_exclude,
)
peft_config = LoraConfig(
target_modules=target_modules,
r=peft_lora_config.lora_r,
lora_alpha=peft_lora_config.lora_alpha,
lora_dropout=peft_lora_config.lora_dropout,
task_type=peft_lora_config.lora_task_type,
use_rslora=peft_lora_config.use_rslora,
bias="none",
modules_to_save=peft_lora_config.lora_modules_to_save,
)
model = get_peft_model(model, peft_config)
else:
peft_config = None
model.config.tokenizer_padding_side = processor.tokenizer.padding_side
model.config.pad_token_id = processor.tokenizer.pad_token_id
return model, processor, peft_config
@@ -0,0 +1,168 @@
import contextlib
from dataclasses import dataclass, field
from typing import Any, Literal
from omegaconf import OmegaConf
from transformers import HfArgumentParser, TrainingArguments
@dataclass
class DataConfig:
train_json_list: list[str] = field(default_factory=lambda: ["/path/to/dataset/meta_data.json"])
val_json_list: list[str] = field(default_factory=lambda: ["/path/to/dataset/meta_data.json"])
test_json_list: list[str] = field(default_factory=lambda: ["/path/to/dataset/meta_data.json"])
soft_label: bool = False
confidence_threshold: float | None = None
max_pixels: int | None = 256 * 28 * 28 # Default max pixels
min_pixels: int | None = 256 * 28 * 28
with_instruction: bool = True
tied_threshold: float | None = None
@dataclass
class TrainingConfig(TrainingArguments):
max_grad_norm: float | None = 1.0
dataset_num_proc: int | None = None
center_rewards_coefficient: float | None = None
disable_flash_attn2: bool = field(default=False)
disable_dropout: bool = field(default=False)
vision_lr: float | None = None
merger_lr: float | None = None
rm_head_lr: float | None = None
special_token_lr: float | None = None
conduct_eval: bool | None = True
load_from_pretrained: str = None
load_from_pretrained_step: int = None
logging_epochs: float | None = None
eval_epochs: float | None = None
save_epochs: float | None = None
remove_unused_columns: bool | None = False
save_full_model: bool | None = False
# Visualization parameters
visualization_steps: int | None = 100
max_viz_samples: int | None = 4
@dataclass
class PEFTLoraConfig:
lora_enable: bool = False
vision_lora: bool = False
lora_r: int = 16
lora_alpha: int = 32
lora_dropout: float = 0.05
lora_target_modules: list[str] | None = None
lora_namespan_exclude: list[str] | None = None
lora_modules_to_save: list[str] | None = None
lora_task_type: str = "CAUSAL_LM"
use_rslora: bool = False
num_lora_modules: int = -1
def __post_init__(self):
if isinstance(self.lora_target_modules, list) and len(self.lora_target_modules) == 1:
self.lora_target_modules = self.lora_target_modules[0]
if isinstance(self.lora_namespan_exclude, list) and len(self.lora_namespan_exclude) == 1:
self.lora_namespan_exclude = self.lora_namespan_exclude[0]
@dataclass
class ModelConfig:
model_name_or_path: str | None = None
model_revision: str = "main"
rm_head_type: str = "default"
rm_head_kwargs: dict | None = None
output_dim: int = 1
use_special_tokens: bool = False
freeze_vision_tower: bool = field(default=False)
freeze_llm: bool = field(default=False)
tune_merger: bool = field(default=False)
trainable_visual_layers: int | None = -1
torch_dtype: Literal["auto", "bfloat16", "float16", "float32"] | None = None
trust_remote_code: bool = False
attn_implementation: str | None = None
load_in_8bit: bool = False
load_in_4bit: bool = False
bnb_4bit_quant_type: Literal["fp4", "nf4"] = "nf4"
use_bnb_nested_quant: bool = False
reward_token: Literal["last", "mean", "special"] = "last"
loss_type: Literal["bt", "reg", "btt", "margin", "constant_margin", "scaled"] = "regular"
loss_hyperparameters: dict = field(default_factory=dict)
checkpoint_path: str | None = None
def __post_init__(self):
if self.load_in_8bit and self.load_in_4bit:
raise ValueError("You can't use 8 bit and 4 bit precision at the same time")
# if isinstance(self.lora_target_modules, list) and len(self.lora_target_modules) == 1:
# self.lora_target_modules = self.lora_target_modules[0]
# if isinstance(self.lora_namespan_exclude, list) and len(self.lora_namespan_exclude) == 1:
# self.lora_namespan_exclude = self.lora_namespan_exclude[0]
########## Functions for get trainable modules' parameters ##########
def parse_args_with_yaml(
dataclass_types: tuple[type, ...],
config_path: str = None,
allow_extra_keys: bool = True,
is_train: bool = True,
) -> tuple[Any, ...]:
"""
Parse arguments using HfArgumentParser with OmegaConf for YAML support.
Args:
dataclass_types: Tuple of dataclass types for HfArgumentParser
args: Optional arguments (if None, will read from sys.argv)
allow_extra_keys: Whether to allow extra keys in config
Returns:
Tuple of parsed dataclass instances
"""
# Read arguments from command line or provided args
# Load YAML config and merge with command line overrides
args = OmegaConf.to_container(OmegaConf.load(config_path))
@contextlib.contextmanager
def _disable_accelerate_state_reset(enabled: bool):
if not enabled:
yield
return
try:
from accelerate.state import AcceleratorState, PartialState
except Exception:
# If accelerate is unavailable, just continue.
yield
return
orig_acc_reset = AcceleratorState._reset_state
orig_partial_reset = PartialState._reset_state
def _no_reset_state(*_args, **_kwargs):
return None
AcceleratorState._reset_state = staticmethod(_no_reset_state)
PartialState._reset_state = staticmethod(_no_reset_state)
try:
yield
finally:
AcceleratorState._reset_state = orig_acc_reset
PartialState._reset_state = orig_partial_reset
# Parse with HfArgumentParser
parser = HfArgumentParser(dataclass_types)
with _disable_accelerate_state_reset(enabled=not is_train):
return parser.parse_dict(args, allow_extra_keys=allow_extra_keys), config_path
if __name__ == "__main__":
data_config, training_args, model_config, peft_lora_config = parse_args_with_yaml(
(DataConfig, TrainingConfig, ModelConfig, PEFTLoraConfig)
)
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2025 Kling Team, Kuaishou Technology
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+120
View File
@@ -0,0 +1,120 @@
<h1 align="center"> Improving Video Generation with Human Feedback </h1>
<div align="center">
<!-- <a href='LICENSE'><img src='https://img.shields.io/badge/license-MIT-yellow'></a> -->
<a href='https://arxiv.org/abs/2501.13918'><img src='https://img.shields.io/badge/arXiv-VideoAlign-red'></a> &nbsp;
<a href='https://gongyeliu.github.io/videoalign/'><img src='https://img.shields.io/badge/Project-VideoAlign-green'></a> &nbsp;
<a href="https://github.com/KwaiVGI/VideoAlign"><img src="https://img.shields.io/badge/GitHub-VideoAlign-9E95B7?logo=github"></a> &nbsp;
<a href='https://huggingface.co/KwaiVGI/VideoReward'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20Model-VideoReward-blue'></a> &nbsp;
<br>
<a href='https://huggingface.co/datasets/KwaiVGI/VideoGen-RewardBench'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20Eval%20Dataset-VideoGen--RewardBench-blue'></a> &nbsp;
<a href='https://huggingface.co/spaces/KwaiVGI/VideoGen-RewardBench'><img src='https://img.shields.io/badge/Space-VideoGen--RewardBench-orange.svg?logo=data:image/svg+xml;charset=utf-8;base64,PHN2ZyB0PSIxNzM5MjA0MzY2MDEwIiBjbGFzcz0iaWNvbiIgdmlld0JveD0iMCAwIDEwMjQgMTAyNCIgdmVyc2lvbj0iMS4xIiB4bWxucz0iaHR0cDovL3d3dy53My5vcmcvMjAwMC9zdmciIHAtaWQ9IjQzNDYiIHdpZHRoPSIyMDAiIGhlaWdodD0iMjAwIj48cGF0aCBkPSJNNjgyLjY2NjY2NyA0NjkuMzMzMzMzVjEyOEgzNDEuMzMzMzMzdjI1Nkg4NS4zMzMzMzN2NTEyaDg1My4zMzMzMzRWNDY5LjMzMzMzM2gtMjU2eiBtLTI1Ni0yNTZoMTcwLjY2NjY2NnY1OTcuMzMzMzM0aC0xNzAuNjY2NjY2VjIxMy4zMzMzMzN6IG0tMjU2IDI1NmgxNzAuNjY2NjY2djM0MS4zMzMzMzRIMTcwLjY2NjY2N3YtMzQxLjMzMzMzNHogbTY4Mi42NjY2NjYgMzQxLjMzMzMzNGgtMTcwLjY2NjY2NnYtMjU2aDE3MC42NjY2NjZ2MjU2eiIgcC1pZD0iNDM0NyIgZmlsbD0iIzhhOGE4YSI+PC9wYXRoPjwvc3ZnPg=='></a> &nbsp;
<br>
</div>
## 📖 Introduction
This repository open-sources the **VideoReward** component -- our VLM-based reward model introduced in the paper [Improving Video Generation with Human Feedback](https://arxiv.org/abs/2501.13918). For Flow-DPO, we provide an implementation for text-to-image tasks [here](https://github.com/yifan123/flow_grpo/blob/main/scripts/single_node/dpo.sh).
VideoReward evaluates generated videos across three critical dimensions:
* Visual Quality (VQ): The clarity, aesthetics, and single-frame reasonableness.
* Motion Quality (MQ): The dynamic stability, dynamic reasonableness, naturalness, and dynamic degress.
* Text Alignment (TA): The relevance between the generated video and the text prompt.
This versatile reward model can be used for data filtering, guidance, reject sampling, DPO, and other RL methods. <br>
<img src=https://gongyeliu.github.io/videoalign/pics/overview.png width="100%"/>
## 📝 Updates
- __[2025.08.14]__: 🔥 We provide the prompt sets used to evaluate video generation performance in this paper, including VBench, VideoGen-Eval, and TA-Hard. See [`./datasets/video_eval_prompts`](./datasets/video_eval_prompts/README.md) for details.
- __[2025.07.17]__: 🔥 Release the [Flow-DPO](https://github.com/yifan123/flow_grpo/blob/main/scripts/single_node/dpo.sh).
- __[2025.02.08]__: 🔥 Release the [VideoGen-RewardBench](https://huggingface.co/datasets/KwaiVGI/VideoGen-RewardBench) and [Leaderboard](https://huggingface.co/spaces/KwaiVGI/VideoGen-RewardBench).
- __[2025.02.08]__: 🔥 Release the [Code](#) and [Checkpoints](https://huggingface.co/KwaiVGI/VideoReward) of VideoReward.
- __[2025.01.23]__: Release the [Paper](https://arxiv.org/abs/2501.13918) and [Project Page](https://gongyeliu.github.io/videoalign/).
## 🚀 Quick Started
### 1. Environment Set Up
Clone this repository and install packages.
```bash
git clone https://github.com/KwaiVGI/VideoAlign
cd VideoAlign
conda env create -f environment.yaml
conda activate VideoReward
pip install flash-attn==2.5.8 --no-build-isolation
```
### 2. Download Pretrained Weights
Please download our checkpoints from [Huggingface](https://huggingface.co/KwaiVGI/VideoReward) and put it in `./checkpoints/`.
```bash
cd checkpoints
git lfs install
git clone https://huggingface.co/KwaiVGI/VideoReward
cd ..
```
### 3. Scoring for a single prompt-video item.
```bash
python inference.py
```
## ✨ Eval the Performance on VideoGen-RewardBench
### 1. Download the VideoGen-RewardBench and put it in `./datasets/`.
```bash
cd dataset
git lfs install
git clone https://huggingface.co/datasets/KwaiVGI/VideoGen-RewardBench
cd ..
```
### 2. Start inference
```bash
python eval_videogen_rewardbench.py
```
## 🏁 Train RM on Your Own Data
### 1. Prepare your own data as the [instruction](./datasets/train/README.md) stated.
### 2. Start training!
```bash
sh train.sh
```
## 🤗 Acknowledgments
Our reward model is based on [QWen2-VL-2B-Instruct](https://huggingface.co/Qwen/Qwen2-VL-2B-Instruct), and our code is build upon [TRL](https://github.com/huggingface/trl) and [Qwen2-VL-Finetune](https://github.com/2U1/Qwen2-VL-Finetune), thanks to all the contributors!
## ⭐ Citation
Please leave us a star ⭐ if you find our work helpful.
```bibtex
@article{liu2025improving,
title={Improving video generation with human feedback},
author={Liu, Jie and Liu, Gongye and Liang, Jiajun and Yuan, Ziyang and Liu, Xiaokun and Zheng, Mingwu and Wu, Xiele and Wang, Qiulin and Qin, Wenyu and Xia, Menghan and others},
journal={arXiv preprint arXiv:2501.13918},
year={2025}
}
```
```bibtex
@article{liu2025flow,
title={Flow-grpo: Training flow matching models via online rl},
author={Liu, Jie and Liu, Gongye and Liang, Jiajun and Li, Yangguang and Liu, Jiaheng and Wang, Xintao and Wan, Pengfei and Zhang, Di and Ouyang, Wanli},
journal={arXiv preprint arXiv:2505.05470},
year={2025}
}
```
+18
View File
@@ -0,0 +1,18 @@
from dataclasses import dataclass
@dataclass
class DataConfig:
meta_data: str = "/path/to/dataset/meta_data.csv"
data_dir: str = "/path/to/dataset"
meta_data_test: str = None
max_frame_pixels: int = 240 * 320
num_frames: float = None
fps: float = 2.0
p_shuffle_frames: float = 0.0
p_color_jitter: float = 0.0
eval_dim: str | list[str] = "VQ"
prompt_template_type: str = "none"
add_noise: bool = False
sample_type: str = "uniform"
use_tied_data: bool = True
+238
View File
@@ -0,0 +1,238 @@
import json
import os
from collections.abc import Mapping
import torch
from data import DataConfig
from prompt_template import build_prompt
from runtime import ModelConfig, PEFTLoraConfig, TrainingConfig, create_model_and_processor, load_model_from_checkpoint
from vision_process import process_vision_info
def load_configs_from_json(config_path):
with open(config_path) as f:
config_dict = json.load(f)
# del config_dict["training_args"]["_n_gpu"]
del config_dict["data_config"]["meta_data"]
del config_dict["data_config"]["data_dir"]
return (
config_dict["data_config"],
None,
config_dict["model_config"],
config_dict["peft_lora_config"],
config_dict["inference_config"] if "inference_config" in config_dict else None,
)
class VideoVLMRewardInference:
def __init__(
self,
load_from_pretrained,
load_from_pretrained_step=-1,
device="cuda",
dtype=torch.bfloat16,
):
config_path = os.path.join(load_from_pretrained, "model_config.json")
(
data_config,
_,
model_config,
peft_lora_config,
inference_config,
) = load_configs_from_json(config_path)
data_config = DataConfig(**data_config)
model_config = ModelConfig(**model_config)
peft_lora_config = PEFTLoraConfig(**peft_lora_config)
training_args = TrainingConfig(
load_from_pretrained=load_from_pretrained,
load_from_pretrained_step=load_from_pretrained_step,
gradient_checkpointing=False,
disable_flash_attn2=False,
bf16=True if dtype == torch.bfloat16 else False,
fp16=True if dtype == torch.float16 else False,
output_dir="",
)
model, processor, peft_config = create_model_and_processor(
model_config=model_config,
peft_lora_config=peft_lora_config,
training_args=training_args,
)
self.device = device
model, checkpoint_step = load_model_from_checkpoint(model, load_from_pretrained, load_from_pretrained_step)
model.eval()
self.model = model
self.processor = processor
self.model.to(self.device)
self.data_config = data_config
self.inference_config = inference_config
def _norm(self, reward):
if self.inference_config is None:
return reward
reward["VQ"] = (reward["VQ"] - self.inference_config["VQ_mean"]) / self.inference_config["VQ_std"]
reward["MQ"] = (reward["MQ"] - self.inference_config["MQ_mean"]) / self.inference_config["MQ_std"]
reward["TA"] = (reward["TA"] - self.inference_config["TA_mean"]) / self.inference_config["TA_std"]
return reward
def _pad_sequence(self, sequences, attention_mask, max_len, padding_side="right"):
"""
Pad the sequences to the maximum length.
"""
assert padding_side in ["right", "left"]
if sequences.shape[1] >= max_len:
return sequences, attention_mask
pad_len = max_len - sequences.shape[1]
padding = (0, pad_len) if padding_side == "right" else (pad_len, 0)
sequences_padded = torch.nn.functional.pad(
sequences, padding, "constant", self.processor.tokenizer.pad_token_id
)
attention_mask_padded = torch.nn.functional.pad(attention_mask, padding, "constant", 0)
return sequences_padded, attention_mask_padded
def _prepare_input(self, data):
"""
Prepare `inputs` before feeding them to the model, converting them to tensors if they are not already and
handling potential state.
"""
if isinstance(data, Mapping):
return type(data)({k: self._prepare_input(v) for k, v in data.items()})
if isinstance(data, (tuple, list)):
return type(data)(self._prepare_input(v) for v in data)
if isinstance(data, torch.Tensor):
kwargs = {"device": self.device}
return data.to(**kwargs)
return data
def _prepare_inputs(self, inputs):
"""
Prepare `inputs` before feeding them to the model, converting them to tensors if they are not already and
handling potential state.
"""
inputs = self._prepare_input(inputs)
if len(inputs) == 0:
raise ValueError
return inputs
def prepare_batch(
self,
video_paths,
prompts,
fps=None,
num_frames=None,
max_pixels=None,
):
fps = self.data_config.fps if fps is None else fps
num_frames = self.data_config.num_frames if num_frames is None else num_frames
max_pixels = self.data_config.max_frame_pixels if max_pixels is None else max_pixels
if num_frames is None:
chat_data = [
[
{
"role": "user",
"content": [
{
"type": "video",
"video": f"file://{video_path}",
"max_pixels": max_pixels,
"fps": fps,
"sample_type": self.data_config.sample_type,
},
{
"type": "text",
"text": build_prompt(
prompt,
self.data_config.eval_dim,
self.data_config.prompt_template_type,
),
},
],
},
]
for video_path, prompt in zip(video_paths, prompts)
]
else:
chat_data = [
[
{
"role": "user",
"content": [
{
"type": "video",
"video": f"file://{video_path}",
"max_pixels": max_pixels,
"nframes": num_frames,
"sample_type": self.data_config.sample_type,
},
{
"type": "text",
"text": build_prompt(
prompt,
self.data_config.eval_dim,
self.data_config.prompt_template_type,
),
},
],
},
]
for video_path, prompt in zip(video_paths, prompts)
]
image_inputs, video_inputs = process_vision_info(chat_data)
batch = self.processor(
text=self.processor.apply_chat_template(chat_data, tokenize=False, add_generation_prompt=True),
images=image_inputs,
videos=video_inputs,
padding=True,
return_tensors="pt",
videos_kwargs={"do_rescale": True},
)
batch = self._prepare_inputs(batch)
return batch
def reward(
self,
video_paths,
prompts,
fps=None,
num_frames=None,
max_pixels=None,
use_norm=True,
):
"""
Inputs:
video_paths: List[str], B paths of the videos.
prompts: List[str], B prompts for the videos.
eval_dims: List[str], N evaluation dimensions.
fps: float, sample rate of the videos. If None, use the default value in the config.
num_frames: int, number of frames of the videos. If None, use the default value in the config.
max_pixels: int, maximum pixels of the videos. If None, use the default value in the config.
use_norm: bool, whether to rescale the output rewards
Outputs:
Rewards: List[dict], N + 1 rewards of the B videos.
"""
assert fps is None or num_frames is None, "fps and num_frames cannot be set at the same time."
batch = self.prepare_batch(video_paths, prompts, fps, num_frames, max_pixels)
rewards = self.model(return_dict=True, **batch)["logits"]
rewards = [{"VQ": reward[0].item(), "MQ": reward[1].item(), "TA": reward[2].item()} for reward in rewards]
for i in range(len(rewards)):
if use_norm:
rewards[i] = self._norm(rewards[i])
rewards[i]["Overall"] = rewards[i]["VQ"] + rewards[i]["MQ"] + rewards[i]["TA"]
return rewards
@@ -0,0 +1,141 @@
VIDEOSCORE_QUERY_PROMPT = """
Suppose you are an expert in judging and evaluating the quality of AI-generated videos,
please watch the frames of a given video and see the text prompt for generating the video,
then give scores based on its {dimension_name}, i.e., {dimension_description}.
Output a float number from 1.0 to 5.0 for this dimension,
the higher the number is, the better the video performs in that sub-score,
the lowest 1.0 means Bad, the highest 5.0 means Perfect/Real (the video is like a real video).
The text prompt used for generation is "{text_prompt}".
"""
DIMENSION_DESCRIPTIONS = {
"VQ": [
"visual quality",
"the quality of the video in terms of clearness, resolution, brightness, and color",
],
"TA": [
"text-to-video alignment",
"the alignment between the text prompt and the video content and motion",
],
"MQ": [
"motion quality",
"the quality of the motion in terms of consistency, smoothness, and completeness",
],
"Overall": [
"Overall Performance",
"the overall performance of the video in terms of visual quality, text-to-video alignment, and motion quality",
],
}
SIMPLE_PROMPT = """
Please evaluate the {dimension_name} of a generated video. Consider {dimension_description}.
The text prompt used for generation is "{text_prompt}".
"""
DETAILED_PROMPT_WITH_SPECIAL_TOKEN = """
You are tasked with evaluating a generated video based on three distinct criteria: Visual Quality, Motion Quality, and Text Alignment. Please provide a rating from 0 to 10 for each of the three categories, with 0 being the worst and 10 being the best. Each evaluation should be independent of the others.
**Visual Quality:**
Evaluate the overall visual quality of the video, with a focus on static factors. The following sub-dimensions should be considered:
- **Reasonableness:** The video should not contain any significant biological or logical errors, such as abnormal body structures or nonsensical environmental setups.
- **Clarity:** Evaluate the sharpness and visibility of the video. The image should be clear and easy to interpret, with no blurring or indistinct areas.
- **Detail Richness:** Consider the level of detail in textures, materials, lighting, and other visual elements (e.g., hair, clothing, shadows).
- **Aesthetic and Creativity:** Assess the artistic aspects of the video, including the color scheme, composition, atmosphere, depth of field, and the overall creative appeal. The scene should convey a sense of harmony and balance.
- **Safety:** The video should not contain harmful or inappropriate content, such as political, violent, or adult material. If such content is present, the image quality and satisfaction score should be the lowest possible.
Please provide the ratings of Visual Quality: <|VQ_reward|>
END
**Motion Quality:**
Assess the dynamic aspects of the video, with a focus on dynamic factors. Consider the following sub-dimensions:
- **Stability:** Evaluate the continuity and stability between frames. There should be no sudden, unnatural jumps, and the video should maintain stable attributes (e.g., no fluctuating colors, textures, or missing body parts).
- **Naturalness:** The movement should align with physical laws and be realistic. For example, clothing should flow naturally with motion, and facial expressions should change appropriately (e.g., blinking, mouth movements).
- **Aesthetic Quality:** The movement should be smooth and fluid. The transitions between different motions or camera angles should be seamless, and the overall dynamic feel should be visually pleasing.
- **Fusion:** Ensure that elements in motion (e.g., edges of the subject, hair, clothing) blend naturally with the background, without obvious artifacts or the feeling of cut-and-paste effects.
- **Clarity of Motion:** The video should be clear and smooth in motion. Pay attention to any areas where the video might have blurry or unsteady sections that hinder visual continuity.
- **Amplitude:** If the video is largely static or has little movement, assign a low score for motion quality.
Please provide the ratings of Motion Quality: <|MQ_reward|>
END
**Text Alignment:**
Assess how well the video matches the textual prompt across the following sub-dimensions:
- **Subject Relevance** Evaluate how accurately the subject(s) in the video (e.g., person, animal, object) align with the textual description. The subject should match the description in terms of number, appearance, and behavior.
- **Motion Relevance:** Evaluate if the dynamic actions (e.g., gestures, posture, facial expressions like talking or blinking) align with the described prompt. The motion should match the prompt in terms of type, scale, and direction.
- **Environment Relevance:** Assess whether the background and scene fit the prompt. This includes checking if real-world locations or scenes are accurately represented, though some stylistic adaptation is acceptable.
- **Style Relevance:** If the prompt specifies a particular artistic or stylistic style, evaluate how well the video adheres to this style.
- **Camera Movement Relevance:** Check if the camera movements (e.g., following the subject, focus shifts) are consistent with the expected behavior from the prompt.
Textual prompt - {text_prompt}
Please provide the ratings of Text Alignment: <|TA_reward|>
END
"""
DETAILED_PROMPT = """
You are tasked with evaluating a generated video based on three distinct criteria: Visual Quality, Motion Quality, and Text Alignment. Please provide a rating from 0 to 10 for each of the three categories, with 0 being the worst and 10 being the best. Each evaluation should be independent of the others.
**Visual Quality:**
Evaluate the overall visual quality of the video, with a focus on static factors. The following sub-dimensions should be considered:
- **Reasonableness:** The video should not contain any significant biological or logical errors, such as abnormal body structures or nonsensical environmental setups.
- **Clarity:** Evaluate the sharpness and visibility of the video. The image should be clear and easy to interpret, with no blurring or indistinct areas.
- **Detail Richness:** Consider the level of detail in textures, materials, lighting, and other visual elements (e.g., hair, clothing, shadows).
- **Aesthetic and Creativity:** Assess the artistic aspects of the video, including the color scheme, composition, atmosphere, depth of field, and the overall creative appeal. The scene should convey a sense of harmony and balance.
- **Safety:** The video should not contain harmful or inappropriate content, such as political, violent, or adult material. If such content is present, the image quality and satisfaction score should be the lowest possible.
**Motion Quality:**
Assess the dynamic aspects of the video, with a focus on dynamic factors. Consider the following sub-dimensions:
- **Stability:** Evaluate the continuity and stability between frames. There should be no sudden, unnatural jumps, and the video should maintain stable attributes (e.g., no fluctuating colors, textures, or missing body parts).
- **Naturalness:** The movement should align with physical laws and be realistic. For example, clothing should flow naturally with motion, and facial expressions should change appropriately (e.g., blinking, mouth movements).
- **Aesthetic Quality:** The movement should be smooth and fluid. The transitions between different motions or camera angles should be seamless, and the overall dynamic feel should be visually pleasing.
- **Fusion:** Ensure that elements in motion (e.g., edges of the subject, hair, clothing) blend naturally with the background, without obvious artifacts or the feeling of cut-and-paste effects.
- **Clarity of Motion:** The video should be clear and smooth in motion. Pay attention to any areas where the video might have blurry or unsteady sections that hinder visual continuity.
- **Amplitude:** If the video is largely static or has little movement, assign a low score for motion quality.
**Text Alignment:**
Assess how well the video matches the textual prompt across the following sub-dimensions:
- **Subject Relevance** Evaluate how accurately the subject(s) in the video (e.g., person, animal, object) align with the textual description. The subject should match the description in terms of number, appearance, and behavior.
- **Motion Relevance:** Evaluate if the dynamic actions (e.g., gestures, posture, facial expressions like talking or blinking) align with the described prompt. The motion should match the prompt in terms of type, scale, and direction.
- **Environment Relevance:** Assess whether the background and scene fit the prompt. This includes checking if real-world locations or scenes are accurately represented, though some stylistic adaptation is acceptable.
- **Style Relevance:** If the prompt specifies a particular artistic or stylistic style, evaluate how well the video adheres to this style.
- **Camera Movement Relevance:** Check if the camera movements (e.g., following the subject, focus shifts) are consistent with the expected behavior from the prompt.
Textual prompt - {text_prompt}
Please provide the ratings of Visual Quality, Motion Quality, and Text Alignment.
"""
SIMPLE_PROMPT_NO_PROMPT = """
Please evaluate the {dimension_name} of a generated video. Consider {dimension_description}.
"""
def build_prompt(prompt, dimension, template_type):
if isinstance(dimension, list) and len(dimension) > 1:
dimension_name = ", ".join([DIMENSION_DESCRIPTIONS[d][0] for d in dimension])
dimension_name = f"overall performance({dimension_name})"
dimension_description = "the overall performance of the video"
else:
if isinstance(dimension, list):
dimension = dimension[0]
dimension_name = DIMENSION_DESCRIPTIONS[dimension][0]
dimension_description = DIMENSION_DESCRIPTIONS[dimension][1]
if template_type == "none":
return prompt
if template_type == "simple":
return SIMPLE_PROMPT.format(
dimension_name=dimension_name,
dimension_description=dimension_description,
text_prompt=prompt,
)
if template_type == "video_score":
return VIDEOSCORE_QUERY_PROMPT.format(
dimension_name=dimension_name,
dimension_description=dimension_description,
text_prompt=prompt,
)
if template_type == "detailed_special":
return DETAILED_PROMPT_WITH_SPECIAL_TOKEN.format(text_prompt=prompt)
if template_type == "detailed":
return DETAILED_PROMPT.format(text_prompt=prompt)
raise ValueError("Invalid template type")
@@ -0,0 +1,117 @@
import torch
from torch import nn
from transformers import Qwen2VLForConditionalGeneration
class Qwen2VLRewardModelBT(Qwen2VLForConditionalGeneration):
def __init__(self, config, output_dim=4, reward_token="last", special_token_ids=None, **kwargs):
del kwargs
super().__init__(config)
self.output_dim = output_dim
hidden_size = getattr(config, "hidden_size", None)
if hidden_size is None and hasattr(config, "text_config"):
hidden_size = getattr(config.text_config, "hidden_size", None)
if hidden_size is None:
raise AttributeError("Qwen2VL reward model config must define hidden_size or text_config.hidden_size")
self.rm_head = nn.Linear(hidden_size, output_dim, bias=False)
self.reward_token = reward_token
self.special_token_ids = special_token_ids
if self.special_token_ids is not None:
self.reward_token = "special"
def forward(
self,
input_ids: torch.LongTensor = None,
attention_mask: torch.Tensor | None = None,
position_ids: torch.LongTensor | None = None,
past_key_values: list[torch.FloatTensor] | None = None,
inputs_embeds: torch.FloatTensor | None = None,
labels: torch.LongTensor | None = None,
use_cache: bool | None = None,
output_attentions: bool | None = None,
output_hidden_states: bool | None = None,
return_dict: bool | None = None,
pixel_values: torch.Tensor | None = None,
pixel_values_videos: torch.FloatTensor | None = None,
image_grid_thw: torch.LongTensor | None = None,
video_grid_thw: torch.LongTensor | None = None,
rope_deltas: torch.LongTensor | None = None,
**kwargs,
):
del kwargs, labels, rope_deltas
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if inputs_embeds is None:
inputs_embeds = self.model.embed_tokens(input_ids)
if pixel_values is not None:
pixel_values = pixel_values.type(self.visual.get_dtype())
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
image_mask = (input_ids == self.config.image_token_id).unsqueeze(-1).expand_as(inputs_embeds)
image_embeds = image_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)
if pixel_values_videos is not None:
pixel_values_videos = pixel_values_videos.type(self.visual.get_dtype())
video_embeds = self.visual(pixel_values_videos, grid_thw=video_grid_thw)
video_mask = (input_ids == self.config.video_token_id).unsqueeze(-1).expand_as(inputs_embeds)
video_embeds = video_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)
if attention_mask is not None:
attention_mask = attention_mask.to(inputs_embeds.device)
outputs = self.model(
input_ids=None,
position_ids=position_ids,
attention_mask=attention_mask,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
hidden_states = outputs[0]
logits = self.rm_head(hidden_states)
if input_ids is not None:
batch_size = input_ids.shape[0]
else:
batch_size = inputs_embeds.shape[0]
if self.config.pad_token_id is None and batch_size != 1:
raise ValueError("Cannot handle batch sizes > 1 if no padding token is defined.")
if self.config.pad_token_id is None:
sequence_lengths = -1
elif input_ids is not None:
sequence_lengths = torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1
sequence_lengths = sequence_lengths % input_ids.shape[-1]
sequence_lengths = sequence_lengths.to(logits.device)
else:
sequence_lengths = -1
if self.reward_token == "last":
pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths]
elif self.reward_token == "mean":
valid_lengths = torch.clamp(sequence_lengths, min=0, max=logits.size(1) - 1)
pooled_logits = torch.stack([logits[i, :valid_lengths[i]].mean(dim=0) for i in range(batch_size)])
elif self.reward_token == "special":
special_token_mask = torch.zeros_like(input_ids, dtype=torch.bool)
for special_token_id in self.special_token_ids:
special_token_mask = special_token_mask | (input_ids == special_token_id)
pooled_logits = logits[special_token_mask, ...]
pooled_logits = pooled_logits.view(batch_size, 3, -1)
if self.output_dim == 3:
pooled_logits = pooled_logits.diagonal(dim1=1, dim2=2)
pooled_logits = pooled_logits.view(batch_size, -1)
else:
raise ValueError("Invalid reward_token")
return {"logits": pooled_logits}
+247
View File
@@ -0,0 +1,247 @@
import glob
from importlib import util
import os
from dataclasses import dataclass, field
from typing import Literal
import safetensors
import torch
from peft import LoraConfig, get_peft_model
from reward_model import Qwen2VLRewardModelBT
from transformers import AutoProcessor, TrainingArguments
@dataclass
class TrainingConfig(TrainingArguments):
max_length: int | None = None
dataset_num_proc: int | None = None
center_rewards_coefficient: float | None = None
disable_flash_attn2: bool = field(default=False)
vision_lr: float | None = None
merger_lr: float | None = None
special_token_lr: float | None = None
conduct_eval: bool | None = True
load_from_pretrained: str = None
load_from_pretrained_step: int = None
logging_epochs: float | None = None
eval_epochs: float | None = None
save_epochs: float | None = None
remove_unused_columns: bool | None = False
save_full_model: bool | None = False
@dataclass
class PEFTLoraConfig:
lora_enable: bool = False
vision_lora: bool = False
lora_r: int = 16
lora_alpha: int = 32
lora_dropout: float = 0.05
lora_target_modules: list[str] | None = None
lora_namespan_exclude: list[str] | None = None
lora_modules_to_save: list[str] | None = None
lora_task_type: str = "CAUSAL_LM"
use_rslora: bool = False
num_lora_modules: int = -1
def __post_init__(self):
if isinstance(self.lora_target_modules, list) and len(self.lora_target_modules) == 1:
self.lora_target_modules = self.lora_target_modules[0]
if isinstance(self.lora_namespan_exclude, list) and len(self.lora_namespan_exclude) == 1:
self.lora_namespan_exclude = self.lora_namespan_exclude[0]
@dataclass
class ModelConfig:
model_name_or_path: str | None = None
model_revision: str = "main"
output_dim: int = 1
use_special_tokens: bool = False
freeze_vision_tower: bool = field(default=False)
freeze_llm: bool = field(default=False)
tune_merger: bool = field(default=False)
torch_dtype: Literal["auto", "bfloat16", "float16", "float32"] | None = None
trust_remote_code: bool = False
attn_implementation: str | None = None
load_in_8bit: bool = False
load_in_4bit: bool = False
bnb_4bit_quant_type: Literal["fp4", "nf4"] = "nf4"
use_bnb_nested_quant: bool = False
reward_token: Literal["last", "mean", "special"] = "last"
loss_type: Literal["bt", "reg", "btt", "margin", "constant_margin", "scaled"] = "regular"
def __post_init__(self):
if self.load_in_8bit and self.load_in_4bit:
raise ValueError("You can't use 8 bit and 4 bit precision at the same time")
def find_target_linear_names(model, num_lora_modules=-1, lora_namespan_exclude=None):
linear_cls = torch.nn.Linear
embedding_cls = torch.nn.Embedding
excluded = lora_namespan_exclude or []
lora_module_names = []
for name, module in model.named_modules():
if any(ex_keyword in name for ex_keyword in excluded):
continue
if isinstance(module, (linear_cls, embedding_cls)):
lora_module_names.append(name)
if num_lora_modules > 0:
lora_module_names = lora_module_names[-num_lora_modules:]
return lora_module_names
def _get_quantization_config(model_config):
if not model_config.load_in_8bit and not model_config.load_in_4bit:
return None
from transformers import BitsAndBytesConfig
if model_config.load_in_8bit:
return BitsAndBytesConfig(load_in_8bit=True)
return BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type=model_config.bnb_4bit_quant_type,
bnb_4bit_use_double_quant=model_config.use_bnb_nested_quant,
)
def create_model_and_processor(
model_config,
peft_lora_config,
training_args,
cache_dir=None,
):
torch_dtype = (
model_config.torch_dtype
if model_config.torch_dtype in ["auto", None]
else getattr(torch, model_config.torch_dtype)
)
quantization_config = _get_quantization_config(model_config)
model_kwargs = dict(
revision=model_config.model_revision,
device_map="auto" if quantization_config is not None else None,
quantization_config=quantization_config,
use_cache=True if training_args.gradient_checkpointing else False,
)
processor = AutoProcessor.from_pretrained(
model_config.model_name_or_path, padding_side="right", cache_dir=cache_dir
)
special_token_ids = None
if model_config.use_special_tokens:
special_tokens = ["<|VQ_reward|>", "<|MQ_reward|>", "<|TA_reward|>"]
processor.tokenizer.add_special_tokens({"additional_special_tokens": special_tokens})
special_token_ids = processor.tokenizer.convert_tokens_to_ids(special_tokens)
has_flash_attn = util.find_spec("flash_attn") is not None
model = Qwen2VLRewardModelBT.from_pretrained(
model_config.model_name_or_path,
output_dim=model_config.output_dim,
reward_token=model_config.reward_token,
special_token_ids=special_token_ids,
torch_dtype=torch_dtype,
attn_implementation=(
"flash_attention_2" if not training_args.disable_flash_attn2 and has_flash_attn else "sdpa"
),
cache_dir=cache_dir,
**model_kwargs,
)
if model_config.use_special_tokens:
model.resize_token_embeddings(len(processor.tokenizer))
if training_args.bf16:
model.to(torch.bfloat16)
if training_args.fp16:
model.to(torch.float16)
if peft_lora_config.lora_enable:
target_modules = find_target_linear_names(
model,
num_lora_modules=peft_lora_config.num_lora_modules,
lora_namespan_exclude=peft_lora_config.lora_namespan_exclude,
)
peft_config = LoraConfig(
target_modules=target_modules,
r=peft_lora_config.lora_r,
lora_alpha=peft_lora_config.lora_alpha,
lora_dropout=peft_lora_config.lora_dropout,
task_type=peft_lora_config.lora_task_type,
use_rslora=peft_lora_config.use_rslora,
bias="none",
modules_to_save=peft_lora_config.lora_modules_to_save,
)
model = get_peft_model(model, peft_config)
else:
peft_config = None
model.config.tokenizer_padding_side = processor.tokenizer.padding_side
model.config.pad_token_id = processor.tokenizer.pad_token_id
return model, processor, peft_config
def _insert_adapter_name_into_state_dict(
state_dict: dict[str, torch.Tensor], adapter_name: str, parameter_prefix: str
) -> dict[str, torch.Tensor]:
peft_model_state_dict = {}
for key, val in state_dict.items():
if parameter_prefix in key:
suffix = key.split(parameter_prefix)[1]
if "." in suffix:
suffix_to_replace = ".".join(suffix.split(".")[1:])
key = key.replace(suffix_to_replace, f"{adapter_name}.{suffix_to_replace}")
else:
key = f"{key}.{adapter_name}"
peft_model_state_dict[key] = val
else:
peft_model_state_dict[key] = val
return peft_model_state_dict
def load_model_from_checkpoint(model, checkpoint_dir, checkpoint_step):
checkpoint_paths = glob.glob(os.path.join(checkpoint_dir, "checkpoint-*"))
checkpoint_paths.sort(key=lambda x: int(x.split("-")[-1]), reverse=True)
if checkpoint_step is None or checkpoint_step == -1:
if checkpoint_paths:
checkpoint_path = checkpoint_paths[0]
else:
checkpoint_path = checkpoint_dir
else:
checkpoint_path = os.path.join(checkpoint_dir, f"checkpoint-{checkpoint_step}")
if checkpoint_path not in checkpoint_paths:
checkpoint_path = checkpoint_paths[0] if checkpoint_paths else checkpoint_dir
checkpoint_step = checkpoint_path.split("checkpoint-")[-1].split("/")[0]
full_ckpt = os.path.join(checkpoint_path, "model.pth")
lora_ckpt = os.path.join(checkpoint_path, "adapter_model.safetensors")
non_lora_ckpt = os.path.join(checkpoint_path, "non_lora_state_dict.pth")
if os.path.exists(full_ckpt):
model_state_dict = torch.load(full_ckpt, map_location="cpu")
model.load_state_dict(model_state_dict)
else:
lora_state_dict = safetensors.torch.load_file(lora_ckpt)
non_lora_state_dict = torch.load(non_lora_ckpt, map_location="cpu")
lora_state_dict = _insert_adapter_name_into_state_dict(
lora_state_dict, adapter_name="default", parameter_prefix="lora_"
)
model_state_dict = model.state_dict()
model_state_dict.update(non_lora_state_dict)
model_state_dict.update(lora_state_dict)
model.load_state_dict(model_state_dict)
return model, checkpoint_step
@@ -0,0 +1,426 @@
## This file is modified from https://github.com/kq-chen/qwen-vl-utils/blob/main/src/qwen_vl_utils/vision_process.py
from __future__ import annotations
import base64
import logging
import math
import os
import sys
import time
import warnings
from functools import lru_cache
from io import BytesIO
import requests
import torch
import torchvision
from packaging import version
from PIL import Image
from torchvision import io, transforms
from torchvision.transforms import InterpolationMode
logger = logging.getLogger(__name__)
IMAGE_FACTOR = 28
MIN_PIXELS = 4 * 28 * 28
MAX_PIXELS = 16384 * 28 * 28
MAX_RATIO = 200
VIDEO_MIN_PIXELS = 128 * 28 * 28
VIDEO_MAX_PIXELS = 768 * 28 * 28
VIDEO_TOTAL_PIXELS = 24576 * 28 * 28
FRAME_FACTOR = 2
FPS = 2.0
FPS_MIN_FRAMES = 4
FPS_MAX_FRAMES = 768
def round_by_factor(number: int, factor: int) -> int:
"""Returns the closest integer to 'number' that is divisible by 'factor'."""
return round(number / factor) * factor
def ceil_by_factor(number: int, factor: int) -> int:
"""Returns the smallest integer greater than or equal to 'number' that is divisible by 'factor'."""
return math.ceil(number / factor) * factor
def floor_by_factor(number: int, factor: int) -> int:
"""Returns the largest integer less than or equal to 'number' that is divisible by 'factor'."""
return math.floor(number / factor) * factor
def smart_resize(
height: int,
width: int,
factor: int = IMAGE_FACTOR,
min_pixels: int = MIN_PIXELS,
max_pixels: int = MAX_PIXELS,
) -> tuple[int, int]:
"""
Rescales the image so that the following conditions are met:
1. Both dimensions (height and width) are divisible by 'factor'.
2. The total number of pixels is within the range ['min_pixels', 'max_pixels'].
3. The aspect ratio of the image is maintained as closely as possible.
"""
if max(height, width) / min(height, width) > MAX_RATIO:
raise ValueError(
f"absolute aspect ratio must be smaller than {MAX_RATIO}, got {max(height, width) / min(height, width)}"
)
h_bar = max(factor, round_by_factor(height, factor))
w_bar = max(factor, round_by_factor(width, factor))
if h_bar * w_bar > max_pixels:
beta = math.sqrt((height * width) / max_pixels)
h_bar = floor_by_factor(height / beta, factor)
w_bar = floor_by_factor(width / beta, factor)
elif h_bar * w_bar < min_pixels:
beta = math.sqrt(min_pixels / (height * width))
h_bar = ceil_by_factor(height * beta, factor)
w_bar = ceil_by_factor(width * beta, factor)
return h_bar, w_bar
def fetch_image(ele: dict[str, str | Image.Image], size_factor: int = IMAGE_FACTOR) -> Image.Image:
if "image" in ele:
image = ele["image"]
else:
image = ele["image_url"]
image_obj = None
if isinstance(image, Image.Image):
image_obj = image
elif image.startswith("http://") or image.startswith("https://"):
image_obj = Image.open(requests.get(image, stream=True).raw)
elif image.startswith("file://"):
image_obj = Image.open(image[7:])
elif image.startswith("data:image"):
if "base64," in image:
_, base64_data = image.split("base64,", 1)
data = base64.b64decode(base64_data)
image_obj = Image.open(BytesIO(data))
else:
image_obj = Image.open(image)
if image_obj is None:
raise ValueError(f"Unrecognized image input, support local path, http url, base64 and PIL.Image, got {image}")
image = image_obj.convert("RGB")
## resize
if "resized_height" in ele and "resized_width" in ele:
resized_height, resized_width = smart_resize(
ele["resized_height"],
ele["resized_width"],
factor=size_factor,
)
else:
width, height = image.size
min_pixels = ele.get("min_pixels", MIN_PIXELS)
max_pixels = ele.get("max_pixels", MAX_PIXELS)
resized_height, resized_width = smart_resize(
height,
width,
factor=size_factor,
min_pixels=min_pixels,
max_pixels=max_pixels,
)
image = image.resize((resized_width, resized_height))
return image
def smart_nframes(
ele: dict,
total_frames: int,
video_fps: float,
) -> int:
"""calculate the number of frames for video used for model inputs.
Args:
ele (dict): a dict contains the configuration of video.
support either `fps` or `nframes`:
- nframes: the number of frames to extract for model inputs.
- fps: the fps to extract frames for model inputs.
- min_frames: the minimum number of frames of the video, only used when fps is provided.
- max_frames: the maximum number of frames of the video, only used when fps is provided.
total_frames (int): the original total number of frames of the video.
video_fps (int | float): the original fps of the video.
Raises:
ValueError: nframes should in interval [FRAME_FACTOR, total_frames].
Returns:
int: the number of frames for video used for model inputs.
"""
assert not ("fps" in ele and "nframes" in ele), "Only accept either `fps` or `nframes`"
if "nframes" in ele:
nframes = round_by_factor(ele["nframes"], FRAME_FACTOR)
else:
fps = ele.get("fps", FPS)
min_frames = ceil_by_factor(ele.get("min_frames", FPS_MIN_FRAMES), FRAME_FACTOR)
max_frames = floor_by_factor(ele.get("max_frames", min(FPS_MAX_FRAMES, total_frames)), FRAME_FACTOR)
nframes = total_frames / video_fps * fps
nframes = min(max(nframes, min_frames), max_frames)
nframes = round_by_factor(nframes, FRAME_FACTOR)
nframes = min(nframes, total_frames)
if not (nframes >= FRAME_FACTOR and nframes <= total_frames):
raise ValueError(f"nframes should in interval [{FRAME_FACTOR}, {total_frames}], but got {nframes}.")
return nframes
def _get_video_fps_fallback(video_path: str) -> float:
"""Get video fps using PyAV or OpenCV as fallback when torchvision info doesn't have it."""
try:
# Try PyAV first (since torchvision uses pyav backend)
import av
container = av.open(video_path)
video_stream = container.streams.video[0]
fps = float(video_stream.average_rate)
container.close()
if fps > 0:
return fps
except Exception:
pass
try:
# Try OpenCV as backup
import cv2
cap = cv2.VideoCapture(video_path)
fps = cap.get(cv2.CAP_PROP_FPS)
cap.release()
if fps > 0:
return float(fps)
except Exception:
pass
logger.error("Error getting video fps using PyAV or OpenCV, using default fallback 30.0 fps.")
# Default fallback
return 30.0
def _read_video_torchvision(
ele: dict,
) -> torch.Tensor:
"""read video using torchvision.io.read_video
Args:
ele (dict): a dict contains the configuration of video.
support keys:
- video: the path of video. support "file://", "http://", "https://" and local path.
- video_start: the start time of video.
- video_end: the end time of video.
Returns:
torch.Tensor: the video tensor with shape (T, C, H, W).
"""
video_path = ele["video"]
# Remove file:// prefix - torchvision doesn't support it (especially for relative paths)
if video_path.startswith("file://"):
video_path = video_path[7:]
if version.parse(torchvision.__version__) < version.parse("0.19.0"):
if "http://" in video_path or "https://" in video_path:
warnings.warn("torchvision < 0.19.0 does not support http/https video path, please upgrade to 0.19.0.")
st = time.time()
video, audio, info = io.read_video(
video_path,
start_pts=ele.get("video_start", 0.0),
end_pts=ele.get("video_end"),
pts_unit="sec",
output_format="TCHW",
)
total_frames = video.size(0)
if total_frames == 0:
raise ValueError(
f"No frames were read from video: {video_path}. "
f"This may be caused by invalid video_start ({ele.get('video_start', 0.0)}) "
f"or video_end ({ele.get('video_end')}) parameters, or the video file may be corrupted."
)
# Try to get video_fps from info, use fallback methods if not available
if "video_fps" in info:
video_fps = info["video_fps"]
else:
# Fallback: use PyAV or OpenCV to get real fps
video_fps = _get_video_fps_fallback(video_path)
# logger.info(f"torchvision: {video_path=}, {total_frames=}, {video_fps=}, time={time.time() - st:.3f}s")
if ele["sample_type"] == "uniform":
nframes = smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
elif ele["sample_type"] == "multi_pts":
frames_each_pts = 6
num_pts = 4
fps = 8
nframes = int(total_frames * fps // video_fps)
frames_idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
start_pt = int(frames_each_pts // 2)
end_pt = int(nframes - frames_each_pts // 2 - 1)
pts = torch.linspace(start_pt, end_pt, num_pts).round().long().tolist()
idx = []
for pt in pts:
idx.extend(frames_idx[pt - frames_each_pts // 2 : pt + frames_each_pts // 2])
video = video[idx]
return video
def is_decord_available() -> bool:
import importlib.util
return importlib.util.find_spec("decord") is not None
def _read_video_decord(
ele: dict,
) -> torch.Tensor:
"""read video using decord.VideoReader
Args:
ele (dict): a dict contains the configuration of video.
support keys:
- video: the path of video. support "file://", "http://", "https://" and local path.
- video_start: the start time of video.
- video_end: the end time of video.
Returns:
torch.Tensor: the video tensor with shape (T, C, H, W).
"""
import decord
video_path = ele["video"]
st = time.time()
vr = decord.VideoReader(video_path)
# TODO: support start_pts and end_pts
if "video_start" in ele or "video_end" in ele:
raise NotImplementedError("not support start_pts and end_pts in decord for now.")
total_frames, video_fps = len(vr), vr.get_avg_fps()
# logger.info(f"decord: {video_path=}, {total_frames=}, {video_fps=}, time={time.time() - st:.3f}s")
if ele["sample_type"] == "uniform":
nframes = smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
# nframes = max(nframes, 8)
# import pdb; pdb.set_trace()
idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
elif ele["sample_type"] == "multi_pts":
frames_each_pts = 6
num_pts = 4
fps = 8
nframes = int(total_frames * fps // video_fps)
frames_idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
start_pt = int(frames_each_pts // 2)
end_pt = int(nframes - frames_each_pts // 2 - 1)
pts = torch.linspace(start_pt, end_pt, num_pts).round().long().tolist()
idx = []
for pt in pts:
idx.extend(frames_idx[pt - frames_each_pts // 2 : pt + frames_each_pts // 2])
video = vr.get_batch(idx).asnumpy()
video = torch.tensor(video).permute(0, 3, 1, 2) # Convert to TCHW format
return video
VIDEO_READER_BACKENDS = {
"decord": _read_video_decord,
"torchvision": _read_video_torchvision,
}
FORCE_QWENVL_VIDEO_READER = os.getenv("FORCE_QWENVL_VIDEO_READER", None)
@lru_cache(maxsize=1)
def get_video_reader_backend() -> str:
if FORCE_QWENVL_VIDEO_READER is not None:
video_reader_backend = FORCE_QWENVL_VIDEO_READER
elif is_decord_available():
video_reader_backend = "decord"
else:
video_reader_backend = "torchvision"
print(f"qwen-vl-utils using {video_reader_backend} to read video.", file=sys.stderr)
return video_reader_backend
def fetch_video(ele: dict, image_factor: int = IMAGE_FACTOR) -> torch.Tensor | list[Image.Image]:
if isinstance(ele["video"], str):
video_reader_backend = get_video_reader_backend()
video = VIDEO_READER_BACKENDS[video_reader_backend](ele)
# import pdb; pdb.set_trace()
nframes, _, height, width = video.shape
min_pixels = ele.get("min_pixels", VIDEO_MIN_PIXELS)
total_pixels = ele.get("total_pixels", VIDEO_TOTAL_PIXELS)
max_pixels = max(
min(VIDEO_MAX_PIXELS, total_pixels / nframes * FRAME_FACTOR),
int(min_pixels * 1.05),
)
max_pixels = ele.get("max_pixels", max_pixels)
if "resized_height" in ele and "resized_width" in ele:
resized_height, resized_width = smart_resize(
ele["resized_height"],
ele["resized_width"],
factor=image_factor,
)
else:
resized_height, resized_width = smart_resize(
height,
width,
factor=image_factor,
min_pixels=min_pixels,
max_pixels=max_pixels,
)
video = transforms.functional.resize(
video,
[resized_height, resized_width],
interpolation=InterpolationMode.BICUBIC,
antialias=True,
).float()
return video
assert isinstance(ele["video"], (list, tuple))
process_info = ele.copy()
process_info.pop("type", None)
process_info.pop("video", None)
images = [
fetch_image({"image": video_element, **process_info}, size_factor=image_factor)
for video_element in ele["video"]
]
nframes = ceil_by_factor(len(images), FRAME_FACTOR)
if len(images) < nframes:
images.extend([images[-1]] * (nframes - len(images)))
return images
def extract_vision_info(conversations: list[dict] | list[list[dict]]) -> list[dict]:
vision_infos = []
if isinstance(conversations[0], dict):
conversations = [conversations]
for conversation in conversations:
for message in conversation:
if isinstance(message["content"], list):
for ele in message["content"]:
if (
"image" in ele
or "image_url" in ele
or "video" in ele
or ele["type"] in ("image", "image_url", "video")
):
vision_infos.append(ele)
return vision_infos
def process_vision_info(
conversations: list[dict] | list[list[dict]],
) -> tuple[list[Image.Image] | None, list[torch.Tensor | list[Image.Image]] | None]:
vision_infos = extract_vision_info(conversations)
## Read images or videos
image_inputs = []
video_inputs = []
for vision_info in vision_infos:
if "image" in vision_info or "image_url" in vision_info:
image_inputs.append(fetch_image(vision_info))
elif "video" in vision_info:
video_inputs.append(fetch_video(vision_info))
else:
raise ValueError("image, image_url or video should in content.")
if len(image_inputs) == 0:
image_inputs = None
if len(video_inputs) == 0:
video_inputs = None
return image_inputs, video_inputs
+75 -17
View File
@@ -11,10 +11,12 @@ import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler, )
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.models.base import ModelBase
SchedulerName = Literal["flow_match_euler", "model_default"]
SchedulerName = Literal["flow_match_euler", "flow_unipc", "model_default"]
TrajectoryName = Literal["ode", "sde_reflow"]
@@ -26,6 +28,7 @@ class SamplingConfig:
scheduler: SchedulerName = "model_default"
trajectory: TrajectoryName = "ode"
flow_shift: float | None = None
guidance_scale: float = 1.0
timesteps: list[float] | None = None
sigmas: list[float] | None = None
@@ -37,6 +40,7 @@ class SamplingConfig:
raise ValueError(f"method.sampling must be a mapping, got {type(raw).__name__}")
supported_keys = {
"flow_shift",
"guidance_scale",
"num_steps",
"scheduler",
"sigmas",
@@ -48,9 +52,9 @@ class SamplingConfig:
raise ValueError(f"Unsupported method.sampling key(s): {unsupported_keys}. "
f"Supported keys: {sorted(supported_keys)}")
scheduler = str(raw.get("scheduler", "model_default") or "model_default").strip().lower()
if scheduler not in {"flow_match_euler", "model_default"}:
if scheduler not in {"flow_match_euler", "flow_unipc", "model_default"}:
raise ValueError("method.sampling.scheduler must be one of "
"{flow_match_euler, model_default}, got "
"{flow_match_euler, flow_unipc, model_default}, got "
f"{raw.get('scheduler')!r}")
trajectory = str(raw.get("trajectory", "ode") or "ode").strip().lower()
if trajectory not in {"ode", "sde_reflow"}:
@@ -69,14 +73,21 @@ class SamplingConfig:
sigmas = [float(s) for s in sigmas]
if timesteps is not None and sigmas is not None and len(timesteps) != len(sigmas):
raise ValueError("method.sampling.timesteps and method.sampling.sigmas must have the same length")
if scheduler == "flow_unipc" and timesteps is not None:
raise ValueError("method.sampling.timesteps is not supported with flow_unipc; "
"use num_steps or sigmas instead")
num_steps = int(raw.get("num_steps", 25) or 25)
if num_steps <= 0:
raise ValueError("method.sampling.num_steps must be positive")
guidance_scale = float(raw.get("guidance_scale", 1.0) or 1.0)
if guidance_scale < 0.0:
raise ValueError("method.sampling.guidance_scale must be non-negative")
return cls(
num_steps=num_steps,
scheduler=scheduler, # type: ignore[arg-type]
trajectory=trajectory, # type: ignore[arg-type]
flow_shift=(None if raw.get("flow_shift", None) in (None, "inherit") else float(raw["flow_shift"])),
guidance_scale=guidance_scale,
timesteps=timesteps,
sigmas=sigmas,
)
@@ -130,13 +141,7 @@ class DiffusionSampler:
for timestep in timesteps:
model_timestep = self._model_timestep(timestep, current)
batch.timesteps = model_timestep
pred_noise = model.predict_noise(
current,
model_timestep,
batch,
conditional=True,
attn_kind="dense",
)
pred_noise = self._predict_with_cfg(model, current, model_timestep, batch)
current = scheduler.step(
pred_noise.flatten(0, 1),
timestep,
@@ -170,6 +175,11 @@ class DiffusionSampler:
if shift is None:
shift = float(getattr(model.noise_scheduler, "shift", 1.0))
scheduler = FlowMatchEulerDiscreteScheduler(shift=float(shift))
elif self.config.scheduler == "flow_unipc":
shift = self.config.flow_shift
if shift is None:
shift = float(getattr(model.noise_scheduler, "shift", 1.0))
scheduler = FlowUniPCMultistepScheduler(shift=float(shift))
else:
scheduler = copy.deepcopy(model.noise_scheduler)
kwargs: dict[str, Any] = {"device": device}
@@ -182,6 +192,8 @@ class DiffusionSampler:
if "num_inference_steps" not in kwargs:
kwargs["num_inference_steps"] = self.config.num_steps
scheduler.set_timesteps(**kwargs)
if hasattr(scheduler, "set_begin_index"):
scheduler.set_begin_index(0)
return scheduler
def _sample_sde_reflow(
@@ -197,13 +209,7 @@ class DiffusionSampler:
for step_idx, timestep in enumerate(timesteps):
timestep_tensor = self._model_timestep(timestep, current)
batch.timesteps = timestep_tensor
pred_clean = model.predict_x0(
current,
timestep_tensor,
batch,
conditional=True,
attn_kind="dense",
)
pred_clean = self._predict_x0_with_cfg(model, current, timestep_tensor, batch)
if step_idx < len(timesteps) - 1:
next_timestep = timesteps[step_idx + 1].reshape(1).to(device=current.device)
noise = torch.randn(
@@ -215,6 +221,58 @@ class DiffusionSampler:
current = model.add_noise(pred_clean, noise, next_timestep)
return pred_clean
def _predict_with_cfg(
self,
model: ModelBase,
current: torch.Tensor,
timestep: torch.Tensor,
batch: TrainingBatch,
) -> torch.Tensor:
cond = model.predict_noise(
current,
timestep,
batch,
conditional=True,
attn_kind="dense",
)
guidance_scale = float(self.config.guidance_scale)
if guidance_scale == 1.0:
return cond
uncond = model.predict_noise(
current,
timestep,
batch,
conditional=False,
attn_kind="dense",
)
return uncond + guidance_scale * (cond - uncond)
def _predict_x0_with_cfg(
self,
model: ModelBase,
current: torch.Tensor,
timestep: torch.Tensor,
batch: TrainingBatch,
) -> torch.Tensor:
cond = model.predict_x0(
current,
timestep,
batch,
conditional=True,
attn_kind="dense",
)
guidance_scale = float(self.config.guidance_scale)
if guidance_scale == 1.0:
return cond
uncond = model.predict_x0(
current,
timestep,
batch,
conditional=False,
attn_kind="dense",
)
return uncond + guidance_scale * (cond - uncond)
@staticmethod
def _model_timestep(
timestep: torch.Tensor,
@@ -17,8 +17,11 @@ class RLValidationConfig:
num_prompts: int = 16
batch_size: int = 16
log_samples: bool = True
max_samples: int | None = None
fps: int = 1
seed: int = 42
data_path: str | None = None
num_latent_t: int | None = None
sampling: dict[str, Any] | None = None
@classmethod
@@ -31,14 +34,25 @@ class RLValidationConfig:
sampling = raw.get("sampling", None)
if sampling is not None and not isinstance(sampling, dict):
raise ValueError(f"method.validation.sampling must be a mapping, got {type(sampling).__name__}")
max_samples = raw.get("max_samples", None)
if max_samples is not None:
max_samples = max(0, int(max_samples))
num_latent_t = raw.get("num_latent_t", None)
if num_latent_t is not None:
num_latent_t = int(num_latent_t)
if num_latent_t <= 0:
raise ValueError("method.validation.num_latent_t must be positive when set")
return cls(
every_steps=max(0, int(raw.get("every_steps", 0) or 0)),
num_steps=max(1, int(raw.get("num_steps", 40) or 40)),
num_prompts=max(1, int(raw.get("num_prompts", 16) or 16)),
batch_size=max(1, int(raw.get("batch_size", 16) or 16)),
log_samples=bool(raw.get("log_samples", True)),
max_samples=max_samples,
fps=max(1, int(raw.get("fps", 1) or 1)),
seed=int(raw.get("seed", 42) or 42),
data_path=(None if data_path in (None, "") else str(data_path)),
num_latent_t=num_latent_t,
sampling=(dict(sampling) if sampling is not None else None),
)
+46 -16
View File
@@ -21,7 +21,11 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.models.base import ModelBase
from fastvideo.train.methods.rl.rewards import build_multi_reward_scorer
from fastvideo.train.methods.rl.rewards import (
GENRL_REWARD_NAMES,
build_multi_reward_scorer,
normalize_reward_weights,
)
from fastvideo.train.methods.rl.common import (
DiffusionSampler,
RLValidationConfig,
@@ -134,14 +138,18 @@ class DiffusionNFTMethod(TrainingMethod):
"{all, positive_only, negative_only, one_only, binary}")
reward_fn = self.method_config.get("reward_fn", None)
if not isinstance(reward_fn, dict) or not reward_fn:
raise ValueError("method.reward_fn must be a non-empty mapping, "
"for example {pickscore: 1.0, clipscore: 1.0}")
self._reward_fn_config = {str(k): float(v) for k, v in reward_fn.items()}
unsupported = sorted(set(self._reward_fn_config) - {"pickscore", "clipscore"})
if unsupported:
raise ValueError(f"Unsupported DiffusionNFT reward(s): {unsupported}. "
"Only pickscore and clipscore are currently ported.")
self._reward_fn_config, reward_backend = normalize_reward_weights(reward_fn)
self._reward_backend = str(
self.method_config.get(
"reward_backend",
reward_backend or "auto",
) or "auto").strip().lower()
if self._reward_backend not in {"auto", "diffusion_nft", "genrl"}:
raise ValueError("method.reward_backend must be one of auto, diffusion_nft, "
f"or genrl, got {self._reward_backend!r}")
if self._reward_backend == "genrl" and not any(name in GENRL_REWARD_NAMES for name in self._reward_fn_config):
raise ValueError("method.reward_backend='genrl' requires at least one GenRL reward "
f"from {sorted(GENRL_REWARD_NAMES)}")
self._reward_scorer: Any | None = None
self._init_optimizer_and_scheduler()
@@ -225,6 +233,7 @@ class DiffusionNFTMethod(TrainingMethod):
self._reward_scorer = build_multi_reward_scorer(
self._reward_fn_config,
device=self.student.device,
backend=self._reward_backend,
)
def _init_optimizer_and_scheduler(self) -> None:
@@ -364,11 +373,7 @@ class DiffusionNFTMethod(TrainingMethod):
batch_items = items[start:start + config.batch_size]
raw_batch = self._collate_validation_rows([item[2] for item in batch_items])
prompts = self._extract_prompts(raw_batch)
batch = self.student.prepare_batch(
raw_batch,
generator=prepare_generator,
latents_source="zeros",
)
batch = self._prepare_validation_batch(raw_batch, prepare_generator)
sampling_result = self._validation_sampler.sample(
self.student,
batch,
@@ -420,6 +425,18 @@ class DiffusionNFTMethod(TrainingMethod):
self._log_progress(f"[DiffusionNFT] validation step {iteration}: finished")
return metrics
def _prepare_validation_batch(
self,
raw_batch: dict[str, Any],
generator: torch.Generator,
) -> TrainingBatch:
return self.student.prepare_batch(
raw_batch,
generator=generator,
latents_source="zeros",
num_latent_t=self._validation_config.num_latent_t,
)
def _get_validation_items(self) -> list[tuple[int, bool, dict[str, Any]]]:
if self._validation_items is not None:
return self._validation_items
@@ -482,11 +499,15 @@ class DiffusionNFTMethod(TrainingMethod):
return
artifacts = []
for item in sorted(logs, key=lambda x: int(x["index"])):
fps = int(self._validation_config.fps)
sorted_logs = sorted(logs, key=lambda x: int(x["index"]))
if self._validation_config.max_samples is not None:
sorted_logs = sorted_logs[:self._validation_config.max_samples]
for item in sorted_logs:
artifact = tracker.video(
media_to_video_array(item["media"]),
caption=validation_caption(str(item["prompt"]), item["rewards"]),
fps=1,
fps=fps,
)
if artifact is not None:
artifacts.append(artifact)
@@ -544,6 +565,7 @@ class DiffusionNFTMethod(TrainingMethod):
effective_grad_accum *= max(1, num_train_timesteps)
current_accum = 0
optimizer_steps = 0
partial_step_micro_steps = 0
loss_terms: dict[str, list[torch.Tensor]] = defaultdict(list)
num_batches = max(1, total_samples // max(1, self._train_batch_size))
training_batch_size = max(1, total_samples // num_batches)
@@ -607,6 +629,11 @@ class DiffusionNFTMethod(TrainingMethod):
progress.update(1)
if current_accum % effective_grad_accum != 0:
partial_step_micro_steps = current_accum % effective_grad_accum
self._log_progress("[DiffusionNFT] final optimizer step uses a partial "
f"gradient accumulation window "
f"({partial_step_micro_steps}/{effective_grad_accum} "
"timestep micro-steps)")
self._clip_student_grads()
self._student_optimizer.step()
self._student_lr_scheduler.step()
@@ -627,6 +654,9 @@ class DiffusionNFTMethod(TrainingMethod):
"nft/iteration": float(iteration),
"nft/num_inner_epochs": float(self._num_inner_epochs),
"nft/inner_micro_steps": float(current_accum),
"nft/effective_grad_accum_micro_steps": float(effective_grad_accum),
"nft/partial_optimizer_step_micro_steps": float(partial_step_micro_steps),
"nft/partial_optimizer_step_ratio": float(partial_step_micro_steps) / float(effective_grad_accum),
"nft/optimizer_steps": float(optimizer_steps),
"ema/update_count": float(self._ema_update_count),
}
+86 -5
View File
@@ -1,6 +1,16 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable reward models for training methods."""
from typing import Any
from fastvideo.train.methods.rl.rewards.diffusion_nft import (
BUILTIN_DEBUG_REWARD_SCORERS,
ExternalDiffusionNFTScorer,
JpegCompressibilityScorer,
JpegIncompressibilityScorer,
MeanLuminanceScorer,
normalize_reward_weights,
)
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
@@ -8,30 +18,101 @@ from fastvideo.train.methods.rl.rewards.frame_rewards import (
from fastvideo.train.methods.rl.rewards.media import (
MultiRewardScorer,
RewardScorer,
media_to_uint8_array,
select_first_frame,
)
GENRL_REWARD_NAMES = frozenset({
"hpsv3_general",
"hpsv3_percentile",
"videoalign_mq",
"videoalign_ta",
"videoalign_vq",
})
_NATIVE_SCORER_CLASSES: dict[str, Any] = {
"pickscore": PickScoreScorer,
"clipscore": ClipScoreScorer,
**BUILTIN_DEBUG_REWARD_SCORERS,
}
def _build_lazy_genrl_scorer(
name: str,
*,
device,
) -> RewardScorer:
if name == "hpsv3_general":
from fastvideo.train.methods.rl.rewards.hpsv3 import HPSv3GeneralScorer
return HPSv3GeneralScorer(device=device)
if name == "hpsv3_percentile":
from fastvideo.train.methods.rl.rewards.hpsv3 import HPSv3PercentileScorer
return HPSv3PercentileScorer(device=device)
if name == "videoalign_mq":
from fastvideo.train.methods.rl.rewards.videoalign import VideoAlignMotionQualityScorer
return VideoAlignMotionQualityScorer(device=device)
if name == "videoalign_ta":
from fastvideo.train.methods.rl.rewards.videoalign import VideoAlignTextAlignmentScorer
return VideoAlignTextAlignmentScorer(device=device)
if name == "videoalign_vq":
from fastvideo.train.methods.rl.rewards.videoalign import VideoAlignVisualQualityScorer
return VideoAlignVisualQualityScorer(device=device)
raise ValueError(f"Unsupported GenRL reward {name!r}. "
f"Available GenRL rewards: {sorted(GENRL_REWARD_NAMES)}")
def build_multi_reward_scorer(
reward_weights,
*,
device="cuda",
backend: str = "auto",
scorers: dict[str, RewardScorer] | None = None,
) -> MultiRewardScorer:
reward_weights, reward_backend = normalize_reward_weights(reward_weights)
backend = reward_backend or str(backend or "auto").strip().lower()
if backend not in {"auto", "diffusion_nft", "genrl"}:
raise ValueError("method.reward_backend must be one of auto, diffusion_nft, or genrl, "
f"got {backend!r}")
available: dict[str, RewardScorer] = dict(scorers or {})
if not available:
available = {
"pickscore": PickScoreScorer(device=device),
"clipscore": ClipScoreScorer(device=device),
}
for name in reward_weights:
if name in available:
continue
if backend == "diffusion_nft" and name not in _NATIVE_SCORER_CLASSES:
available[name] = ExternalDiffusionNFTScorer(name, device=device)
continue
if name in GENRL_REWARD_NAMES:
available[name] = _build_lazy_genrl_scorer(name, device=device)
continue
scorer_cls = _NATIVE_SCORER_CLASSES.get(name)
if scorer_cls is None:
if backend == "genrl":
raise ValueError(f"Unsupported GenRL reward {name!r}. "
f"Available GenRL rewards: {sorted(GENRL_REWARD_NAMES)}")
available[name] = ExternalDiffusionNFTScorer(name, device=device)
elif name in BUILTIN_DEBUG_REWARD_SCORERS:
available[name] = scorer_cls()
else:
available[name] = scorer_cls(device=device)
return MultiRewardScorer(reward_weights, scorers=available)
__all__ = [
"ClipScoreScorer",
"ExternalDiffusionNFTScorer",
"JpegCompressibilityScorer",
"JpegIncompressibilityScorer",
"MeanLuminanceScorer",
"MultiRewardScorer",
"PickScoreScorer",
"RewardScorer",
"build_multi_reward_scorer",
"media_to_uint8_array",
"normalize_reward_weights",
"select_first_frame",
]
@@ -0,0 +1,160 @@
# SPDX-License-Identifier: Apache-2.0
"""DiffusionNFT-compatible reward adapters."""
from __future__ import annotations
import os
import sys
from collections.abc import Mapping, Sequence
from io import BytesIO
from pathlib import Path
from typing import Any
import numpy as np
from PIL import Image
import torch
from fastvideo.train.methods.rl.rewards.media import media_to_uint8_array
def normalize_reward_weights(
reward_config: Any,
) -> tuple[dict[str, float], str | None]:
"""Accept flat or DiffusionNFT-style nested reward mappings."""
backend = None
raw_rewards = reward_config
if isinstance(raw_rewards, Mapping) and "backend" in raw_rewards:
backend = str(raw_rewards["backend"]).strip().lower()
if isinstance(raw_rewards, Mapping) and "rewards" in raw_rewards:
raw_rewards = raw_rewards["rewards"]
if not isinstance(raw_rewards, Mapping) or not raw_rewards:
raise ValueError("method.reward_fn must be a non-empty mapping, "
"for example {pickscore: 1.0, clipscore: 1.0} "
"or {rewards: {videoalign_vq: 1.0}}")
return {str(key).strip().lower(): float(value) for key, value in raw_rewards.items()}, backend
class JpegIncompressibilityScorer:
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
del prompts
arr = media_to_uint8_array(media)
sizes: list[float] = []
for sample in arr:
frames = sample[np.newaxis] if sample.ndim == 3 else sample
frame_sizes: list[float] = []
for frame in frames:
buffer = BytesIO()
Image.fromarray(frame).save(buffer, format="JPEG", quality=95)
frame_sizes.append(buffer.tell() / 1000.0)
sizes.append(float(np.mean(frame_sizes)))
return torch.tensor(sizes, device=media.device, dtype=torch.float32)
class JpegCompressibilityScorer:
def __init__(self) -> None:
self._incompressibility = JpegIncompressibilityScorer()
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
return -self._incompressibility(media, prompts) / 500.0
class MeanLuminanceScorer:
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
del prompts
reduce_dims = tuple(range(1, media.ndim))
return media.detach().float().mean(dim=reduce_dims)
def _ensure_local_diffusion_nft_on_path() -> None:
explicit_root = os.environ.get("DIFFUSION_NFT_ROOT")
candidates: list[Path] = []
if explicit_root:
candidates.append(Path(explicit_root))
candidates.append(Path("/cache/DiffusionNFT"))
for parent in Path(__file__).resolve().parents:
candidates.append(parent / "DiffusionNFT")
for root_path in candidates:
if (root_path / "flow_grpo").is_dir():
root = str(root_path)
if root not in sys.path:
sys.path.insert(0, root)
return
def _flatten_video_for_image_reward(
media: torch.Tensor,
prompts: Sequence[str],
) -> tuple[torch.Tensor, list[str], int | None]:
if media.ndim != 5:
return media, list(prompts), None
batch, channels, frames, height, width = media.shape
flat = media.permute(0, 2, 1, 3, 4).reshape(batch * frames, channels, height, width)
flat_prompts = [prompt for prompt in prompts for _ in range(frames)]
return flat, flat_prompts, frames
class ExternalDiffusionNFTScorer:
"""Adapter for rewards from a local DiffusionNFT/Flow-GRPO checkout."""
def __init__(
self,
name: str,
*,
device: torch.device | str = "cuda",
) -> None:
self.name = str(name).strip().lower()
self.device = torch.device(device)
_ensure_local_diffusion_nft_on_path()
try:
from flow_grpo import rewards as flow_rewards
except ImportError as exc:
raise ImportError(f"Reward {self.name!r} requires DiffusionNFT's "
"Flow-GRPO reward package. Set DIFFUSION_NFT_ROOT "
"to a checkout containing flow_grpo, or choose a "
"built-in reward.") from exc
self._score_fn = flow_rewards.multi_score(self.device, {self.name: 1.0})
@torch.no_grad()
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
reward_media, reward_prompts, frames = _flatten_video_for_image_reward(media, prompts)
metadata: list[dict[str, Any]] = [{} for _ in reward_prompts]
scores, _ = self._score_fn(
reward_media,
reward_prompts,
metadata,
only_strict=True,
)
value = torch.as_tensor(scores[self.name], device=self.device, dtype=torch.float32)
if frames is not None:
if value.numel() % frames != 0:
raise RuntimeError(f"Reward {self.name!r} returned {value.numel()} scores, "
f"which is not divisible by num_frames={frames}.")
value = value.reshape(-1, frames).mean(dim=1)
return value
BUILTIN_DEBUG_REWARD_SCORERS = {
"jpeg_incompressibility": JpegIncompressibilityScorer,
"jpeg_compressibility": JpegCompressibilityScorer,
"mean_luminance": MeanLuminanceScorer,
}
+212
View File
@@ -0,0 +1,212 @@
# SPDX-License-Identifier: Apache-2.0
"""HPSv3 reward scorers for RL training."""
from __future__ import annotations
import os
import sys
import tempfile
from pathlib import Path
from typing import Any
import numpy as np
import torch
from fastvideo.train.methods.rl.rewards.media import media_to_uint8_array
_HPSV3_ROOT = Path(__file__).resolve().parents[4] / "third_party" / "rl_rewards" / "HPSv3"
if _HPSV3_ROOT.exists() and str(_HPSV3_ROOT) not in sys.path:
sys.path.insert(0, str(_HPSV3_ROOT))
_HPSV3_INFERENCERS: dict[str, Any] = {}
_HPSV3_LOAD_PATCHED = False
def _patch_transformers_video_input_alias() -> None:
from transformers import image_utils
if not hasattr(image_utils, "VideoInput"):
image_utils.VideoInput = image_utils.ImageInput
def _remap_hpsv3_state_dict(state_dict: dict[str, Any]) -> dict[str, Any]:
remapped = {}
for key, value in state_dict.items():
if key.startswith("visual."):
key = f"model.{key}"
elif key.startswith("model.layers.") or key.startswith("model.embed_tokens.") or key.startswith("model.norm."):
key = f"model.language_model.{key[len('model.'):]}"
key = key.replace("base_model.model.visual.", "base_model.model.model.visual.", 1)
key = key.replace("base_model.model.model.layers.", "base_model.model.model.language_model.layers.", 1)
key = key.replace("base_model.model.model.embed_tokens.",
"base_model.model.model.language_model.embed_tokens.", 1)
key = key.replace("base_model.model.model.norm.", "base_model.model.model.language_model.norm.", 1)
remapped[key] = value
return remapped
def _walk_model_graph(model: Any):
stack = [model]
seen = set()
while stack:
current = stack.pop()
if current is None or id(current) in seen:
continue
seen.add(id(current))
yield current
for attr in ("base_model", "model"):
child = getattr(current, attr, None)
if child is not None:
stack.append(child)
def _patch_load_state_dict(cls: Any) -> None:
if getattr(cls, "_fastvideo_qwen2vl_key_remap", False):
return
original_load_state_dict = cls.load_state_dict
def load_state_dict_with_key_remap(self, state_dict, strict=True, assign=False):
state_dict = _remap_hpsv3_state_dict(state_dict)
return original_load_state_dict(self, state_dict, strict=strict, assign=assign)
cls.load_state_dict = load_state_dict_with_key_remap
cls._fastvideo_qwen2vl_key_remap = True
def _patch_hpsv3_state_dict_loader() -> None:
global _HPSV3_LOAD_PATCHED
if _HPSV3_LOAD_PATCHED:
return
from hpsv3.model.reward_model import Qwen2VLRewardModelBT
_patch_load_state_dict(Qwen2VLRewardModelBT)
try:
from peft import PeftModel
except ImportError:
PeftModel = None
if PeftModel is not None:
_patch_load_state_dict(PeftModel)
_HPSV3_LOAD_PATCHED = True
def _patch_hpsv3_runtime_model(model: Any) -> None:
for candidate in _walk_model_graph(model):
language_model = getattr(candidate, "language_model", None)
if language_model is not None and not hasattr(candidate, "embed_tokens") and hasattr(language_model,
"embed_tokens"):
candidate.embed_tokens = language_model.embed_tokens
def _normalize_device(device: torch.device | str) -> str:
return str(torch.device(device))
def _move_hpsv3_inferencer(inferencer: Any, device: torch.device | str) -> None:
device_str = _normalize_device(device)
model = getattr(inferencer, "model", None)
if model is not None and hasattr(model, "to"):
model.to(device)
inferencer.device = device_str
def set_hpsv3_device(device: torch.device | str) -> None:
key = _normalize_device(device)
if key in _HPSV3_INFERENCERS:
return
for old_key, inferencer in list(_HPSV3_INFERENCERS.items()):
if old_key != key:
_move_hpsv3_inferencer(inferencer, device)
_HPSV3_INFERENCERS[key] = inferencer
del _HPSV3_INFERENCERS[old_key]
return
def _get_hpsv3_inferencer(device: torch.device | str) -> Any:
key = _normalize_device(device)
if key not in _HPSV3_INFERENCERS:
try:
_patch_transformers_video_input_alias()
from hpsv3 import HPSv3RewardInferencer
_patch_hpsv3_state_dict_loader()
except ImportError as exc:
raise ImportError("HPSv3 rewards require the HPSv3 package or the vendored "
"fastvideo/third_party/rl_rewards/HPSv3 directory.") from exc
inferencer = HPSv3RewardInferencer(device=device)
_patch_hpsv3_runtime_model(inferencer.model)
_HPSV3_INFERENCERS[key] = inferencer
return _HPSV3_INFERENCERS[key]
def _save_frame_to_temp(frame: np.ndarray) -> str:
from PIL import Image
fd, path = tempfile.mkstemp(suffix=".png")
os.close(fd)
Image.fromarray(frame).save(path)
return path
def _extract_reward_scalar(result: Any) -> float:
if isinstance(result, torch.Tensor):
return float(result.item())
if isinstance(result, float | int):
return float(result)
if isinstance(result, list | np.ndarray):
return float(np.mean(result))
return float(result)
class HPSv3GeneralScorer:
"""Score every frame with a generic quality prompt and average."""
def __init__(self, *, device: torch.device | str = "cuda") -> None:
self.device = torch.device(device)
@torch.no_grad()
def __call__(self, media: torch.Tensor, prompts) -> torch.Tensor:
del prompts
inferencer = _get_hpsv3_inferencer(self.device)
images_np = media_to_uint8_array(media)
batch_scores = []
for sample in images_np:
frames = sample[np.newaxis] if sample.ndim == 3 else sample
frame_scores = []
for frame in frames:
path = _save_frame_to_temp(frame)
try:
rewards = inferencer.reward(["A high-quality image"], [path])
frame_scores.append(_extract_reward_scalar(rewards[0][0]))
finally:
os.remove(path)
batch_scores.append(float(np.mean(frame_scores)))
return torch.tensor(batch_scores, device=self.device, dtype=torch.float32)
class HPSv3PercentileScorer:
"""Score frames with per-prompt text and average the top 30 percent."""
def __init__(self, *, device: torch.device | str = "cuda") -> None:
self.device = torch.device(device)
@torch.no_grad()
def __call__(self, media: torch.Tensor, prompts) -> torch.Tensor:
inferencer = _get_hpsv3_inferencer(self.device)
images_np = media_to_uint8_array(media)
batch_scores = []
for sample_idx, sample in enumerate(images_np):
frames = sample[np.newaxis] if sample.ndim == 3 else sample
prompt = prompts[sample_idx] if sample_idx < len(prompts) else "A high-quality image"
frame_scores = []
for frame in frames:
path = _save_frame_to_temp(frame)
try:
rewards = inferencer.reward([prompt], [path])
frame_scores.append(_extract_reward_scalar(rewards[0][0]))
finally:
os.remove(path)
if not frame_scores:
batch_scores.append(0.0)
continue
top_k = max(1, int(len(frame_scores) * 0.3))
batch_scores.append(float(np.mean(sorted(frame_scores, reverse=True)[:top_k])))
return torch.tensor(batch_scores, device=self.device, dtype=torch.float32)
@@ -5,6 +5,7 @@ from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
import numpy as np
import torch
RewardScorer = Callable[[torch.Tensor, Sequence[str]], torch.Tensor]
@@ -27,6 +28,32 @@ def select_first_frame(media: torch.Tensor) -> torch.Tensor:
f"got {tuple(media.shape)}")
def media_to_uint8_array(media: torch.Tensor | np.ndarray) -> np.ndarray:
"""Convert image/video media to uint8 NHWC or NFHWC arrays."""
if isinstance(media, torch.Tensor):
media = media.detach().float().clamp(0, 1).cpu().numpy()
media = np.asarray(media)
if media.ndim == 4:
if media.shape[-1] in (1, 3):
pass
elif media.shape[1] in (1, 3):
media = media.transpose(0, 2, 3, 1)
elif media.ndim == 5:
if media.shape[-1] in (1, 3):
pass
elif media.shape[2] in (1, 3):
media = media.transpose(0, 1, 3, 4, 2)
elif media.shape[1] in (1, 3):
media = media.transpose(0, 2, 3, 4, 1)
else:
raise ValueError("media must have shape [B, C, H, W], [B, H, W, C], "
"[B, C, T, H, W], [B, T, C, H, W], or "
f"[B, T, H, W, C], got {tuple(media.shape)}")
if media.dtype in (np.float16, np.float32, np.float64):
media = np.clip(media * 255.0, 0, 255).round().astype(np.uint8)
return media
class MultiRewardScorer:
"""Weighted sum of reusable media reward scorers.
@@ -0,0 +1,310 @@
# SPDX-License-Identifier: Apache-2.0
"""VideoAlign reward scorers for motion, visual quality, and text alignment."""
from __future__ import annotations
import os
import sys
import tempfile
from importlib import import_module
from pathlib import Path
from typing import Any
import cv2
import numpy as np
import torch
from fastvideo.train.methods.rl.rewards.media import media_to_uint8_array
_VIDEOALIGN_ROOT = Path(__file__).resolve().parents[4] / "third_party" / "rl_rewards" / "VideoAlign"
if _VIDEOALIGN_ROOT.is_dir() and str(_VIDEOALIGN_ROOT) not in sys.path:
sys.path.insert(0, str(_VIDEOALIGN_ROOT))
_VIDEOALIGN_INFERENCERS: dict[str, Any] = {}
_VIDEOALIGN_PATCHED = False
def _normalize_device_str(device: torch.device | str) -> str:
return str(torch.device(device))
def _move_videoalign_inferencer(inferencer: Any, device: torch.device | str) -> None:
device_str = _normalize_device_str(device)
model = getattr(inferencer, "model", None)
if model is not None and hasattr(model, "to"):
model.to(device)
inferencer.device = device_str
def _remap_qwen2vl_state_dict_keys(state_dict: dict[str, Any]) -> dict[str, Any]:
remapped = {}
for key, value in state_dict.items():
if key.startswith("visual."):
key = f"model.{key}"
elif key.startswith("model.layers.") or key.startswith("model.embed_tokens.") or key.startswith("model.norm."):
key = f"model.language_model.{key[len('model.'):]}"
key = key.replace("base_model.model.visual.", "base_model.model.model.visual.", 1)
key = key.replace("base_model.model.model.layers.", "base_model.model.model.language_model.layers.", 1)
key = key.replace("base_model.model.model.embed_tokens.",
"base_model.model.model.language_model.embed_tokens.", 1)
key = key.replace("base_model.model.model.norm.", "base_model.model.model.language_model.norm.", 1)
remapped[key] = value
return remapped
def _walk_model_graph(model: Any):
stack = [model]
seen = set()
while stack:
current = stack.pop()
if current is None or id(current) in seen:
continue
seen.add(id(current))
yield current
for attr in ("base_model", "model"):
child = getattr(current, attr, None)
if child is not None:
stack.append(child)
def _patch_load_state_dict(cls: Any) -> None:
if getattr(cls, "_fastvideo_qwen2vl_key_remap", False):
return
original_load_state_dict = cls.load_state_dict
def load_state_dict_with_key_remap(self, state_dict, strict=True, assign=False):
state_dict = _remap_qwen2vl_state_dict_keys(state_dict)
if not assign:
try:
assign = any(getattr(param, "is_meta", False) for param in self.parameters())
except Exception:
assign = False
return original_load_state_dict(self, state_dict, strict=strict, assign=assign)
cls.load_state_dict = load_state_dict_with_key_remap
cls._fastvideo_qwen2vl_key_remap = True
def _select_videoalign_frame_indices(
vision_mod: Any,
ele: dict[str, Any],
total_frames: int,
video_fps: float,
) -> list[int]:
sample_type = ele.get("sample_type", "uniform")
if sample_type == "uniform":
nframes = vision_mod.smart_nframes(ele, total_frames=total_frames, video_fps=video_fps)
return torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
if sample_type == "multi_pts":
frames_each_pts = 6
num_pts = 4
fps = 8
nframes = max(frames_each_pts, int(total_frames * fps // video_fps))
frame_idx = torch.linspace(0, total_frames - 1, nframes).round().long().tolist()
start_pt = int(frames_each_pts // 2)
end_pt = int(nframes - frames_each_pts // 2 - 1)
pts = torch.linspace(start_pt, end_pt, num_pts).round().long().tolist()
idx = []
for pt in pts:
idx.extend(frame_idx[pt - frames_each_pts // 2:pt + frames_each_pts // 2])
return idx
raise ValueError(f"Unsupported VideoAlign sample_type: {sample_type}")
def _read_video_opencv(vision_mod: Any, ele: dict[str, Any]) -> torch.Tensor:
video_path = ele["video"]
if video_path.startswith("file://"):
video_path = video_path[7:]
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise ValueError(f"Could not open video: {video_path}")
video_fps = float(cap.get(cv2.CAP_PROP_FPS) or 30.0)
total_file_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT) or 0)
start_frame = max(0, int(round(float(ele.get("video_start", 0.0) or 0.0) * video_fps)))
end_sec = ele.get("video_end")
if end_sec is None:
end_frame = total_file_frames if total_file_frames > 0 else None
else:
end_frame = int(round(float(end_sec) * video_fps))
if total_file_frames > 0:
end_frame = min(end_frame, total_file_frames)
cap.set(cv2.CAP_PROP_POS_FRAMES, start_frame)
frames = []
current_frame = start_frame
while end_frame is None or current_frame < end_frame:
ok, frame = cap.read()
if not ok:
break
frames.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB))
current_frame += 1
cap.release()
if not frames:
raise ValueError(f"No frames were read from video: {video_path}.")
idx = _select_videoalign_frame_indices(vision_mod, ele, total_frames=len(frames), video_fps=video_fps)
video = np.stack([frames[i] for i in idx], axis=0)
return torch.from_numpy(video).permute(0, 3, 1, 2)
def _torchvision_read_video_available() -> bool:
try:
torchvision_io = import_module("torchvision.io")
except Exception:
return False
return hasattr(torchvision_io, "read_video")
def _patch_videoalign_video_reader() -> None:
vision_mod = import_module("vision_process")
if "opencv" not in vision_mod.VIDEO_READER_BACKENDS:
def read_video_opencv(ele):
return _read_video_opencv(vision_mod, ele)
vision_mod.VIDEO_READER_BACKENDS["opencv"] = read_video_opencv
if _torchvision_read_video_available():
return
vision_mod.__dict__["FORCE_QWENVL_VIDEO_READER"] = "opencv"
if hasattr(vision_mod.get_video_reader_backend, "cache_clear"):
vision_mod.get_video_reader_backend.cache_clear()
def _patch_videoalign_modules() -> Any:
global _VIDEOALIGN_PATCHED
inference_mod = import_module("inference")
if _VIDEOALIGN_PATCHED:
return inference_mod
reward_model_mod = import_module("reward_model")
_patch_videoalign_video_reader()
_patch_load_state_dict(reward_model_mod.Qwen2VLRewardModelBT)
try:
peft_mod = import_module("peft")
except ImportError:
peft_mod = None
if peft_mod is not None:
_patch_load_state_dict(peft_mod.PeftModel)
_VIDEOALIGN_PATCHED = True
return inference_mod
def _patch_videoalign_runtime_model(model: Any) -> None:
for candidate in _walk_model_graph(model):
language_model = getattr(candidate, "language_model", None)
if language_model is not None and not hasattr(candidate, "embed_tokens") and hasattr(language_model,
"embed_tokens"):
candidate.embed_tokens = language_model.embed_tokens
def set_videoalign_device(device: torch.device | str) -> None:
key = _normalize_device_str(device)
for old_key, inferencer in list(_VIDEOALIGN_INFERENCERS.items()):
if old_key != key and old_key.split(":")[-1] != key:
new_key = f"{inferencer._key_prefix}:{key}"
_move_videoalign_inferencer(inferencer, device)
_VIDEOALIGN_INFERENCERS[new_key] = inferencer
del _VIDEOALIGN_INFERENCERS[old_key]
def _get_inferencer(
device: torch.device | str,
checkpoint_path: str | None = None,
) -> Any:
if checkpoint_path is None:
checkpoint_path = os.environ.get(
"VIDEOALIGN_CHECKPOINT_PATH",
str(_VIDEOALIGN_ROOT / "checkpoints"),
)
checkpoint_path = os.path.abspath(checkpoint_path)
key = _normalize_device_str(device)
cache_key = f"{checkpoint_path}:{key}"
if cache_key not in _VIDEOALIGN_INFERENCERS:
try:
inference_mod = _patch_videoalign_modules()
video_reward_cls = inference_mod.VideoVLMRewardInference
except ImportError as exc:
raise ImportError("VideoAlign rewards require the VideoAlign runtime files under "
"fastvideo/third_party/rl_rewards/VideoAlign and a "
"VIDEOALIGN_CHECKPOINT_PATH checkpoint.") from exc
inferencer = video_reward_cls(load_from_pretrained=checkpoint_path, device=device)
_patch_videoalign_runtime_model(inferencer.model)
inferencer._key_prefix = checkpoint_path or "default"
_VIDEOALIGN_INFERENCERS[cache_key] = inferencer
return _VIDEOALIGN_INFERENCERS[cache_key]
def _convert_to_grayscale(frames: np.ndarray) -> np.ndarray:
if frames.ndim == 4 and frames.shape[-1] == 3:
gray = np.mean(frames, axis=-1, keepdims=True)
return np.repeat(gray.astype(np.uint8), 3, axis=-1)
return frames
def _save_video_to_temp(frames: np.ndarray, fps: int = 8) -> str:
fd, path = tempfile.mkstemp(suffix=".mp4")
os.close(fd)
height, width = frames.shape[1], frames.shape[2]
writer = cv2.VideoWriter(path, cv2.VideoWriter.fourcc(*"mp4v"), fps, (width, height))
for frame in frames:
writer.write(cv2.cvtColor(frame, cv2.COLOR_RGB2BGR))
writer.release()
return path
class _VideoAlignScorer:
score_key: str
def __init__(
self,
*,
device: torch.device | str = "cuda",
checkpoint_path: str | None = None,
) -> None:
self.device = torch.device(device)
self.checkpoint_path = checkpoint_path
def _prompt(self, prompts, index: int) -> str:
return prompts[index] if prompts and index < len(prompts) else ""
def _frames(self, frames: np.ndarray) -> np.ndarray:
return frames
@torch.no_grad()
def __call__(self, media: torch.Tensor, prompts) -> torch.Tensor:
inferencer = _get_inferencer(self.device, self.checkpoint_path)
images_np = media_to_uint8_array(media)
batch_scores = []
for sample_idx, sample in enumerate(images_np):
frames = sample[np.newaxis] if sample.ndim == 3 else sample
path = _save_video_to_temp(self._frames(frames))
try:
results = inferencer.reward([path], [self._prompt(prompts, sample_idx)], use_norm=True)
batch_scores.append(float(results[0].get(self.score_key, 0.0)))
finally:
os.remove(path)
return torch.tensor(batch_scores, device=self.device, dtype=torch.float32)
class VideoAlignMotionQualityScorer(_VideoAlignScorer):
score_key = "MQ"
def _frames(self, frames: np.ndarray) -> np.ndarray:
return _convert_to_grayscale(frames)
def _prompt(self, prompts, index: int) -> str:
del prompts, index
return ""
class VideoAlignVisualQualityScorer(_VideoAlignScorer):
score_key = "VQ"
def _prompt(self, prompts, index: int) -> str:
del prompts, index
return ""
class VideoAlignTextAlignmentScorer(_VideoAlignScorer):
score_key = "TA"
+6 -1
View File
@@ -126,8 +126,13 @@ class ModelBase(ABC):
*,
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
num_latent_t: int | None = None,
) -> TrainingBatch:
"""Convert a dataloader batch into forward primitives."""
"""Convert a dataloader batch into forward primitives.
``num_latent_t`` may override the configured temporal length when
creating zero latents for sampling-only paths such as validation.
"""
@abstractmethod
def add_noise(
+7 -1
View File
@@ -81,6 +81,7 @@ class CosmosModel(WanModel):
*,
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
num_latent_t: int | None = None,
) -> TrainingBatch:
"""Same flow as Wan, but uses Cosmos VAE
normalisation."""
@@ -97,6 +98,9 @@ class CosmosModel(WanModel):
infos = raw_batch.get("info_list")
if latents_source == "zeros":
resolved_num_latent_t = tc.data.num_latent_t if num_latent_t is None else int(num_latent_t)
if resolved_num_latent_t <= 0:
raise ValueError("num_latent_t must be positive when creating zero latents")
batch_size = encoder_hidden_states.shape[0]
vae_config = (
tc.pipeline_config.vae_config # type: ignore[union-attr]
@@ -112,13 +116,15 @@ class CosmosModel(WanModel):
latents = torch.zeros(
batch_size,
num_channels,
tc.data.num_latent_t,
resolved_num_latent_t,
latent_height,
latent_width,
device=device,
dtype=dtype,
)
elif latents_source == "data":
if num_latent_t is not None:
raise ValueError("num_latent_t can only override latents_source='zeros'")
if "vae_latent" not in raw_batch:
raise ValueError("vae_latent not found in batch "
"and latents_source='data'")
+7 -1
View File
@@ -75,6 +75,7 @@ class HunyuanModel(WanModel):
*,
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
num_latent_t: int | None = None,
) -> TrainingBatch:
"""Same flow as Wan, but uses Hunyuan VAE normalisation."""
self.ensure_negative_conditioning()
@@ -90,6 +91,9 @@ class HunyuanModel(WanModel):
infos = raw_batch.get("info_list")
if latents_source == "zeros":
resolved_num_latent_t = tc.data.num_latent_t if num_latent_t is None else int(num_latent_t)
if resolved_num_latent_t <= 0:
raise ValueError("num_latent_t must be positive when creating zero latents")
batch_size = encoder_hidden_states.shape[0]
vae_config = (
tc.pipeline_config.vae_config # type: ignore[union-attr]
@@ -105,13 +109,15 @@ class HunyuanModel(WanModel):
latents = torch.zeros(
batch_size,
num_channels,
tc.data.num_latent_t,
resolved_num_latent_t,
latent_height,
latent_width,
device=device,
dtype=dtype,
)
elif latents_source == "data":
if num_latent_t is not None:
raise ValueError("num_latent_t can only override latents_source='zeros'")
if "vae_latent" not in raw_batch:
raise ValueError("vae_latent not found in batch "
"and latents_source='data'")
@@ -54,6 +54,7 @@ class MatrixGame2Model(WanModel):
*,
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
num_latent_t: int | None = None,
) -> TrainingBatch:
assert self.training_config is not None
tc = self.training_config
@@ -65,8 +66,16 @@ class MatrixGame2Model(WanModel):
batch_size = self._infer_batch_size(raw_batch)
if latents_source == "zeros":
latents = self._make_zero_latents(batch_size=batch_size)
resolved_num_latent_t = tc.data.num_latent_t if num_latent_t is None else int(num_latent_t)
if resolved_num_latent_t <= 0:
raise ValueError("num_latent_t must be positive when creating zero latents")
latents = self._make_zero_latents(
batch_size=batch_size,
num_latent_t=resolved_num_latent_t,
)
elif latents_source == "data":
if num_latent_t is not None:
raise ValueError("num_latent_t can only override latents_source='zeros'")
latents = raw_batch["vae_latent"][:, :, :tc.data.num_latent_t]
latents = latents.to(device=device, dtype=dtype)
else:
@@ -427,7 +436,12 @@ class MatrixGame2Model(WanModel):
return int(raw_batch["clip_feature"].shape[0])
raise ValueError("Unable to infer batch size from Matrix-Game 2.0 batch")
def _make_zero_latents(self, *, batch_size: int) -> torch.Tensor:
def _make_zero_latents(
self,
*,
batch_size: int,
num_latent_t: int,
) -> torch.Tensor:
assert self.training_config is not None
vae_config = self.training_config.pipeline_config.vae_config.arch_config # type: ignore[union-attr]
num_channels = vae_config.z_dim
@@ -437,7 +451,7 @@ class MatrixGame2Model(WanModel):
return torch.zeros(
batch_size,
num_channels,
self.training_config.data.num_latent_t,
num_latent_t,
latent_height,
latent_width,
device=self.device,
+7 -1
View File
@@ -235,6 +235,7 @@ class WanModel(ModelBase):
*,
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
num_latent_t: int | None = None,
) -> TrainingBatch:
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
@@ -250,6 +251,9 @@ class WanModel(ModelBase):
infos = raw_batch.get("info_list")
if latents_source == "zeros":
resolved_num_latent_t = tc.data.num_latent_t if num_latent_t is None else int(num_latent_t)
if resolved_num_latent_t <= 0:
raise ValueError("num_latent_t must be positive when creating zero latents")
batch_size = encoder_hidden_states.shape[0]
vae_config = (
tc.pipeline_config.vae_config.arch_config # type: ignore[union-attr]
@@ -261,13 +265,15 @@ class WanModel(ModelBase):
latents = torch.zeros(
batch_size,
num_channels,
tc.data.num_latent_t,
resolved_num_latent_t,
latent_height,
latent_width,
device=device,
dtype=dtype,
)
elif latents_source == "data":
if num_latent_t is not None:
raise ValueError("num_latent_t can only override latents_source='zeros'")
if "vae_latent" not in raw_batch:
raise ValueError("vae_latent not found in batch "
"and latents_source='data'")
+467
View File
@@ -0,0 +1,467 @@
"""Run DiffusionNFT Wan RL training on Modal.
By default, video generation uses the Microsoft World-R1 enhanced prompt
dataset and VideoAlign/VideoReward rewards. ``num_frames`` is exposed as a
Modal argument so the same video-policy path can run cheap one-frame debug jobs
or multi-frame video generation. Pass a DiffusionNFT dataset name to
``dataset`` to load prompts from the cached external DiffusionNFT checkout.
Expected Modal resources in the hao-ai-lab workspace:
- Volume `fastvideo-data` mounted under the repo for text-only parquet prompts.
- Volume `fastvideo-runs` mounted under the repo for checkpoints/logs.
- Volume `fastvideo-cache` mounted under the repo for Hugging Face/DiffusionNFT caches.
- Secrets `wandb-adamlee00` and `hf-adamlee00`.
"""
from __future__ import annotations
import modal
app = modal.App("fastvideo-diffusion-nft-wan")
CONFIG_PATH = "examples/train/configs/rl/wan/diffusion_nft_videoalign.yaml"
PROJECT_ROOT = "/root/FastVideo"
MODAL_DATA_ROOT = f"{PROJECT_ROOT}/.modal_data"
MODAL_CACHE_ROOT = f"{PROJECT_ROOT}/.modal_cache"
DIFFUSION_NFT_ROOT = f"{MODAL_CACHE_ROOT}/DiffusionNFT"
OUTPUT_DIR_BASE = f"{PROJECT_ROOT}/outputs/diffusion_nft_wan"
WANDB_ENTITY = "adamlee00"
WANDB_SECRET_NAME = "wandb-adamlee00"
HF_SECRET_NAME = "hf-adamlee00"
DEFAULT_GPU_TYPE = "A100-80GB"
DEFAULT_NUM_GPUS = 4
DEFAULT_HSDP_REPLICATE_DIM = 1
DEFAULT_HSDP_SHARD_DIM = DEFAULT_NUM_GPUS
DEFAULT_MAX_TRAIN_STEPS = 50
DEFAULT_NUM_SAMPLES_PER_PROMPT = 4
DEFAULT_NUM_BATCHES_PER_EPOCH = 1
DEFAULT_COLLECTION_BATCH_SIZE = 4
DEFAULT_INNER_EPOCHS = 1
DEFAULT_TRAIN_BATCH_SIZE = 4
DEFAULT_GRADIENT_ACCUMULATION_STEPS = 1
DEFAULT_LEARNING_RATE = -1.0
DEFAULT_SAMPLE_NUM_STEPS = 50
DEFAULT_SAMPLE_FLOW_SHIFT = -1.0
DEFAULT_SAMPLE_GUIDANCE_SCALE = -1.0
DEFAULT_VALIDATION_NUM_STEPS = 50
DEFAULT_VALIDATION_NUM_PROMPTS = 16
DEFAULT_VALIDATION_BATCH_SIZE = 4
DEFAULT_VALIDATION_NUM_FRAMES = 0
DEFAULT_NUM_FRAMES = 17
DEFAULT_NUM_LATENT_T = 0
DEFAULT_LOG_SAMPLE_MAX_VIDEOS = 16
DEFAULT_PREPROCESS_BATCH_SIZE = 128
DEFAULT_PREPROCESS_NUM_GPUS = 1
DEFAULT_DATASET = "world-r1-enhanced-dynamic"
DEFAULT_REWARD = "videoalign"
DEFAULT_MAX_PROMPTS = "512"
VIDEOALIGN_CKPT_ROOT = f"{MODAL_CACHE_ROOT}/VideoReward"
MODAL_CONTEXT_IGNORE = [
".git/fsmonitor--daemon.ipc",
".cache",
".cache/**",
"__pycache__",
"**/__pycache__/**",
"*.pyc",
".mypy_cache",
".ruff_cache",
".pytest_cache",
"logs",
"logs/**",
"outputs",
"outputs/**",
"wandb",
"wandb/**",
]
data_vol = modal.Volume.from_name("fastvideo-data")
runs_vol = modal.Volume.from_name("fastvideo-runs")
cache_vol = modal.Volume.from_name("fastvideo-cache")
image = (modal.Image.from_registry(
"nvidia/cuda:12.8.1-devel-ubuntu22.04",
add_python="3.12",
).entrypoint([]).apt_install(
"git",
"git-lfs",
"ffmpeg",
"libgl1",
"libglib2.0-0",
"build-essential",
"ninja-build",
"cmake",
).pip_install("uv").add_local_dir(
".",
PROJECT_ROOT,
copy=True,
ignore=MODAL_CONTEXT_IGNORE,
).run_commands(
"cd /root/FastVideo && uv pip install --system --prerelease=allow -e .",
"uv pip install --system --prerelease=allow "
"--index-url https://download.pytorch.org/whl/cu128 "
"--upgrade torch torchvision torchaudio",
"uv pip install --system --no-cache-dir "
"https://github.com/mjun0812/flash-attention-prebuild-wheels/"
"releases/download/v0.7.16/"
"flash_attn-2.8.3+cu128torch2.10-cp312-cp312-linux_x86_64.whl",
"cd /root/FastVideo && uv pip install --system "
"-r examples/train/requirements-diffusion-nft.txt",
"python -c 'import flash_attn; print(\"flash_attn ok\")'",
"python -c 'import cv2, imageio; print(\"video io ok\")'",
"python -c 'import cloudpickle, pyarrow, torchdata; "
"print(\"training deps ok\")'",
"python -c 'import datasets, peft, qwen_vl_utils, "
"safetensors, timm; "
"from transformers import Qwen2VLForConditionalGeneration; "
"print(\"reward deps ok\")'",
).env({
"WANDB_MODE": "online",
"WANDB_ENTITY": WANDB_ENTITY,
"TOKENIZERS_PARALLELISM": "false",
"NUM_GPUS": str(DEFAULT_NUM_GPUS),
"HF_HOME": f"{MODAL_CACHE_ROOT}/huggingface",
"HF_HUB_CACHE": f"{MODAL_CACHE_ROOT}/huggingface",
"TRANSFORMERS_CACHE": f"{MODAL_CACHE_ROOT}/huggingface",
"DIFFUSION_NFT_ROOT": DIFFUSION_NFT_ROOT,
"VIDEOALIGN_CHECKPOINT_PATH": VIDEOALIGN_CKPT_ROOT,
"FASTVIDEO_ATTENTION_BACKEND": "FLASH_ATTN",
"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True",
}))
@app.function(
image=image,
gpu=f"{DEFAULT_GPU_TYPE}:{DEFAULT_NUM_GPUS}",
timeout=8 * 60 * 60,
volumes={
MODAL_DATA_ROOT: data_vol,
f"{PROJECT_ROOT}/outputs": runs_vol,
MODAL_CACHE_ROOT: cache_vol,
},
secrets=[
modal.Secret.from_name(WANDB_SECRET_NAME),
modal.Secret.from_name(HF_SECRET_NAME),
],
)
def train(
max_train_steps: int = DEFAULT_MAX_TRAIN_STEPS,
num_samples_per_prompt: int = DEFAULT_NUM_SAMPLES_PER_PROMPT,
num_batches_per_epoch: int = DEFAULT_NUM_BATCHES_PER_EPOCH,
collection_batch_size: int = DEFAULT_COLLECTION_BATCH_SIZE,
inner_epochs: int = DEFAULT_INNER_EPOCHS,
train_batch_size: int = DEFAULT_TRAIN_BATCH_SIZE,
gradient_accumulation_steps: int = DEFAULT_GRADIENT_ACCUMULATION_STEPS,
learning_rate: float = DEFAULT_LEARNING_RATE,
sample_num_steps: int = DEFAULT_SAMPLE_NUM_STEPS,
sample_flow_shift: float = DEFAULT_SAMPLE_FLOW_SHIFT,
sample_guidance_scale: float = DEFAULT_SAMPLE_GUIDANCE_SCALE,
validation_num_steps: int = DEFAULT_VALIDATION_NUM_STEPS,
validation_num_prompts: int = DEFAULT_VALIDATION_NUM_PROMPTS,
validation_batch_size: int = DEFAULT_VALIDATION_BATCH_SIZE,
validation_num_frames: int = DEFAULT_VALIDATION_NUM_FRAMES,
num_frames: int = DEFAULT_NUM_FRAMES,
num_latent_t: int = DEFAULT_NUM_LATENT_T,
log_sample_max_videos: int = DEFAULT_LOG_SAMPLE_MAX_VIDEOS,
preprocess_batch_size: int = DEFAULT_PREPROCESS_BATCH_SIZE,
preprocess_num_gpus: int = DEFAULT_PREPROCESS_NUM_GPUS,
dataset: str = DEFAULT_DATASET,
reward: str = DEFAULT_REWARD,
max_prompts: str = DEFAULT_MAX_PROMPTS,
check_rewards: bool = True,
):
from datetime import datetime, timezone
import json
import subprocess
import sys
from pathlib import Path
repo = Path(PROJECT_ROOT)
dataset = dataset.strip().lower()
reward = reward.strip().lower()
preprocess_batch_size = int(preprocess_batch_size)
preprocess_num_gpus = int(preprocess_num_gpus)
log_sample_max_videos = int(log_sample_max_videos)
num_batches_per_epoch = int(num_batches_per_epoch)
learning_rate = float(learning_rate)
sample_num_steps = int(sample_num_steps)
sample_flow_shift = float(sample_flow_shift)
sample_guidance_scale = float(sample_guidance_scale)
validation_num_steps = int(validation_num_steps)
validation_num_prompts = int(validation_num_prompts)
validation_batch_size = int(validation_batch_size)
validation_num_frames = int(validation_num_frames)
if preprocess_batch_size <= 0:
raise ValueError("--preprocess-batch-size must be positive")
if learning_rate < 0.0 and learning_rate != DEFAULT_LEARNING_RATE:
raise ValueError("--learning-rate must be >= 0")
if sample_num_steps < 0:
raise ValueError("--sample-num-steps must be >= 0")
if num_batches_per_epoch < 0:
raise ValueError("--num-batches-per-epoch must be >= 0")
if sample_flow_shift < 0.0 and sample_flow_shift != DEFAULT_SAMPLE_FLOW_SHIFT:
raise ValueError("--sample-flow-shift must be >= 0")
if (sample_guidance_scale < 0.0
and sample_guidance_scale != DEFAULT_SAMPLE_GUIDANCE_SCALE):
raise ValueError("--sample-guidance-scale must be >= 0")
if validation_num_steps < 0:
raise ValueError("--validation-num-steps must be >= 0")
if validation_num_prompts < 0:
raise ValueError("--validation-num-prompts must be >= 0")
if validation_batch_size < 0:
raise ValueError("--validation-batch-size must be >= 0")
if validation_num_frames < 0:
raise ValueError("--validation-num-frames must be >= 0")
if validation_num_frames > 0 and (validation_num_frames - 1) % 4 != 0:
raise ValueError("Wan validation frame counts must satisfy "
"num_frames = (num_latent_t - 1) * 4 + 1; "
f"got {validation_num_frames}")
if preprocess_num_gpus != 1:
raise ValueError("FastVideo text preprocessing currently supports "
"--preprocess-num-gpus 1 only.")
if DEFAULT_HSDP_REPLICATE_DIM * DEFAULT_HSDP_SHARD_DIM != DEFAULT_NUM_GPUS:
raise ValueError(
"Invalid HSDP mesh: replicate_dim * shard_dim must equal "
f"num_gpus ({DEFAULT_HSDP_REPLICATE_DIM} * "
f"{DEFAULT_HSDP_SHARD_DIM} != {DEFAULT_NUM_GPUS}).")
if (DEFAULT_NUM_GPUS * collection_batch_size) % num_samples_per_prompt != 0:
raise ValueError("DiffusionNFT K-repeat sampling requires "
"num_gpus * collection_batch_size to be divisible by "
"--num-samples-per-prompt "
f"({DEFAULT_NUM_GPUS} * {collection_batch_size} vs "
f"{num_samples_per_prompt}).")
if (DEFAULT_NUM_GPUS * train_batch_size) % num_samples_per_prompt != 0:
raise ValueError("DiffusionNFT training batches should keep full prompt "
"repeat groups: num_gpus * train_batch_size must be "
"divisible by --num-samples-per-prompt "
f"({DEFAULT_NUM_GPUS} * {train_batch_size} vs "
f"{num_samples_per_prompt}).")
output_dir = (f"{OUTPUT_DIR_BASE}_"
f"{datetime.now(timezone.utc).strftime('%Y%m%d_%H%M%S')}")
prep_cmd = [
"python",
"examples/train/prepare_diffusion_nft_assets.py",
"--repo-root",
str(repo),
"--config",
CONFIG_PATH,
"--data-root",
MODAL_DATA_ROOT,
"--cache-root",
MODAL_CACHE_ROOT,
"--output-dir",
output_dir,
"--run-config-dir",
f"{PROJECT_ROOT}/outputs/diffusion_nft_run_configs",
"--diffusion-nft-root",
DIFFUSION_NFT_ROOT,
"--videoalign-checkpoint-path",
VIDEOALIGN_CKPT_ROOT,
"--dataset",
dataset,
"--reward",
reward,
"--max-prompts",
str(max_prompts),
"--num-frames",
str(num_frames),
"--num-latent-t",
str(num_latent_t),
"--num-gpus",
str(DEFAULT_NUM_GPUS),
"--hsdp-replicate-dim",
str(DEFAULT_HSDP_REPLICATE_DIM),
"--hsdp-shard-dim",
str(DEFAULT_HSDP_SHARD_DIM),
"--max-train-steps",
str(max_train_steps),
"--gradient-accumulation-steps",
str(gradient_accumulation_steps),
"--num-samples-per-prompt",
str(num_samples_per_prompt),
"--collection-batch-size",
str(collection_batch_size),
"--inner-epochs",
str(inner_epochs),
"--train-batch-size",
str(train_batch_size),
"--log-sample-max-videos",
str(log_sample_max_videos),
"--preprocess-batch-size",
str(preprocess_batch_size),
"--preprocess-num-gpus",
str(preprocess_num_gpus),
"--json",
]
if check_rewards:
prep_cmd.append("--check-rewards")
if learning_rate >= 0.0:
prep_cmd.extend(["--learning-rate", str(learning_rate)])
if sample_num_steps > 0:
prep_cmd.extend(["--sample-num-steps", str(sample_num_steps)])
if sample_flow_shift >= 0.0:
prep_cmd.extend(["--sample-flow-shift", str(sample_flow_shift)])
if sample_guidance_scale >= 0.0:
prep_cmd.extend(["--sample-guidance-scale", str(sample_guidance_scale)])
print("Preparing DiffusionNFT assets with tracked example CLI:", flush=True)
completed = subprocess.run(
prep_cmd,
cwd=repo,
stdout=subprocess.PIPE,
stderr=sys.stderr,
text=True,
check=True,
)
print(completed.stdout, end="", flush=True)
summary = json.loads(completed.stdout.strip().splitlines()[-1])
run_config_path = Path(summary["run_config"])
parquet_dir = Path(summary["parquet_dir"])
output_dir = summary["output_dir"]
resolved_num_frames = int(summary["num_frames"])
resolved_num_latent_t = int(summary["num_latent_t"])
resolved_validation_num_frames = validation_num_frames or resolved_num_frames
resolved_validation_num_latent_t = ((resolved_validation_num_frames - 1) // 4 + 1)
cache_vol.commit()
data_vol.commit()
print(f"Output dir: {output_dir}", flush=True)
print(
"DiffusionNFT probe settings: "
f"max_train_steps={max_train_steps} "
f"num_samples_per_prompt={num_samples_per_prompt} "
f"num_batches_per_epoch={num_batches_per_epoch if num_batches_per_epoch > 0 else 'config'} "
f"collection_batch_size={collection_batch_size} "
f"inner_epochs={inner_epochs} "
f"train_batch_size={train_batch_size} "
f"gradient_accumulation_steps={gradient_accumulation_steps} "
f"learning_rate={learning_rate if learning_rate >= 0 else 'config'} "
f"sample_num_steps={sample_num_steps if sample_num_steps > 0 else 'config'} "
f"sample_flow_shift={sample_flow_shift if sample_flow_shift >= 0 else 'config'} "
f"sample_guidance_scale={sample_guidance_scale if sample_guidance_scale >= 0 else 'config'} "
f"validation_num_steps={validation_num_steps if validation_num_steps > 0 else 'config'} "
f"validation_num_prompts={validation_num_prompts if validation_num_prompts > 0 else 'config'} "
f"validation_batch_size={validation_batch_size if validation_batch_size > 0 else 'config'} "
f"validation_num_frames={resolved_validation_num_frames} "
f"validation_num_latent_t={resolved_validation_num_latent_t} "
f"num_frames={resolved_num_frames} "
f"num_latent_t={resolved_num_latent_t} "
f"log_sample_max_videos={log_sample_max_videos} "
f"preprocess_batch_size={preprocess_batch_size} "
f"preprocess_num_gpus={preprocess_num_gpus} "
f"dataset={dataset} "
f"reward={reward} "
f"max_prompts={max_prompts} "
f"resolved_prompts={summary['prompt_count']}",
flush=True,
)
cmd = [
"bash",
"examples/train/run.sh",
str(run_config_path),
"--training.data.data_path",
str(parquet_dir),
"--training.checkpoint.output_dir",
output_dir,
"--training.tracker.run_name",
Path(output_dir).name,
"--training.loop.max_train_steps",
str(int(max_train_steps)),
"--training.distributed.num_gpus",
str(DEFAULT_NUM_GPUS),
"--training.distributed.hsdp_replicate_dim",
str(DEFAULT_HSDP_REPLICATE_DIM),
"--training.distributed.hsdp_shard_dim",
str(DEFAULT_HSDP_SHARD_DIM),
"--training.loop.gradient_accumulation_steps",
str(int(gradient_accumulation_steps)),
"--training.data.num_frames",
str(resolved_num_frames),
"--training.data.num_latent_t",
str(resolved_num_latent_t),
"--method.num_video_per_prompt",
str(int(num_samples_per_prompt)),
"--method.sample_train_batch_size",
str(int(collection_batch_size)),
"--method.num_inner_epochs",
str(int(inner_epochs)),
"--method.train_batch_size",
str(int(train_batch_size)),
]
if validation_num_steps > 0:
cmd.extend(["--method.validation.num_steps", str(validation_num_steps)])
if validation_num_prompts > 0:
cmd.extend(["--method.validation.num_prompts", str(validation_num_prompts)])
if validation_batch_size > 0:
cmd.extend(["--method.validation.batch_size", str(validation_batch_size)])
cmd.extend(["--method.validation.num_latent_t", str(resolved_validation_num_latent_t)])
if num_batches_per_epoch > 0:
cmd.extend(["--method.num_batches_per_epoch", str(num_batches_per_epoch)])
subprocess.run(["nvidia-smi"], check=True)
subprocess.run(
cmd,
cwd=repo,
stdout=sys.stdout,
stderr=sys.stderr,
check=True,
)
runs_vol.commit()
cache_vol.commit()
@app.local_entrypoint()
def main(
max_train_steps: int = DEFAULT_MAX_TRAIN_STEPS,
num_samples_per_prompt: int = DEFAULT_NUM_SAMPLES_PER_PROMPT,
num_batches_per_epoch: int = DEFAULT_NUM_BATCHES_PER_EPOCH,
collection_batch_size: int = DEFAULT_COLLECTION_BATCH_SIZE,
inner_epochs: int = DEFAULT_INNER_EPOCHS,
train_batch_size: int = DEFAULT_TRAIN_BATCH_SIZE,
gradient_accumulation_steps: int = DEFAULT_GRADIENT_ACCUMULATION_STEPS,
learning_rate: float = DEFAULT_LEARNING_RATE,
sample_num_steps: int = DEFAULT_SAMPLE_NUM_STEPS,
sample_flow_shift: float = DEFAULT_SAMPLE_FLOW_SHIFT,
sample_guidance_scale: float = DEFAULT_SAMPLE_GUIDANCE_SCALE,
validation_num_steps: int = DEFAULT_VALIDATION_NUM_STEPS,
validation_num_prompts: int = DEFAULT_VALIDATION_NUM_PROMPTS,
validation_batch_size: int = DEFAULT_VALIDATION_BATCH_SIZE,
validation_num_frames: int = DEFAULT_VALIDATION_NUM_FRAMES,
num_frames: int = DEFAULT_NUM_FRAMES,
num_latent_t: int = DEFAULT_NUM_LATENT_T,
log_sample_max_videos: int = DEFAULT_LOG_SAMPLE_MAX_VIDEOS,
preprocess_batch_size: int = DEFAULT_PREPROCESS_BATCH_SIZE,
preprocess_num_gpus: int = DEFAULT_PREPROCESS_NUM_GPUS,
dataset: str = DEFAULT_DATASET,
reward: str = DEFAULT_REWARD,
max_prompts: str = DEFAULT_MAX_PROMPTS,
check_rewards: bool = True,
):
train.spawn(
max_train_steps=max_train_steps,
num_samples_per_prompt=num_samples_per_prompt,
num_batches_per_epoch=num_batches_per_epoch,
collection_batch_size=collection_batch_size,
inner_epochs=inner_epochs,
train_batch_size=train_batch_size,
gradient_accumulation_steps=gradient_accumulation_steps,
learning_rate=learning_rate,
sample_num_steps=sample_num_steps,
sample_flow_shift=sample_flow_shift,
sample_guidance_scale=sample_guidance_scale,
validation_num_steps=validation_num_steps,
validation_num_prompts=validation_num_prompts,
validation_batch_size=validation_batch_size,
validation_num_frames=validation_num_frames,
num_frames=num_frames,
num_latent_t=num_latent_t,
log_sample_max_videos=log_sample_max_videos,
preprocess_batch_size=preprocess_batch_size,
preprocess_num_gpus=preprocess_num_gpus,
dataset=dataset,
reward=reward,
max_prompts=max_prompts,
check_rewards=check_rewards,
)
+337
View File
@@ -0,0 +1,337 @@
#!/usr/bin/env bash
set -euo pipefail
# Submit/run the Modal DiffusionNFT Wan video RL equivalent on a Slurm cluster.
#
# From the login node:
# PARTITION=all NUM_GPUS=4 bash scripts/train/train_diffusion_nft_wan_videoalign_slurm.sh
#
# Useful overrides:
# NUM_FRAMES=17 MAX_TRAIN_STEPS=10 CHECK_REWARDS=0 ...
# bash scripts/train/train_diffusion_nft_wan_videoalign_slurm.sh --method.validation.num_prompts 4
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../.." && pwd)"
ENV_FILE="${ENV_FILE:-${REPO_ROOT}/.env}"
if [[ -f "${ENV_FILE}" ]]; then
set -a
# shellcheck source=/dev/null
source "${ENV_FILE}"
set +a
fi
CONFIG="${CONFIG:-examples/train/configs/rl/wan/diffusion_nft_videoalign.yaml}"
PARTITION="${PARTITION:-all}"
NUM_NODES="${NUM_NODES:-1}"
NUM_GPUS="${NUM_GPUS:-4}"
GRES="${GRES:-gpu:nvidia_g:${NUM_GPUS}}"
CPUS_PER_TASK="${CPUS_PER_TASK:-32}"
MEM="${MEM:-0}"
TIME="${TIME:-08:00:00}"
JOB_NAME="${JOB_NAME:-diffusion_nft_wan_videoalign}"
SLURM_LOG_DIR="${SLURM_LOG_DIR:-logs/slurm}"
MASTER_PORT="${MASTER_PORT:-29531}"
ACCOUNT="${ACCOUNT:-}"
QOS="${QOS:-}"
EXCLUDE="${EXCLUDE:-}"
DRY_RUN="${DRY_RUN:-0}"
NUM_FRAMES="${NUM_FRAMES:-29}"
NUM_LATENT_T="${NUM_LATENT_T:-0}"
if (( (NUM_FRAMES - 1) % 4 != 0 )); then
echo "NUM_FRAMES must satisfy Wan's 4n + 1 rule; got ${NUM_FRAMES}." >&2
echo "29 is the closest supported value to 30." >&2
exit 1
fi
DERIVED_NUM_LATENT_T="$(( (NUM_FRAMES - 1) / 4 + 1 ))"
if (( NUM_LATENT_T > 0 && NUM_LATENT_T != DERIVED_NUM_LATENT_T )); then
echo "NUM_FRAMES=${NUM_FRAMES} implies NUM_LATENT_T=${DERIVED_NUM_LATENT_T}; got ${NUM_LATENT_T}." >&2
exit 1
fi
if [[ "${_FASTVIDEO_DIFFUSION_NFT_SLURM_WORKER:-0}" != "1" ]]; then
mkdir -p "${REPO_ROOT}/${SLURM_LOG_DIR}"
sbatch_args=(
--job-name="${JOB_NAME}"
--partition="${PARTITION}"
--nodes="${NUM_NODES}"
--ntasks="${NUM_NODES}"
--ntasks-per-node=1
--gres="${GRES}"
--cpus-per-task="${CPUS_PER_TASK}"
--mem="${MEM}"
--time="${TIME}"
--output="${REPO_ROOT}/${SLURM_LOG_DIR}/${JOB_NAME}_%j.out"
--error="${REPO_ROOT}/${SLURM_LOG_DIR}/${JOB_NAME}_%j.err"
)
if [[ -n "${ACCOUNT}" ]]; then
sbatch_args+=(--account="${ACCOUNT}")
fi
if [[ -n "${QOS}" ]]; then
sbatch_args+=(--qos="${QOS}")
fi
if [[ -n "${EXCLUDE}" ]]; then
sbatch_args+=(--exclude="${EXCLUDE}")
fi
echo "=== DiffusionNFT Wan Slurm submission ==="
echo "repo: ${REPO_ROOT}"
echo "config: ${CONFIG}"
echo "partition: ${PARTITION}"
echo "nodes: ${NUM_NODES}"
echo "gpus/node: ${NUM_GPUS}"
echo "gres: ${GRES}"
echo "job name: ${JOB_NAME}"
echo "extra args: $*"
echo "========================================="
if [[ "${DRY_RUN}" == "1" ]]; then
printf 'sbatch'
printf ' %q' "${sbatch_args[@]}"
printf ' --export=ALL,_FASTVIDEO_DIFFUSION_NFT_SLURM_WORKER=1 %q' "$0"
if (($# > 0)); then
printf ' %q' "$@"
fi
echo
exit 0
fi
sbatch "${sbatch_args[@]}" \
--export=ALL,_FASTVIDEO_DIFFUSION_NFT_SLURM_WORKER=1 \
"$0" "$@"
exit 0
fi
cd "${REPO_ROOT}"
if [[ "${CONFIG}" = /* ]]; then
CONFIG_PATH="${CONFIG}"
else
CONFIG_PATH="${REPO_ROOT}/${CONFIG}"
fi
if [[ ! -f "${CONFIG_PATH}" ]]; then
echo "Training config not found: ${CONFIG_PATH}" >&2
exit 1
fi
if [[ -n "${CONDA_ENV_PATH:-}" ]]; then
if [[ ! -x "${CONDA_ENV_PATH}/bin/python" ]]; then
echo "CONDA_ENV_PATH was set, but no python was found at ${CONDA_ENV_PATH}/bin/python" >&2
exit 1
fi
export PATH="${CONDA_ENV_PATH}/bin:${PATH}"
elif [[ -n "${CONDA_ROOT:-}" ]]; then
if [[ ! -f "${CONDA_ROOT}/etc/profile.d/conda.sh" ]]; then
echo "CONDA_ROOT was set, but conda.sh was not found under ${CONDA_ROOT}" >&2
exit 1
fi
# shellcheck source=/dev/null
source "${CONDA_ROOT}/etc/profile.d/conda.sh"
conda activate "${CONDA_ENV:-fastvideo}"
fi
if [[ "${INSTALL_DEPS:-0}" == "1" ]]; then
uv pip install --prerelease=allow -e .
uv pip install --prerelease=allow -r examples/train/requirements-diffusion-nft.txt
fi
RUN_ID="${RUN_ID:-$(date -u +%Y%m%d_%H%M%S)}"
OUTPUT_ROOT="${OUTPUT_ROOT:-${REPO_ROOT}/outputs/diffusion_nft_wan}"
OUTPUT_DIR="${OUTPUT_DIR:-${OUTPUT_ROOT}_${RUN_ID}}"
DATA_ROOT="${DATA_ROOT:-${REPO_ROOT}/.slurm_data}"
CACHE_PARENT="${SCRATCH:-${REPO_ROOT}/.slurm_cache}"
CACHE_ROOT="${CACHE_ROOT:-${CACHE_PARENT}/diffusion_nft}"
RUN_CONFIG_DIR="${RUN_CONFIG_DIR:-${REPO_ROOT}/outputs/diffusion_nft_run_configs/${RUN_ID}}"
DIFFUSION_NFT_ROOT="${DIFFUSION_NFT_ROOT:-${CACHE_ROOT}/DiffusionNFT}"
VIDEOALIGN_CHECKPOINT_PATH="${VIDEOALIGN_CHECKPOINT_PATH:-${CACHE_ROOT}/VideoReward}"
DATASET="${DATASET:-world-r1-enhanced-dynamic}"
REWARD="${REWARD:-videoalign}"
MAX_PROMPTS="${MAX_PROMPTS:-512}"
MAX_TRAIN_STEPS="${MAX_TRAIN_STEPS:-50}"
NUM_SAMPLES_PER_PROMPT="${NUM_SAMPLES_PER_PROMPT:-4}"
NUM_BATCHES_PER_EPOCH="${NUM_BATCHES_PER_EPOCH:-1}"
COLLECTION_BATCH_SIZE="${COLLECTION_BATCH_SIZE:-4}"
INNER_EPOCHS="${INNER_EPOCHS:-1}"
TRAIN_BATCH_SIZE="${TRAIN_BATCH_SIZE:-4}"
GRADIENT_ACCUMULATION_STEPS="${GRADIENT_ACCUMULATION_STEPS:-1}"
LEARNING_RATE="${LEARNING_RATE:--1}"
SAMPLE_NUM_STEPS="${SAMPLE_NUM_STEPS:-50}"
SAMPLE_FLOW_SHIFT="${SAMPLE_FLOW_SHIFT:--1}"
SAMPLE_GUIDANCE_SCALE="${SAMPLE_GUIDANCE_SCALE:--1}"
VALIDATION_NUM_STEPS="${VALIDATION_NUM_STEPS:-50}"
VALIDATION_NUM_PROMPTS="${VALIDATION_NUM_PROMPTS:-16}"
VALIDATION_BATCH_SIZE="${VALIDATION_BATCH_SIZE:-4}"
LOG_SAMPLE_MAX_VIDEOS="${LOG_SAMPLE_MAX_VIDEOS:-16}"
PREPROCESS_BATCH_SIZE="${PREPROCESS_BATCH_SIZE:-128}"
PREPROCESS_NUM_GPUS="${PREPROCESS_NUM_GPUS:-1}"
PREPROCESS_MASTER_PORT="${PREPROCESS_MASTER_PORT:-29541}"
DATALOADER_NUM_WORKERS="${DATALOADER_NUM_WORKERS:-0}"
PYTHON_BIN="${PYTHON_BIN:-python3}"
SP_SIZE="${SP_SIZE:-1}"
TP_SIZE="${TP_SIZE:-1}"
HSDP_REPLICATE_DIM="${HSDP_REPLICATE_DIM:-1}"
HSDP_SHARD_DIM="${HSDP_SHARD_DIM:-$((NUM_NODES * NUM_GPUS))}"
PROJECT_NAME="${PROJECT_NAME:-diffusion_nft_wan}"
RUN_NAME="${RUN_NAME:-wan2.1_diffusion_nft_videoalign_${RUN_ID}}"
CHECK_REWARDS="${CHECK_REWARDS:-1}"
REWARD_DEVICE="${REWARD_DEVICE:-auto}"
TOTAL_GPUS=$((NUM_NODES * NUM_GPUS))
if (( HSDP_REPLICATE_DIM * HSDP_SHARD_DIM != TOTAL_GPUS )); then
echo "Invalid HSDP mesh: HSDP_REPLICATE_DIM * HSDP_SHARD_DIM must equal total GPUs." >&2
echo "Got ${HSDP_REPLICATE_DIM} * ${HSDP_SHARD_DIM} != ${TOTAL_GPUS}" >&2
exit 1
fi
if (( (TOTAL_GPUS * COLLECTION_BATCH_SIZE) % NUM_SAMPLES_PER_PROMPT != 0 )); then
echo "K-repeat sampling requires TOTAL_GPUS * COLLECTION_BATCH_SIZE divisible by NUM_SAMPLES_PER_PROMPT." >&2
exit 1
fi
if (( (TOTAL_GPUS * TRAIN_BATCH_SIZE) % NUM_SAMPLES_PER_PROMPT != 0 )); then
echo "Training batches should keep full prompt groups: TOTAL_GPUS * TRAIN_BATCH_SIZE must be divisible by NUM_SAMPLES_PER_PROMPT." >&2
exit 1
fi
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export WANDB_MODE="${WANDB_MODE:-online}"
export WANDB_BASE_URL="${WANDB_BASE_URL:-https://api.wandb.ai}"
export FASTVIDEO_ATTENTION_BACKEND="${FASTVIDEO_ATTENTION_BACKEND:-FLASH_ATTN}"
export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-/tmp/triton_cache_diffusion_nft_wan_${SLURM_JOB_ID:-manual}}"
export HF_HOME="${HF_HOME:-${CACHE_ROOT}/huggingface}"
export HF_HUB_CACHE="${HF_HUB_CACHE:-${HF_HOME}}"
export TRANSFORMERS_CACHE="${TRANSFORMERS_CACHE:-${HF_HOME}}"
export DIFFUSION_NFT_ROOT
export VIDEOALIGN_CHECKPOINT_PATH
mkdir -p "${OUTPUT_DIR}" "${DATA_ROOT}" "${CACHE_ROOT}" "${RUN_CONFIG_DIR}" logs/train
echo "=== DiffusionNFT Wan Slurm job ==="
echo "host: $(hostname)"
echo "slurm job: ${SLURM_JOB_ID:-manual}"
echo "nodes / gpus per node: ${NUM_NODES} / ${NUM_GPUS}"
echo "total gpus: ${TOTAL_GPUS}"
echo "output dir: ${OUTPUT_DIR}"
echo "data root: ${DATA_ROOT}"
echo "cache root: ${CACHE_ROOT}"
echo "dataset / reward: ${DATASET} / ${REWARD}"
echo "frames / latent T: ${NUM_FRAMES} / ${DERIVED_NUM_LATENT_T} (29 frames is the closest Wan-supported value to 30)"
echo "validation steps/prompts/batch: ${VALIDATION_NUM_STEPS} / ${VALIDATION_NUM_PROMPTS} / ${VALIDATION_BATCH_SIZE}"
echo "=================================="
nvidia-smi || true
prep_cmd=(
"${PYTHON_BIN}" examples/train/prepare_diffusion_nft_assets.py
--repo-root "${REPO_ROOT}"
--config "${CONFIG_PATH}"
--data-root "${DATA_ROOT}"
--cache-root "${CACHE_ROOT}"
--output-dir "${OUTPUT_DIR}"
--run-config-dir "${RUN_CONFIG_DIR}"
--diffusion-nft-root "${DIFFUSION_NFT_ROOT}"
--videoalign-checkpoint-path "${VIDEOALIGN_CHECKPOINT_PATH}"
--dataset "${DATASET}"
--reward "${REWARD}"
--max-prompts "${MAX_PROMPTS}"
--num-frames "${NUM_FRAMES}"
--num-latent-t "${DERIVED_NUM_LATENT_T}"
--num-gpus "${TOTAL_GPUS}"
--sp-size "${SP_SIZE}"
--tp-size "${TP_SIZE}"
--hsdp-replicate-dim "${HSDP_REPLICATE_DIM}"
--hsdp-shard-dim "${HSDP_SHARD_DIM}"
--max-train-steps "${MAX_TRAIN_STEPS}"
--gradient-accumulation-steps "${GRADIENT_ACCUMULATION_STEPS}"
--num-samples-per-prompt "${NUM_SAMPLES_PER_PROMPT}"
--collection-batch-size "${COLLECTION_BATCH_SIZE}"
--inner-epochs "${INNER_EPOCHS}"
--train-batch-size "${TRAIN_BATCH_SIZE}"
--log-sample-max-videos "${LOG_SAMPLE_MAX_VIDEOS}"
--preprocess-batch-size "${PREPROCESS_BATCH_SIZE}"
--preprocess-num-gpus "${PREPROCESS_NUM_GPUS}"
--preprocess-master-port "${PREPROCESS_MASTER_PORT}"
--dataloader-num-workers "${DATALOADER_NUM_WORKERS}"
--project-name "${PROJECT_NAME}"
--run-name "${RUN_NAME}"
--json
)
if [[ "${CHECK_REWARDS}" == "1" ]]; then
prep_cmd+=(--check-rewards --reward-device "${REWARD_DEVICE}")
fi
if awk "BEGIN {exit !(${LEARNING_RATE} >= 0)}"; then
prep_cmd+=(--learning-rate "${LEARNING_RATE}")
fi
if (( SAMPLE_NUM_STEPS > 0 )); then
prep_cmd+=(--sample-num-steps "${SAMPLE_NUM_STEPS}")
fi
if awk "BEGIN {exit !(${SAMPLE_FLOW_SHIFT} >= 0)}"; then
prep_cmd+=(--sample-flow-shift "${SAMPLE_FLOW_SHIFT}")
fi
if awk "BEGIN {exit !(${SAMPLE_GUIDANCE_SCALE} >= 0)}"; then
prep_cmd+=(--sample-guidance-scale "${SAMPLE_GUIDANCE_SCALE}")
fi
echo "Preparing DiffusionNFT assets:"
printf ' %q' "${prep_cmd[@]}"
echo
"${prep_cmd[@]}"
RUN_CONFIG_PATH="${RUN_CONFIG_DIR}/diffusion_nft_wan_run.yaml"
if [[ ! -f "${RUN_CONFIG_PATH}" ]]; then
echo "Prepared run config not found: ${RUN_CONFIG_PATH}" >&2
exit 1
fi
nodes=( $(scontrol show hostnames "${SLURM_JOB_NODELIST:-$(hostname)}") )
MASTER_ADDR="${MASTER_ADDR:-${nodes[0]}}"
export MASTER_ADDR MASTER_PORT
train_args=(
--config "${RUN_CONFIG_PATH}"
--training.checkpoint.output_dir "${OUTPUT_DIR}"
--training.tracker.run_name "${RUN_NAME}"
--training.loop.max_train_steps "${MAX_TRAIN_STEPS}"
--training.distributed.num_gpus "${TOTAL_GPUS}"
--training.distributed.sp_size "${SP_SIZE}"
--training.distributed.tp_size "${TP_SIZE}"
--training.distributed.hsdp_replicate_dim "${HSDP_REPLICATE_DIM}"
--training.distributed.hsdp_shard_dim "${HSDP_SHARD_DIM}"
--training.loop.gradient_accumulation_steps "${GRADIENT_ACCUMULATION_STEPS}"
--training.data.num_frames "${NUM_FRAMES}"
--training.data.num_latent_t "${DERIVED_NUM_LATENT_T}"
--method.num_video_per_prompt "${NUM_SAMPLES_PER_PROMPT}"
--method.sample_train_batch_size "${COLLECTION_BATCH_SIZE}"
--method.num_inner_epochs "${INNER_EPOCHS}"
--method.train_batch_size "${TRAIN_BATCH_SIZE}"
)
if (( VALIDATION_NUM_STEPS > 0 )); then
train_args+=(--method.validation.num_steps "${VALIDATION_NUM_STEPS}")
fi
if (( VALIDATION_NUM_PROMPTS > 0 )); then
train_args+=(--method.validation.num_prompts "${VALIDATION_NUM_PROMPTS}")
fi
if (( VALIDATION_BATCH_SIZE > 0 )); then
train_args+=(--method.validation.batch_size "${VALIDATION_BATCH_SIZE}")
fi
if (( NUM_BATCHES_PER_EPOCH > 0 )); then
train_args+=(--method.num_batches_per_epoch "${NUM_BATCHES_PER_EPOCH}")
fi
train_cmd=(
srun
--ntasks "${NUM_NODES}"
--ntasks-per-node 1
bash -c
'torchrun --nnodes "$1" --nproc_per_node "$2" --node_rank "$SLURM_PROCID" --rdzv_backend c10d --rdzv_endpoint "$3" -m fastvideo.train.entrypoint.train "${@:4}"'
bash
"${NUM_NODES}"
"${NUM_GPUS}"
"${MASTER_ADDR}:${MASTER_PORT}"
"${train_args[@]}"
)
echo "Launching training:"
printf ' %q' "${train_cmd[@]}" "$@"
echo
"${train_cmd[@]}" "$@"
@@ -15,6 +15,34 @@ class _FakeEMA:
self.updates += 1
class _FakeTracker:
def __init__(self):
self.videos = []
self.artifacts = []
def video(self, data, *, caption=None, fps=None):
self.videos.append({
"data": data,
"caption": caption,
"fps": fps,
})
return f"video-{len(self.videos)}"
def log_artifacts(self, artifacts, step):
self.artifacts.append((artifacts, step))
class _FakeValidationStudent:
def __init__(self):
self.kwargs = None
def prepare_batch(self, raw_batch, **kwargs):
self.kwargs = kwargs
return raw_batch
def test_reward_diagnostic_metrics_match_per_prompt_groups():
method = object.__new__(DiffusionNFTMethod)
method._trained_prompt_hashes = set()
@@ -56,6 +84,57 @@ def test_update_ema_honors_update_after_step():
assert method._ema_update_count == 2
def test_log_validation_samples_caps_video_count_and_uses_configured_fps():
method = object.__new__(DiffusionNFTMethod)
method._validation_config = SimpleNamespace(fps=16, max_samples=1)
tracker = _FakeTracker()
method.tracker = tracker
method._log_validation_samples(
[
{
"index": 1,
"prompt": "second",
"media": torch.ones(3, 2, 4, 5),
"rewards": {
"avg": 0.2,
},
},
{
"index": 0,
"prompt": "first",
"media": torch.ones(3, 2, 4, 5),
"rewards": {
"avg": 0.1,
},
},
],
iteration=10,
)
assert len(tracker.videos) == 1
assert tracker.videos[0]["fps"] == 16
assert "first" in tracker.videos[0]["caption"]
assert tracker.artifacts == [({"validation/videos": ["video-1"]}, 10)]
def test_prepare_validation_batch_uses_validation_only_temporal_length():
method = object.__new__(DiffusionNFTMethod)
method.student = _FakeValidationStudent()
method._validation_config = SimpleNamespace(num_latent_t=13)
generator = torch.Generator().manual_seed(0)
raw_batch = {"prompt": ["test"]}
result = method._prepare_validation_batch(raw_batch, generator)
assert result is raw_batch
assert method.student.kwargs == {
"generator": generator,
"latents_source": "zeros",
"num_latent_t": 13,
}
def test_num_train_timesteps_uses_explicit_schedule_length():
method = object.__new__(DiffusionNFTMethod)
method._sample_steps = 25
+57 -2
View File
@@ -1,7 +1,13 @@
import torch
import pytest
import torch
from fastvideo.train.methods.rl.rewards import MultiRewardScorer, select_first_frame
from fastvideo.train.methods.rl.rewards import (
MultiRewardScorer,
build_multi_reward_scorer,
media_to_uint8_array,
normalize_reward_weights,
select_first_frame,
)
def test_select_first_frame_for_video_tensor():
@@ -47,6 +53,55 @@ def test_multi_reward_weighted_sum_with_injected_scorers():
torch.testing.assert_close(scores["avg"], torch.tensor([3.5, 8.5]))
def test_build_multi_reward_accepts_nested_diffusion_nft_config():
scorer = build_multi_reward_scorer(
{"rewards": {
"pickscore": 2.0,
}},
device="cpu",
scorers={"pickscore": lambda media, prompts: torch.tensor([1.5, 2.5])},
)
scores = scorer(torch.zeros(2, 3, 4, 5, 6), ["a", "b"])
torch.testing.assert_close(scores["avg"], torch.tensor([3.0, 5.0]))
def test_normalize_reward_weights_reports_nested_backend():
weights, backend = normalize_reward_weights({
"backend": "genrl",
"rewards": {
"videoalign_vq": 0.75,
"hpsv3_general": 0.25,
},
})
assert backend == "genrl"
assert weights == {
"videoalign_vq": 0.75,
"hpsv3_general": 0.25,
}
def test_media_to_uint8_array_converts_video_tensor_to_nfhwc():
media = torch.zeros(2, 3, 4, 5, 6)
media[:, 0] = 1.0
array = media_to_uint8_array(media)
assert array.shape == (2, 4, 5, 6, 3)
assert array.dtype.name == "uint8"
assert array[..., 0].max() == 255
def test_build_multi_reward_instantiates_debug_reward_without_device_arg():
scorer = build_multi_reward_scorer({"mean_luminance": 1.0}, device="cpu")
scores = scorer(torch.ones(2, 3, 4, 5, 6), ["a", "b"])
torch.testing.assert_close(scores["avg"], torch.ones(2))
def test_multi_reward_validates_score_shape():
scorer = MultiRewardScorer(
{"pickscore": 1.0},
+91 -4
View File
@@ -4,6 +4,7 @@ import pytest
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.methods.rl.common import (
DiffusionSampler,
RLValidationConfig,
SamplingConfig,
distributed_k_repeat_indices,
media_to_video_array,
@@ -50,18 +51,23 @@ class _FakeModel:
self.noise_scheduler = _FakeScheduler()
self.add_noise_calls = 0
self.timestep_shapes = []
self.conditional_calls = []
def predict_noise(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
del attn_kind
self.conditional_calls.append(bool(conditional))
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return torch.zeros_like(noisy_latents)
value = 1.0 if conditional else -1.0
return torch.full_like(noisy_latents, value)
def predict_x0(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
del attn_kind
self.conditional_calls.append(bool(conditional))
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return noisy_latents
value = 1.0 if conditional else -1.0
return noisy_latents + value
def add_noise(self, clean_latents, noise, timestep):
del timestep
@@ -119,6 +125,17 @@ def test_sampling_config_rejects_unknown_keys():
SamplingConfig.from_mapping({"solver": "dpm2"})
def test_sampling_config_accepts_flow_unipc_and_rejects_explicit_timesteps():
cfg = SamplingConfig.from_mapping({"scheduler": "flow_unipc", "num_steps": 50, "flow_shift": 8})
assert cfg.scheduler == "flow_unipc"
assert cfg.num_steps == 50
assert cfg.flow_shift == 8.0
with pytest.raises(ValueError, match="timesteps is not supported with flow_unipc"):
SamplingConfig.from_mapping({"scheduler": "flow_unipc", "timesteps": [1000, 500, 0]})
def test_sampler_restores_original_batch_timestep_after_sampling():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=2))
@@ -139,6 +156,31 @@ def test_euler_sampler_does_not_renoise_between_steps():
assert model.add_noise_calls == 0
assert model.timestep_shapes == [(2,), (2,), (2,), (2,)]
assert model.conditional_calls == [True, True, True, True]
def test_euler_sampler_applies_cfg_guidance_when_requested():
baseline_model = _FakeModel()
guided_model = _FakeModel()
baseline_sampler = DiffusionSampler(SamplingConfig(num_steps=1))
guided_sampler = DiffusionSampler(SamplingConfig(num_steps=1, guidance_scale=3.0))
baseline = baseline_sampler.sample(
baseline_model,
_batch(),
generator=torch.Generator().manual_seed(0),
)
guided = guided_sampler.sample(
guided_model,
_batch(),
generator=torch.Generator().manual_seed(0),
)
# Fake cond prediction is +1 and uncond is -1, so guidance=3 changes
# the one-step scheduler update from +1 to +5.
assert torch.allclose(guided.latents - baseline.latents, torch.full_like(guided.latents, 4.0))
assert baseline_model.conditional_calls == [True]
assert guided_model.conditional_calls == [True, False]
def test_sde_reflow_sampler_renoises_between_steps():
@@ -174,6 +216,39 @@ def test_diffusion_nft_config_uses_rl_sampler_not_dmd_pipeline():
assert cfg.method["validation"]["log_samples"] is True
def test_diffusion_nft_video_config_uses_genrl_rewards_in_clean_layout():
config_path = "examples/train/configs/rl/wan/diffusion_nft_videoalign.yaml"
cfg = load_run_config(config_path)
raw_text = open(config_path, encoding="utf-8").read()
assert cfg.method["_target_"] == "fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod"
assert cfg.method["reward_backend"] == "genrl"
assert cfg.method["reward_fn"]["rewards"] == {
"videoalign_vq": 1.0,
"videoalign_mq": 1.0,
"videoalign_ta": 1.0,
}
assert cfg.method["sampling"]["scheduler"] == "flow_unipc"
assert cfg.method["sampling"]["trajectory"] == "ode"
assert cfg.method["sampling"]["num_steps"] == 50
assert cfg.method["sampling"]["flow_shift"] == 8.0
assert cfg.method["sampling"]["guidance_scale"] == 6.0
assert cfg.method["validation"]["num_steps"] == 50
assert cfg.method["validation"]["num_prompts"] == 16
assert cfg.method["validation"]["batch_size"] == 4
assert cfg.method["validation"]["max_samples"] == 4
assert cfg.method["validation"]["fps"] == 16
assert cfg.method["beta"] == 0.1
assert cfg.method["kl_beta"] == 0.0001
assert cfg.training.loop.gradient_accumulation_steps == 24
assert cfg.training.data.num_latent_t == 20
assert cfg.training.data.num_frames == 77
assert "rl/reward/" not in raw_text
assert "WanDMDPipeline" not in raw_text
assert "solver" not in cfg.method["sampling"]
def test_validation_shard_indices_are_stable_and_padded():
rank0 = validation_shard_indices(5, rank=0, world_size=2)
rank1 = validation_shard_indices(5, rank=1, world_size=2)
@@ -182,6 +257,18 @@ def test_validation_shard_indices_are_stable_and_padded():
assert rank1 == [(1, True), (3, True), (0, False)]
def test_rl_validation_config_parses_video_shape_and_fps():
config = RLValidationConfig.from_mapping({
"fps": 16,
"max_samples": 2,
"num_latent_t": 13,
})
assert config.fps == 16
assert config.max_samples == 2
assert config.num_latent_t == 13
def test_distributed_k_repeat_indices_repeats_prompts_globally():
rank0 = distributed_k_repeat_indices(
dataset_length=100,