Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
753eed7def | ||
|
|
0f1877b138 | ||
|
|
a033a164f2 | ||
|
|
11aca73fc2 | ||
|
|
0e93752d5d | ||
|
|
9cb5dd6e93 | ||
|
|
a0595973da | ||
|
|
5721beb40a | ||
|
|
ae56104ba0 | ||
|
|
7c8aab0930 | ||
|
|
f0439797a7 | ||
|
|
79edcc7c1e | ||
|
|
aa7bc51507 |
@@ -31,6 +31,7 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
!examples/train/requirements-diffusion-nft.txt
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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}
|
||||
@@ -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)
|
||||
)
|
||||
@@ -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.
|
||||
@@ -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>
|
||||
<a href='https://gongyeliu.github.io/videoalign/'><img src='https://img.shields.io/badge/Project-VideoAlign-green'></a>
|
||||
<a href="https://github.com/KwaiVGI/VideoAlign"><img src="https://img.shields.io/badge/GitHub-VideoAlign-9E95B7?logo=github"></a>
|
||||
<a href='https://huggingface.co/KwaiVGI/VideoReward'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20Model-VideoReward-blue'></a>
|
||||
<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>
|
||||
<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>
|
||||
<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}
|
||||
}
|
||||
```
|
||||
@@ -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
|
||||
@@ -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}
|
||||
@@ -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
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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(
|
||||
|
||||
@@ -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'")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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'")
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user