Files
aigc-apps-VideoX-Fun/scripts/get_rewards_1video.py
T

261 lines
10 KiB
Python

# ---------------------------------------------------------------------------
# This script computes reward scores for video-prompt pairs using any reward
# model defined in videox_fun.reward.reward_fn (e.g. HPSReward, MPSReward,
# VideoAlignReward, HPSv3Reward, etc.).
#
# Argument style is kept identical to train_reward_lora.py so that users can
# directly copy --reward_fn and --reward_fn_kwargs from their training shell
# scripts without modification.
#
# Typical usage (single model):
# 1) AestheticReward (v2.5, custom predictor)
# python scripts/get_rewards_1video.py \
# --reward_fn AestheticReward \
# --reward_fn_kwargs '{"version": "v2.5"}' \
# --video_path asset/1.mp4 --prompts "1girl"
# # Optional: override encoder and predictor paths
# # --reward_fn_kwargs '{"version": "v2.5", "encoder_path": "models/Personalized_Model/CLIP-ViT-H-14-laion2B-s32B-b79K-ViT-H-14-laion2B-s32B-b79K", "predictor_path": "models/Personalized_Model/predictor"}'
#
# 2) HPSReward
# python scripts/get_rewards_1video.py \
# --reward_fn HPSReward \
# --reward_fn_kwargs '{"version": "v2.1"}' \
# --video_path asset/1.mp4 --prompts "1girl"
# # Optional: override model checkpoint path
# # --reward_fn_kwargs '{"version": "v2.1", "model_path": "models/Personalized_Model/HPS_v2.1_compressed.pt"}'
#
# 3) PickScoreReward
# python scripts/get_rewards_1video.py \
# --reward_fn PickScoreReward \
# --reward_fn_kwargs '{}' \
# --video_path asset/1.mp4 --prompts "1girl"
# # Optional: override model and processor paths
# # --reward_fn_kwargs '{"model_path": "models/Personalized_Model/PickScore_v1", "processor_name_or_path": "models/Personalized_Model/CLIP-ViT-H-14-laion2B-s32B-b79K-processor"}'
#
# 4) MPSReward
# python scripts/get_rewards_1video.py \
# --reward_fn MPSReward \
# --reward_fn_kwargs '{}' \
# --video_path asset/1.mp4 --prompts "1girl"
# # Optional: override model and processor paths simultaneously
# # --reward_fn_kwargs '{"model_path": "models/Personalized_Model/MPS_overall.pth", "processor_name_or_path": "models/Personalized_Model/CLIP-ViT-H-14-laion2B-s32B-b79K-processor"}'
#
# 5) HPSv3Reward
# python scripts/get_rewards_1video.py \
# --reward_fn HPSv3Reward \
# --reward_fn_kwargs '{"checkpoint_path": "models/Personalized_Model/HPSv3/HPSv3.safetensors"}' \
# --video_path asset/1.mp4 --prompts "1girl"
# # Optional: override base model and checkpoint paths simultaneously
# # --reward_fn_kwargs '{"checkpoint_path": "models/Personalized_Model/HPSv3/HPSv3.safetensors", "model_name_or_path": "models/Personalized_Model/Qwen2-VL-7B-Instruct"}'
#
# 6) VideoAlignReward
# python scripts/get_rewards_1video.py \
# --reward_fn VideoAlignReward \
# --reward_fn_kwargs '{"model_path": "models/Personalized_Model/VideoReward/", "fps": 16, "reward_dim": "Overall"}' \
# --video_path asset/1.mp4 --prompts "1girl"
# # Optional: override checkpoint dir and base model path simultaneously
# # --reward_fn_kwargs '{"model_path": "models/Personalized_Model/VideoReward/", "model_name_or_path": "models/Personalized_Model/Qwen2-VL-2B-Instruct", "fps": 16, "reward_dim": "Overall"}'
# #
# # reward_dim options:
# # - "VQ" : Visual Quality (clearness, resolution, brightness, color)
# # - "MQ" : Motion Quality (consistency, smoothness, completeness)
# # - "TA" : Text-to-Video Alignment (prompt-content & motion match)
# # - "Overall": Overall Performance = VQ + MQ + TA (sum of the three)
#
# Notes:
# * --num_sampled_frames defaults to -1, meaning ALL frames are fed to the
# reward model. For long videos this can be very slow / VRAM-heavy.
# * Videos are processed one-by-one (batch_size == 1) to keep peak VRAM low.
# * Both --video_path and --prompts accept multiple values; they must have
# the same length and are matched by position.
# ---------------------------------------------------------------------------
import argparse
import json
import os
import sys
import numpy as np
import torch
import torchvision.transforms as transforms
try:
from decord import VideoReader
except ImportError:
from videox_fun.data.utils import AVVideoReader as VideoReader
from einops import rearrange
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, project_root)
import videox_fun.reward.reward_fn as reward_fn
def parse_args():
parser = argparse.ArgumentParser(
description="Compute rewards for videos using a specified reward model. "
"Example: python scripts/get_rewards.py --reward_fn HPSReward "
'--reward_fn_kwargs \'{"version": "v2.1"}\' '
'--video_path video1.mp4 video2.mp4 --prompts "A cat plays with a ball" "A dog runs on grass"'
)
parser.add_argument(
"--reward_fn",
type=str,
default="HPSReward",
help='Reward model class name. Options: AestheticReward, HPSReward, PickScoreReward, '
'MPSReward, HPSv3Reward, VideoAlignReward.',
)
parser.add_argument(
"--reward_fn_kwargs",
type=str,
default=None,
help='JSON string of keyword arguments passed to the reward model constructor. '
'E.g., \'{"version": "v2.1", "device": "cuda"}\'',
)
parser.add_argument(
"--video_path",
type=str,
nargs="+",
required=True,
help="List of video file paths.",
)
parser.add_argument(
"--prompts",
type=str,
nargs="+",
required=True,
help="List of text prompts corresponding to each video.",
)
parser.add_argument(
"--device",
type=str,
default="cuda",
help="Device to run the reward model on (e.g., cuda, cuda:0, cpu).",
)
parser.add_argument(
"--dtype",
type=str,
default="bfloat16",
choices=["float16", "bfloat16", "float32"],
help="Torch dtype for the reward model and input tensors.",
)
parser.add_argument(
"--num_sampled_frames",
type=int,
default=-1,
help="Number of uniformly sampled frames from each video. Set to -1 (default) to use all frames.",
)
args = parser.parse_args()
return args
def parse_video_prompt_pairs(video_paths, prompts):
"""Zip video paths and prompts into (video_path, prompt) tuples."""
if len(video_paths) != len(prompts):
raise ValueError(
f"Number of --video_path ({len(video_paths)}) and --prompts ({len(prompts)}) must be equal."
)
pairs = []
for video_path, prompt in zip(video_paths, prompts):
if not os.path.exists(video_path):
raise FileNotFoundError(f"Video file not found: {video_path}")
pairs.append((video_path, prompt))
return pairs
def sample_frames_from_video(video_path, num_sampled_frames):
"""Sample frames from a video and return a [T, C, H, W] float32 tensor in [0, 1].
If num_sampled_frames is -1, all frames are used.
"""
vr = VideoReader(video_path)
total_frames = len(vr)
if total_frames == 0:
raise ValueError(f"Video has no frames: {video_path}")
if num_sampled_frames == -1:
sampled_frame_indices = np.arange(total_frames, dtype=int)
else:
sampled_frame_indices = np.linspace(0, total_frames, num_sampled_frames, endpoint=False, dtype=int)
sampled_frames = vr.get_batch(sampled_frame_indices).asnumpy() # [T, H, W, C]
to_tensor = transforms.ToTensor()
frames_tensor = torch.stack([to_tensor(frame) for frame in sampled_frames], dim=0) # [T, C, H, W]
return frames_tensor
def main():
args = parse_args()
dtype_map = {
"float16": torch.float16,
"bfloat16": torch.bfloat16,
"float32": torch.float32,
}
dtype = dtype_map[args.dtype]
# Parse reward constructor kwargs
reward_fn_kwargs = {}
if args.reward_fn_kwargs is not None:
reward_fn_kwargs = json.loads(args.reward_fn_kwargs)
# Ensure device and dtype are passed (user can override via reward_fn_kwargs)
reward_fn_kwargs.setdefault("device", args.device)
reward_fn_kwargs.setdefault("dtype", dtype)
# Parse video:prompt pairs
pairs = parse_video_prompt_pairs(args.video_path, args.prompts)
video_paths = [p[0] for p in pairs]
prompts = [p[1] for p in pairs]
# Sample frames from each video
frames_list = []
for video_path in video_paths:
frames = sample_frames_from_video(video_path, args.num_sampled_frames)
frames_list.append(frames)
# Stack to [B, T, C, H, W] then rearrange to [B, C, T, H, W] as required by reward_fn
batch_frames = torch.stack(frames_list, dim=0)
batch_frames = rearrange(batch_frames, "b t c h w -> b c t h w")
# Instantiate reward model
if not hasattr(reward_fn, args.reward_fn):
raise ValueError(
f"Reward function '{args.reward_fn}' not found in videox_fun.reward.reward_fn. "
f"Available: {reward_fn.__all__}"
)
reward_model_cls = getattr(reward_fn, args.reward_fn)
print(f"Loading reward model: {args.reward_fn} with kwargs {reward_fn_kwargs}")
reward_model = reward_model_cls(**reward_fn_kwargs)
# Process one video at a time to minimize VRAM usage
batch_size = 1
num_batches = len(pairs)
all_per_sample_rewards = []
total_loss = 0.0
total_reward = 0.0
with torch.no_grad():
for b in range(num_batches):
start = b * batch_size
end = min(start + batch_size, len(pairs))
batch_frames_slice = batch_frames[start:end].to(args.device, dtype=dtype)
batch_prompts_slice = prompts[start:end]
loss, reward = reward_model(batch_frames_slice, batch_prompts_slice)
per_sample_rewards = reward_model.get_reward(batch_frames_slice, batch_prompts_slice)
total_loss += loss.item() * (end - start)
total_reward += reward.item() * (end - start)
all_per_sample_rewards.append(per_sample_rewards.cpu())
# Aggregate results
all_per_sample_rewards = torch.cat(all_per_sample_rewards, dim=0)
avg_loss = total_loss / len(pairs)
avg_reward = total_reward / len(pairs)
print("Per-sample rewards:")
for idx, (vp, prompt) in enumerate(pairs):
print(f" [{idx}] Reward: {all_per_sample_rewards[idx].item():.6f} | Video: {vp} | Prompt: {prompt}")
print("=" * 60)
if __name__ == "__main__":
main()