Compare commits

...
22 Commits
Author SHA1 Message Date
Davids048 a520eceed7 document vendored reward runtimes in agent memory 2026-06-08 20:22:13 +00:00
Davids048 b2c75534ac fix mypy issues for vendored reward runtimes 2026-06-08 20:07:46 +00:00
Davids048 0b9c325984 fix mypy issues in RL integration files 2026-06-08 19:54:35 +00:00
Davids048 2afa837d96 handle yapf and ruff issues 2026-06-08 19:22:11 +00:00
Davids048 83bb602836 apply yapf formatting to reward runtimes 2026-06-06 20:59:50 +00:00
Davids048 606b42b6f9 fix spelling and markdown pre-commit issues 2026-06-06 20:58:49 +00:00
Davids048 e9a7a128c2 patch: resolve VideoAlign checkpoint and FA2 fallback
Resolve the default VideoAlign checkpoint path by downloading the KlingTeam/VideoReward Hugging Face snapshot and passing the returned local snapshot directory into VideoVLMRewardInference. Explicit checkpoint_path values still work as local path overrides.

Make the VideoAlign FlashAttention fallback check for classic FlashAttention-2 metadata/API instead of only checking for a flash_attn namespace. FastVideo may have FlashAttention-4/CuTe installed, but Transformers' flash_attention_2 path requires classic FlashAttention-2. When classic FA2 is unavailable, warn and use SDPA for the VideoAlign reward model.
2026-06-05 08:54:37 +00:00
Davids048 78b6995e13 WIP: vendor HPSv3 and VideoAlign reward runtimes
Adds vendored runtime code under fastvideo/train/methods/rl/reward/HPSv3 and fastvideo/train/methods/rl/reward/VideoAlign, replacing the broken gitlinks with normal tracked files.

The provenance and porting rules live in the vendor package markers: fastvideo/train/methods/rl/reward/HPSv3/__init__.py, fastvideo/train/methods/rl/reward/HPSv3/hpsv3/__init__.py, and fastvideo/train/methods/rl/reward/VideoAlign/__init__.py.

The FastVideo wrapper scripts hpsv3.py and videoalign.py depend on these vendored packages through explicit package imports. The copied runtime files are ported faithfully from upstream, with only import-path changes needed for package importability; this means some third-party code is unrelated to FastVideo internals but remains faithful to the source.

Known follow-ups: this vendor code does not currently pass pre-commit. VideoAlign also still assumes the checkpoint artifacts already exist under its checkpoints path or the configured VIDEOALIGN_CHECKPOINT_PATH; the missing/downloaded checkpoint resolution is not addressed in this commit.
2026-06-05 07:34:58 +00:00
Davids048 d858270708 patch formatting based on ruff hints 2026-06-05 01:46:09 +00:00
Davids048 298845ccd9 apply pre-commit formatting 2026-06-05 01:45:03 +00:00
Adam 9c2fa5aa9d [feat] GenRL: add explicit HPSv3 VideoAlign recipes (#1405) 2026-06-05 01:14:43 +00:00
Adam 73e0fa727f [feat] GenRL: add Wan LoRA adapter support (#1404) 2026-06-05 01:14:43 +00:00
AdamandDavids048 1fbe5a868c [feat] GenRL: fix PPO loop cadence and diagnostics (#1403)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
2026-06-05 01:14:43 +00:00
AdamandDavids048 e1be068c46 [feat] GenRL: add runtime and memory stability helpers (#1402)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
2026-06-05 01:14:43 +00:00
Adam cd69575095 [feat] GenRL: keep repeated prompt samples on one rank (#1401) 2026-06-05 01:14:43 +00:00
AdamandDavids048 2c19cbdebd [feat] GenRL: stabilize reward model compatibility (#1400)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
2026-06-05 01:14:43 +00:00
Peiyuan Zhang 5ae633a551 mv to utils 2026-06-05 01:14:42 +00:00
Peiyuan Zhang 6e663214a7 sampled video looks correct 2026-06-05 01:14:42 +00:00
Peiyuan Zhang 8b0ba24deb time profiling and log sampled videos 2026-06-05 01:14:42 +00:00
Peiyuan Zhang ef83574e44 gen RL runing 2026-06-05 01:14:14 +00:00
Peiyuan Zhang 9f2fc2f303 improve ema 2026-06-05 01:13:50 +00:00
Peiyuan Zhang c52de7e9df first edit 2026-06-05 01:07:44 +00:00
58 changed files with 9576 additions and 8 deletions
+24
View File
@@ -117,6 +117,30 @@ FastVideo-WorldModel/
| `TOKENIZERS_PARALLELISM` | Set `false` to avoid fork warnings |
| `HF_HOME` | HuggingFace cache directory |
## RL Reward Runtime Vendoring
Commit `78b6995e` (`WIP: vendor HPSv3 and VideoAlign reward runtimes`) added
vendored runtime code under `fastvideo/train/methods/rl/reward/HPSv3` and
`fastvideo/train/methods/rl/reward/VideoAlign`, replacing broken gitlinks with
normal tracked files.
The provenance and porting rules live in the vendor package markers:
`fastvideo/train/methods/rl/reward/HPSv3/__init__.py`,
`fastvideo/train/methods/rl/reward/HPSv3/hpsv3/__init__.py`, and
`fastvideo/train/methods/rl/reward/VideoAlign/__init__.py`.
The FastVideo wrapper scripts `hpsv3.py` and `videoalign.py` depend on these
vendored packages through explicit package imports. The copied runtime files
are ported faithfully from upstream, with only import-path changes needed for
package importability; this means some third-party code is unrelated to
FastVideo internals but remains faithful to the source.
Known follow-ups from commit `78b6995e`: the vendored code was not expected to
pass pre-commit at that commit. VideoAlign also assumed checkpoint artifacts
already existed under its checkpoints path or the configured
`VIDEOALIGN_CHECKPOINT_PATH`; later work resolved the default checkpoint path by
downloading the `KlingTeam/VideoReward` Hugging Face snapshot.
## Build & Test Commands
```bash
@@ -0,0 +1,164 @@
# GenRL / Video GRPO: Wan 2.1 T2V 1.3B with HPSv3 and VideoAlign rewards.
#
# Ported from GenRL/config/longcat.yaml.
#
# - Student: trainable full-parameter model by default.
# - LoRA is still available via models.student.use_lora=true.
# - Full fine-tuning with beta > 0 requires models.reference and much
# more memory; keep beta at 0.0 for the 4xH100 probe run.
#
# Usage:
# torchrun --nnodes=1 --nproc_per_node=4 \
# -m fastvideo.train.entrypoint.train \
# --config examples/train/configs/genrl_wan2.1_t2v_1.3B_hpsv3_videoalign.yaml
models:
student:
_target_: fastvideo.train.models.wan.wan_genrl.GenRLWanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
# Set true for LoRA. LoRA can use method.beta with disable_adapter().
use_lora: false
lora_r: 128
lora_alpha: 64
lora_init_weights: gaussian
lora_path: null
lora_target_modules:
- to_k
- to_out
- to_q
- to_v
- ffn.fc_in
- ffn.fc_out
enable_gradient_checkpointing_type: full
method:
_target_: fastvideo.train.methods.rl.genrl.GenRLMethod
# ---- Reward functions ----
reward_fn:
hpsv3_general: 1.0
hpsv3_percentile: 1.0
videoalign_mq: 1.0
videoalign_ta: 1.0
reward_module: null
reward_on_gpu: true
# ---- Data ----
prompt_dataset_path: GenRL/datasets/filtered_prompts
prompt_fn: filtered_prompts
# ---- Sampling ----
sample_batch_size: 4
eval_batch_size: 2
# Sample multiple rollout microbatches, average their PPO losses, then
# apply one optimizer update. This reduces reward/advantage variance.
num_batches_per_epoch: 4
accumulate_ppo_microbatches: true
eval_every_steps: 20
eval_num_batches: 1
eval_num_steps: 16
eval_guidance_scale: 4.5
num_inference_steps: 16
guidance_scale: 4.5
num_video_per_prompt: 4
noise_level: 1.0
sde_type: flow_sde
sde_window_size: 1
sde_window_range: [0, 6]
diffusion_clip: true
diffusion_clip_value: 0.45
kl_reward: 0
same_latent: true
# ---- Video dimensions ----
height: 480
width: 832
num_frames: 81
# ---- PPO training ----
train_batch_size: 4
num_inner_epochs: 1
clip_range: 1.0e-4
adv_clip_max: 5.0
# Official LoRA LongCat uses beta: 3.0e-4 with disable_adapter().
# For full fine-tuning on 4 H100s, avoid a second frozen Wan copy.
beta: 0.0
use_cfg: true
# Flash-GRPO-style temporal gradient rectification: avoid the large
# LongCat sigma/dt multiplier while debugging full-FT stability.
loss_reweighting: flash_tgr
loss_reweighting_clip: null
weight_advantages: true
# Match official GenRL PPO cadence. With sde_window_size: 1 this is
# equivalent to one optimizer step per sampled trajectory timestep.
optimizer_step_per_timestep: true
log_post_update_kl: true
max_grad_norm: 1.0
seed: 42
# ---- Advantage computation ----
per_prompt_stat_tracking: true
global_std: false
max_group_std: true
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
# Full fine-tuning needs FSDP/HSDP sharding across all 4 GPUs.
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
# Not used by GenRL (prompt dataloaders are in method config)
# but required by the config parser.
data_path: ""
train_batch_size: 1
seed: 42
num_height: 480
num_width: 832
num_frames: 81
optimizer:
# Full-parameter visual GRPO is much more sensitive than LoRA.
# DanceGRPO reports 5e-6 to 2e-5 as the practical range.
learning_rate: 1.0e-5
betas: [0.9, 0.999]
weight_decay: 1.0e-4
lr_scheduler: constant
lr_warmup_steps: 0
loop:
# max_train_steps == num_epochs in GenRL terms.
max_train_steps: 100000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/genrl_longcat
training_state_checkpointing_steps: 100
checkpoints_total_limit: 3
tracker:
project_name: VideoRL
# Leave blank so W&B auto-generates a unique display name per run.
run_name: ""
model:
enable_gradient_checkpointing_type: full
callbacks:
# Gradnorm call back Disabled; GenRLMethod clips internally.
ema:
decay: 0.9
start_iter: 0
update_interval: 8
log_rl_samples:
every_steps: 1
max_videos: 4
fps: 16
pipeline:
flow_shift: 3.0
@@ -0,0 +1,113 @@
# GenRL / Video GRPO: Wan 2.1 T2V 1.3B — OCR reward, full finetune, 4 GPUs.
#
# Usage:
# torchrun --nnodes=1 --nproc_per_node=4 \
# -m fastvideo.train.entrypoint.train \
# --config examples/train/configs/genrl_wan2.1_t2v_1.3B_ocr.yaml
models:
student:
_target_: fastvideo.train.models.wan.wan_genrl.GenRLWanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
enable_gradient_checkpointing_type: full
method:
_target_: fastvideo.train.methods.rl.genrl.GenRLMethod
# ---- Reward functions ----
reward_fn:
video_ocr: 1.0
reward_module: null
# ---- Data ----
prompt_dataset_path: GenRL/datasets/ocr
prompt_fn: general_ocr
# ---- Sampling ----
sample_batch_size: 4
eval_batch_size: 2
num_batches_per_epoch: 1
num_inference_steps: 16
guidance_scale: 4.5
num_video_per_prompt: 4
noise_level: 1.0
sde_type: flow_sde
sde_window_size: 1
sde_window_range: [0, 6]
diffusion_clip: true
diffusion_clip_value: 0.45
kl_reward: 0
same_latent: true
# ---- Video dimensions ----
height: 480
width: 832
num_frames: 81
# ---- PPO training ----
train_batch_size: 4
num_inner_epochs: 1
clip_range: 1.0e-3
adv_clip_max: 5.0
# No frozen reference model is configured in this launch.
beta: 0.0
use_cfg: true
loss_reweighting: longcat
weight_advantages: false
max_grad_norm: 1.0
seed: 42
# ---- Advantage computation ----
per_prompt_stat_tracking: true
global_std: false
max_group_std: true
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: ""
train_batch_size: 1
seed: 42
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 1.0e-4
betas: [0.9, 0.999]
weight_decay: 1.0e-4
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 100000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/genrl_ocr
training_state_checkpointing_steps: 100
checkpoints_total_limit: 3
tracker:
project_name: VideoRL
run_name: wan_2_1_t2v_1_3b_ocr
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 0.0 # Disabled; GenRLMethod clips internally.
ema:
decay: 0.9
start_iter: 0
pipeline:
flow_shift: 3.0
+1 -1
View File
@@ -107,7 +107,7 @@ srun torchrun \\
--node_rank \$SLURM_PROCID \\
--rdzv_backend=c10d \\
--rdzv_endpoint="\$MASTER_ADDR:\$MASTER_PORT" \\
fastvideo/train/entrypoint/train.py \\
-m fastvideo.train.entrypoint.train \\
--config ${CONFIG} \\
--training.distributed.num_gpus ${TOTAL_GPUS} \\
${EXTRA_ARGS[*]:-}
+27 -5
View File
@@ -317,8 +317,25 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
input_parallel = splitted_input[tp_rank].contiguous()
output_parallel = self.base_layer.quant_method.apply(self.base_layer, input_parallel)
if self.set_lora:
output_parallel = self.apply_lora(output_parallel, input_parallel)
if not self.merged and not self.disable_lora:
lora_A = self.lora_A
lora_B = self.lora_B
if lora_A is None or lora_B is None:
raise RuntimeError("LoRA weights (lora_A, lora_B) must be initialized "
"before forward pass when LoRA is enabled.")
if isinstance(lora_A, DTensor):
lora_A = lora_A.to_local()
if isinstance(lora_B, DTensor):
lora_B = lora_B.to_local()
lora_A_sliced = self.slice_lora_a_weights(lora_A.to(input_parallel, non_blocking=True))
lora_B_sliced = self.slice_lora_b_weights(lora_B.to(output_parallel, non_blocking=True))
delta = input_parallel @ lora_A_sliced.T @ lora_B_sliced.T
if self.lora_alpha != self.lora_rank:
delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore
)
output_parallel = output_parallel + delta
if self.base_layer.reduce_results and self.base_layer.tp_size > 1:
output_ = tensor_model_parallel_all_reduce(output_parallel)
@@ -334,12 +351,17 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
return output, output_bias
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
tp_rank = get_tp_rank()
shard_size = self.base_layer.input_size_per_partition
# LoRA A gets its input size from base_layer.weight.shape[1].
# If that size is already input_size_per_partition, A is already
# sharded; otherwise it is a global tensor and needs TP slicing.
if A.shape[1] == shard_size:
return A.contiguous()
tp_rank = get_tp_rank()
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
A = A[:, start_idx:end_idx].contiguous()
return A
return A[:, start_idx:end_idx].contiguous()
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
return B
+3
View File
@@ -8,6 +8,8 @@ from fastvideo.train.callbacks.ema import (
EMACallback, )
from fastvideo.train.callbacks.grad_clip import (
GradNormClipCallback, )
from fastvideo.train.callbacks.log_rl_samples import (
LogRLSamplesCallback, )
from fastvideo.train.callbacks.validation import (
ValidationCallback, )
@@ -16,5 +18,6 @@ __all__ = [
"CallbackDict",
"EMACallback",
"GradNormClipCallback",
"LogRLSamplesCallback",
"ValidationCallback",
]
+2
View File
@@ -24,6 +24,7 @@ _BUILTIN_CALLBACKS: dict[str, str] = {
"grad_clip": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback",
"validation": "fastvideo.train.callbacks.validation.ValidationCallback",
"ema": "fastvideo.train.callbacks.ema.EMACallback",
"log_rl_samples": "fastvideo.train.callbacks.log_rl_samples.LogRLSamplesCallback",
}
@@ -63,6 +64,7 @@ class Callback:
self,
method: TrainingMethod,
iteration: int = 0,
outputs: dict[str, Any] | None = None,
) -> None:
pass
+6 -1
View File
@@ -47,9 +47,11 @@ class EMACallback(Callback):
*,
decay: float = 0.9999,
start_iter: int = 0,
update_interval: int = 1,
) -> None:
self._decay = float(decay)
self._start_iter = int(start_iter)
self._update_interval = max(1, int(update_interval))
self._ema_started = False
self.student_ema: EMA_FSDP | None = None
@@ -78,9 +80,10 @@ class EMACallback(Callback):
)
logger.info(
"EMA callback enabled (decay=%s, "
"start_iter=%d).",
"start_iter=%d, update_interval=%d).",
self._decay,
self._start_iter,
self._update_interval,
)
def on_training_step_end(
@@ -94,6 +97,8 @@ class EMACallback(Callback):
if iteration < self._start_iter:
return
if (iteration - self._start_iter) % self._update_interval != 0:
return
if not self._ema_started:
logger.info(
"Starting EMA updates at iteration %d "
+2 -1
View File
@@ -8,7 +8,7 @@ Optionally logs per-module grad norms to the tracker.
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import Any, TYPE_CHECKING
from fastvideo.logger import init_logger
from fastvideo.train.callbacks.callback import Callback
@@ -41,6 +41,7 @@ class GradNormClipCallback(Callback):
self,
method: TrainingMethod,
iteration: int = 0,
outputs: dict[str, Any] | None = None,
) -> None:
max_norm = self._max_grad_norm
if max_norm <= 0.0:
+116
View File
@@ -0,0 +1,116 @@
# SPDX-License-Identifier: Apache-2.0
"""Callback to log sampled RL videos to the tracker."""
from __future__ import annotations
import contextlib
import os
import tempfile
from typing import Any, TYPE_CHECKING
import imageio
import numpy as np
import torch
from fastvideo.logger import init_logger
from fastvideo.train.callbacks.callback import Callback
if TYPE_CHECKING:
from fastvideo.train.methods.base import TrainingMethod
logger = init_logger(__name__)
class LogRLSamplesCallback(Callback):
"""Log RL-sampled videos to the experiment tracker.
Expects ``outputs`` to contain:
- ``sample_videos``: uint8 tensor (B, 3, T, H, W).
- ``sample_prompts``: list of prompt strings.
Configuration (YAML ``callbacks.log_rl_samples``):
.. code-block:: yaml
callbacks:
log_rl_samples:
every_steps: 5
max_videos: 4
fps: 16
"""
def __init__(
self,
*,
every_steps: int = 1,
max_videos: int = 4,
fps: int = 16,
) -> None:
self._every_steps = int(every_steps)
self._max_videos = int(max_videos)
self._fps = int(fps)
def on_before_optimizer_step(
self,
method: TrainingMethod,
iteration: int = 0,
outputs: dict[str, Any] | None = None,
) -> None:
if outputs is None:
return
if (self._every_steps > 0 and iteration % self._every_steps != 0):
return
videos = outputs.get("sample_videos")
prompts = outputs.get("sample_prompts")
if videos is None:
return
tracker = getattr(method, "tracker", None)
if tracker is None:
return
self._log_videos(tracker, videos, prompts, iteration)
def _log_videos(
self,
tracker: Any,
videos: torch.Tensor,
prompts: list[str] | None,
step: int,
) -> None:
n = min(len(videos), self._max_videos)
tmp_dir = tempfile.mkdtemp(prefix="rl_samples_")
video_logs = []
try:
for i in range(n):
# (3, T, H, W) uint8 -> (T, H, W, 3) numpy.
v = videos[i].permute(1, 2, 3, 0)
frames = v.numpy().astype(np.uint8)
fname = os.path.join(tmp_dir, f"sample_{step}_{i}.mp4")
imageio.mimsave(fname, frames, fps=self._fps)
caption = (prompts[i] if prompts and i < len(prompts) else None)
art = tracker.video(fname, caption=caption, fps=self._fps)
if art is not None:
video_logs.append(art)
if video_logs:
tracker.log_artifacts(
{"rl_sample_videos": video_logs},
step,
)
logger.info(
"Logged %d RL sample videos at step %d",
len(video_logs),
step,
)
finally:
# Clean up temp files.
for f in os.listdir(tmp_dir):
with contextlib.suppress(OSError):
os.remove(os.path.join(tmp_dir, f))
with contextlib.suppress(OSError):
os.rmdir(tmp_dir)
+4
View File
@@ -9,6 +9,7 @@ __all__ = [
"KDMethod",
"SelfForcingMethod",
"DiffusionForcingSFTMethod",
"GenRLMethod",
]
@@ -28,4 +29,7 @@ def __getattr__(name: str) -> object:
if name == "DiffusionForcingSFTMethod":
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
return DiffusionForcingSFTMethod
if name == "GenRLMethod":
from fastvideo.train.methods.rl.genrl import GenRLMethod
return GenRLMethod
raise AttributeError(name)
+2
View File
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
"""Reinforcement learning methods for video generation."""
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,17 @@
"""Vendored runtime subset of HPSv3.
Source: https://github.com/MizzenAI/HPSv3
Commit: bd0c5fcb5f587617b0169c07222ab78d01e2f3c2
Purpose: Runtime reward inference integration for FastVideo GenRL.
This is temporary minimal vendoring for PR integration. It is expected to be
cleaned up and normalized later.
Porting rules:
- Include only files required by the runtime import closure used by FastVideo.
- When an upstream file is required, copy the entire file faithfully.
- Only adjust imports as needed to make the vendored code import through package
paths instead of sys.path mutation or ambiguous top-level imports.
- Do not perform style, typing, or behavioral cleanup as part of this vendoring
step.
"""
+21
View File
@@ -0,0 +1,21 @@
"""Vendored HPSv3 runtime package.
Source: https://github.com/MizzenAI/HPSv3
Commit: bd0c5fcb5f587617b0169c07222ab78d01e2f3c2
Purpose: Runtime reward inference integration for FastVideo GenRL.
This is temporary minimal vendoring for PR integration. It is expected to be
cleaned up and normalized later.
Porting rules:
- Include only files required by the runtime import closure used by FastVideo.
- When an upstream file is required, copy the entire file faithfully.
- Only adjust imports as needed to make the vendored code import through package
paths instead of sys.path mutation or ambiguous top-level imports.
- Do not perform style, typing, or behavioral cleanup as part of this vendoring
step.
"""
from .inference import HPSv3RewardInferencer
__all__ = ["HPSv3RewardInferencer"]
@@ -0,0 +1,60 @@
# 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"
deepspeed: hpsv3/config/ds_config/zero2.json
save_only_model: True
save_full_model: True
dataloader_num_workers: 8
@@ -0,0 +1 @@
"""Vendored HPSv3 dataset runtime helpers."""
@@ -0,0 +1,185 @@
import numpy as np
import torch
from .utils import process_vision_info
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.
**Visual Quality:**
Evaluate the overall visual quality of the image. The following sub-dimensions should be considered:
- **Reasonableness:** The image 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 image. 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 image, 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 image 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.
**Text Alignment:**
Assess how well the image matches the textual prompt across the following sub-dimensions:
- **Subject Relevance** Evaluate how accurately the subject(s) in the image (e.g., person, animal, object) align with the textual description. The subject should match the description in terms of number, appearance, and behavior.
- **Style Relevance:** If the prompt specifies a particular artistic or stylistic style, evaluate how well the image adheres to this style.
- **Contextual Consistency**: Assess whether the background, setting, and surrounding elements in the image logically fit the scenario described in the prompt. The environment should support and enhance the subject without contradictions.
- **Attribute Fidelity**: Check if specific attributes mentioned in the prompt (e.g., colors, clothing, accessories, expressions, actions) are faithfully represented in the image. Minor deviations may be acceptable, but critical attributes should be preserved.
- **Semantic Coherence**: Evaluate whether the overall meaning and intent of the prompt are captured in the image. The generated content should not introduce elements that conflict with or distort the original description.
Textual prompt - {text_prompt}
"""
INSTRUCTION_debug = """
{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 QWen2VLDataCollator:
def __init__(
self,
processor,
with_instruction=True,
max_pixels=256 * 28 * 28, # Default max pixels
min_pixels=256 * 28 * 28, # Default min pixels
use_special_tokens=True,
):
self.processor = processor
self.with_instruction = with_instruction
self.max_pixels = max_pixels
self.min_pixels = min_pixels
self.use_special_tokens = use_special_tokens
def _clean_message(
self,
texts,
images,
max_pixels=256 * 28 * 28,
min_pixels=256 * 28 * 28,
with_instruction=True,
use_special_tokens=True,
):
"""
remove unnecessary keys from message(very very necessary)
"""
message_list = []
for text, image in zip(texts, images, strict=False):
out_message = [{
"role":
"user",
"content": [
{
"type": "image",
"image": image,
"min_pixels": min_pixels,
"max_pixels": max_pixels,
},
{
"type":
"text",
"text": (INSTRUCTION.format(text_prompt=text) +
prompt_with_special_token if use_special_tokens else prompt_without_special_token),
},
],
}]
message_list.append(out_message)
return message_list
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 __call__(self, inputs, with_instruction=True):
"""
Preprocess inputs to token sequences and return a batch
"""
images_1, images_2, texts_1, texts_2 = [], [], [], []
for idx, batch in enumerate(inputs):
texts_1.append(batch["text_1"])
texts_2.append(batch["text_2"])
images_1.append(batch["image_1"])
images_2.append(batch["image_2"])
messages_batch_1 = self._clean_message(
texts_1,
images_1,
max_pixels=self.max_pixels,
min_pixels=self.min_pixels,
with_instruction=self.with_instruction,
use_special_tokens=self.use_special_tokens,
)
messages_batch_2 = self._clean_message(
texts_2,
images_2,
max_pixels=self.max_pixels,
min_pixels=self.min_pixels,
with_instruction=self.with_instruction,
use_special_tokens=self.use_special_tokens,
)
# import pdb; pdb.set_trace()
image_inputs_1, _ = process_vision_info(messages_batch_1)
image_inputs_2, _ = process_vision_info(messages_batch_2)
image_inputs_1 = [np.array(image_inputs_1[i]) / 255.0 for i in range(len(image_inputs_1))]
image_inputs_2 = [np.array(image_inputs_2[i]) / 255.0 for i in range(len(image_inputs_2))]
do_rescale = False
batch_1 = self.processor(
text=self.processor.apply_chat_template(messages_batch_1, tokenize=False, add_generation_prompt=True),
images=image_inputs_1,
videos=None,
padding=True,
return_tensors="pt",
images_kwargs={"do_rescale": do_rescale},
)
batch_2 = self.processor(
text=self.processor.apply_chat_template(messages_batch_2, tokenize=False, add_generation_prompt=True),
images=image_inputs_2,
videos=None,
padding=True,
return_tensors="pt",
images_kwargs={"do_rescale": do_rescale},
)
# pdb.set_trace()
max_len = max(batch_1["input_ids"].shape[1], batch_2["input_ids"].shape[1])
batch_1["input_ids"], batch_1["attention_mask"] = self._pad_sequence(batch_1["input_ids"],
batch_1["attention_mask"], max_len,
"right")
batch_2["input_ids"], batch_2["attention_mask"] = self._pad_sequence(batch_2["input_ids"],
batch_2["attention_mask"], max_len,
"right")
batch = {
"batch_1": batch_1,
"batch_2": batch_2,
"choice_dist": torch.stack([batch["choice_dist"] for batch in inputs]),
# Store original text prompts for visualization
"text_1": texts_1,
"text_2": texts_2,
"image_1": image_inputs_1,
"image_2": image_inputs_2,
}
return batch
@@ -0,0 +1,75 @@
import torch
from torch.utils.data import Dataset
import random
import json
import os
from tqdm import tqdm
class PairwiseOriginalDataset(Dataset):
def __init__(
self,
json_list,
soft_label=False,
confidence_threshold=None,
):
self.samples = []
for json_file in json_list:
with open(json_file) as f:
data = json.load(f)
self.samples.extend(data)
self.soft_label = soft_label
self.confidence_threshold = confidence_threshold
if confidence_threshold is not None:
new_samples = []
for sample in tqdm(self.samples, desc="Filtering samples according to confidence threshold"):
if sample.get("confidence", float("inf")) >= confidence_threshold:
new_samples.append(sample)
self.samples = new_samples
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
while True:
index = idx
try:
return self.get_single_item(index)
except Exception as e:
print(f"Error processing sample at index {idx}: {e}")
import traceback
traceback.print_exc()
index = random.randint(0, len(self.samples) - 1)
if index == idx:
continue
idx = index
def get_single_item(self, idx):
sample = self.samples[idx]
# Load image paths
image_1 = sample["path1"]
image_2 = sample["path2"]
assert os.path.exists(image_1) and os.path.exists(image_2), f'{image_1} or {image_2}'
text_1 = sample["prompt"]
text_2 = sample["prompt"]
# Process Label
if self.soft_label:
choice_dist = sorted(sample["choice_dist"], reverse=True)
assert (torch.sum(torch.tensor(choice_dist)) > 0), "Choice distribution cannot be zero."
label = torch.tensor(choice_dist[0]) / torch.sum(torch.tensor(choice_dist))
else:
label = torch.tensor(1).float()
# breakpoint()
return {
"image_1": image_1,
"image_2": image_2,
"text_1": text_1,
"text_2": text_2,
"label": label,
"confidence": sample.get("confidence", 1.0),
"choice_dist": torch.tensor(sample.get("choice_dist", [1.0, 0.0])),
}
@@ -0,0 +1,414 @@
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 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:
image = ele["image"] if "image" in ele else ele["image_url"]
image_obj = None
if isinstance(image, Image.Image | torch.Tensor):
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}")
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: int | 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)
if nframes > total_frames:
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.",
stacklevel=2,
)
if "file://" in video_path:
video_path = video_path[7:]
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"]
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
else:
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
+155
View File
@@ -0,0 +1,155 @@
import os
from collections.abc import Mapping
import torch
import huggingface_hub
from .dataset.utils import process_vision_info
from .dataset.data_collator_qwen import prompt_with_special_token, prompt_without_special_token, INSTRUCTION
from .utils.parser import ModelConfig, PEFTLoraConfig, TrainingConfig, DataConfig, parse_args_with_yaml
from .train import create_model_and_processor
from pathlib import Path
_MODEL_CONFIG_PATH = Path(__file__).parent / "config/"
class HPSv3RewardInferencer:
def __init__(self, config_path=None, checkpoint_path=None, device='cuda', differentiable=False):
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,
differentiable=differentiable,
)
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()})
elif isinstance(data, tuple | list):
return type(data)(self._prepare_input(v) for v in data)
elif 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, image_paths, prompts):
max_pixels = 256 * 28 * 28
min_pixels = 256 * 28 * 28
message_list = []
for text, image in zip(prompts, image_paths, strict=False):
out_message = [{
"role":
"user",
"content": [
{
"type": "image",
"image": image,
"min_pixels": min_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
if __name__ == "__main__":
config_path = 'config/inference/HPSv3_7B.yaml'
checkpoint_path = 'checkpoints/HPSv3_7B.pth'
device = 'cuda'
dtype = torch.bfloat16
inferencer = HPSv3RewardInferencer(config_path, checkpoint_path, device=device)
image_paths = ["assets/example1.png", "assets/example2.png"]
prompts = [
"cute chibi anime cartoon fox, smiling wagging tail with a small cartoon heart above sticker",
"cute chibi anime cartoon fox, smiling wagging tail with a small cartoon heart above sticker"
]
rewards = inferencer.reward(image_paths, prompts)
print(rewards[0][0].item()) # miu and sigma. we select miu as the final output
print(rewards[1][0].item())
@@ -0,0 +1 @@
"""Vendored HPSv3 reward model runtime helpers."""
@@ -0,0 +1,615 @@
# Copyright 2024 The Qwen team, Alibaba Group and the HuggingFace Inc. team. All rights reserved.
#
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
# and OPT implementations in this library. It has been modified from its
# original forms to accommodate minor architectural differences compared
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Image processor class for Qwen2-VL.
This module provides both differentiable and non-differentiable image processing methods:
1. For DIFFERENTIABLE processing (torch.autograd compatible):
- Pass torch.Tensor to _preprocess() method
- Use preprocess_tensor() method directly
- All operations use PyTorch functions (F.interpolate, tensor operations, etc.)
2. For NON-DIFFERENTIABLE processing (original functionality):
- Pass PIL images or numpy arrays to preprocess() method
- Uses PIL/transformers image processing functions and numpy operations
The differentiable path supports:
- Bilinear interpolation for resizing (instead of PIL resampling)
- Tensor-based rescaling and normalization
- Differentiable patch extraction and reshaping
"""
import math
import numpy as np
import torch
import torch.nn.functional as F
from transformers.image_processing_utils import BaseImageProcessor, BatchFeature
from transformers.image_transforms import (
convert_to_rgb,
resize,
to_channel_dimension_format,
)
from transformers.image_utils import (
OPENAI_CLIP_MEAN,
OPENAI_CLIP_STD,
ChannelDimension,
ImageInput,
PILImageResampling,
VideoInput,
get_image_size,
infer_channel_dimension_format,
is_scaled_image,
is_valid_image,
make_list_of_images,
to_numpy_array,
valid_images,
validate_preprocess_arguments,
)
from transformers.utils import TensorType, is_vision_available, logging
logger = logging.get_logger(__name__)
if is_vision_available():
from PIL import Image
def make_batched_images(images) -> list[list[ImageInput]]:
"""
Accepts images in list or nested list format, and makes a list of images for preprocessing.
Args:
images (`Union[List[List[ImageInput]], List[ImageInput], ImageInput]`):
The input image.
Returns:
list: A list of images.
"""
if isinstance(images, list | tuple) and isinstance(images[0], list | tuple) and is_valid_image(images[0][0]):
return [img for img_list in images for img in img_list]
elif isinstance(images, list | tuple) and is_valid_image(images[0]):
return images
elif is_valid_image(images):
return [images]
raise ValueError(f"Could not make batched images from {images}")
# Copied from transformers.models.llava_next_video.image_processing_llava_next_video.make_batched_videos
def make_batched_videos(videos) -> list[VideoInput]:
if isinstance(videos, list | tuple) and isinstance(videos[0], list | tuple) and is_valid_image(videos[0][0]):
return videos
elif isinstance(videos, list | tuple) and is_valid_image(videos[0]):
if isinstance(videos[0], Image.Image):
return [videos]
elif len(videos[0].shape) == 4:
return [list(video) for video in videos]
elif is_valid_image(videos) and len(videos.shape) == 4:
return [list(videos)]
raise ValueError(f"Could not make batched video from {videos}")
def smart_resize(height: int,
width: int,
factor: int = 28,
min_pixels: int = 56 * 56,
max_pixels: int = 14 * 14 * 4 * 1280):
"""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 height < factor or width < factor:
raise ValueError(f"height:{height} or width:{width} must be larger than factor:{factor}")
elif max(height, width) / min(height, width) > 200:
raise ValueError(
f"absolute aspect ratio must be smaller than 200, got {max(height, width) / min(height, width)}")
h_bar = round(height / factor) * factor
w_bar = round(width / factor) * factor
if h_bar * w_bar > max_pixels:
beta = math.sqrt((height * width) / max_pixels)
h_bar = math.floor(height / beta / factor) * factor
w_bar = math.floor(width / beta / factor) * factor
elif h_bar * w_bar < min_pixels:
beta = math.sqrt(min_pixels / (height * width))
h_bar = math.ceil(height * beta / factor) * factor
w_bar = math.ceil(width * beta / factor) * factor
return h_bar, w_bar
class Qwen2VLImageProcessor(BaseImageProcessor):
r"""
Constructs a Qwen2-VL image processor that dynamically resizes images based on the original images.
Args:
do_resize (`bool`, *optional*, defaults to `True`):
Whether to resize the image's (height, width) dimensions.
resample (`PILImageResampling`, *optional*, defaults to `Resampling.BICUBIC`):
Resampling filter to use when resizing the image.
do_rescale (`bool`, *optional*, defaults to `True`):
Whether to rescale the image by the specified scale `rescale_factor`.
rescale_factor (`int` or `float`, *optional*, defaults to `1/255`):
Scale factor to use if rescaling the image.
do_normalize (`bool`, *optional*, defaults to `True`):
Whether to normalize the image.
image_mean (`float` or `List[float]`, *optional*, defaults to `[0.48145466, 0.4578275, 0.40821073]`):
Mean to use if normalizing the image. This is a float or list of floats for each channel in the image.
image_std (`float` or `List[float]`, *optional*, defaults to `[0.26862954, 0.26130258, 0.27577711]`):
Standard deviation to use if normalizing the image. This is a float or list of floats for each channel in the image.
do_convert_rgb (`bool`, *optional*, defaults to `True`):
Whether to convert the image to RGB.
min_pixels (`int`, *optional*, defaults to `56 * 56`):
The min pixels of the image to resize the image.
max_pixels (`int`, *optional*, defaults to `28 * 28 * 1280`):
The max pixels of the image to resize the image.
patch_size (`int`, *optional*, defaults to 14):
The spatial patch size of the vision encoder.
temporal_patch_size (`int`, *optional*, defaults to 2):
The temporal patch size of the vision encoder.
merge_size (`int`, *optional*, defaults to 2):
The merge size of the vision encoder to llm encoder.
"""
model_input_names = ["pixel_values", "image_grid_thw", "pixel_values_videos", "video_grid_thw"]
def __init__(
self,
do_resize: bool = True,
resample: PILImageResampling = PILImageResampling.BICUBIC,
do_rescale: bool = True,
rescale_factor: int | float = 1 / 255,
do_normalize: bool = True,
image_mean: float | list[float] | None = None,
image_std: float | list[float] | None = None,
do_convert_rgb: bool = True,
min_pixels: int = 56 * 56,
max_pixels: int = 28 * 28 * 1280,
patch_size: int = 14,
temporal_patch_size: int = 2,
merge_size: int = 2,
**kwargs,
) -> None:
super().__init__(**kwargs)
self.do_resize = do_resize
self.resample = resample
self.do_rescale = do_rescale
self.rescale_factor = rescale_factor
self.do_normalize = do_normalize
self.image_mean = image_mean if image_mean is not None else OPENAI_CLIP_MEAN
self.image_std = image_std if image_std is not None else OPENAI_CLIP_STD
self.min_pixels = min_pixels
self.max_pixels = max_pixels
self.patch_size = patch_size
self.temporal_patch_size = temporal_patch_size
self.merge_size = merge_size
self.size = {"min_pixels": min_pixels, "max_pixels": max_pixels}
self.do_convert_rgb = do_convert_rgb
def _preprocess_differentiable(
self,
images: torch.Tensor,
do_resize: bool = None,
do_rescale: bool = None,
rescale_factor: float = None,
do_normalize: bool = None,
image_mean: float | list[float] | None = None,
image_std: float | list[float] | None = None,
):
"""
Differentiable version of image preprocessing using torch operations.
Args:
images: torch.Tensor of shape (B, C, H, W) or (C, H, W)
Returns:
flatten_patches: torch.Tensor - flattened patches
grid_thw: tuple - (grid_t, grid_h, grid_w)
"""
if images.dim() == 3:
images = images.unsqueeze(0) # Add batch dimension
batch_size, channels, height, width = images.shape
processed_images = []
resized_height, resized_width = height, width
for i in range(batch_size):
image = images[i] # (C, H, W)
if do_resize:
resized_height, resized_width = smart_resize(
height,
width,
factor=self.patch_size * self.merge_size,
min_pixels=self.min_pixels,
max_pixels=self.max_pixels,
)
# Use differentiable interpolation
image = F.interpolate(image.unsqueeze(0),
size=(resized_height, resized_width),
mode='bilinear',
align_corners=False).squeeze(0)
if do_rescale:
image = image * rescale_factor
if do_normalize:
if isinstance(image_mean, list | tuple):
mean = torch.tensor(image_mean, device=image.device, dtype=image.dtype).view(-1, 1, 1)
std = torch.tensor(image_std, device=image.device, dtype=image.dtype).view(-1, 1, 1)
else:
mean = image_mean
std = image_std
image = (image - mean) / std
processed_images.append(image)
# Stack all processed images
patches = torch.stack(processed_images) # (B, C, H, W)
# Handle temporal dimension
if patches.shape[0] == 1:
patches = patches.repeat(self.temporal_patch_size, 1, 1, 1)
# Reshape for patch extraction
batch_size, channel, resized_height, resized_width = patches.shape
grid_t = batch_size // self.temporal_patch_size
grid_h, grid_w = resized_height // self.patch_size, resized_width // self.patch_size
# Differentiable patch extraction and reshaping
patches = patches.view(
grid_t,
self.temporal_patch_size,
channel,
grid_h // self.merge_size,
self.merge_size,
self.patch_size,
grid_w // self.merge_size,
self.merge_size,
self.patch_size,
)
patches = patches.permute(0, 3, 6, 4, 7, 2, 1, 5, 8)
flatten_patches = patches.reshape(grid_t * grid_h * grid_w,
channel * self.temporal_patch_size * self.patch_size * self.patch_size)
return flatten_patches, (grid_t, grid_h, grid_w)
def _preprocess(
self,
images: ImageInput | VideoInput,
do_resize: bool = None,
resample: PILImageResampling = None,
do_rescale: bool = None,
rescale_factor: float = None,
do_normalize: bool = None,
image_mean: float | list[float] | None = None,
image_std: float | list[float] | None = None,
do_convert_rgb: bool = None,
data_format: ChannelDimension | None = ChannelDimension.FIRST,
input_data_format: str | ChannelDimension | None = None,
):
"""
Preprocess an image or batch of images. Copy of the `preprocess` method from `CLIPImageProcessor`.
Args:
images (`ImageInput`):
Image or batch of images to preprocess. Expects pixel values ranging from 0 to 255. If pixel values range from 0 to 1, set `do_rescale=False`.
vision_info (`List[Dict]`, *optional*):
Optional list of dictionaries containing additional information about vision inputs.
do_resize (`bool`, *optional*, defaults to `self.do_resize`):
Whether to resize the image.
resample (`PILImageResampling`, *optional*, defaults to `self.resample`):
Resampling filter to use if resizing the image. This can be one of the `PILImageResampling` enums.
do_rescale (`bool`, *optional*, defaults to `self.do_rescale`):
Whether to rescale the image.
rescale_factor (`float`, *optional*, defaults to `self.rescale_factor`):
Scale factor to use if rescaling the image.
do_normalize (`bool`, *optional*, defaults to `self.do_normalize`):
Whether to normalize the image.
image_mean (`float` or `List[float]`, *optional*, defaults to `self.image_mean`):
Mean to use if normalizing the image. Can be a float or a list of floats corresponding to the number of channels in the image.
image_std (`float` or `List[float]`, *optional*, defaults to `self.image_std`):
Standard deviation to use if normalizing the image. Can be a float or a list of floats corresponding to the number of channels in the image.
do_convert_rgb (`bool`, *optional*, defaults to `self.do_convert_rgb`):
Whether to convert the image to RGB.
data_format (`ChannelDimension`, *optional*, defaults to `ChannelDimension.FIRST`):
The channel dimension format for the output image. Can be one of:
- `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
- `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
- Unset: Use the channel dimension format of the input image.
input_data_format (`ChannelDimension` or `str`, *optional*):
The channel dimension format for the input image. Can be one of:
- `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
- `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
- `"none"` or `ChannelDimension.NONE`: image in (height, width) format. - `"none"` or `ChannelDimension.NONE`: image in (height, width) format.
"""
# Check if input is already a torch tensor (differentiable path)
if isinstance(images, torch.Tensor):
return self._preprocess_differentiable(
images,
do_resize=do_resize,
do_rescale=do_rescale,
rescale_factor=rescale_factor,
do_normalize=do_normalize,
image_mean=image_mean,
image_std=image_std,
)
# Original non-differentiable path for backward compatibility
images = make_list_of_images(images)
if do_convert_rgb:
images = [convert_to_rgb(image) for image in images]
# All transformations expect numpy arrays.
images = [to_numpy_array(image) for image in images]
if is_scaled_image(images[0]) and do_rescale:
logger.warning_once(
"It looks like you are trying to rescale already rescaled images. If the input"
" images have pixel values between 0 and 1, set `do_rescale=False` to avoid rescaling them again.")
if input_data_format is None:
# We assume that all images have the same channel dimension format.
input_data_format = infer_channel_dimension_format(images[0])
height, width = get_image_size(images[0], channel_dim=input_data_format)
resized_height, resized_width = height, width
processed_images = []
for image in images:
if do_resize:
resized_height, resized_width = smart_resize(
height,
width,
factor=self.patch_size * self.merge_size,
min_pixels=self.min_pixels,
max_pixels=self.max_pixels,
)
image = resize(image,
size=(resized_height, resized_width),
resample=resample,
input_data_format=input_data_format)
if do_rescale:
image = self.rescale(image, scale=rescale_factor, input_data_format=input_data_format)
if do_normalize:
image = self.normalize(image=image, mean=image_mean, std=image_std, input_data_format=input_data_format)
image = to_channel_dimension_format(image, data_format, input_channel_dim=input_data_format)
processed_images.append(image)
# NOTE: The following operations use numpy and are NOT differentiable
# For differentiable operations, pass torch.Tensor as input to use _preprocess_differentiable
patches = np.array(processed_images)
if data_format == ChannelDimension.LAST:
patches = patches.transpose(0, 3, 1, 2)
if patches.shape[0] == 1:
patches = np.tile(patches, (self.temporal_patch_size, 1, 1, 1))
channel = patches.shape[1]
grid_t = patches.shape[0] // self.temporal_patch_size
grid_h, grid_w = resized_height // self.patch_size, resized_width // self.patch_size
patches = patches.reshape(
grid_t,
self.temporal_patch_size,
channel,
grid_h // self.merge_size,
self.merge_size,
self.patch_size,
grid_w // self.merge_size,
self.merge_size,
self.patch_size,
)
patches = patches.transpose(0, 3, 6, 4, 7, 2, 1, 5, 8)
flatten_patches = patches.reshape(grid_t * grid_h * grid_w,
channel * self.temporal_patch_size * self.patch_size * self.patch_size)
return flatten_patches, (grid_t, grid_h, grid_w)
def preprocess_tensor(
self,
images: torch.Tensor,
do_resize: bool = None,
do_rescale: bool = None,
rescale_factor: float = None,
do_normalize: bool = None,
image_mean: float | list[float] | None = None,
image_std: float | list[float] | None = None,
):
"""
Differentiable preprocessing method for torch tensors.
Args:
images: torch.Tensor of shape (B, C, H, W) or (C, H, W)
Returns:
dict containing:
- pixel_values: torch.Tensor - processed patches
- image_grid_thw: torch.Tensor - grid dimensions
"""
do_resize = do_resize if do_resize is not None else self.do_resize
do_rescale = do_rescale if do_rescale is not None else self.do_rescale
rescale_factor = rescale_factor if rescale_factor is not None else self.rescale_factor
do_normalize = do_normalize if do_normalize is not None else self.do_normalize
image_mean = image_mean if image_mean is not None else self.image_mean
image_std = image_std if image_std is not None else self.image_std
patches, image_grid_thw = self._preprocess_differentiable(
images,
do_resize=do_resize,
do_rescale=do_rescale,
rescale_factor=rescale_factor,
do_normalize=do_normalize,
image_mean=image_mean,
image_std=image_std,
)
return {"pixel_values": patches, "image_grid_thw": torch.tensor(image_grid_thw, device=patches.device)}
def preprocess(
self,
images: ImageInput,
videos: VideoInput = None,
do_resize: bool = None,
size: dict[str, int] = None,
resample: PILImageResampling = None,
do_rescale: bool = None,
rescale_factor: float = None,
do_normalize: bool = None,
image_mean: float | list[float] | None = None,
image_std: float | list[float] | None = None,
do_convert_rgb: bool = None,
return_tensors: str | TensorType | None = None,
data_format: ChannelDimension | None = ChannelDimension.FIRST,
input_data_format: str | ChannelDimension | None = None,
):
"""
Args:
images (`ImageInput`):
Image to preprocess. Expects a single or batch of images with pixel values ranging from 0 to 255. If
passing in images with pixel values between 0 and 1, set `do_rescale=False`.
videos (`VideoInput`):
Video to preprocess. Expects a single or batch of videos with pixel values ranging from 0 to 255. If
passing in videos with pixel values between 0 and 1, set `do_rescale=False`.
do_resize (`bool`, *optional*, defaults to `self.do_resize`):
Whether to resize the image.
size (`Dict[str, int]`, *optional*, defaults to `self.size`):
Size of the image after resizing. Shortest edge of the image is resized to size["shortest_edge"], with
the longest edge resized to keep the input aspect ratio.
resample (`int`, *optional*, defaults to `self.resample`):
Resampling filter to use if resizing the image. This can be one of the enum `PILImageResampling`. Only
has an effect if `do_resize` is set to `True`.
do_rescale (`bool`, *optional*, defaults to `self.do_rescale`):
Whether to rescale the image.
rescale_factor (`float`, *optional*, defaults to `self.rescale_factor`):
Rescale factor to rescale the image by if `do_rescale` is set to `True`.
do_normalize (`bool`, *optional*, defaults to `self.do_normalize`):
Whether to normalize the image.
image_mean (`float` or `List[float]`, *optional*, defaults to `self.image_mean`):
Image mean to use for normalization. Only has an effect if `do_normalize` is set to `True`.
image_std (`float` or `List[float]`, *optional*, defaults to `self.image_std`):
Image standard deviation to use for normalization. Only has an effect if `do_normalize` is set to
`True`.
do_convert_rgb (`bool`, *optional*, defaults to `self.do_convert_rgb`):
Whether to convert the image to RGB.
return_tensors (`str` or `TensorType`, *optional*):
The type of tensors to return. Can be one of:
- Unset: Return a list of `np.ndarray`.
- `TensorType.TENSORFLOW` or `'tf'`: Return a batch of type `tf.Tensor`.
- `TensorType.PYTORCH` or `'pt'`: Return a batch of type `torch.Tensor`.
- `TensorType.NUMPY` or `'np'`: Return a batch of type `np.ndarray`.
- `TensorType.JAX` or `'jax'`: Return a batch of type `jax.numpy.ndarray`.
data_format (`ChannelDimension` or `str`, *optional*, defaults to `ChannelDimension.FIRST`):
The channel dimension format for the output image. Can be one of:
- `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
- `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
- Unset: Use the channel dimension format of the input image.
input_data_format (`ChannelDimension` or `str`, *optional*):
The channel dimension format for the input image. If unset, the channel dimension format is inferred
from the input image. Can be one of:
- `"channels_first"` or `ChannelDimension.FIRST`: image in (num_channels, height, width) format.
- `"channels_last"` or `ChannelDimension.LAST`: image in (height, width, num_channels) format.
- `"none"` or `ChannelDimension.NONE`: image in (height, width) format.
"""
do_resize = do_resize if do_resize is not None else self.do_resize
size = size if size is not None else self.size
resample = resample if resample is not None else self.resample
do_rescale = do_rescale if do_rescale is not None else self.do_rescale
rescale_factor = rescale_factor if rescale_factor is not None else self.rescale_factor
do_normalize = do_normalize if do_normalize is not None else self.do_normalize
image_mean = image_mean if image_mean is not None else self.image_mean
image_std = image_std if image_std is not None else self.image_std
do_convert_rgb = do_convert_rgb if do_convert_rgb is not None else self.do_convert_rgb
if images is not None:
images = make_batched_images(images)
if videos is not None:
videos = make_batched_videos(videos)
if images is not None and not valid_images(images):
raise ValueError("Invalid image type. Must be of type PIL.Image.Image, numpy.ndarray, "
"torch.Tensor, tf.Tensor or jax.ndarray.")
validate_preprocess_arguments(
rescale_factor=rescale_factor,
do_normalize=do_normalize,
image_mean=image_mean,
image_std=image_std,
do_resize=do_resize,
size=size,
resample=resample,
)
if images is not None:
pixel_values, vision_grid_thws = [], []
for image in images:
patches, image_grid_thw = self._preprocess(
image,
do_resize=do_resize,
resample=resample,
do_rescale=do_rescale,
rescale_factor=rescale_factor,
do_normalize=do_normalize,
image_mean=image_mean,
image_std=image_std,
data_format=data_format,
do_convert_rgb=do_convert_rgb,
input_data_format=input_data_format,
)
pixel_values.extend(patches)
vision_grid_thws.append(image_grid_thw)
if not isinstance(pixel_values[0], torch.Tensor):
pixel_values = np.array(pixel_values)
else:
pixel_values = torch.stack(pixel_values)
vision_grid_thws = np.array(vision_grid_thws)
data = {"pixel_values": pixel_values, "image_grid_thw": vision_grid_thws}
if videos is not None:
pixel_values, vision_grid_thws = [], []
for images in videos:
patches, video_grid_thw = self._preprocess(
images,
do_resize=do_resize,
resample=resample,
do_rescale=do_rescale,
rescale_factor=rescale_factor,
do_normalize=do_normalize,
image_mean=image_mean,
image_std=image_std,
data_format=data_format,
do_convert_rgb=do_convert_rgb,
input_data_format=input_data_format,
)
pixel_values.extend(patches)
vision_grid_thws.append(video_grid_thw)
pixel_values = np.array(pixel_values)
vision_grid_thws = np.array(vision_grid_thws)
data = {"pixel_values_videos": pixel_values, "video_grid_thw": vision_grid_thws}
return BatchFeature(data=data, tensor_type=return_tensors)
@@ -0,0 +1,817 @@
import os
import math
import matplotlib.pyplot as plt
import safetensors
import numpy as np
import torch
import torch.nn as nn
import datasets
from torch.utils.data import Dataset, DataLoader
from peft import PeftModel
from transformers import Qwen2VLForConditionalGeneration
from transformers.modeling_utils import PreTrainedModel
from transformers.trainer import TrainerCallback
from transformers.trainer import (
is_sagemaker_mp_enabled,
is_peft_available,
is_datasets_available,
WEIGHTS_NAME,
TRAINING_ARGS_NAME,
SAFE_WEIGHTS_NAME,
PREFIX_CHECKPOINT_DIR,
logger,
)
from transformers.trainer_pt_utils import nested_detach
from trl import RewardTrainer
from ..utils.training_utils import get_peft_state_non_lora_maybe_zero_3
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,
):
super().__init__(config)
# pdb.set_trace()
self.output_dim = output_dim
if rm_head_type == "default":
self.rm_head = nn.Linear(config.hidden_size, output_dim, bias=False)
elif rm_head_type == "ranknet":
if rm_head_kwargs is not None:
for layer in range(rm_head_kwargs.get("num_layers", 3)):
if layer == 0:
self.rm_head = nn.Sequential(
nn.Linear(config.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:
self.rm_head = nn.Sequential(
nn.Linear(config.hidden_size, 1024),
nn.ReLU(),
nn.Dropout(0.05),
nn.Linear(1024, 16),
nn.ReLU(),
nn.Linear(16, output_dim),
)
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,
):
## modified from the origin class Qwen2VLForConditionalGeneration
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)
# pdb.set_trace()
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] # [B, L, D]
with torch.autocast(device_type='cuda', dtype=torch.float32):
logits = self.rm_head(hidden_states) # [B, L, N]
batch_size = input_ids.shape[0] if input_ids is not None else inputs_embeds.shape[0]
## get sequence length
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
else:
if input_ids is not None:
# if no pad token found, use modulo instead of reverse indexing for ONNX compatibility
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
## get the last token's logits
if self.reward_token == "last":
pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths]
elif self.reward_token == "mean":
## get the mean of all valid tokens' logits
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_ids = self.tokenizer.convert_tokens_to_ids(self.special_tokens)
# create a mask for special tokens
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) # [B, 3, N] assert 3 attributes
pooled_logits = pooled_logits.view(batch_size, -1)
# pdb.set_trace()
else:
raise ValueError("Invalid reward_token")
return {"logits": pooled_logits}
def _convert_A_B_to_chosen_rejected(
rewards_A,
rewards_B,
tied_threshold=None,
choice_dist=None,
):
"""
Inputs:
rewards_A: [B, 1]
rewards_B: [B, 1]
Outputs:
rewards_chosen: [B, 1]
rewards_rejected: [B, 1]
nontied_mask: [B, 1] (preference labels that is not tied)
"""
chosen_label = torch.ones_like(rewards_A, dtype=torch.int64).to(rewards_A.device) # [B, 1]
rewards_chosen = rewards_A
rewards_rejected = rewards_B
if tied_threshold is None:
nontied_mask = torch.ones_like(chosen_label, dtype=torch.float32).to(rewards_A.device)
else:
nontied_mask = (torch.abs((choice_dist[:, 0] - choice_dist[:, 1]) / torch.sum(choice_dist, dim=-1))
> tied_threshold)
print(nontied_mask)
return (
rewards_chosen,
rewards_rejected,
nontied_mask,
)
class PartialEmbeddingUpdateCallback(TrainerCallback):
"""
Callback to update the embedding of special tokens
Only the special tokens are updated, the rest of the embeddings are kept fixed
"""
def __init__(self, special_token_ids):
super().__init__()
self.special_token_ids = special_token_ids
self.orig_embeds_params = None
def on_train_begin(self, args, state, control, **kwargs):
model = kwargs.get("model")
self.orig_embeds_params = model.get_input_embeddings().weight.clone().detach()
def on_step_end(self, args, state, control, **kwargs):
# pdb.set_trace()
model = kwargs.get("model")
tokenizer = kwargs.get("tokenizer")
index_no_updates = torch.ones((len(tokenizer), ), dtype=torch.bool)
index_no_updates[self.special_token_ids] = False
with torch.no_grad():
model.get_input_embeddings().weight[index_no_updates] = (self.orig_embeds_params[index_no_updates])
class VLMRewardTrainer(RewardTrainer):
def __init__(self,
loss_type="regular",
loss_hyperparameters=None,
tied_threshold=None,
visualization_steps=500,
max_viz_samples=4,
*args,
**kwargs):
super().__init__(*args, **kwargs)
self.loss_type = loss_type
self.tied_threshold = tied_threshold
self.rewards_chosen_accumulated = []
self.rewards_rejected_accumulated = []
self.loss_hyperparameters = loss_hyperparameters if loss_hyperparameters is not None else {}
self.visualization_steps = visualization_steps
self.max_viz_samples = max_viz_samples
def get_eval_dataloader(self, eval_dataset: str | Dataset | None = None) -> DataLoader:
"""
Returns the evaluation [`~torch.utils.data.DataLoader`].
Subclass and override this method if you want to inject some custom behavior.
Args:
eval_dataset (`str` or `torch.utils.data.Dataset`, *optional*):
If a `str`, will use `self.eval_dataset[eval_dataset]` as the evaluation dataset. If a `Dataset`, will override `self.eval_dataset` and must implement `__len__`. If it is a [`~datasets.Dataset`], columns not accepted by the `model.forward()` method are automatically removed.
"""
if eval_dataset is None and self.eval_dataset is None:
raise ValueError("Trainer: evaluation requires an eval_dataset.")
# If we have persistent workers, don't do a fork bomb especially as eval datasets
# don't change during training
dataloader_key = eval_dataset if isinstance(eval_dataset, str) else "eval"
if (hasattr(self, "_eval_dataloaders") and dataloader_key in self._eval_dataloaders
and self.args.dataloader_persistent_workers):
return self.accelerator.prepare(self._eval_dataloaders[dataloader_key])
eval_dataset = (self.eval_dataset[eval_dataset] if isinstance(eval_dataset, str) else
eval_dataset if eval_dataset is not None else self.eval_dataset)
data_collator = self.data_collator
if is_datasets_available() and isinstance(eval_dataset, datasets.Dataset):
eval_dataset = self._remove_unused_columns(eval_dataset, description="evaluation")
else:
data_collator = self._get_collator_with_removed_columns(data_collator, description="evaluation")
dataloader_params = {
"batch_size": self.args.eval_batch_size,
"collate_fn": data_collator,
"num_workers": self.args.dataloader_num_workers,
"pin_memory": self.args.dataloader_pin_memory,
"persistent_workers": self.args.dataloader_persistent_workers,
}
if not isinstance(eval_dataset, torch.utils.data.IterableDataset):
dataloader_params["sampler"] = self._get_eval_sampler(eval_dataset)
dataloader_params["drop_last"] = self.args.dataloader_drop_last
dataloader_params["prefetch_factor"] = self.args.dataloader_prefetch_factor
# accelerator.free_memory() will destroy the references, so
# we need to store the non-prepared version
eval_dataloader = DataLoader(eval_dataset, **dataloader_params)
if self.args.dataloader_persistent_workers:
if hasattr(self, "_eval_dataloaders"):
self._eval_dataloaders[dataloader_key] = eval_dataloader
else:
self._eval_dataloaders = {dataloader_key: eval_dataloader}
return self.accelerator.prepare(eval_dataloader)
def create_optimizer(self):
"""
Setup the optimizer.
We provide a reasonable default that works well. If you want to use something else, you can pass a tuple in the
Trainer's init through `optimizers`, or subclass and override this method in a subclass.
"""
if is_sagemaker_mp_enabled():
return super().create_optimizer()
opt_model = self.model
if self.optimizer is None:
decay_parameters = self.get_decay_parameter_names(opt_model)
decay_parameters = [name for name in decay_parameters if "bias" not in name]
lr_mapper = {}
visual_parameters = []
merger_parameters = []
rm_head_parameters = []
if self.args.vision_lr is not None:
lr_mapper["visual"] = self.args.vision_lr
visual_parameters = [
name for name, _ in opt_model.named_parameters() if "visual" in name and "merger" not in name
]
if self.args.merger_lr is not None:
lr_mapper["merger"] = self.args.merger_lr
merger_parameters = [name for name, _ in opt_model.named_parameters() if "merger" in name]
if self.args.rm_head_lr is not None:
lr_mapper["rm_head"] = self.args.rm_head_lr
rm_head_parameters = [name for name, _ in opt_model.named_parameters() if "rm_head" in name]
if len(lr_mapper) > 0:
special_lr_parameters = merger_parameters + visual_parameters + rm_head_parameters
optimizer_grouped_parameters = [
{
"params": [
p for n, p in opt_model.named_parameters()
if (n in decay_parameters and n not in special_lr_parameters and p.requires_grad)
],
"weight_decay":
self.args.weight_decay,
},
{
"params": [
p for n, p in opt_model.named_parameters()
if (n not in decay_parameters and n not in special_lr_parameters and p.requires_grad)
],
"weight_decay":
0.0,
},
]
if visual_parameters:
optimizer_grouped_parameters.extend([
{
"params": [
p for n, p in opt_model.named_parameters()
if (n in decay_parameters and n in visual_parameters and p.requires_grad)
],
"weight_decay":
self.args.weight_decay,
"lr":
self.args.vision_lr,
},
{
"params": [
p for n, p in opt_model.named_parameters()
if (n not in decay_parameters and n in visual_parameters and p.requires_grad)
],
"weight_decay":
0.0,
"lr":
self.args.vision_lr,
},
])
if merger_parameters:
optimizer_grouped_parameters.extend([
{
"params": [
p for n, p in opt_model.named_parameters()
if (n in decay_parameters and n in merger_parameters and p.requires_grad)
],
"weight_decay":
self.args.weight_decay,
"lr":
self.args.merger_lr,
},
{
"params": [
p for n, p in opt_model.named_parameters()
if (n not in decay_parameters and n in merger_parameters and p.requires_grad)
],
"weight_decay":
0.0,
"lr":
self.args.merger_lr,
},
])
if rm_head_parameters:
optimizer_grouped_parameters.extend([
{
"params": [
p for n, p in opt_model.named_parameters()
if (n in decay_parameters and n in rm_head_parameters and p.requires_grad)
],
"weight_decay":
self.args.weight_decay,
"lr":
self.args.rm_head_lr,
},
{
"params": [
p for n, p in opt_model.named_parameters()
if (n not in decay_parameters and n in rm_head_parameters and p.requires_grad)
],
"weight_decay":
0.0,
"lr":
self.args.rm_head_lr,
},
])
else:
optimizer_grouped_parameters = [
{
"params":
[p for n, p in opt_model.named_parameters() if (n in decay_parameters and p.requires_grad)],
"weight_decay":
self.args.weight_decay,
},
{
"params":
[p for n, p in opt_model.named_parameters() if (n not in decay_parameters and p.requires_grad)],
"weight_decay":
0.0,
},
]
if self.model.special_token_ids:
special_token_embeddings = opt_model.get_input_embeddings().weight
special_token_embeddings.requires_grad = True
optimizer_grouped_parameters.extend([
{
# "params": [p for n, p in opt_model.get_input_embeddings().named_parameters() if (p.requires_grad)],
"params": [special_token_embeddings],
"lr": self.args.special_token_lr,
"weight_decay": 0.0,
},
])
optimizer_cls, optimizer_kwargs = self.get_optimizer_cls_and_kwargs(self.args, opt_model)
self.optimizer = optimizer_cls(optimizer_grouped_parameters, **optimizer_kwargs)
return self.optimizer
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
rewards_A = model(return_dict=True, **inputs["batch_1"])["logits"]
rewards_B = model(return_dict=True, **inputs["batch_2"])["logits"]
# Log to TensorBoard for visualization
if (hasattr(self.state, 'global_step') and self.state.global_step % self.visualization_steps == 0
and self.state.global_step > 0):
# Pass the original inputs which should contain the text prompts
self._log_training_visualization(inputs, rewards_A, rewards_B)
# calculate loss, optionally modulate with margin
# get chosen and rejected rewards from the chosen label
(
rewards_chosen,
rewards_rejected,
nontied_mask,
) = _convert_A_B_to_chosen_rejected(
rewards_A,
rewards_B,
tied_threshold=self.tied_threshold,
choice_dist=inputs["choice_dist"],
)
loss_dict = {}
if self.loss_type == "bt":
# Bradley-Terry model
loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected)
out_mask = nontied_mask
loss = loss * out_mask
loss = loss.mean()
elif self.loss_type == "likelihood_displacement":
# Bradley-Terry model
loss = -nn.functional.logsigmoid(rewards_chosen - self.loss_hyperparameters['tau'] * rewards_rejected)
out_mask = nontied_mask
loss = loss * out_mask
loss = loss.mean()
elif self.loss_type == "constant_margin":
# Bradley-Terry model with constant margin
loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - 0.57)
out_mask = nontied_mask
loss = loss * out_mask
loss = loss.mean()
elif self.loss_type == "btt":
# Bradley-Terry-With-Ties model
k = 5.0
log_k = math.log(k)
log_k2_sub_1 = math.log(k**2 - 1)
bt_loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - log_k)
same_loss = (-nn.functional.logsigmoid(rewards_chosen - rewards_rejected - log_k) -
nn.functional.logsigmoid(rewards_rejected - rewards_chosen - log_k) - log_k2_sub_1)
loss = bt_loss * nontied_mask.float() + same_loss * (1 - nontied_mask.float())
out_mask = torch.ones_like(nontied_mask, dtype=torch.float32).to(rewards_A.device) # [B, 1]
loss = loss * out_mask
loss = loss.mean()
elif self.loss_type == "hpsv2":
device = rewards_A.device
rewards = torch.nn.functional.softmax(torch.cat([rewards_A, rewards_B], dim=-1), dim=-1)
text_0_logits, text_1_logits = rewards[:, 0], rewards[:, 1]
label_0, label_1 = torch.ones_like(text_0_logits), torch.zeros_like(text_0_logits)
text_logits = torch.stack([text_0_logits, text_1_logits], dim=-1)
text_0_labels = torch.zeros(text_logits.shape[0], device=device, dtype=torch.long)
text_1_labels = text_0_labels + 1
text_0_loss = torch.nn.functional.cross_entropy(text_logits, text_0_labels, reduction="none")
text_1_loss = torch.nn.functional.cross_entropy(text_logits, text_1_labels, reduction="none")
loss = label_0 * text_0_loss + label_1 * text_1_loss
# absolute_example_weight = 1 / num_per_prompt
# denominator = absolute_example_weight.sum()
# weight_per_example = absolute_example_weight / denominator
# text_loss *= weight_per_example
loss = loss.sum()
elif self.loss_type == "uncertainty":
batch_size = rewards_A.shape[0]
mean_chosen = rewards_A[:, 0]
mean_rejected = rewards_B[:, 0]
sigma_chosen = torch.exp(rewards_A[:, 1])
sigma_rejected = torch.exp(rewards_B[:, 1])
mean_z = mean_chosen - mean_rejected
sigma_z = torch.sqrt(sigma_chosen**2 + sigma_rejected**2)
z_samples = torch.randn(batch_size, 1000).to(sigma_z.device).to(
torch.float16) * sigma_z.unsqueeze(1).repeat(1, 1000) + mean_z.unsqueeze(1).repeat(1, 1000)
loss = -torch.nn.functional.logsigmoid(z_samples).mean()
else:
raise NotImplementedError(f"Loss type {self.loss_type} not implemented.")
loss_dict.update({"loss": loss.item()})
if return_outputs:
## return rewards_A/B instead of chosen/rejected
## easier to calculate metrics for multi-attribute
return loss, {
"rewards_A": rewards_A,
"rewards_B": rewards_B,
}
return loss
def prediction_step(
self,
model,
inputs,
prediction_loss_only,
ignore_keys=None,
):
model.eval()
inputs = self._prepare_inputs(inputs)
if ignore_keys is None:
if hasattr(self.model, "config"):
ignore_keys = getattr(self.model.config, "keys_to_ignore_at_inference", [])
else:
ignore_keys = []
with torch.no_grad():
loss, logits_dict = self.compute_loss(model, inputs, return_outputs=True)
if prediction_loss_only:
return (loss, None, None)
loss = loss.detach()
logits = tuple(v for k, v in logits_dict.items() if k not in ignore_keys)
logits = nested_detach(logits)
if self.loss_type != "uncertainty":
logits = torch.cat(logits, dim=1) # [B, 2]
else:
logits = torch.cat([p[:, [0]] for p in logits], dim=1)
labels = torch.ones((logits.shape[0], 1)).to(logits.device)
return loss, logits, labels
def _log_training_visualization(self, inputs, rewards_A, rewards_B):
"""Log training samples and predictions to TensorBoard"""
try:
# Get tensorboard writer from trainer
writer = None
if hasattr(self, 'log_metrics') and hasattr(self.args,
'report_to') and 'tensorboard' in self.args.report_to:
# Try to get the writer from the logger
from torch.utils.tensorboard import SummaryWriter
if not hasattr(self, '_tb_writer'):
self._tb_writer = SummaryWriter(log_dir=self.args.logging_dir)
writer = self._tb_writer
if writer is None:
return
step = self.state.global_step
batch_size = min(len(rewards_A), self.max_viz_samples)
# Log scalar metrics
for i in range(batch_size):
score_A = rewards_A[i].float().detach().cpu().numpy()
score_B = rewards_B[i].float().detach().cpu().numpy()
# Convert to float for logging
score_A_val = float(score_A.mean()) if score_A.ndim > 0 else float(score_A)
score_B_val = float(score_B.mean()) if score_B.ndim > 0 else float(score_B)
score_diff = score_A_val - score_B_val
writer.add_scalar(f'train_viz/sample_{i}/score_A', score_A_val, step)
writer.add_scalar(f'train_viz/sample_{i}/score_B', score_B_val, step)
writer.add_scalar(f'train_viz/sample_{i}/score_diff', score_diff, step)
try:
# Get image data from inputs
image_A = inputs['image_1'][i] if 'image_1' in inputs else None
image_B = inputs['image_2'][i] if 'image_2' in inputs else None
# Get prompt text from the original batch (now properly stored)
prompt_A = inputs.get('text_1', ['Unknown prompt'])[i] if 'text_1' in inputs else 'Unknown prompt'
fig, axes = plt.subplots(nrows=1, ncols=2, figsize=(12, 8))
fig.text(0.05,
0.05,
f'Prompt:\n{prompt_A[:200]}{"..." if len(prompt_A) > 200 else ""}',
ha='left',
va='bottom',
fontsize=8,
wrap=True,
bbox=dict(boxstyle="round,pad=0.3", facecolor="lightblue", alpha=0.7))
img_A_np = np.array(image_A)
if img_A_np.ndim == 3 and img_A_np.shape[0] == 3: # CHW format
img_A_np = np.transpose(img_A_np, (1, 2, 0))
img_A_np = np.clip(img_A_np, 0, 1) # Ensure values are in [0,1]
axes[0].imshow(img_A_np)
axes[0].set_title(f'Image A - Score: {score_A_val:.3f}')
axes[0].axis('off')
img_B_np = np.array(image_B)
if img_B_np.ndim == 3 and img_B_np.shape[0] == 3: # CHW format
img_B_np = np.transpose(img_B_np, (1, 2, 0))
img_B_np = np.clip(img_B_np, 0, 1) # Ensure values are in [0,1]
axes[1].imshow(img_B_np)
axes[1].set_title(f'Image B - Score: {score_B_val:.3f}')
axes[1].axis('off')
# Add prediction info
winner = "A" if score_diff > 0 else "B"
plt.suptitle(
f'Step {step} - Sample {i} | Predicted Winner: Image {winner} | Diff: {score_diff:.3f}',
fontsize=14)
plt.tight_layout()
# Log figure to tensorboard
writer.add_figure(f'train_viz/sample_{i}_comparison', fig, step)
plt.close(fig)
except Exception as viz_error:
print(f"Warning: Could not extract images for visualization: {viz_error}")
continue
# Log aggregate statistics
all_scores_A = rewards_A.float().detach().cpu().numpy()
all_scores_B = rewards_B.float().detach().cpu().numpy()
writer.add_histogram('train_viz/all_scores_A', all_scores_A, step)
writer.add_histogram('train_viz/all_scores_B', all_scores_B, step)
writer.add_scalar('train_viz/mean_score_A', float(all_scores_A.mean()), step)
writer.add_scalar('train_viz/mean_score_B', float(all_scores_B.mean()), step)
writer.add_scalar('train_viz/mean_score_diff', float((all_scores_A - all_scores_B).mean()), step)
except Exception as e:
print(f"Error in training visualization: {e}")
def _save_checkpoint(self, model, trial, metrics=None):
if isinstance(self.model, PeftModel):
checkpoint_folder = f"{PREFIX_CHECKPOINT_DIR}-{self.state.global_step}"
if self.hp_search_backend is None and trial is None:
self.store_flos()
run_dir = self._get_output_dir(trial=trial)
output_dir = os.path.join(run_dir, checkpoint_folder)
os.makedirs(output_dir, exist_ok=True)
# TODO: Just Temp
self.save_model(output_dir, _internal_call=True)
# pdb.set_trace()
if not self.args.save_full_model:
non_lora_weights = get_peft_state_non_lora_maybe_zero_3(self.model.named_parameters(),
require_grad_only=True)
torch.save(
non_lora_weights,
os.path.join(output_dir, "non_lora_state_dict.pth"),
)
# safetensors.torch.save(non_lora_weights, os.path.join(output_dir, "non_lora_model.safetensors"))
if not self.args.save_only_model:
# Save optimizer and scheduler
self._save_optimizer_and_scheduler(output_dir)
# Save RNG state
self._save_rng_state(output_dir)
else:
super(RewardTrainer, self)._save_checkpoint(model, trial, metrics)
def _save(self, output_dir: str | None = None, state_dict=None):
# If we are executing this function, we are the process zero, so we don't check for that.
output_dir = output_dir if output_dir is not None else self.args.output_dir
os.makedirs(output_dir, exist_ok=True)
logger.info(f"Saving model checkpoint to {output_dir}")
# pdb.set_trace()
supported_classes = ((PreTrainedModel, ) if not is_peft_available() else (PreTrainedModel, PeftModel))
# Save a trained model and configuration using `save_pretrained()`.
# They can then be reloaded using `from_pretrained()`
if not isinstance(self.model, supported_classes):
if state_dict is None:
state_dict = self.model.state_dict()
if isinstance(self.accelerator.unwrap_model(self.model), supported_classes):
self.accelerator.unwrap_model(self.model).save_pretrained(
output_dir,
state_dict=state_dict,
safe_serialization=self.args.save_safetensors,
)
else:
logger.info("Trainer.model is not a `PreTrainedModel`, only saving its state dict.")
if self.args.save_safetensors:
safetensors.torch.save_file(
state_dict,
os.path.join(output_dir, SAFE_WEIGHTS_NAME),
metadata={"format": "pt"},
)
else:
torch.save(state_dict, os.path.join(output_dir, WEIGHTS_NAME))
else:
if not self.args.save_full_model:
state_dict = {k: v for k, v in state_dict.items() if "wte" not in k}
self.model.save_pretrained(
output_dir,
state_dict=state_dict,
safe_serialization=self.args.save_safetensors,
)
else:
torch.save(state_dict, os.path.join(output_dir, "model.pth"))
if self.tokenizer is not None:
os.makedirs(os.path.join(output_dir, "tokenizer"), exist_ok=True)
self.tokenizer.save_pretrained(os.path.join(output_dir, "tokenizer"))
# Good practice: save your training arguments together with the trained model
torch.save(self.args, os.path.join(output_dir, TRAINING_ARGS_NAME))
# pdb.set_trace()
def compute_multi_attr_accuracy(eval_pred, metainfo_idxs=None) -> dict[str, float]:
predictions, labels = eval_pred
metrics = {}
pred_curr = predictions
label_curr = labels.squeeze(1)
total_count = np.sum(label_curr != 0)
rewards_chosen = pred_curr[:, 0]
rewards_rejected = pred_curr[:, 1]
rewards_chosen_avg = np.sum(rewards_chosen) / total_count
rewards_rejected_avg = np.sum(rewards_rejected) / total_count
accuracy = np.sum(rewards_chosen > rewards_rejected) / total_count
metrics.update({
"Acc": accuracy,
"R_chosen_avg": rewards_chosen_avg,
"R_rejected_avg": rewards_rejected_avg,
})
return metrics
+286
View File
@@ -0,0 +1,286 @@
import json
import os
import fire
from dataclasses import asdict
from functools import partial
import torch
from .model.qwen2vl_trainer import (
Qwen2VLRewardModelBT,
VLMRewardTrainer,
compute_multi_attr_accuracy,
PartialEmbeddingUpdateCallback,
)
from .dataset.pairwise_dataset import PairwiseOriginalDataset
from .dataset.data_collator_qwen import QWen2VLDataCollator
from .utils.parser import ModelConfig, PEFTLoraConfig, TrainingConfig, DataConfig
from .utils.training_utils import load_model_from_checkpoint, find_target_linear_names
from .utils.parser import parse_args_with_yaml
from transformers import AutoProcessor
from peft import LoraConfig, get_peft_model
from trl import get_kbit_device_map, get_quantization_config
from .model.differentiable_image_processor import Qwen2VLImageProcessor
try:
import flash_attn
except ImportError:
flash_attn = None
print("Flash Attention is not installed. Falling to SDPA.")
def create_model_and_processor(
model_config,
peft_lora_config,
training_args,
cache_dir=None,
differentiable=False,
):
# create model
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=get_kbit_device_map() if quantization_config is not None else None,
quantization_config=quantization_config,
use_cache=False)
# create processor and set padding
processor = AutoProcessor.from_pretrained(model_config.model_name_or_path,
padding_side="right",
cache_dir=cache_dir)
if differentiable:
processor.image_processor = Qwen2VLImageProcessor()
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)
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 flash_attn is not None 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)
# create lora and peft model
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 save_configs_to_json(data_config, training_args, model_config, peft_lora_config):
"""
Save all configurations to a JSON file.
"""
config_dict = {
"data_config": asdict(data_config),
"training_args": asdict(training_args),
"model_config": asdict(model_config),
"peft_lora_config": asdict(peft_lora_config),
}
# del information about local device
del config_dict["training_args"]["local_rank"]
del config_dict["training_args"]["_n_gpu"]
save_path = os.path.join(training_args.output_dir, "model_config.json")
os.makedirs(training_args.output_dir, exist_ok=True)
print(training_args.output_dir)
with open(save_path, "w") as f:
json.dump(config_dict, f, indent=4)
def set_requires_grad(parameters, requires_grad):
for p in parameters:
p.requires_grad = requires_grad
def train(config, local_rank=0, debug=False):
## ===> Step 1: Parse arguments
(data_config, training_args, model_config, peft_lora_config), config_path = (parse_args_with_yaml(
(DataConfig, TrainingConfig, ModelConfig, PEFTLoraConfig), config, is_train=True))
training_args.output_dir = os.path.join(training_args.output_dir, config.split("/")[-1].split(".")[0])
training_args.logging_dir = training_args.output_dir
# check valid (lora config)
assert not (peft_lora_config.lora_enable and model_config.freeze_llm
), "When using LoRA, the LLM should not be frozen. If you want to freeze the LLM, please disable LoRA."
if not peft_lora_config.lora_enable:
assert (not peft_lora_config.vision_lora
), "Error: model_config.lora_enable is not enabled, but model_config.vision_lora is enabled."
else:
if peft_lora_config.lora_namespan_exclude is None:
peft_lora_config.lora_namespan_exclude = []
if not peft_lora_config.vision_lora:
peft_lora_config.lora_namespan_exclude += ["visual"]
## ===> Step 2: Load model and configure
model, processor, peft_config = create_model_and_processor(
model_config=model_config,
peft_lora_config=peft_lora_config,
training_args=training_args,
)
## load model
if training_args.load_from_pretrained is not None:
model, checkpoint_step = load_model_from_checkpoint(
model,
training_args.load_from_pretrained,
training_args.load_from_pretrained_step,
)
model.train()
if peft_lora_config.lora_enable:
model_to_configure = model.model
else:
model_to_configure = model
# set requires_grad for LLM
set_requires_grad(model_to_configure.model.parameters(), not model_config.freeze_llm)
set_requires_grad(model_to_configure.model.embed_tokens.parameters(), False)
if not peft_lora_config.vision_lora:
# set requires_grad for visual encoder and merger
set_requires_grad(model_to_configure.visual.parameters(), not model_config.freeze_vision_tower)
set_requires_grad(model_to_configure.visual.merger.parameters(), model_config.tune_merger)
if model_config.trainable_visual_layers: # This is inverse order to index of model.visual.blocks, set -1 to unfreeze all layers
assert model_config.trainable_visual_layers <= len(
model_to_configure.visual.blocks
), "trainable_visual_layers should be less than or equal to the number of visual blocks"
freeze_layer_num = len(
model_to_configure.visual.blocks
) - model_config.trainable_visual_layers if model_config.trainable_visual_layers > 0 else 0
for index, layer in enumerate(model_to_configure.visual.blocks):
if index < freeze_layer_num:
set_requires_grad(layer.parameters(), False)
else:
set_requires_grad(layer.parameters(), True)
# set requires_grad for regression head
set_requires_grad(model_to_configure.rm_head.parameters(), True)
## ===> Step 3: Load Dataset and configure
train_dataset = PairwiseOriginalDataset(
data_config.train_json_list,
data_config.soft_label,
data_config.confidence_threshold,
)
test_set_dict = {}
for item in data_config.test_json_list:
test_set_dict[item[0]] = PairwiseOriginalDataset(
item[1],
data_config.soft_label,
data_config.confidence_threshold,
)
print(f"===> Selected {len(train_dataset)} samples for training.")
for key, value in test_set_dict.items():
print(f"===> Selected {len(value)} samples for {key} testing.")
num_gpu = int(os.environ.get("WORLD_SIZE", 1))
data_collator = QWen2VLDataCollator(
processor,
max_pixels=data_config.max_pixels,
min_pixels=data_config.min_pixels,
with_instruction=data_config.with_instruction,
use_special_tokens=model_config.use_special_tokens,
)
compute_metrics = partial(compute_multi_attr_accuracy)
actual_batch_size = (training_args.per_device_train_batch_size * training_args.gradient_accumulation_steps *
num_gpu)
total_steps = (training_args.num_train_epochs * len(train_dataset) // actual_batch_size)
if training_args.save_epochs is not None:
training_args.save_steps = round(training_args.save_epochs * len(train_dataset) / actual_batch_size)
if training_args.eval_epochs is not None:
training_args.eval_steps = round(training_args.eval_epochs * len(train_dataset) / actual_batch_size)
if training_args.logging_epochs is not None:
training_args.logging_steps = round(training_args.logging_epochs * len(train_dataset) / actual_batch_size)
if training_args.local_rank == -1 or training_args.local_rank == 0:
print(f"===> Using {num_gpu} GPUs.")
print(f"===> Total Batch Size: {actual_batch_size}")
print(f"===> Training Epochs: {training_args.num_train_epochs}")
print(f"===> Total Steps: {total_steps}")
print(f"===> Save Steps: {training_args.save_steps}")
print(f"===> Eval Steps: {training_args.eval_steps}")
print(f"===> Logging Steps: {training_args.logging_steps}")
## ===> Step 4: Save configs for re-check
if training_args.local_rank == -1 or training_args.local_rank == 0:
save_configs_to_json(data_config, training_args, model_config, peft_lora_config)
print(train_dataset)
## ===> Step 5: Start Training!
special_token_ids = model.special_token_ids
callbacks = []
if special_token_ids is not None:
callbacks.append(PartialEmbeddingUpdateCallback(special_token_ids))
trainer = VLMRewardTrainer(
model=model,
compute_metrics=compute_metrics,
data_collator=data_collator,
args=training_args,
train_dataset=train_dataset,
eval_dataset=(test_set_dict if training_args.conduct_eval else None),
peft_config=peft_config,
callbacks=callbacks,
loss_type=model_config.loss_type,
loss_hyperparameters=model_config.loss_hyperparameters,
tokenizer=processor.tokenizer,
tied_threshold=data_config.tied_threshold,
visualization_steps=training_args.visualization_steps,
max_viz_samples=training_args.max_viz_samples,
)
trainer.train()
if training_args.local_rank == -1 or training_args.local_rank == 0:
model_state_dict = model.state_dict()
torch.save(model_state_dict, os.path.join(training_args.output_dir, "final_model.pth"))
model.config.save_pretrained(training_args.output_dir)
if __name__ == "__main__":
fire.Fire(train)
@@ -0,0 +1 @@
"""Vendored HPSv3 runtime utilities."""
@@ -0,0 +1,142 @@
from typing import Any, Literal
from omegaconf import OmegaConf
from transformers import HfArgumentParser
from dataclasses import dataclass, field
from transformers import 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=lambda: {})
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))
if not is_train:
args.pop('deepspeed', None)
# Parse with HfArgumentParser
parser = HfArgumentParser(dataclass_types)
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,144 @@
import torch
import os
import glob
import safetensors
def maybe_zero_3(param, ignore_status=False, name=None):
from deepspeed import zero
from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus
if hasattr(param, "ds_id"):
if param.ds_status == ZeroParamStatus.NOT_AVAILABLE and not ignore_status:
print(f"Parameter {name} is not available in ZeRO-3, please check the ZeRO-3 status.")
with zero.GatheredParameters([param]):
param = param.data.detach().cpu().clone()
else:
param = param.detach().cpu().clone()
return param
# Borrowed from peft.utils.get_peft_model_state_dict
def get_peft_state_maybe_zero_3(named_params, bias):
if bias == "none":
to_return = {k: t for k, t in named_params if "lora_" in k}
elif bias == "all":
to_return = {k: t for k, t in named_params if "lora_" in k or "bias" in k}
elif bias == "lora_only":
to_return = {}
maybe_lora_bias = {}
lora_bias_names = set()
for k, t in named_params:
if "lora_" in k:
to_return[k] = t
bias_name = k.split("lora_")[0] + "bias"
lora_bias_names.add(bias_name)
elif "bias" in k:
maybe_lora_bias[k] = t
for k, t in maybe_lora_bias:
if bias_name in lora_bias_names:
to_return[bias_name] = t
else:
raise NotImplementedError
to_return = {k: maybe_zero_3(v, ignore_status=True) for k, v in to_return.items()}
return to_return
def get_peft_state_non_lora_maybe_zero_3(named_params, require_grad_only=True):
to_return = {k: t for k, t in named_params if "lora_" not in k}
if require_grad_only:
to_return = {k: t for k, t in to_return.items() if t.requires_grad}
to_return = {k: maybe_zero_3(v, ignore_status=True).cpu() for k, v in to_return.items()}
return to_return
def _insert_adapter_name_into_state_dict(state_dict: dict[str, torch.Tensor], adapter_name: str,
parameter_prefix: str) -> dict[str, torch.Tensor]:
"""Utility function to remap the state_dict keys to fit the PEFT model by inserting the adapter name."""
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 save_video(tensor, path):
from torchvision.io import write_video
tensor = tensor * 255.0
tensor = tensor.permute(0, 2, 3, 1)
tensor = tensor.clamp(0, 255).byte()
write_video(path, tensor, 4, video_codec="h264")
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:
# get the latest checkpoint
checkpoint_path = checkpoint_paths[0]
print(f"===> Checkpoint step is not provided, using the latest checkpoint: {checkpoint_path}")
else:
checkpoint_path = os.path.join(checkpoint_dir, f"checkpoint-{checkpoint_step}")
if checkpoint_path not in checkpoint_paths:
checkpoint_path = checkpoint_paths[0]
print(f"===> Checkpoint step {checkpoint_step} not found, using the latest checkpoint: {checkpoint_path}")
else:
print(f"===> Checkpoint step {checkpoint_step} found, using the specified checkpoint: {checkpoint_path}")
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
def find_target_linear_names(model, num_lora_modules=-1, lora_namespan_exclude=None, verbose=False):
"""
Find the target linear modules for LoRA.
"""
linear_cls = torch.nn.Linear
embedding_cls = torch.nn.Embedding
if lora_namespan_exclude is None:
lora_namespan_exclude = []
lora_module_names = []
for name, module in model.named_modules():
if any(ex_keyword in name for ex_keyword in lora_namespan_exclude):
# print(f"Excluding module: {name}")
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:]
if verbose:
print(f"Found {len(lora_module_names)} lora modules: {lora_module_names}")
return lora_module_names
@@ -0,0 +1,17 @@
"""Vendored runtime subset of VideoAlign.
Source: https://github.com/KlingAIResearch/VideoAlign
Commit: 219ab9db64c045e5181a2202d11f686439351292
Purpose: Runtime reward inference integration for FastVideo GenRL.
This is temporary minimal vendoring for PR integration. It is expected to be
cleaned up and normalized later.
Porting rules:
- Include only files required by the runtime import closure used by FastVideo.
- When an upstream file is required, copy the entire file faithfully.
- Only adjust imports as needed to make the vendored code import through package
paths instead of sys.path mutation or ambiguous top-level imports.
- Do not perform style, typing, or behavioral cleanup as part of this vendoring
step.
"""
@@ -0,0 +1,8 @@
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 ..
```
@@ -0,0 +1,280 @@
from dataclasses import dataclass
import torch
from .prompt_template import build_prompt
# from qwen_vl_utils import process_vision_info
from .vision_process import process_vision_info
@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
def convert_GSB_csv_to_reward_data(example,
data_dir,
eval_dims=None,
max_pixels=448 * 448,
fps=2.0,
num_frames=None,
prompt_template_type="none",
sample_type="uniform"):
"""
Convert Good/Same/Bad csv data to reward data.
Args:
example (dict): A dataframe containing the GSB csv data.
data_dir (str): The directory path to the video files.
eval_dim (str): The dimension to evaluate ("VQ"/"MQ"/"TA").
max_pixels (int): The maximum number of pixels allowed for videos.
num_frames (float): Number of frames.
prompt_template_type (str): The type of prompt template to use ("none"/"simple"/"video_score").
Returns:
dict: A dictionary containing the reward data.
"""
if eval_dims is None:
eval_dims = ["VQ"]
A_data = [{
"role":
"user",
"content": [
{
"type": "video",
"video": f"file://{data_dir}/{example['path_A']}",
"max_pixels": max_pixels,
"fps": fps if num_frames is None else None,
"nframes": min(num_frames, example["num_frames_A"]) if num_frames is not None else None,
"sample_type": sample_type,
},
{
"type": "text",
"text": build_prompt(example["prompt"], eval_dims, prompt_template_type)
},
],
}]
B_data = [{
"role":
"user",
"content": [
{
"type": "video",
"video": f"file://{data_dir}/{example['path_B']}",
"max_pixels": max_pixels,
"fps": fps if num_frames is None else None,
"nframes": min(num_frames, example["num_frames_B"]) if num_frames is not None else None,
"sample_type": sample_type,
},
{
"type": "text",
"text": build_prompt(example["prompt"], eval_dims, prompt_template_type)
},
],
}]
chosen_labels = []
A_scores = []
B_scores = []
for eval_dim in eval_dims:
### chosen_label: 1 if A is chosen, -1 if B is chosen, 0 if tied.
### 22 if invalid. ooaaeeaa o.O
try:
if example[f"{eval_dim}"] is not None:
if example[f"{eval_dim}"] == "A":
chosen_label = 1
elif example[f"{eval_dim}"] == "B":
chosen_label = -1
elif example[f"{eval_dim}"] == "same":
chosen_label = 0
elif example[f"{eval_dim}"] == "invalid":
chosen_label = 22
else:
chosen_label = 22
else:
chosen_label = 22
except Exception:
chosen_label = 22
chosen_labels.append(chosen_label)
if f"MOS_A_{eval_dim}" in example and f"MOS_B_{eval_dim}" in example:
try:
A_score = example[f"MOS_A_{eval_dim}"] if example[f"MOS_A_{eval_dim}"] is not None else 0.0
B_score = example[f"MOS_B_{eval_dim}"] if example[f"MOS_B_{eval_dim}"] is not None else 0.0
except Exception:
A_score = 0.0
B_score = 0.0
A_scores.append(A_score)
B_scores.append(B_score)
else:
A_scores.append(0.0)
B_scores.append(0.0)
chosen_labels = torch.tensor(chosen_labels, dtype=torch.long)
A_scores = torch.tensor(A_scores, dtype=torch.float)
B_scores = torch.tensor(B_scores, dtype=torch.float)
metainfo_idx = None
if 'metainfo_idx' in example:
metainfo_idx = example['metainfo_idx']
return {
"A_data": A_data,
"B_data": B_data,
"A_scores": A_scores,
"B_scores": B_scores,
"chosen_label": chosen_labels,
"metainfo_idx": metainfo_idx,
}
class QWen2VLDataCollator:
def __init__(self, processor, add_noise=False, p_shuffle_frames=0.0, p_color_jitter=0.0):
self.processor = processor
self.add_noise = add_noise
self.set_noise_step = None
self.p_shuffle_frames = p_shuffle_frames
self.p_color_jitter = p_color_jitter
self.noise_adder = None
def _clean_message(self, message):
"""
remove unnecessary keys from message(very very necessary)
"""
message_content = message[0]["content"][0]
out_message = [{
"role":
"user",
"content": [
{
"type": "video",
"video": message_content["video"],
"max_pixels": message_content["max_pixels"],
"fps": message_content.get("fps", None),
"nframes": message_content.get("nframes", None),
"sample_type": message_content.get("sample_type", "uniform"),
},
{
"type": "text",
"text": message[0]["content"][1]["text"]
},
],
}]
if out_message[0]["content"][0]["fps"] is None:
out_message[0]["content"][0].pop("fps")
if out_message[0]["content"][0]["nframes"] is None:
out_message[0]["content"][0].pop("nframes")
return out_message
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 __call__(self, features, enable_noise=True):
"""
Preprocess inputs to token sequences and return a batch
"""
# try:
features_A = []
features_B = []
# check if we have a margin. If we do, we need to batch it as well
# has_margin = "margin" in features[0]
has_idx = "metainfo_idx" in features[0] and features[0]["metainfo_idx"] is not None
for idx, feature in enumerate(features):
features_A.append(self._clean_message(feature["A_data"]))
features_B.append(self._clean_message(feature["B_data"]))
# import pdb; pdb.set_trace()
image_inputs_A, video_inputs_A = process_vision_info(features_A)
image_inputs_B, video_inputs_B = process_vision_info(features_B)
video_inputs_A = [video_inputs_A[i].float() / 255.0 for i in range(len(video_inputs_A))]
video_inputs_B = [video_inputs_B[i].float() / 255.0 for i in range(len(video_inputs_B))]
do_rescale = False
# print(f"{video_inputs_A[0].shape}, {video_inputs_B[0].shape}")
# if not enable_noise:
# print("Not training, no noise added.")
batch_A = self.processor(
text=self.processor.apply_chat_template(features_A, tokenize=False, add_generation_prompt=True),
images=image_inputs_A,
videos=video_inputs_A,
padding=True,
return_tensors="pt",
videos_kwargs={"do_rescale": do_rescale},
)
batch_B = self.processor(
text=self.processor.apply_chat_template(features_B, tokenize=False, add_generation_prompt=True),
images=image_inputs_B,
videos=video_inputs_B,
padding=True,
return_tensors="pt",
videos_kwargs={"do_rescale": do_rescale},
)
# pdb.set_trace()
max_len = max(batch_A["input_ids"].shape[1], batch_B["input_ids"].shape[1])
batch_A["input_ids"], batch_A["attention_mask"] = self._pad_sequence(batch_A["input_ids"],
batch_A["attention_mask"], max_len,
"right")
batch_B["input_ids"], batch_B["attention_mask"] = self._pad_sequence(batch_B["input_ids"],
batch_B["attention_mask"], max_len,
"right")
# print(f"Batch A: {batch_A['input_ids'].shape}, Batch B: {batch_B['input_ids'].shape}")
chosen_label = torch.stack([torch.tensor(feature["chosen_label"]) for feature in features])
A_scores = torch.stack([torch.tensor(feature["A_scores"]) for feature in features])
B_scores = torch.stack([torch.tensor(feature["B_scores"]) for feature in features])
batch = {
"A": batch_A,
"B": batch_B,
"return_loss": True,
"chosen_label": chosen_label,
"A_scores": A_scores,
"B_scores": B_scores,
}
if has_idx:
metainfo_idx = torch.stack([torch.tensor(feature["metainfo_idx"]) for feature in features])
batch["metainfo_idx"] = metainfo_idx
# pdb.set_trace()
return batch
# except Exception as e:
# print(f"Error processing batch: {e} in reading.")
# # get next batch
# return None
@@ -0,0 +1,237 @@
import json
import os
from collections.abc import Mapping
import torch
from .vision_process import process_vision_info
from .data import DataConfig
from .utils import ModelConfig, PEFTLoraConfig, TrainingConfig
from .utils import load_model_from_checkpoint
from .train_reward import create_model_and_processor
from .prompt_template import build_prompt
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.get("inference_config", 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=dtype == torch.bfloat16,
fp16=dtype == torch.float16,
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
else:
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()})
elif isinstance(data, tuple | list):
return type(data)(self._prepare_input(v) for v in data)
elif isinstance(data, torch.Tensor):
kwargs = {"device": self.device}
## TODO: Maybe need to add dtype
# if self.is_deepspeed_enabled and (torch.is_floating_point(data) or torch.is_complex(data)):
# # NLP models inputs are int/uint and those get adjusted to the right dtype of the
# # embedding. Other models such as wav2vec2's inputs are already float and thus
# # may need special handling to match the dtypes of the model
# kwargs.update({"dtype": self.accelerator.state.deepspeed_plugin.hf_ds_config.dtype()})
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, strict=False)]
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, strict=False)]
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
if __name__ == "__main__":
load_from_pretrained = "./checkpoints"
device = "cuda:0"
dtype = torch.bfloat16
inferencer = VideoVLMRewardInference(load_from_pretrained, device=device, dtype=dtype)
video_paths = [
"datasets/train/videos/example_1_A.mp4",
"datasets/train/videos/example_1_B.mp4",
"datasets/train/videos/example_2_A.mp4",
]
prompts = [
"The camera remains still, a girl with braided hair and wearing a pink dress approached the chair in the room and sat on it, the background is a cozy bedroom, warm indoor lighting.",
"The camera remains still, a girl with braided hair and wearing a pink dress approached the chair in the room and sat on it, the background is a cozy bedroom, warm indoor lighting.",
"The camera follows a young explorer through an abandoned urban building at night, exploring hidden corridors and forgotten spaces, with a mix of light and shadow creating a mysterious atmosphere.",
]
with torch.no_grad():
rewards = inferencer.reward(video_paths, prompts, use_norm=True)
print(rewards)
@@ -0,0 +1,129 @@
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
elif template_type == "simple":
return SIMPLE_PROMPT.format(dimension_name=dimension_name,
dimension_description=dimension_description,
text_prompt=prompt)
elif template_type == "video_score":
return VIDEOSCORE_QUERY_PROMPT.format(dimension_name=dimension_name,
dimension_description=dimension_description,
text_prompt=prompt)
elif template_type == "detailed_special":
return DETAILED_PROMPT_WITH_SPECIAL_TOKEN.format(text_prompt=prompt)
elif template_type == "detailed":
return DETAILED_PROMPT.format(text_prompt=prompt)
else:
raise ValueError("Invalid template type")
@@ -0,0 +1,313 @@
import ast
import json
import os
import random
from dataclasses import asdict
from functools import partial
import torch
from datasets import load_dataset
from peft import LoraConfig, get_peft_model
from transformers import AutoProcessor, HfArgumentParser
from trl import get_kbit_device_map, get_quantization_config
from .trainer import Qwen2VLRewardModelBT, VideoVLMRewardTrainer, compute_multi_attr_accuracy, PartialEmbeddingUpdateCallback
from .data import DataConfig, QWen2VLDataCollator, convert_GSB_csv_to_reward_data
from .utils import ModelConfig, PEFTLoraConfig, TrainingConfig
from .utils import load_model_from_checkpoint
def save_configs_to_json(data_config, training_args, model_config, peft_lora_config):
"""
Save all configurations to a JSON file.
"""
config_dict = {
"data_config": asdict(data_config),
"training_args": asdict(training_args),
"model_config": asdict(model_config),
"peft_lora_config": asdict(peft_lora_config),
}
# del information about local device
del config_dict["training_args"]["local_rank"]
del config_dict["training_args"]["_n_gpu"]
save_path = os.path.join(training_args.output_dir, "model_config.json")
os.makedirs(training_args.output_dir, exist_ok=True)
print(training_args.output_dir)
with open(save_path, "w") as f:
json.dump(config_dict, f, indent=4)
def find_target_linear_names(model, num_lora_modules=-1, lora_namespan_exclude=None, verbose=False):
"""
Find the target linear modules for LoRA.
"""
linear_cls = torch.nn.Linear
embedding_cls = torch.nn.Embedding
if lora_namespan_exclude is None:
lora_namespan_exclude = []
lora_module_names = []
for name, module in model.named_modules():
if any(ex_keyword in name for ex_keyword in lora_namespan_exclude):
# print(f"Excluding module: {name}")
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:]
if verbose:
print(f"Found {len(lora_module_names)} lora modules: {lora_module_names}")
return lora_module_names
def set_requires_grad(parameters, requires_grad):
for p in parameters:
p.requires_grad = requires_grad
def create_model_and_processor(
model_config,
peft_lora_config,
training_args,
cache_dir=None,
):
# create model
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=get_kbit_device_map() if quantization_config is not None else None,
quantization_config=quantization_config,
use_cache=bool(training_args.gradient_checkpointing),
)
# pdb.set_trace()
# create processor and set padding
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)
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 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)
# create lora and peft model
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 create_dataset(data_config, meta_file=None):
if meta_file is None:
meta_file = data_config.meta_data
dataset = load_dataset('csv', data_files=meta_file)
def add_idx(example, idx):
example['metainfo_idx'] = idx
return example
dataset['train'] = dataset['train'].map(lambda example, idx: add_idx(example, idx), with_indices=True)
if not data_config.use_tied_data:
filter_func = lambda example: any(example[f"{dim}"] != "same" for dim in data_config.eval_dim)
dataset = dataset.filter(filter_func)
# convert data to reward data
convert_func = lambda example: convert_GSB_csv_to_reward_data(
example,
data_config.data_dir,
data_config.eval_dim,
data_config.max_frame_pixels,
data_config.fps,
data_config.num_frames,
data_config.prompt_template_type,
sample_type=data_config.sample_type,
)
dataset = dataset.map(convert_func, remove_columns=dataset['train'].column_names, load_from_cache_file=False)
dataset = dataset['train']
# pdb.set_trace()
return dataset
def train():
## ===> Step 1: Parse arguments
parser = HfArgumentParser((DataConfig, TrainingConfig, ModelConfig, PEFTLoraConfig))
data_config, training_args, model_config, peft_lora_config = parser.parse_args_into_dataclasses()
# pdb.set_trace()
# check valid (lora config)
assert not (peft_lora_config.lora_enable and model_config.freeze_llm
), 'When using LoRA, the LLM should not be frozen. If you want to freeze the LLM, please disable LoRA.'
if not peft_lora_config.lora_enable:
assert not peft_lora_config.vision_lora, \
"Error: model_config.lora_enable is not enabled, but model_config.vision_lora is enabled."
else:
if peft_lora_config.lora_namespan_exclude is not None:
peft_lora_config.lora_namespan_exclude = ast.literal_eval(peft_lora_config.lora_namespan_exclude)
else:
peft_lora_config.lora_namespan_exclude = []
if not peft_lora_config.vision_lora:
peft_lora_config.lora_namespan_exclude += ["visual"]
# pdb.set_trace()
## ===> Step 2: Load model and configure
model, processor, peft_config = create_model_and_processor(
model_config=model_config,
peft_lora_config=peft_lora_config,
training_args=training_args,
)
## load model
if training_args.load_from_pretrained is not None:
model, checkpoint_step = load_model_from_checkpoint(model, training_args.load_from_pretrained,
training_args.load_from_pretrained_step)
model.train()
if peft_lora_config.lora_enable:
model_to_configure = model.model
else:
model_to_configure = model
# set requires_grad for LLM
set_requires_grad(model_to_configure.model.parameters(), not model_config.freeze_llm)
if not peft_lora_config.vision_lora:
# set requires_grad for visual encoder and merger
set_requires_grad(model_to_configure.visual.parameters(), not model_config.freeze_vision_tower)
set_requires_grad(model_to_configure.visual.merger.parameters(), model_config.tune_merger)
# set requires_grad for regression head
set_requires_grad(model_to_configure.rm_head.parameters(), True)
## ===> Step 3: Load Dataset and configure
if isinstance(data_config.eval_dim, str):
data_config.eval_dim = [data_config.eval_dim]
# datasets = create_dataset(data_config)
# train_dataset = concatenate_datasets([datasets[dim] for dim in data_config.eval_dim])
train_dataset = create_dataset(data_config)
train_dataset = train_dataset.shuffle(seed=42)
if training_args.conduct_eval:
if data_config.meta_data_test is not None:
random.seed(42)
valid_dataset = create_dataset(data_config, meta_file=data_config.meta_data_test)
# indices = random.sample(range(len(valid_dataset)), 1000)
# valid_dataset = valid_dataset.select(indices)
else:
dataset = train_dataset.train_test_split(test_size=0.02)
train_dataset = dataset['train']
valid_dataset = dataset['test']
else:
valid_dataset = None
print(f"===> Selected {len(train_dataset)} samples for training.")
print(f"===> Selected {len(valid_dataset)} samples for testing.")
num_gpu = int(os.environ.get("WORLD_SIZE", 1))
data_collator = QWen2VLDataCollator(
processor,
add_noise=data_config.add_noise,
p_shuffle_frames=data_config.p_shuffle_frames,
p_color_jitter=data_config.p_color_jitter,
)
compute_metrics = partial(compute_multi_attr_accuracy, eval_dims=data_config.eval_dim)
actual_batch_size = training_args.per_device_train_batch_size * training_args.gradient_accumulation_steps * num_gpu
total_steps = training_args.num_train_epochs * len(train_dataset) // actual_batch_size
if training_args.save_epochs is not None:
training_args.save_steps = round(training_args.save_epochs * len(train_dataset) / actual_batch_size)
if training_args.eval_epochs is not None:
training_args.eval_steps = round(training_args.eval_epochs * len(train_dataset) / actual_batch_size)
if training_args.logging_epochs is not None:
training_args.logging_steps = round(training_args.logging_epochs * len(train_dataset) / actual_batch_size)
if training_args.local_rank == -1 or training_args.local_rank == 0:
print(f"===> Using {num_gpu} GPUs.")
print(f"===> Total Batch Size: {actual_batch_size}")
print(f"===> Training Epochs: {training_args.num_train_epochs}")
print(f"===> Total Steps: {total_steps}")
print(f"===> Save Steps: {training_args.save_steps}")
print(f"===> Eval Steps: {training_args.eval_steps}")
print(f"===> Logging Steps: {training_args.logging_steps}")
# pdb.set_trace()
## ===> Step 4: Save configs for re-check
if training_args.local_rank == -1 or training_args.local_rank == 0:
save_configs_to_json(data_config, training_args, model_config, peft_lora_config)
print(train_dataset)
## ===> Step 5: Start Training!
special_token_ids = model.special_token_ids
callbacks = []
if special_token_ids is not None:
callbacks.append(PartialEmbeddingUpdateCallback(special_token_ids))
trainer = VideoVLMRewardTrainer(
model=model,
compute_metrics=compute_metrics,
data_collator=data_collator,
args=training_args,
train_dataset=train_dataset,
eval_dataset=valid_dataset if training_args.conduct_eval else None,
peft_config=peft_config,
callbacks=callbacks,
loss_type=model_config.loss_type,
tokenizer=processor.tokenizer,
)
trainer.train()
if training_args.local_rank == -1 or training_args.local_rank == 0:
model_state_dict = model.state_dict()
torch.save(model_state_dict, os.path.join(training_args.output_dir, 'final_model.pth'))
model.config.save_pretrained(training_args.output_dir)
if __name__ == "__main__":
train()
@@ -0,0 +1,633 @@
import os
import math
# from training.train_utils import get_peft_state_maybe_zero_3, get_peft_state_non_lora_maybe_zero_3
import pandas as pd
import safetensors
import numpy as np
import torch
import torch.nn as nn
import datasets
from torch.utils.data import Dataset, DataLoader
from peft import PeftModel
from transformers import Qwen2VLForConditionalGeneration
from transformers.modeling_utils import PreTrainedModel
from transformers.trainer import TrainerCallback
from transformers.trainer import (
is_sagemaker_mp_enabled,
is_peft_available,
is_datasets_available,
WEIGHTS_NAME,
TRAINING_ARGS_NAME,
SAFE_WEIGHTS_NAME,
PREFIX_CHECKPOINT_DIR,
logger,
is_torch_xla_available,
)
from transformers.trainer_pt_utils import nested_detach
from trl import RewardTrainer
from .utils import get_peft_state_non_lora_maybe_zero_3
if is_torch_xla_available():
pass
else:
IS_XLA_FSDPV2_POST_2_2 = False
class Qwen2VLRewardModelBT(Qwen2VLForConditionalGeneration):
def __init__(self, config, output_dim=4, reward_token="last", special_token_ids=None):
super().__init__(config)
# pdb.set_trace()
self.output_dim = output_dim
self.rm_head = nn.Linear(config.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,
):
## modified from the origin class Qwen2VLForConditionalGeneration
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
# pdb.set_trace()
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] # [B, L, D]
logits = self.rm_head(hidden_states) # [B, L, N]
batch_size = input_ids.shape[0] if input_ids is not None else inputs_embeds.shape[0]
## get sequence length
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
else:
if input_ids is not None:
# if no pad token found, use modulo instead of reverse indexing for ONNX compatibility
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
## get the last token's logits
if self.reward_token == "last":
pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths]
elif self.reward_token == "mean":
## get the mean of all valid tokens' logits
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_ids = self.tokenizer.convert_tokens_to_ids(self.special_tokens)
# create a mask for special tokens
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) # [B, 3, N] assert 3 attributes
if self.output_dim == 3:
pooled_logits = pooled_logits.diagonal(dim1=1, dim2=2)
pooled_logits = pooled_logits.view(batch_size, -1)
# pdb.set_trace()
else:
raise ValueError("Invalid reward_token")
return {"logits": pooled_logits}
def _convert_A_B_to_chosen_rejected(rewards_A, rewards_B, scores_A, scores_B, chosen_label, label_dim=None):
"""
Inputs:
rewards_A: [B, N]
rewards_B: [B, N]
scores_A: [B, N]
scores_B: [B, N]
chosen_label: [B, N]
Outputs:
rewards_chosen: [B, N]
rewards_rejected: [B, N]
scores_chosen: [B, N]
scores_rejected: [B, N]
nontied_mask: [B, N] (preference labels that is not tied)
valid_mask: [B, N] (all valid labels)
"""
chosen_mask = (chosen_label == 1)
# rejected_mask = (chosen_label == -1)
rejected_mask = (chosen_label != 1)
if label_dim is not None:
N = chosen_label.size(1)
chosen_mask = chosen_mask[:, label_dim].unsqueeze(1).expand(-1, N)
rejected_mask = rejected_mask[:, label_dim].unsqueeze(1).expand(-1, N)
rewards_chosen = torch.where(chosen_mask, rewards_A, rewards_B)
rewards_rejected = torch.where(rejected_mask, rewards_A, rewards_B)
scores_chosen = torch.where(chosen_mask, scores_A, scores_B)
scores_rejected = torch.where(rejected_mask, scores_A, scores_B)
nontied_mask = ((chosen_label == 1) | (chosen_label == -1)).float()
if label_dim is not None:
nontied_mask = nontied_mask[:, label_dim].unsqueeze(1).expand(-1, N)
valid_mask = (chosen_label != 22).float()
if label_dim is not None:
valid_mask = valid_mask[:, label_dim].unsqueeze(1).expand(-1, N)
# rewards_chosen = rewards_chosen * valid_mask
# rewards_rejected = rewards_rejected * valid_mask
return rewards_chosen, rewards_rejected, scores_chosen, scores_rejected, nontied_mask, valid_mask
class PartialEmbeddingUpdateCallback(TrainerCallback):
"""
Callback to update the embedding of special tokens
Only the special tokens are updated, the rest of the embeddings are kept fixed
"""
def __init__(self, special_token_ids):
super().__init__()
self.special_token_ids = special_token_ids
self.orig_embeds_params = None
def on_train_begin(self, args, state, control, **kwargs):
model = kwargs.get("model")
self.orig_embeds_params = model.get_input_embeddings().weight.clone().detach()
def on_step_end(self, args, state, control, **kwargs):
# pdb.set_trace()
model = kwargs.get("model")
tokenizer = kwargs.get("tokenizer")
index_no_updates = torch.ones((len(tokenizer), ), dtype=torch.bool)
index_no_updates[self.special_token_ids] = False
with torch.no_grad():
model.get_input_embeddings().weight[index_no_updates] = self.orig_embeds_params[index_no_updates]
class VideoVLMRewardTrainer(RewardTrainer):
def __init__(self, loss_type="regular", enable_noise_in_eval=False, *args, **kwargs):
super().__init__(*args, **kwargs)
self.loss_type = loss_type
self.enable_noise_in_eval = enable_noise_in_eval
self.rewards_chosen_accumulated = []
self.rewards_rejected_accumulated = []
self.scores_chosen_accumulated = []
self.scores_rejected_accumulated = []
def get_eval_dataloader(self, eval_dataset: str | Dataset | None = None) -> DataLoader:
"""
Returns the evaluation [`~torch.utils.data.DataLoader`].
Subclass and override this method if you want to inject some custom behavior.
Args:
eval_dataset (`str` or `torch.utils.data.Dataset`, *optional*):
If a `str`, will use `self.eval_dataset[eval_dataset]` as the evaluation dataset. If a `Dataset`, will override `self.eval_dataset` and must implement `__len__`. If it is a [`~datasets.Dataset`], columns not accepted by the `model.forward()` method are automatically removed.
"""
if eval_dataset is None and self.eval_dataset is None:
raise ValueError("Trainer: evaluation requires an eval_dataset.")
# If we have persistent workers, don't do a fork bomb especially as eval datasets
# don't change during training
dataloader_key = eval_dataset if isinstance(eval_dataset, str) else "eval"
if (hasattr(self, "_eval_dataloaders") and dataloader_key in self._eval_dataloaders
and self.args.dataloader_persistent_workers):
return self.accelerator.prepare(self._eval_dataloaders[dataloader_key])
eval_dataset = (self.eval_dataset[eval_dataset] if isinstance(eval_dataset, str) else
eval_dataset if eval_dataset is not None else self.eval_dataset)
data_collator = lambda features: self.data_collator(features, enable_noise=self.enable_noise_in_eval)
if is_datasets_available() and isinstance(eval_dataset, datasets.Dataset):
eval_dataset = self._remove_unused_columns(eval_dataset, description="evaluation")
else:
data_collator = self._get_collator_with_removed_columns(data_collator, description="evaluation")
dataloader_params = {
"batch_size": self.args.eval_batch_size,
"collate_fn": data_collator,
"num_workers": self.args.dataloader_num_workers,
"pin_memory": self.args.dataloader_pin_memory,
"persistent_workers": self.args.dataloader_persistent_workers,
}
if not isinstance(eval_dataset, torch.utils.data.IterableDataset):
dataloader_params["sampler"] = self._get_eval_sampler(eval_dataset)
dataloader_params["drop_last"] = self.args.dataloader_drop_last
dataloader_params["prefetch_factor"] = self.args.dataloader_prefetch_factor
# accelerator.free_memory() will destroy the references, so
# we need to store the non-prepared version
eval_dataloader = DataLoader(eval_dataset, **dataloader_params)
if self.args.dataloader_persistent_workers:
if hasattr(self, "_eval_dataloaders"):
self._eval_dataloaders[dataloader_key] = eval_dataloader
else:
self._eval_dataloaders = {dataloader_key: eval_dataloader}
return self.accelerator.prepare(eval_dataloader)
def create_optimizer(self):
"""
Setup the optimizer.
We provide a reasonable default that works well. If you want to use something else, you can pass a tuple in the
Trainer's init through `optimizers`, or subclass and override this method in a subclass.
"""
if is_sagemaker_mp_enabled():
return super().create_optimizer()
opt_model = self.model
if self.optimizer is None:
decay_parameters = self.get_decay_parameter_names(opt_model)
decay_parameters = [name for name in decay_parameters if "bias" not in name]
lr_mapper = {}
visual_parameters = []
merger_parameters = []
if self.args.vision_lr is not None:
lr_mapper["visual"] = self.args.vision_lr
visual_parameters = [
name for name, _ in opt_model.named_parameters() if "visual" in name and "merger" not in name
]
if self.args.merger_lr is not None:
lr_mapper["merger"] = self.args.merger_lr
merger_parameters = [name for name, _ in opt_model.named_parameters() if "merger" in name]
if len(lr_mapper) > 0:
special_lr_parameters = merger_parameters + visual_parameters
optimizer_grouped_parameters = [
{
"params": [
p for n, p in opt_model.named_parameters()
if (n in decay_parameters and n not in special_lr_parameters and p.requires_grad)
],
"weight_decay":
self.args.weight_decay,
},
{
"params": [
p for n, p in opt_model.named_parameters()
if (n not in decay_parameters and n not in special_lr_parameters and p.requires_grad)
],
"weight_decay":
0.0,
},
]
if visual_parameters:
optimizer_grouped_parameters.extend([
{
"params": [
p for n, p in opt_model.named_parameters()
if (n in decay_parameters and n in visual_parameters and p.requires_grad)
],
"weight_decay":
self.args.weight_decay,
"lr":
self.args.vision_lr,
},
{
"params": [
p for n, p in opt_model.named_parameters()
if (n not in decay_parameters and n in visual_parameters and p.requires_grad)
],
"weight_decay":
0.0,
"lr":
self.args.vision_lr,
},
])
if merger_parameters:
optimizer_grouped_parameters.extend([
{
"params": [
p for n, p in opt_model.named_parameters()
if (n in decay_parameters and n in merger_parameters and p.requires_grad)
],
"weight_decay":
self.args.weight_decay,
"lr":
self.args.merger_lr,
},
{
"params": [
p for n, p in opt_model.named_parameters()
if (n not in decay_parameters and n in merger_parameters and p.requires_grad)
],
"weight_decay":
0.0,
"lr":
self.args.merger_lr,
},
])
else:
optimizer_grouped_parameters = [
{
"params":
[p for n, p in opt_model.named_parameters() if (n in decay_parameters and p.requires_grad)],
"weight_decay":
self.args.weight_decay,
},
{
"params":
[p for n, p in opt_model.named_parameters() if (n not in decay_parameters and p.requires_grad)],
"weight_decay":
0.0,
},
]
if self.model.special_token_ids:
special_token_embeddings = opt_model.get_input_embeddings().weight
special_token_embeddings.requires_grad = True
optimizer_grouped_parameters.extend([
{
# "params": [p for n, p in opt_model.get_input_embeddings().named_parameters() if (p.requires_grad)],
"params": [special_token_embeddings],
"lr": self.args.special_token_lr,
"weight_decay": 0.0,
},
])
optimizer_cls, optimizer_kwargs = self.get_optimizer_cls_and_kwargs(self.args, opt_model)
self.optimizer = optimizer_cls(optimizer_grouped_parameters, **optimizer_kwargs)
return self.optimizer
# def training_step(self, model, inputs, num_items_in_batch=None):
# pdb.set_trace()
# return super(VideoVLMRewardTrainer, self).training_step(model, inputs, num_items_in_batch)
def compute_loss(
self,
model,
inputs,
return_outputs=False,
):
rewards_A = model(return_dict=True, **inputs['A'])["logits"]
rewards_B = model(return_dict=True, **inputs['B'])["logits"]
# calculate loss, optionally modulate with margin
# get chosen and rejected rewards from the chosen label
rewards_chosen, rewards_rejected, scores_chosen, scores_rejected, nontied_mask, valid_mask = _convert_A_B_to_chosen_rejected(
rewards_A, rewards_B, inputs["A_scores"], inputs["B_scores"], inputs["chosen_label"])
# pdb.set_trace()
inputs["margin"] = scores_chosen - scores_rejected
loss_dict = {}
if self.loss_type == "bt":
# Bradley-Terry model
loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected)
out_mask = nontied_mask
elif self.loss_type == "margin":
# Bradley-Terry model with margin
loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - inputs["margin"])
out_mask = nontied_mask
elif self.loss_type == "constant_margin":
# Bradley-Terry model with constant margin
loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - 0.57)
out_mask = nontied_mask
elif self.loss_type == "scaled":
# Bradley-Terry model with scaled margin
loss = (-(inputs["margin"] + 0.0) * nn.functional.logsigmoid(rewards_chosen - rewards_rejected))
out_mask = nontied_mask
elif self.loss_type == "reg":
# regression loss
rewards = torch.stack([rewards_A, rewards_B], dim=1)
scores = torch.stack([inputs["A_scores"], inputs["B_scores"]], dim=1)
out_mask = scores != 0.0
scores = (scores - 3.0) # rescale
# pdb.set_trace()
loss = nn.functional.mse_loss(rewards, scores, reduction="none")
elif self.loss_type == "btt":
# Bradley-Terry-With-Ties model
k = 5.0
log_k = math.log(k)
log_k2_sub_1 = math.log(k**2 - 1)
bt_loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - log_k)
same_loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - log_k) \
-nn.functional.logsigmoid(rewards_rejected - rewards_chosen - log_k) \
-log_k2_sub_1
loss = bt_loss * nontied_mask + same_loss * (1 - nontied_mask)
out_mask = valid_mask
else:
raise NotImplementedError(f"Loss type {self.loss_type} not implemented.")
loss = loss * out_mask
loss = loss.mean()
loss_dict.update({"loss": loss.item()})
if return_outputs:
## return rewards_A/B instead of chosen/rejected
## easier to calculate metrics for multi-attribute
return loss, {
"rewards_A": rewards_A,
"rewards_B": rewards_B,
}
return loss
def prediction_step(
self,
model,
inputs,
prediction_loss_only,
ignore_keys=None,
):
inputs = self._prepare_inputs(inputs)
if ignore_keys is None:
if hasattr(self.model, "config"):
ignore_keys = getattr(self.model.config, "keys_to_ignore_at_inference", [])
else:
ignore_keys = []
with torch.no_grad():
loss, logits_dict = self.compute_loss(model, inputs, return_outputs=True)
if prediction_loss_only:
return (loss, None, None)
loss = loss.detach()
logits = tuple(v for k, v in logits_dict.items() if k not in ignore_keys)
logits = nested_detach(logits)
logits = torch.stack(logits).permute(1, 0, 2) # [B, 2, N]
labels = inputs["chosen_label"] # [B, N], values in {-1, 0, 1}
return loss, logits, labels
def _save_checkpoint(self, model, trial, metrics=None):
if isinstance(self.model, PeftModel):
checkpoint_folder = f"{PREFIX_CHECKPOINT_DIR}-{self.state.global_step}"
if self.hp_search_backend is None and trial is None:
self.store_flos()
run_dir = self._get_output_dir(trial=trial)
output_dir = os.path.join(run_dir, checkpoint_folder)
os.makedirs(output_dir, exist_ok=True)
# TODO: Just Temp
self.save_model(output_dir, _internal_call=True)
# pdb.set_trace()
if not self.args.save_full_model:
non_lora_weights = get_peft_state_non_lora_maybe_zero_3(self.model.named_parameters(),
require_grad_only=True)
torch.save(non_lora_weights, os.path.join(output_dir, "non_lora_state_dict.pth"))
# safetensors.torch.save(non_lora_weights, os.path.join(output_dir, "non_lora_model.safetensors"))
if not self.args.save_only_model:
# Save optimizer and scheduler
self._save_optimizer_and_scheduler(output_dir)
# Save RNG state
self._save_rng_state(output_dir)
else:
super()._save_checkpoint(model, trial, metrics)
def _save(self, output_dir: str | None = None, state_dict=None):
# If we are executing this function, we are the process zero, so we don't check for that.
output_dir = output_dir if output_dir is not None else self.args.output_dir
os.makedirs(output_dir, exist_ok=True)
logger.info(f"Saving model checkpoint to {output_dir}")
# pdb.set_trace()
supported_classes = (PreTrainedModel, ) if not is_peft_available() else (PreTrainedModel, PeftModel)
# Save a trained model and configuration using `save_pretrained()`.
# They can then be reloaded using `from_pretrained()`
if not isinstance(self.model, supported_classes):
if state_dict is None:
state_dict = self.model.state_dict()
if isinstance(self.accelerator.unwrap_model(self.model), supported_classes):
self.accelerator.unwrap_model(self.model).save_pretrained(output_dir,
state_dict=state_dict,
safe_serialization=self.args.save_safetensors)
else:
logger.info("Trainer.model is not a `PreTrainedModel`, only saving its state dict.")
if self.args.save_safetensors:
safetensors.torch.save_file(state_dict,
os.path.join(output_dir, SAFE_WEIGHTS_NAME),
metadata={"format": "pt"})
else:
torch.save(state_dict, os.path.join(output_dir, WEIGHTS_NAME))
else:
if not self.args.save_full_model:
state_dict = {k: v for k, v in state_dict.items() if "wte" not in k}
self.model.save_pretrained(output_dir,
state_dict=state_dict,
safe_serialization=self.args.save_safetensors)
else:
torch.save(state_dict, os.path.join(output_dir, 'model.pth'))
if self.tokenizer is not None:
os.makedirs(os.path.join(output_dir, "tokenizer"), exist_ok=True)
self.tokenizer.save_pretrained(os.path.join(output_dir, "tokenizer"))
# Good practice: save your training arguments together with the trained model
torch.save(self.args, os.path.join(output_dir, TRAINING_ARGS_NAME))
# pdb.set_trace()
def compute_multi_attr_accuracy(eval_pred, metainfo_idxs=None, eval_dims=None, save_path=None) -> dict[str, float]:
predictions, labels = eval_pred
metrics = {}
for idx, eval_dim in enumerate(eval_dims):
pred_curr = predictions[:, :, idx]
label_curr = labels[:, idx]
# pdb.set_trace()
## calculate the average scores of rewards_chosen and rewards_rejected
valid_mask = (label_curr != 0)
rewards_chosen = np.where(label_curr == 1, pred_curr[:, 0], pred_curr[:, 1])
rewards_rejected = np.where(label_curr == -1, pred_curr[:, 0], pred_curr[:, 1])
rewards_chosen_avg = np.sum(rewards_chosen * valid_mask) / np.sum(valid_mask)
rewards_rejected_avg = np.sum(rewards_rejected * valid_mask) / np.sum(valid_mask)
pred_curr = np.argmax(pred_curr, axis=1)
pred_curr = np.where(pred_curr == 0, 1, -1)
accuracy = np.array(pred_curr == label_curr, dtype=float)
accuracy = np.sum(accuracy * valid_mask) / np.sum(valid_mask)
metrics.update({
f"accuracy_{eval_dim}": accuracy,
f"rewards_chosen_avg_{eval_dim}": rewards_chosen_avg,
f"rewards_rejected_avg_{eval_dim}": rewards_rejected_avg,
})
if save_path is not None and metainfo_idxs is not None:
df = pd.DataFrame(metainfo_idxs, columns=["metainfo_idx"])
for idx, eval_dim in enumerate(eval_dims):
rewards_A = predictions[:, 0, idx]
rewards_B = predictions[:, 1, idx]
df[f"reward_A_{eval_dim}"] = rewards_A
df[f"reward_B_{eval_dim}"] = rewards_B
df.to_csv(save_path, index=False)
print(f"===> Inference results saved to {save_path}")
# pdb.set_trace()
return metrics
@@ -0,0 +1,208 @@
import os
import glob
import logging
from dataclasses import dataclass, field
from typing import Literal
import safetensors
import torch
from transformers import TrainingArguments
########## DataClass For Configure ##########
@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")
# 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 maybe_zero_3(param, ignore_status=False, name=None):
from deepspeed import zero
from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus
if hasattr(param, "ds_id"):
if param.ds_status == ZeroParamStatus.NOT_AVAILABLE and not ignore_status:
logging.warning("%s: param.ds_status != ZeroParamStatus.NOT_AVAILABLE: %s", name, param.ds_status)
with zero.GatheredParameters([param]):
param = param.data.detach().cpu().clone()
else:
param = param.detach().cpu().clone()
return param
# Borrowed from peft.utils.get_peft_model_state_dict
def get_peft_state_maybe_zero_3(named_params, bias):
if bias == "none":
to_return = {k: t for k, t in named_params if "lora_" in k}
elif bias == "all":
to_return = {k: t for k, t in named_params if "lora_" in k or "bias" in k}
elif bias == "lora_only":
to_return = {}
maybe_lora_bias = {}
lora_bias_names = set()
for k, t in named_params:
if "lora_" in k:
to_return[k] = t
bias_name = k.split("lora_")[0] + "bias"
lora_bias_names.add(bias_name)
elif "bias" in k:
maybe_lora_bias[k] = t
for k, t in maybe_lora_bias:
if bias_name in lora_bias_names:
to_return[bias_name] = t
else:
raise NotImplementedError
to_return = {k: maybe_zero_3(v, ignore_status=True) for k, v in to_return.items()}
return to_return
def get_peft_state_non_lora_maybe_zero_3(named_params, require_grad_only=True):
to_return = {k: t for k, t in named_params if "lora_" not in k}
if require_grad_only:
to_return = {k: t for k, t in to_return.items() if t.requires_grad}
to_return = {k: maybe_zero_3(v, ignore_status=True).cpu() for k, v in to_return.items()}
return to_return
########## Load Models From Folder ##########
def _insert_adapter_name_into_state_dict(state_dict: dict[str, torch.Tensor], adapter_name: str,
parameter_prefix: str) -> dict[str, torch.Tensor]:
"""Utility function to remap the state_dict keys to fit the PEFT model by inserting the adapter name."""
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 save_video(tensor, path):
from torchvision.io import write_video
tensor = tensor * 255.0
tensor = tensor.permute(0, 2, 3, 1)
tensor = tensor.clamp(0, 255).byte()
write_video(path, tensor, 4, video_codec='h264')
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:
# get the latest checkpoint
checkpoint_path = checkpoint_paths[0]
print(f"===> Checkpoint step is not provided, using the latest checkpoint: {checkpoint_path}")
else:
checkpoint_path = os.path.join(checkpoint_dir, f"checkpoint-{checkpoint_step}")
if checkpoint_path not in checkpoint_paths:
checkpoint_path = checkpoint_paths[0]
print(f"===> Checkpoint step {checkpoint_step} not found, using the latest checkpoint: {checkpoint_path}")
else:
print(f"===> Checkpoint step {checkpoint_step} found, using the specified checkpoint: {checkpoint_path}")
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,367 @@
## 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 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:
image = ele["image"] if "image" in ele else 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: int | 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)
if nframes > total_frames:
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.",
stacklevel=2,
)
if "file://" in video_path:
video_path = video_path[7:]
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"]
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
else:
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,21 @@
# SPDX-License-Identifier: Apache-2.0
"""Reward functions for RL-based video generation training."""
from fastvideo.train.methods.rl.reward.hpsv3 import (
hpsv3_general_score,
hpsv3_percentile_score,
)
from fastvideo.train.methods.rl.reward.ocr import (
video_ocr_score, )
from fastvideo.train.methods.rl.reward.videoalign import (
videoalign_mq_score,
videoalign_ta_score,
)
__all__ = [
"hpsv3_general_score",
"hpsv3_percentile_score",
"video_ocr_score",
"videoalign_mq_score",
"videoalign_ta_score",
]
+278
View File
@@ -0,0 +1,278 @@
# SPDX-License-Identifier: Apache-2.0
"""HPSv3 reward functions for visual quality assessment."""
from __future__ import annotations
import os
import tempfile
from typing import Any
import numpy as np
import torch
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.reward.utils import (
prepare_images, )
logger = init_logger(__name__)
# Global cache of HPSv3 inferencers keyed by device.
_HPSV3_INFERENCERS: dict[str, Any] = {}
_HPSV3_LOAD_PATCHED = False
def _patch_transformers_video_input_alias() -> None:
"""Keep HPSv3 compatible with newer transformers releases.
HPSv3 imports ``VideoInput`` from ``transformers.image_utils`` for type
annotations. Some transformers versions used by FastVideo no longer
export that alias, even though the runtime image utilities HPSv3 needs are
still present.
"""
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]:
"""Adapt HPSv3 checkpoints saved with older Qwen2-VL key names."""
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):
"""Yield common wrapper/base model objects without importing PEFT."""
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:
"""Patch a model class to accept old Qwen2-VL checkpoint keys."""
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:
"""Patch HPSv3 reward model loading for transformers key drift."""
global _HPSV3_LOAD_PATCHED
if _HPSV3_LOAD_PATCHED:
return
from fastvideo.train.methods.rl.reward.HPSv3.hpsv3.model.qwen2vl_trainer 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:
"""Add aliases expected by HPSv3's older Qwen2-VL forward."""
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.__dict__["embed_tokens"] = language_model.embed_tokens
def _normalize_device(device: torch.device | str) -> str:
if isinstance(device, torch.device):
return str(device)
return str(torch.device(device))
def _move_hpsv3_inferencer(inferencer: Any, device: torch.device | str) -> None:
"""Move an HPSv3 inferencer across devices.
HPSv3RewardInferencer does not expose ``.to()``, but it stores its torch
module on ``.model`` and reads ``.device`` when preparing batches.
"""
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:
"""Move cached HPSv3 inferencer to given device."""
key = _normalize_device(device)
if key in _HPSV3_INFERENCERS:
return
# Move from any existing device.
for old_key, inf in list(_HPSV3_INFERENCERS.items()):
if old_key != key:
_move_hpsv3_inferencer(inf, device)
_HPSV3_INFERENCERS[key] = inf
del _HPSV3_INFERENCERS[old_key]
return
def _get_hpsv3_inferencer(device: torch.device | str) -> Any:
"""Get or create HPSv3 inferencer for device."""
key = _normalize_device(device)
if key not in _HPSV3_INFERENCERS:
try:
_patch_transformers_video_input_alias()
from fastvideo.train.methods.rl.reward.HPSv3.hpsv3 import HPSv3RewardInferencer
_patch_hpsv3_state_dict_loader()
except ImportError as exc:
msg = ("Failed to import HPSv3. Ensure the HPSv3 submodule is "
"checked out under fastvideo/train/methods/rl/reward/HPSv3 "
"and that its transformers dependencies are compatible.")
raise ImportError(msg) from exc
inf = HPSv3RewardInferencer(device=device)
_patch_hpsv3_runtime_model(inf.model)
_HPSV3_INFERENCERS[key] = inf
return _HPSV3_INFERENCERS[key]
def _save_frame_to_temp(frame: np.ndarray) -> str:
"""Save a frame to a temporary PNG file."""
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) -> float:
"""Extract a float from HPSv3 result."""
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)
def hpsv3_general_score(device):
"""Return a reward fn that scores frames with
'A high-quality image' as prompt.
Returns mean score across all frames."""
def _score(images, prompts, metadata, only_strict=False):
inf = _get_hpsv3_inferencer(device)
images_np = prepare_images(images)
batch_scores = []
for b in range(len(images_np)):
frames = images_np[b]
if frames.ndim == 3:
frames = frames[np.newaxis]
frame_scores = []
for frame in frames:
path = _save_frame_to_temp(frame)
try:
rewards = inf.reward(["A high-quality image"], [path])
frame_scores.append(_extract_reward_scalar(rewards[0][0]))
finally:
os.remove(path)
batch_scores.append(np.mean(frame_scores))
reward = torch.tensor(batch_scores, device=device).float()
return {"avg": reward}, {}
return _score
def hpsv3_percentile_score(device):
"""Return a reward fn that scores frames with per-prompt
text. Returns mean of top 30% frame scores."""
def _score(images, prompts, metadata, only_strict=False):
inf = _get_hpsv3_inferencer(device)
images_np = prepare_images(images)
batch_scores = []
for b in range(len(images_np)):
frames = images_np[b]
if frames.ndim == 3:
frames = frames[np.newaxis]
prompt = (prompts[b] if prompts and b < len(prompts) else "A high-quality image")
frame_scores = []
for frame in frames:
path = _save_frame_to_temp(frame)
try:
rewards = inf.reward([prompt], [path])
frame_scores.append(_extract_reward_scalar(rewards[0][0]))
finally:
os.remove(path)
# Top 30% percentile.
if frame_scores:
k = max(1, int(len(frame_scores) * 0.3))
top_k = sorted(frame_scores, reverse=True)[:k]
batch_scores.append(np.mean(top_k))
else:
batch_scores.append(0.0)
reward = torch.tensor(batch_scores, device=device).float()
return {"avg": reward}, {}
return _score
+95
View File
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
"""OCR-based reward for video-text alignment."""
from __future__ import annotations
import re
import numpy as np
import torch
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.reward.utils import (
prepare_images, )
logger = init_logger(__name__)
def _levenshtein_distance(s1: str, s2: str) -> int:
"""Compute Levenshtein edit distance between strings."""
if len(s1) < len(s2):
return _levenshtein_distance(s2, s1)
if len(s2) == 0:
return len(s1)
prev_row = list(range(len(s2) + 1))
for i, c1 in enumerate(s1):
curr_row = [i + 1]
for j, c2 in enumerate(s2):
insertions = prev_row[j + 1] + 1
deletions = curr_row[j] + 1
substitutions = prev_row[j] + (c1 != c2)
curr_row.append(min(insertions, deletions, substitutions))
prev_row = curr_row
return prev_row[-1]
def _extract_text_from_prompt(prompt: str) -> str:
"""Extract expected text from prompt (within quotes)."""
match = re.search(r'["\'](.+?)["\']', prompt)
if match:
return match.group(1)
return prompt
def video_ocr_score():
"""Return an OCR-based reward function (CPU)."""
try:
from paddleocr import PaddleOCR
except ImportError as exc:
msg = ("paddleocr not installed. "
"Install via: pip install paddleocr")
raise ImportError(msg) from exc
ocr = PaddleOCR(use_angle_cls=True, lang="en", use_gpu=False)
def _score(images, prompts, metadata, only_strict=False):
images_np = prepare_images(images)
batch_scores = []
for b in range(len(images_np)):
frames = images_np[b]
if frames.ndim == 3:
frames = frames[np.newaxis]
expected = _extract_text_from_prompt(prompts[b] if prompts else "").lower()
# Sample every 4th frame.
sample_indices = list(range(0, len(frames), 4))
if not sample_indices:
sample_indices = [0]
best_score = 0.0
for idx in sample_indices:
frame = frames[idx]
result = ocr.ocr(frame, cls=True)
detected = ""
if result and result[0]:
texts = [line[1][0] for line in result[0] if line[1]]
detected = " ".join(texts).lower()
if not expected:
score = 1.0 if detected else 0.0
elif not detected:
score = 0.0
else:
dist = _levenshtein_distance(detected, expected)
max_len = max(len(detected), len(expected))
score = 1.0 - (dist / max_len)
best_score = max(best_score, score)
batch_scores.append(best_score)
reward = torch.tensor(batch_scores).float()
return {"avg": reward}, {}
return _score
@@ -0,0 +1,47 @@
# SPDX-License-Identifier: Apache-2.0
"""Utility functions for reward computation."""
from __future__ import annotations
import numpy as np
import torch
def prepare_images(images: torch.Tensor | np.ndarray, ) -> np.ndarray:
"""Convert tensor images to uint8 numpy (NHWC or NFHWC).
Accepts:
- (N, C, H, W) or (N, H, W, C) tensors/arrays
- (N, F, C, H, W) or (N, F, H, W, C) video tensors
Returns:
uint8 numpy array in HWC/FHWC layout.
"""
if isinstance(images, torch.Tensor):
images = images.detach().cpu().numpy()
images = np.asarray(images)
if images.ndim == 4:
# Image batch: (N, C, H, W) or (N, H, W, C)
if images.shape[-1] in (1, 3):
pass
elif images.shape[1] in (1, 3):
images = images.transpose(0, 2, 3, 1)
elif images.ndim == 5:
# Video batch: (N, F, H, W, C), (N, F, C, H, W),
# or (N, C, F, H, W). Check channel-last first because
# one-frame videos have shape[1] == 1.
if images.shape[-1] in (1, 3):
pass
elif images.shape[1] == 3 and images.shape[2] == 1:
# (N, C=3, F=1, H, W) -> (N, F, H, W, C)
images = images.transpose(0, 2, 3, 4, 1)
elif images.shape[2] in (1, 3):
# (N, F, C, H, W) -> (N, F, H, W, C)
images = images.transpose(0, 1, 3, 4, 2)
elif images.shape[1] in (1, 3):
# (N, C, F, H, W) -> (N, F, H, W, C)
images = images.transpose(0, 2, 3, 4, 1)
if images.dtype == np.float32 or images.dtype == np.float64:
images = np.clip(images * 255, 0, 255).astype(np.uint8)
return images
@@ -0,0 +1,457 @@
# SPDX-License-Identifier: Apache-2.0
"""VideoAlign reward functions for motion quality and
text-video alignment."""
from __future__ import annotations
import os
import tempfile
from importlib import import_module, metadata
from typing import Any
from huggingface_hub import snapshot_download
import numpy as np
import torch
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.reward.utils import (
prepare_images, )
logger = init_logger(__name__)
# Global cache of VideoAlign inferencers.
_VIDEOALIGN_INFERENCERS: dict[str, Any] = {}
_VIDEOALIGN_PATCHED = False
def _normalize_device_str(device) -> str:
if isinstance(device, torch.device):
return str(device)
return str(torch.device(device))
def _move_videoalign_inferencer(inferencer: Any, device) -> None:
"""Move a VideoAlign inferencer across devices."""
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]:
"""Adapt checkpoints saved with older Qwen2-VL key names."""
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):
"""Yield common wrapper/base model objects without importing PEFT."""
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:
"""Patch a model class to accept old VideoAlign checkpoint keys."""
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)
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:
"""Read local MP4s without relying on torchvision.io.read_video."""
import cv2
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 _flash_attn2_available() -> bool:
"""Return True when Transformers' classic FlashAttention-2 path is usable."""
try:
metadata.version("flash_attn")
flash_attn_mod = import_module("flash_attn")
import_module("flash_attn.bert_padding")
except (AttributeError, ImportError, metadata.PackageNotFoundError):
available = False
else:
available = hasattr(flash_attn_mod, "flash_attn_func")
if available:
return True
logger.warning("Classic FlashAttention-2 is unavailable; using SDPA for the "
"VideoAlign reward model. Install flash-attn to enable "
"FlashAttention-2.")
return False
def _patch_videoalign_video_reader() -> None:
"""Register an OpenCV reader for torchvision builds without read_video."""
from fastvideo.train.methods.rl.reward.VideoAlign import vision_process as vision_mod
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.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:
"""Patch VideoAlign for the FastVideo dependency set."""
global _VIDEOALIGN_PATCHED
from fastvideo.train.methods.rl.reward.VideoAlign import inference as inference_mod
if _VIDEOALIGN_PATCHED:
return inference_mod
from fastvideo.train.methods.rl.reward.VideoAlign import train_reward as train_reward_mod
from fastvideo.train.methods.rl.reward.VideoAlign import trainer as trainer_mod
_patch_videoalign_video_reader()
# Transformers' ``flash_attention_2`` path expects classic flash-attn.
# FastVideo may have FA4/CuTe installed, which exposes a ``flash_attn``
# namespace but not the FA2 metadata/API that Transformers checks.
if not _flash_attn2_available():
for mod in (train_reward_mod, inference_mod):
original_create = mod.create_model_and_processor
def create_model_and_processor_sdpa(
*args,
_original_create=original_create,
**kwargs,
):
training_args = kwargs.get("training_args")
if training_args is not None:
training_args.disable_flash_attn2 = True
return _original_create(*args, **kwargs)
mod.create_model_and_processor = create_model_and_processor_sdpa
_patch_load_state_dict(trainer_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:
"""Add aliases expected by VideoAlign's older Qwen2-VL forward."""
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.__dict__["embed_tokens"] = language_model.embed_tokens
def set_videoalign_device(device) -> None:
"""Move cached VideoAlign inferencers to device."""
key = _normalize_device_str(device)
for old_key, inf in list(_VIDEOALIGN_INFERENCERS.items()):
if old_key != key and old_key.split(":")[0] != key:
new_key = inf._key_prefix + ":" + key
_move_videoalign_inferencer(inf, device)
_VIDEOALIGN_INFERENCERS[new_key] = inf
del _VIDEOALIGN_INFERENCERS[old_key]
def _resolve_videoalign_checkpoint_path(checkpoint_path: str | None) -> str:
"""Return a local VideoAlign checkpoint directory."""
if checkpoint_path is not None:
return os.path.abspath(checkpoint_path)
return snapshot_download(
repo_id="KlingTeam/VideoReward",
repo_type="model",
allow_patterns=(
"model_config.json",
"checkpoint-*/model.pth",
),
)
def _get_inferencer(
device,
checkpoint_path: str | None = None,
):
"""Get or create VideoAlign inferencer."""
checkpoint_path = _resolve_videoalign_checkpoint_path(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()
VideoVLMRewardInference = inference_mod.VideoVLMRewardInference
except ImportError as exc:
msg = ("VideoAlign not found. Ensure the "
"VideoAlign submodule is checked out "
"under fastvideo/train/methods/rl/"
"reward/VideoAlign")
raise ImportError(msg) from exc
inf = VideoVLMRewardInference(
load_from_pretrained=checkpoint_path,
device=device,
)
_patch_videoalign_runtime_model(inf.model)
inf._key_prefix = checkpoint_path or "default"
_VIDEOALIGN_INFERENCERS[cache_key] = inf
return _VIDEOALIGN_INFERENCERS[cache_key]
def _convert_to_grayscale(frames: np.ndarray, ) -> np.ndarray:
"""Convert FHWC frames to grayscale FHWC."""
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:
"""Save frames to a temporary MP4 file."""
import cv2
fd, path = tempfile.mkstemp(suffix=".mp4")
os.close(fd)
h, w = frames.shape[1], frames.shape[2]
fourcc = cv2.VideoWriter.fourcc(*"mp4v")
writer = cv2.VideoWriter(path, fourcc, fps, (w, h))
for frame in frames:
bgr = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
writer.write(bgr)
writer.release()
return path
def videoalign_mq_score(
device,
checkpoint_path: str | None = None,
):
"""Return Motion Quality reward fn (grayscale)."""
def _score(images, prompts, metadata, only_strict=False):
inf = _get_inferencer(device, checkpoint_path)
images_np = prepare_images(images)
batch_scores = []
for b in range(len(images_np)):
frames = images_np[b]
if frames.ndim == 3:
frames = frames[np.newaxis]
gray_frames = _convert_to_grayscale(frames)
path = _save_video_to_temp(gray_frames)
try:
results = inf.reward([path], [""], use_norm=True)
mq = float(results[0].get("MQ", 0))
batch_scores.append(mq)
finally:
os.remove(path)
reward = torch.tensor(batch_scores, device=device).float()
return {"avg": reward}, {}
return _score
def videoalign_ta_score(
device,
checkpoint_path: str | None = None,
):
"""Return Text-Video Alignment reward fn (color)."""
def _score(images, prompts, metadata, only_strict=False):
inf = _get_inferencer(device, checkpoint_path)
images_np = prepare_images(images)
batch_scores = []
for b in range(len(images_np)):
frames = images_np[b]
if frames.ndim == 3:
frames = frames[np.newaxis]
prompt = (prompts[b] if prompts and b < len(prompts) else "")
path = _save_video_to_temp(frames)
try:
results = inf.reward([path], [prompt], use_norm=True)
ta = float(results[0].get("TA", 0))
batch_scores.append(ta)
finally:
os.remove(path)
reward = torch.tensor(batch_scores, device=device).float()
return {"avg": reward}, {}
return _score
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
"""Utility modules for RL training."""
@@ -0,0 +1,170 @@
# SPDX-License-Identifier: Apache-2.0
"""Advantage computation for RL training."""
from __future__ import annotations
from typing import Any
import numpy as np
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.utils.stat_tracking import (
EPSILON,
PerPromptStatTracker,
)
logger = init_logger(__name__)
def _normalize_rewards(
rewards: np.ndarray,
epsilon: float = EPSILON,
) -> np.ndarray:
"""Normalize rewards to zero mean and unit variance."""
return (rewards - rewards.mean()) / (rewards.std() + epsilon)
def _compute_kl_advantages(
gathered_kl: np.ndarray,
kl_stat_tracker: PerPromptStatTracker | None,
prompts: list[str] | None,
use_per_prompt: bool,
) -> np.ndarray:
"""Compute KL advantages (negative = penalty)."""
if use_per_prompt and kl_stat_tracker is not None:
return kl_stat_tracker.update(prompts, -gathered_kl)
return _normalize_rewards(-gathered_kl)
def calculate_zero_std_ratio(
prompts,
gathered_rewards: dict[str, np.ndarray],
reward_key: str = "avg",
) -> float:
"""Compute fraction of prompts with zero std."""
prompts_arr = np.array(prompts)
rewards = gathered_rewards.get(reward_key)
if rewards is None:
return 0.0
unique = np.unique(prompts_arr)
zero_count = 0
for p in unique:
r = rewards[prompts_arr == p]
if np.std(r) < EPSILON:
zero_count += 1
return zero_count / max(len(unique), 1)
def compute_advantages(
reward_fn_cfg: dict[str, float],
weight_advantages: bool,
per_prompt_stat_tracking: bool,
kl_reward: float,
samples: dict[str, Any],
gathered_rewards: dict[str, np.ndarray],
gathered_kl: np.ndarray,
prompts: list[str] | None,
stat_tracker: PerPromptStatTracker | None,
reward_stat_trackers: (dict[str, PerPromptStatTracker] | None),
kl_stat_tracker: PerPromptStatTracker | None,
) -> tuple[np.ndarray, dict[str, Any]]:
"""Compute advantages from gathered rewards and KL.
Supports two modes:
- Mode 1 (default): Weight rewards, then advantages.
- Mode 2 (weight_advantages=True): Per-reward
advantages, then weight.
Returns:
(advantages, log_dict)
"""
log_dict: dict[str, Any] = {}
if weight_advantages:
if per_prompt_stat_tracking:
if reward_stat_trackers is None:
msg = ("reward_stat_trackers required when "
"weight_advantages=True and "
"per_prompt_stat_tracking=True")
raise ValueError(msg)
weighted_list = []
for reward_name in reward_fn_cfg:
raw_key = f"{reward_name}_raw"
adv = reward_stat_trackers[reward_name].update(prompts, gathered_rewards[raw_key])
weight = reward_fn_cfg[reward_name]
weighted_list.append(adv * weight)
if kl_reward > 0:
if kl_stat_tracker is None:
msg = ("kl_stat_tracker required when "
"weight_advantages=True and "
"kl_reward > 0")
raise ValueError(msg)
kl_adv = _compute_kl_advantages(
gathered_kl,
kl_stat_tracker,
prompts,
use_per_prompt=True,
)
weighted_list.append(kl_adv * kl_reward)
advantages = sum(weighted_list)
first_name = next(iter(reward_fn_cfg))
group_size, trained_num = (reward_stat_trackers[first_name].get_stats())
zero_std_ratios = {}
for rn in reward_fn_cfg:
raw_key = f"{rn}_raw"
zero_std_ratios[f"zero_std_ratio_{rn}"] = calculate_zero_std_ratio(
prompts,
gathered_rewards,
reward_key=f"ori_{raw_key}",
)
log_dict = {
"group_size": group_size,
"trained_prompt_num": trained_num,
**zero_std_ratios,
}
for t in reward_stat_trackers.values():
t.clear()
if kl_stat_tracker is not None:
kl_stat_tracker.clear()
else:
weighted_list = []
for reward_name in reward_fn_cfg:
raw_key = f"{reward_name}_raw"
raw = gathered_rewards[raw_key]
adv = _normalize_rewards(raw)
weight = reward_fn_cfg[reward_name]
weighted_list.append(adv * weight)
if kl_reward > 0:
kl_adv = _compute_kl_advantages(
gathered_kl,
None,
None,
use_per_prompt=False,
)
weighted_list.append(kl_adv * kl_reward)
advantages = sum(weighted_list)
elif per_prompt_stat_tracking:
if stat_tracker is None:
msg = ("stat_tracker required when "
"per_prompt_stat_tracking=True")
raise ValueError(msg)
advantages = stat_tracker.update(prompts, gathered_rewards["avg"])
group_size, trained_num = (stat_tracker.get_stats())
zero_std = calculate_zero_std_ratio(prompts, gathered_rewards)
log_dict = {
"group_size": group_size,
"trained_prompt_num": trained_num,
"zero_std_ratio": zero_std,
}
stat_tracker.clear()
else:
advantages = _normalize_rewards(gathered_rewards["avg"])
return advantages, log_dict
+198
View File
@@ -0,0 +1,198 @@
# SPDX-License-Identifier: Apache-2.0
"""Text prompt datasets and samplers for RL training."""
from __future__ import annotations
import json
import os
import torch
from torch.utils.data import DataLoader, Dataset, Sampler
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class TextPromptDataset(Dataset):
"""Load plain text prompts from train.txt / test.txt."""
def __init__(self, dataset: str, split: str = "train"):
self.file_path = os.path.join(dataset, f"{split}.txt")
with open(self.file_path) as f:
self.prompts = [line.strip() for line in f.readlines()]
def __len__(self) -> int:
return len(self.prompts)
def __getitem__(self, idx: int | tuple[int, int]) -> dict:
epoch_tag = None
if isinstance(idx, tuple):
epoch_tag, idx = idx
return {
"epoch": epoch_tag,
"prompt": self.prompts[idx],
"metadata": {},
}
@staticmethod
def collate_fn(examples: list[dict], ) -> tuple[int | None, list[str], list[dict]]:
epoch_tags = [example.get("epoch") for example in examples]
epoch_tag = (epoch_tags[0] if all(tag == epoch_tags[0] for tag in epoch_tags) else None)
prompts = [example["prompt"] for example in examples]
metadatas = [example["metadata"] for example in examples]
return epoch_tag, prompts, metadatas
class JsonPromptDataset(Dataset):
"""Load prompts from JSONL files."""
def __init__(self, dataset: str, split: str = "train"):
self.file_path = os.path.join(dataset, f"{split}.json")
self._prompts: list[str] = []
self._metadatas: list[dict] = []
self._load_all_prompts()
def _load_all_prompts(self) -> None:
with open(self.file_path, encoding="utf-8") as f:
for raw_line in f:
line = raw_line.strip()
if not line:
continue
try:
item = json.loads(line)
prompt = item.get("prompt", "")
if prompt:
self._prompts.append(prompt)
metadata = {k: v for k, v in item.items() if k != "prompt"}
self._metadatas.append(metadata)
except json.JSONDecodeError as e:
logger.warning(
"Skipping invalid JSON line: %s",
e,
)
def __len__(self) -> int:
return len(self._prompts)
def __getitem__(self, idx: int | tuple[int, int]) -> dict:
epoch_tag = None
if isinstance(idx, tuple):
epoch_tag, idx = idx
return {
"epoch": epoch_tag,
"prompt": self._prompts[idx],
"metadata": (self._metadatas[idx] if self._metadatas else {}),
}
@staticmethod
def collate_fn(examples: list[dict], ) -> tuple[int | None, list[str], list[dict]]:
epoch_tags = [example.get("epoch") for example in examples]
epoch_tag = (epoch_tags[0] if all(tag == epoch_tags[0] for tag in epoch_tags) else None)
prompts = [example["prompt"] for example in examples]
metadatas = [example["metadata"] for example in examples]
return epoch_tag, prompts, metadatas
class DistributedKRepeatSampler(Sampler):
def __init__(
self,
dataset: Dataset,
batch_size: int,
k: int,
num_replicas: int,
rank: int,
seed: int = 0,
):
self.dataset = dataset
self.batch_size = batch_size
self.k = k # Repeats/videos per prompt.
self.num_replicas = num_replicas
self.rank = rank
self.seed = seed
self.total_samples = num_replicas * batch_size
if self.k <= 0:
raise ValueError(f"k must be a positive integer. Got k={k}.")
if self.batch_size % self.k != 0:
raise ValueError("batch_size must be divisible by k so each rank receives "
"whole prompt groups. Got "
f"batch_size={batch_size}, k={k}.")
assert self.total_samples % self.k == 0, (f"k cannot divide n*b: k={k} "
f"num_replicas={num_replicas} "
f"batch_size={batch_size}")
self.m = self.total_samples // self.k # Unique prompts across ranks.
if len(self.dataset) < self.m:
raise ValueError("dataset must contain at least one prompt per global "
"prompt group. Got "
f"dataset_size={len(self.dataset)}, required={self.m}.")
self.groups_per_rank = self.batch_size // self.k
self.epoch = 0
def __iter__(self):
while True:
g = torch.Generator()
g.manual_seed(self.seed + self.epoch)
indices = torch.randperm(len(self.dataset), generator=g)[:self.m].tolist()
start = self.rank * self.groups_per_rank
end = start + self.groups_per_rank
rank_groups = indices[start:end]
yield [(self.epoch, idx) for idx in rank_groups for _ in range(self.k)]
def set_epoch(self, epoch: int):
self.epoch = epoch
def build_prompt_dataloaders(
prompt_dataset_path: str,
prompt_fn: str,
sample_batch_size: int,
eval_batch_size: int,
num_video_per_prompt: int,
num_processes: int,
process_index: int,
seed: int,
) -> tuple[DataLoader, DataLoader, DistributedKRepeatSampler]:
"""Build train/eval prompt dataloaders.
Returns:
(train_dataloader, test_dataloader, train_sampler)
"""
if prompt_fn == "general_ocr":
train_ds = TextPromptDataset(prompt_dataset_path, "train")
test_ds = TextPromptDataset(prompt_dataset_path, "test")
collate = TextPromptDataset.collate_fn
elif prompt_fn == "filtered_prompts":
train_ds = JsonPromptDataset(prompt_dataset_path, "train")
test_ds = JsonPromptDataset(prompt_dataset_path, "test")
collate = JsonPromptDataset.collate_fn
else:
msg = (f"Unsupported prompt_fn: {prompt_fn}. "
"Use 'general_ocr' or 'filtered_prompts'.")
raise NotImplementedError(msg)
train_sampler = DistributedKRepeatSampler(
dataset=train_ds,
batch_size=sample_batch_size,
k=num_video_per_prompt,
num_replicas=num_processes,
rank=process_index,
seed=seed,
)
train_dl = DataLoader(
train_ds,
batch_sampler=train_sampler,
num_workers=1,
collate_fn=collate,
prefetch_factor=1,
persistent_workers=False,
)
test_dl = DataLoader(
test_ds,
batch_size=eval_batch_size,
collate_fn=collate,
shuffle=False,
num_workers=8,
)
return train_dl, test_dl, train_sampler
@@ -0,0 +1,108 @@
# SPDX-License-Identifier: Apache-2.0
"""Single diffusion step for PPO training phase."""
from __future__ import annotations
import torch
from fastvideo.train.methods.rl.utils.sde import (
sde_step_with_logprob, )
def compute_log_prob(
model,
scheduler,
sample: dict[str, torch.Tensor],
j: int,
embeds: torch.Tensor,
negative_embeds: torch.Tensor | None,
guidance_scale: float,
use_cfg: bool,
noise_level: float,
sde_type: str,
diffusion_clip: bool = False,
diffusion_clip_value: float = 0.45,
) -> tuple[
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
float,
]:
"""Run one diffusion step and return log-probability.
Uses model.forward_transformer_raw() for the forward
pass.
Args:
model: WanModel instance.
scheduler: Noise scheduler.
sample: Dict with latents, next_latents, timesteps.
j: Timestep index within the trajectory.
embeds: Conditional text embeddings.
negative_embeds: Unconditional embeddings (or None).
guidance_scale: CFG scale.
use_cfg: Whether to use classifier-free guidance.
noise_level: SDE noise level.
sde_type: 'flow_sde' or 'flow_cps'.
diffusion_clip: Clip SDE variance.
diffusion_clip_value: Clip threshold.
Returns:
(prev_sample, log_prob, prev_sample_mean,
std_dev_t, dt_sqrt, sigma, sigma_max)
"""
dtype = embeds.dtype
latents_j = sample["latents"][:, j]
timestep_j = sample["timesteps"][:, j]
if use_cfg and negative_embeds is not None:
noise_pred_text = model.forward_transformer_raw(
latents_j.to(dtype),
timestep_j,
embeds,
)
noise_pred_uncond = model.forward_transformer_raw(
latents_j.to(dtype),
timestep_j,
negative_embeds,
)
noise_pred = (noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond))
else:
noise_pred = model.forward_transformer_raw(
latents_j.to(dtype),
timestep_j,
embeds,
)
(
prev_sample,
log_prob,
prev_sample_mean,
std_dev_t,
dt_sqrt,
sigma,
sigma_max,
) = sde_step_with_logprob(
scheduler,
noise_pred.float(),
timestep_j,
latents_j.float(),
noise_level=noise_level,
prev_sample=sample["next_latents"][:, j].float(),
sde_type=sde_type,
diffusion_clip=diffusion_clip,
diffusion_clip_value=diffusion_clip_value,
return_sqrt_dt_and_std_dev_t=True,
)
return (
prev_sample,
log_prob,
prev_sample_mean,
std_dev_t,
dt_sqrt,
sigma,
sigma_max,
)
@@ -0,0 +1,43 @@
# SPDX-License-Identifier: Apache-2.0
"""Text embedding utilities for RL training."""
from __future__ import annotations
import torch
def compute_text_embeddings(
prompts: list[str],
text_encoder,
tokenizer,
max_sequence_length: int = 512,
device: torch.device | str = "cuda",
) -> torch.Tensor:
"""Encode text prompts into embeddings using T5.
Args:
prompts: List of text prompts.
text_encoder: T5 text encoder model.
tokenizer: T5 tokenizer.
max_sequence_length: Max token length.
device: Target device.
Returns:
Tensor of shape (B, L, D).
"""
text_inputs = tokenizer(
prompts,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids.to(device)
attention_mask = text_inputs.attention_mask.to(device)
with torch.no_grad():
prompt_embeds = text_encoder(text_input_ids, attention_mask=attention_mask).last_hidden_state
# make padding token 0
prompt_embeds[attention_mask == 0] = 0
return prompt_embeds
@@ -0,0 +1,124 @@
# SPDX-License-Identifier: Apache-2.0
"""Evaluation loop for RL training."""
from __future__ import annotations
from collections.abc import Callable
from typing import Any
import torch
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.utils.embeddings import (
compute_text_embeddings, )
from fastvideo.train.methods.rl.utils.pipeline import (
wan_denoising_with_logprob, )
logger = init_logger(__name__)
def eval_once(
model,
scheduler,
test_dataloader,
text_encoder,
tokenizer,
sample_neg_prompt_embeds: torch.Tensor,
eval_reward_fn: Callable,
global_step: int,
ema_callback,
*,
eval_num_steps: int,
eval_guidance_scale: float,
height: int,
width: int,
num_frames: int,
device: torch.device,
world_size: int,
rank: int,
is_main_process: bool,
tracker: Any | None = None,
max_batches: int | None = None,
seed: int = 0,
) -> dict[str, float]:
"""Run evaluation on test set.
Args:
ema_callback: An ``EMACallback`` instance (or
``None``). Used to temporarily swap EMA
weights into the transformer for evaluation.
Returns:
Dict of aggregated eval metrics.
"""
model.transformer.eval()
all_rewards: dict[str, list[float]] = {}
# Use EMA context manager if available.
if ema_callback is not None:
ctx = ema_callback.ema_context(model.transformer)
else:
from contextlib import nullcontext
ctx = nullcontext()
with ctx:
generator = torch.Generator(device=device)
for batch_idx, (
_epoch_tag,
prompts,
metadata,
) in enumerate(test_dataloader):
if max_batches is not None and batch_idx >= max_batches:
break
prompt_embeds = compute_text_embeddings(
prompts,
text_encoder,
tokenizer,
max_sequence_length=512,
device=device,
)
neg_prompt_embeds = sample_neg_prompt_embeds[:len(prompts)]
with torch.no_grad():
generator.manual_seed(seed + batch_idx)
(
videos,
_latents,
_log_probs,
_kls,
_timesteps,
) = wan_denoising_with_logprob(
model,
scheduler,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=(neg_prompt_embeds),
num_inference_steps=eval_num_steps,
guidance_scale=eval_guidance_scale,
height=height,
width=width,
num_frames=num_frames,
generator=generator,
deterministic=True,
sde_type="flow_sde",
)
rewards, _ = eval_reward_fn(videos, prompts, metadata)
for key, val in rewards.items():
if key not in all_rewards:
all_rewards[key] = []
if isinstance(val, torch.Tensor):
all_rewards[key].extend(val.detach().cpu().tolist())
else:
all_rewards[key].append(float(val))
# Aggregate metrics.
metrics = {}
for key, vals in all_rewards.items():
avg = sum(vals) / max(len(vals), 1)
metrics[f"eval_{key}"] = avg
if is_main_process and tracker is not None:
tracker.log(metrics, global_step)
return metrics
@@ -0,0 +1,321 @@
# SPDX-License-Identifier: Apache-2.0
"""Multi-step denoising with log-probability tracking.
Replaces diffusers' WanPipeline.__call__ for RL training.
Uses WanModel's forward_transformer_raw() instead of
the diffusers pipeline.
"""
from __future__ import annotations
import contextlib
import random
import time
from typing import Any
import torch
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.utils.sde import (
sde_step_with_logprob, )
_pipeline_logger = init_logger(__name__)
def wan_denoising_with_logprob(
model,
scheduler,
prompt_embeds: torch.Tensor,
negative_prompt_embeds: torch.Tensor | None = None,
num_inference_steps: int = 50,
guidance_scale: float = 5.0,
height: int = 480,
width: int = 832,
num_frames: int = 81,
generator: torch.Generator | None = None,
noise_level: float = 0.7,
sde_type: str = "flow_sde",
deterministic: bool = False,
diffusion_clip: bool = False,
diffusion_clip_value: float = 0.45,
sde_window_size: int = 0,
sde_window_range: tuple[int, int] | None = None,
kl_reward: float = 0.0,
ref_transformer: torch.nn.Module | None = None,
lora_model: Any | None = None,
) -> tuple[
torch.Tensor,
list[torch.Tensor],
list[torch.Tensor],
list[torch.Tensor],
list[torch.Tensor],
]:
"""Run full denoising loop, collecting latent
trajectories and log-probabilities at each step.
Args:
model: WanModel (or GenRLWanModel) with
forward_transformer_raw and vae.
scheduler: Noise scheduler (UniPC/Euler).
prompt_embeds: (B, L, D) text embeddings.
negative_prompt_embeds: (B, L, D) or None.
num_inference_steps: Number of denoising steps.
guidance_scale: CFG scale.
height, width, num_frames: Video dimensions.
generator: RNG for reproducibility.
noise_level: SDE noise level.
sde_type: 'flow_sde' or 'flow_cps'.
deterministic: If True, no SDE noise.
diffusion_clip: Clip SDE variance.
diffusion_clip_value: Clip threshold.
sde_window_size: Window size for SDE training.
sde_window_range: (start, end) range for window.
kl_reward: KL penalty weight (>0 enables KL).
ref_transformer: Reference model for KL.
lora_model: LoRA model with disable_adapter().
Returns:
(videos, all_latents, all_log_probs,
all_kl, all_timesteps)
"""
device = model.device
batch_size = prompt_embeds.shape[0]
dtype = prompt_embeds.dtype
do_cfg = (guidance_scale > 1.0 and negative_prompt_embeds is not None)
# Prepare initial noise.
vae_config = model.vae.config
vae_scale_temporal = getattr(vae_config, "temporal_compression_ratio", 4)
vae_scale_spatial = getattr(vae_config, "spatial_compression_ratio", 8)
num_channels = getattr(vae_config, "z_dim", 16)
latent_frames = (num_frames - 1) // vae_scale_temporal + 1
latent_h = height // vae_scale_spatial
latent_w = width // vae_scale_spatial
latent_shape = (
1,
num_channels,
latent_frames,
latent_h,
latent_w,
)
if isinstance(generator, list):
latents = torch.cat([
torch.randn(
*latent_shape,
generator=generator[i],
device=device,
dtype=torch.float32,
) for i in range(batch_size)
])
else:
latents = torch.randn(
batch_size,
num_channels,
latent_frames,
latent_h,
latent_w,
generator=generator,
device=device,
dtype=torch.float32,
)
# Setup scheduler.
scheduler.set_timesteps(num_inference_steps, device=device)
timesteps = scheduler.timesteps
# Window setup.
use_window = (sde_window_size > 0 and sde_window_range is not None)
if use_window:
assert sde_window_range is not None
window_range_start, window_range_end = sde_window_range
if (window_range_end - window_range_start < sde_window_size):
msg = (f"sde_window_range span "
f"({window_range_end - window_range_start}) "
f"must be >= sde_window_size "
f"({sde_window_size})")
raise ValueError(msg)
if generator is not None:
gen = (generator[0] if isinstance(generator, list) else generator)
max_start = (window_range_end - sde_window_size)
start = torch.randint(
window_range_start,
max_start + 1,
(1, ),
generator=gen,
device=device,
).item()
else:
start = random.randint(
window_range_start,
window_range_end - sde_window_size,
)
end = start + sde_window_size
sde_window = (start, end)
else:
sde_window = None
all_latents: list[torch.Tensor] = [] if sde_window is not None else [latents]
all_log_probs: list[torch.Tensor] = []
all_kl: list[torch.Tensor] = []
all_timesteps: list[torch.Tensor] = []
_denoise_fwd_time = 0.0
_denoise_sde_time = 0.0
_denoise_kl_time = 0.0
for i, t in enumerate(timesteps):
latents_ori = latents.clone()
timestep = t.expand(batch_size)
# Conditional prediction.
torch.cuda.synchronize()
_fwd_t0 = time.perf_counter()
noise_pred = model.forward_transformer_raw(
latents.to(dtype),
timestep,
prompt_embeds,
)
noise_pred = noise_pred.to(dtype)
# CFG.
if do_cfg:
noise_uncond = model.forward_transformer_raw(
latents.to(dtype),
timestep,
negative_prompt_embeds,
)
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
torch.cuda.synchronize()
_denoise_fwd_time += time.perf_counter() - _fwd_t0
# Determine noise level for this step.
if sde_window is not None:
window_start, window_end = sde_window
if i < window_start:
cur_noise_level = 0.0
elif i == window_start:
cur_noise_level = noise_level
all_latents.append(latents)
elif window_start < i < window_end:
cur_noise_level = noise_level
else:
cur_noise_level = 0.0
else:
cur_noise_level = noise_level
# SDE step.
_sde_t0 = time.perf_counter()
(
latents,
log_prob,
prev_latents_mean,
std_dev_t,
sigma,
sigma_max,
) = sde_step_with_logprob(
scheduler,
noise_pred.float(),
t.unsqueeze(0),
latents.float(),
noise_level=cur_noise_level,
sde_type=sde_type,
deterministic=deterministic,
diffusion_clip=diffusion_clip,
diffusion_clip_value=diffusion_clip_value,
)
_denoise_sde_time += time.perf_counter() - _sde_t0
prev_latents = latents.clone()
# Record.
in_window = (sde_window is not None and sde_window[0] <= i < sde_window[1])
should_record = (sde_window is None) or in_window
if should_record:
all_latents.append(latents)
all_log_probs.append(log_prob)
all_timesteps.append(t)
# KL computation.
_kl_t0 = time.perf_counter()
if should_record and kl_reward > 0 and not deterministic:
ref_model = ref_transformer
ref_ctx: Any = contextlib.nullcontext()
if ref_model is None and lora_model is not None:
ref_model = lora_model
ref_ctx = lora_model.disable_adapter()
if ref_model is not None:
with ref_ctx:
ref_noise = ref_model(
hidden_states=latents_ori.to(dtype),
timestep=timestep,
encoder_hidden_states=prompt_embeds,
return_dict=False,
)
ref_noise = ref_noise.to(dtype)
if do_cfg:
ref_uncond = ref_model(
hidden_states=latents_ori.to(dtype),
timestep=timestep,
encoder_hidden_states=(negative_prompt_embeds),
return_dict=False,
)
ref_noise = (ref_uncond + guidance_scale * (ref_noise - ref_uncond))
(
_,
_ref_log_prob,
ref_prev_mean,
ref_std,
_ref_sigma,
_ref_sigma_max,
) = sde_step_with_logprob(
scheduler,
ref_noise.float(),
t.unsqueeze(0),
latents_ori.float(),
noise_level=noise_level,
sde_type=sde_type,
prev_sample=prev_latents.float(),
deterministic=deterministic,
diffusion_clip=diffusion_clip,
diffusion_clip_value=diffusion_clip_value,
)
kl = (prev_latents_mean - ref_prev_mean)**2 / (2 * std_dev_t**2)
kl = kl.mean(dim=tuple(range(1, kl.ndim)))
all_kl.append(kl)
else:
all_kl.append(torch.zeros(batch_size, device=device))
elif should_record:
all_kl.append(torch.zeros(batch_size, device=device))
torch.cuda.synchronize()
_denoise_kl_time += time.perf_counter() - _kl_t0
# Decode to video.
torch.cuda.synchronize()
_vae_t0 = time.perf_counter()
videos = model.decode_latents(latents)
torch.cuda.synchronize()
_vae_done = time.perf_counter()
_pipeline_logger.info(
"[denoising] %d steps: "
"transformer_fwd=%.1fs sde_step=%.1fs "
"kl=%.1fs vae_decode=%.1fs",
len(timesteps),
_denoise_fwd_time,
_denoise_sde_time,
_denoise_kl_time,
_vae_done - _vae_t0,
)
return (
videos,
all_latents,
all_log_probs,
all_kl,
all_timesteps,
)
+192
View File
@@ -0,0 +1,192 @@
# SPDX-License-Identifier: Apache-2.0
"""Reward function loading and composition for RL training."""
from __future__ import annotations
import gc
import importlib
import inspect
import time
from collections.abc import Callable
from contextlib import contextmanager
import torch
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.reward import (
hpsv3_general_score,
hpsv3_percentile_score,
video_ocr_score,
videoalign_mq_score,
videoalign_ta_score,
)
from fastvideo.train.methods.rl.reward.hpsv3 import (
_HPSV3_INFERENCERS,
set_hpsv3_device,
)
from fastvideo.train.methods.rl.reward.videoalign import (
_VIDEOALIGN_INFERENCERS,
set_videoalign_device,
)
logger = init_logger(__name__)
_BUILTIN_REWARDS: dict[str, Callable] = {
"video_ocr": video_ocr_score,
"hpsv3_general": hpsv3_general_score,
"hpsv3_percentile": hpsv3_percentile_score,
"videoalign_mq": videoalign_mq_score,
"videoalign_ta": videoalign_ta_score,
}
_GPU_REWARD_NAMES = {
"hpsv3_general",
"hpsv3_percentile",
"videoalign_mq",
"videoalign_ta",
}
def load_reward_fn(
name: str,
device,
module_path: str | None = None,
):
"""Load a reward function by name."""
if module_path:
mod = importlib.import_module(module_path)
fn = getattr(mod, f"{name}_score", None)
if fn is None:
msg = (f"Reward {name}_score not found "
f"in {module_path}")
raise ValueError(msg)
return fn(device) if callable(fn) else fn
if name in _BUILTIN_REWARDS:
fn = _BUILTIN_REWARDS[name]
sig = inspect.signature(fn)
accepts_device = any(p.name in {"device", "dev"} for p in sig.parameters.values())
return fn(device) if accepts_device else fn()
def _zero_fn(images, prompts, metadata, only_strict=False):
batch = (len(prompts) if prompts is not None else 1)
zeros = torch.zeros(batch, device=device)
return {"avg": zeros}, {}
return _zero_fn
def multi_score(
device,
reward_cfg: dict[str, float],
module_path: str | None = None,
return_raw_scores: bool = False,
):
"""Compose multiple reward heads.
Args:
device: Device for reward computation.
reward_cfg: Dict mapping reward name to weight.
module_path: Optional custom module path.
return_raw_scores: If True, include raw scores.
Returns:
A callable (images, prompts, metadata, only_strict)
-> (scores_dict, metadata_dict).
"""
reward_fns = {}
weights = {}
for name, weight in reward_cfg.items():
reward_fns[name] = load_reward_fn(name, device, module_path)
weights[name] = weight
def _fn(images, prompts, metadata, only_strict=True):
scores = {}
for name, fn in reward_fns.items():
out, _meta = fn(images, prompts, metadata)
val = out.get("avg", out.get("reward", out)) if isinstance(out, dict) else out
if return_raw_scores:
scores[f"{name}_raw"] = val
scores[name] = val * weights[name]
stacked = torch.stack([scores[name] for name in reward_cfg], dim=0)
scores["avg"] = stacked.mean(0)
return scores, {}
return _fn
def _has_reward(reward_cfg, names) -> bool:
if not reward_cfg:
return False
return any(name in reward_cfg for name in names)
def _device_type(device) -> str:
if isinstance(device, torch.device):
return device.type
return torch.device(device).type
def move_reward_models(reward_cfg, device) -> None:
"""Move GPU-backed reward models to device."""
if _has_reward(
reward_cfg,
{"hpsv3_general", "hpsv3_percentile"},
):
set_hpsv3_device(device)
if _has_reward(
reward_cfg,
{"videoalign_mq", "videoalign_ta"},
):
set_videoalign_device(device)
def clear_reward_models(reward_cfg) -> None:
"""Drop cached GPU-backed reward models before PPO training."""
cleared = False
if _has_reward(
reward_cfg,
{"hpsv3_general", "hpsv3_percentile"},
):
_HPSV3_INFERENCERS.clear()
cleared = True
if _has_reward(
reward_cfg,
{"videoalign_mq", "videoalign_ta"},
):
_VIDEOALIGN_INFERENCERS.clear()
cleared = True
if cleared:
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
@contextmanager
def reward_models_on_device(reward_cfg, device):
"""Temporarily move reward models to device."""
if _has_reward(reward_cfg, _GPU_REWARD_NAMES):
use_cuda = _device_type(device) == "cuda"
_t0 = time.perf_counter()
move_reward_models(reward_cfg, device)
if use_cuda:
torch.cuda.synchronize()
_t1 = time.perf_counter()
logger.info("[rewards] move_to_device=%.1fs", _t1 - _t0)
try:
yield
finally:
_t2 = time.perf_counter()
move_reward_models(reward_cfg, "cpu")
if use_cuda:
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
_t3 = time.perf_counter()
logger.info(
"[rewards] move_to_cpu+gc=%.1fs",
_t3 - _t2,
)
else:
yield
@@ -0,0 +1,246 @@
# SPDX-License-Identifier: Apache-2.0
"""Sampling epoch for RL training — generate videos and
compute rewards."""
from __future__ import annotations
import hashlib
import time
from collections.abc import Callable
from typing import Any
import torch
from fastvideo.logger import init_logger
from fastvideo.train.methods.rl.utils.embeddings import (
compute_text_embeddings, )
from fastvideo.train.methods.rl.utils.pipeline import (
wan_denoising_with_logprob, )
logger = init_logger(__name__)
SEED_EPOCH_STRIDE = 10_000
def create_generator(
prompts: list[str],
base_seed: int,
device: torch.device,
) -> list[torch.Generator]:
"""Create deterministic generators seeded by prompt."""
generators = []
for prompt in prompts:
prompt_seed = int.from_bytes(
hashlib.blake2b(
prompt.encode("utf-8"),
digest_size=8,
).digest(),
"big",
)
g = torch.Generator(device=device)
g.manual_seed(base_seed + prompt_seed % (2**31))
generators.append(g)
return generators
def sample_epoch(
model,
scheduler,
train_sampler,
train_iter,
reward_fn: Callable,
sample_neg_prompt_embeds: torch.Tensor,
text_encoder,
tokenizer,
executor,
epoch: int,
global_step: int,
*,
# Config values passed explicitly.
sample_batch_size: int,
num_batches_per_epoch: int,
num_inference_steps: int,
guidance_scale: float,
height: int,
width: int,
num_frames: int,
noise_level: float,
sde_type: str,
diffusion_clip: bool,
diffusion_clip_value: float,
sde_window_size: int,
sde_window_range: tuple[int, int] | None,
kl_reward: float,
same_latent: bool,
seed: int,
device: torch.device,
is_main_process: bool,
ref_transformer: torch.nn.Module | None = None,
lora_model: Any | None = None,
tracker: Any | None = None,
async_reward_scoring: bool = True,
) -> tuple[
list[dict[str, Any]],
list[torch.Tensor],
list[list[str]],
]:
"""Run one sampling epoch: generate videos and compute rewards.
Returns:
Tuple of (samples, all_videos, all_prompts):
- samples: list of sample dicts with prompt_ids,
prompt_embeds, latents, log_probs, kl,
timesteps, rewards.
- all_videos: list of decoded video tensors
per batch, each (B, 3, T, H, W) in [0, 1].
- all_prompts: list of prompt string lists
per batch.
"""
samples = []
all_videos: list[torch.Tensor] = []
all_prompts: list[list[str]] = []
for i in range(num_batches_per_epoch):
current_epoch_tag = (epoch * num_batches_per_epoch + i)
train_sampler.set_epoch(current_epoch_tag)
# Drain until epoch tag matches.
while True:
epoch_tag, prompts, prompt_metadata = next(train_iter)
if epoch_tag == current_epoch_tag:
break
torch.cuda.synchronize()
_t_batch_start = time.perf_counter()
_t_embed = time.perf_counter()
prompt_embeds = compute_text_embeddings(
prompts,
text_encoder,
tokenizer,
max_sequence_length=512,
device=device,
)
prompt_ids = tokenizer(
prompts,
padding="max_length",
max_length=512,
truncation=True,
return_tensors="pt",
).input_ids.to(device)
# Generator setup.
gen: torch.Generator | list[torch.Generator]
if same_latent:
gen = create_generator(
prompts,
base_seed=seed + epoch * SEED_EPOCH_STRIDE + i,
device=device,
)
else:
gen = torch.Generator(device=device)
gen.manual_seed(seed + epoch * SEED_EPOCH_STRIDE + i)
torch.cuda.synchronize()
_t_embed_done = time.perf_counter()
with torch.no_grad():
(
videos,
latents_list,
log_probs_list,
kls_list,
timesteps_list,
) = wan_denoising_with_logprob(
model,
scheduler,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=(sample_neg_prompt_embeds),
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
height=height,
width=width,
num_frames=num_frames,
generator=gen,
noise_level=noise_level,
sde_type=sde_type,
diffusion_clip=diffusion_clip,
diffusion_clip_value=diffusion_clip_value,
sde_window_size=sde_window_size,
sde_window_range=sde_window_range,
kl_reward=kl_reward,
ref_transformer=ref_transformer,
lora_model=lora_model,
)
torch.cuda.synchronize()
_t_denoise_done = time.perf_counter()
latents = torch.stack(latents_list, dim=1)
log_probs = torch.stack(log_probs_list, dim=1)
kls = torch.stack(kls_list, dim=1)
kl = kls.detach()
timesteps = (torch.stack(timesteps_list).unsqueeze(0).repeat(sample_batch_size, 1))
videos_cpu = videos.detach().cpu()
# Collect decoded videos and prompts for logging.
all_videos.append(videos_cpu)
all_prompts.append(list(prompts))
if async_reward_scoring:
rewards = executor.submit(
reward_fn,
videos_cpu,
prompts,
prompt_metadata,
True,
)
time.sleep(0)
else:
rewards = (videos_cpu, list(prompts), prompt_metadata)
del videos
logger.info(
"[sample_epoch] batch %d/%d: "
"text_embed=%.1fs denoise=%.1fs "
"batch_total=%.1fs",
i + 1,
num_batches_per_epoch,
_t_embed_done - _t_embed,
_t_denoise_done - _t_embed_done,
_t_denoise_done - _t_batch_start,
)
samples.append({
"prompt_ids": prompt_ids,
"prompt_embeds": prompt_embeds,
"negative_prompt_embeds": (sample_neg_prompt_embeds),
"timesteps": timesteps,
"latents": latents[:, :-1],
"next_latents": latents[:, 1:],
"log_probs": log_probs,
"kl": kl,
"rewards": rewards,
})
# Wait for all rewards.
torch.cuda.synchronize()
_t_reward_wait = time.perf_counter()
for sample in samples:
if async_reward_scoring:
rewards, _ = sample["rewards"].result()
else:
videos_cpu, prompts, prompt_metadata = sample["rewards"]
torch.cuda.empty_cache()
rewards, _ = reward_fn(videos_cpu, prompts, prompt_metadata, True)
sample["rewards"] = {key: torch.as_tensor(value, device=device).float() for key, value in rewards.items()}
_t_reward_done = time.perf_counter()
logger.info(
"[sample_epoch] reward_wait=%.1fs",
_t_reward_done - _t_reward_wait,
)
return samples, all_videos, all_prompts
+151
View File
@@ -0,0 +1,151 @@
# SPDX-License-Identifier: Apache-2.0
"""SDE step with log-probability computation for RL training."""
from __future__ import annotations
import math
import torch
from diffusers.utils.torch_utils import randn_tensor
def sde_step_with_logprob(
scheduler,
model_output: torch.FloatTensor,
timestep: float | torch.FloatTensor,
sample: torch.FloatTensor,
noise_level: float = 0.7,
prev_sample: torch.FloatTensor | None = None,
generator: torch.Generator | None = None,
sde_type: str | None = "flow_sde",
deterministic: bool = False,
return_sqrt_dt_and_std_dev_t: bool = False,
diffusion_clip: bool = False,
diffusion_clip_value: float = 0.45,
):
"""Predict the sample from the previous timestep by reversing
the SDE, returning log-probability of the transition.
Args:
scheduler: Noise scheduler with sigmas and timestep index.
model_output: Predicted noise/velocity.
timestep: Current timestep(s).
sample: Current latents.
noise_level: Noise level for SDE/CPS computation.
prev_sample: Optional precomputed previous sample.
generator: Optional RNG for sampling prev_sample.
sde_type: 'flow_sde' or 'flow_cps'.
deterministic: If True, no noise added.
return_sqrt_dt_and_std_dev_t: If True, return extra terms.
diffusion_clip: If True, clip std_dev_t.
diffusion_clip_value: Clipping threshold.
Returns:
If return_sqrt_dt_and_std_dev_t:
(prev_sample, log_prob, prev_sample_mean,
std_dev_t, sqrt_neg_dt, sigma, sigma_max)
Else:
(prev_sample, log_prob, prev_sample_mean,
std_dev_t * sqrt_neg_dt, sigma, sigma_max)
"""
model_output = model_output.float()
sample = sample.float()
if prev_sample is not None:
prev_sample = prev_sample.float()
step_index = [scheduler.index_for_timestep(t) for t in timestep]
prev_step_index = [step + 1 for step in step_index]
scheduler.sigmas = scheduler.sigmas.to(sample.device)
sigma = scheduler.sigmas[step_index].view(-1, 1, 1, 1, 1)
sigma_prev = (scheduler.sigmas[prev_step_index].view(-1, 1, 1, 1, 1))
sigma_max = scheduler.sigmas[1].item()
dt = sigma_prev - sigma
if sde_type == "flow_sde":
std_dev_t = (torch.sqrt(sigma / (1 - torch.where(
sigma == 1,
torch.tensor(
sigma_max,
device=sigma.device,
dtype=sigma.dtype,
),
sigma,
))) * noise_level)
if diffusion_clip:
max_std_dev_t = (diffusion_clip_value / torch.sqrt(-1 * dt))
std_dev_t = torch.minimum(std_dev_t, max_std_dev_t)
prev_sample_mean = (sample * (1 + std_dev_t**2 / (2 * sigma) * dt) + model_output *
(1 + std_dev_t**2 * (1 - sigma) / (2 * sigma)) * dt)
if prev_sample is None:
variance_noise = randn_tensor(
model_output.shape,
generator=generator,
device=model_output.device,
dtype=model_output.dtype,
)
prev_sample = (prev_sample_mean + std_dev_t * torch.sqrt(-1 * dt) * variance_noise)
if deterministic:
prev_sample = sample + dt * model_output
std_scale = std_dev_t * torch.sqrt(-1 * dt)
if torch.all(std_scale == 0):
log_prob = torch.zeros_like(prev_sample)
else:
std_scale = torch.clamp(
std_scale,
min=math.sqrt(torch.finfo(std_scale.dtype).tiny),
)
log_prob = (-((prev_sample.detach() - prev_sample_mean)**2) / (2 * (std_scale**2)) - torch.log(std_scale) -
torch.log(torch.sqrt(2 * torch.as_tensor(math.pi))))
elif sde_type == "flow_cps":
std_dev_t = sigma_prev * math.sin(noise_level * math.pi / 2)
pred_original_sample = sample - sigma * model_output
noise_estimate = (sample + model_output * (1 - sigma))
prev_sample_mean = pred_original_sample * (1 - sigma_prev) + noise_estimate * torch.sqrt(sigma_prev**2 -
std_dev_t**2)
if prev_sample is None:
variance_noise = randn_tensor(
model_output.shape,
generator=generator,
device=model_output.device,
dtype=model_output.dtype,
)
prev_sample = (prev_sample_mean + std_dev_t * variance_noise)
if deterministic:
prev_sample = (pred_original_sample * (1 - sigma_prev) + noise_estimate * sigma_prev)
log_prob = -((prev_sample.detach() - prev_sample_mean)**2)
else:
msg = (f"Unknown sde_type: {sde_type}. "
"Must be 'flow_sde' or 'flow_cps'.")
raise ValueError(msg)
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
if return_sqrt_dt_and_std_dev_t:
return (
prev_sample,
log_prob,
prev_sample_mean,
std_dev_t,
torch.sqrt(-1 * dt),
sigma,
sigma_max,
)
return (
prev_sample,
log_prob,
prev_sample_mean,
std_dev_t * torch.sqrt(-1 * dt),
sigma,
sigma_max,
)
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-prompt statistics tracking for advantage computation."""
from __future__ import annotations
import numpy as np
import torch
EPSILON = 1e-4
class PerPromptStatTracker:
"""Track per-prompt reward history and compute advantages."""
def __init__(
self,
use_global_std: bool = False,
max_group_std: bool = False,
):
self.use_global_std = use_global_std
self.max_group_std = max_group_std
self.stats: dict[str, np.ndarray] = {}
self.history_prompts: set[int] = set()
def update(
self,
prompts,
rewards,
mode: str = "grpo",
) -> np.ndarray:
"""Update stats and compute advantages.
Args:
prompts: Iterable of prompt strings.
rewards: Array-like rewards aligned with prompts.
mode: Advantage mode: grpo|rwr|sft|dpo.
Returns:
Advantages array aligned with prompts.
"""
prompts = np.array(prompts)
rewards = np.array(rewards, dtype=np.float64)
unique = np.unique(prompts)
advantages = np.empty_like(rewards) * 0.0
for prompt in unique:
prompt_rewards = rewards[prompts == prompt]
if prompt not in self.stats:
self.stats[prompt] = []
self.stats[prompt].extend(prompt_rewards)
self.history_prompts.add(hash(prompt))
self.stats[prompt] = np.stack(self.stats[prompt])
max_std = None
if self.max_group_std and len(unique) > 0:
prompt_stds = []
for prompt in unique:
prompt_std = (np.std(
self.stats[prompt],
axis=0,
keepdims=True,
) + EPSILON)
prompt_stds.append(prompt_std)
max_std_value = max(np.max(std) for std in prompt_stds)
max_std = np.full_like(prompt_stds[0], max_std_value)
for prompt in unique:
prompt_rewards = rewards[prompts == prompt]
mean = np.mean(self.stats[prompt], axis=0, keepdims=True)
if self.use_global_std:
std = (np.std(rewards, axis=0, keepdims=True) + EPSILON)
elif self.max_group_std:
std = max_std
else:
std = (np.std(
self.stats[prompt],
axis=0,
keepdims=True,
) + EPSILON)
if mode == "grpo":
advantages[prompts == prompt] = (prompt_rewards - mean) / std
elif mode == "rwr":
advantages[prompts == prompt] = (prompt_rewards)
elif mode == "sft":
advantages[prompts == prompt] = ((torch.tensor(prompt_rewards) == torch.max(
torch.tensor(prompt_rewards))).float().numpy())
elif mode == "dpo":
pa = torch.tensor(prompt_rewards)
max_idx = torch.argmax(pa)
min_idx = torch.argmin(pa)
if max_idx == min_idx:
min_idx = 0
max_idx = 1
result = torch.zeros_like(pa).float()
result[max_idx] = 1.0
result[min_idx] = -1.0
advantages[prompts == prompt] = (result.numpy())
return advantages
def get_stats(self) -> tuple[float, int]:
"""Return (avg_group_size, num_unique_prompts)."""
avg = (sum(len(v) for v in self.stats.values()) / len(self.stats) if self.stats else 0)
return avg, len(self.history_prompts)
def clear(self):
"""Clear stored statistics."""
self.stats = {}
+2
View File
@@ -5,3 +5,5 @@ from fastvideo.train.models.wan.wan import (
WanModel as WanModel, )
from fastvideo.train.models.wan.wan_causal import (
WanCausalModel as WanCausalModel, )
from fastvideo.train.models.wan.wan_genrl import (
GenRLWanModel as GenRLWanModel, )
+82
View File
@@ -502,6 +502,88 @@ class WanModel(ModelBase):
def _get_transformer(self, timestep: torch.Tensor) -> torch.nn.Module:
return self.transformer
# ------------------------------------------------------------------
# RL pipeline primitives
# ------------------------------------------------------------------
def forward_transformer_raw(
self,
latents: torch.Tensor,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
) -> torch.Tensor:
"""Direct transformer forward for RL pipelines.
Bypasses batch preparation / attention-metadata.
Uses dense attention (attn_metadata=None).
Args:
latents: (B, C, T, H, W) in diffusion space.
timestep: (B,) or scalar timestep.
encoder_hidden_states: (B, L, D) text embeddings.
Returns:
Model output tensor (B, C, T, H, W).
"""
dtype = self._get_training_dtype()
device_type = self.device.type
with (
torch.autocast(device_type, dtype=dtype),
set_forward_context(
current_timestep=timestep,
attn_metadata=None,
),
):
output = self.transformer(
hidden_states=latents.to(dtype),
timestep=timestep,
encoder_hidden_states=(encoder_hidden_states.to(dtype)),
return_dict=False,
)
return output
def decode_latents(
self,
latents: torch.Tensor,
) -> torch.Tensor:
"""Decode latents to pixel-space video.
Denormalizes from flow diffusion space and passes
through the VAE decoder.
Args:
latents: (B, C, T, H, W) denoised latents.
Returns:
Video tensor (B, 3, T_out, H_out, W_out)
in [0, 1] range.
"""
vae = self.vae
vae_config = vae.config
z_dim = getattr(vae_config, "z_dim", 16)
latents_mean = (torch.tensor(vae_config.latents_mean).view(1, z_dim, 1, 1, 1).to(latents.device, latents.dtype))
latents_std_inv = (1.0 / torch.tensor(vae_config.latents_std).view(1, z_dim, 1, 1, 1)).to(
latents.device, latents.dtype)
latents = latents / latents_std_inv + latents_mean
# Decode one sample at a time.
vae_dtype = next(vae.parameters()).dtype
videos = []
with torch.no_grad():
for idx in range(latents.shape[0]):
decoded = vae.decode(latents[idx:idx + 1].to(vae_dtype))
if isinstance(decoded, tuple):
decoded = decoded[0]
videos.append(decoded.float())
video = torch.cat(videos, dim=0)
# Normalize to [0, 1].
video = (video / 2 + 0.5).clamp(0, 1)
return video
# ------------------------------------------------------------------
def _get_uncond_text_dict(
self,
batch: TrainingBatch,
+249
View File
@@ -0,0 +1,249 @@
# SPDX-License-Identifier: Apache-2.0
"""Wan model extended for GenRL (RL training with text prompts).
Overrides ``init_preprocessors`` to load a T5 text encoder
and tokenizer instead of the standard parquet video
dataloader. Provides a trivial dummy dataloader so the
trainer's outer loop has something to iterate over — the
real prompt dataloaders are created by the GenRLMethod.
"""
from __future__ import annotations
from contextlib import contextmanager
from types import MethodType
from typing import Any, TYPE_CHECKING
import torch
from fastvideo.distributed import (
get_sp_group,
get_world_group,
)
from fastvideo.logger import init_logger
from fastvideo.train.models.wan.wan import WanModel
from fastvideo.train.utils.moduleloader import (
load_module_from_path, )
if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
TrainingConfig, )
logger = init_logger(__name__)
def _is_lora_target(
module_name: str,
target_modules: list[str],
) -> bool:
return any(module_name == target or module_name.endswith(f".{target}") for target in target_modules)
def _apply_fastvideo_lora(
transformer: Any,
*,
lora_rank: int,
lora_alpha: int,
target_modules: list[str],
init_weights: str,
) -> int:
from fastvideo.layers.lora.linear import (
get_lora_layer,
replace_submodule,
)
transformer.requires_grad_(False)
converted_count = 0
for name, layer in list(transformer.named_modules()):
if not _is_lora_target(name, target_modules):
continue
lora_layer = get_lora_layer(
layer,
lora_rank=lora_rank,
lora_alpha=lora_alpha,
training_mode=True,
)
if lora_layer is None:
continue
_init_lora_weights(lora_layer, init_weights, lora_rank)
replace_submodule(transformer, name, lora_layer)
converted_count += 1
return converted_count
def _init_lora_weights(
lora_layer: Any,
init_weights: str,
lora_rank: int,
) -> None:
"""Match PEFT's useful LoRA initialization modes."""
init = init_weights.lower()
if init == "default":
return
lora_A = getattr(lora_layer, "lora_A", None)
lora_B = getattr(lora_layer, "lora_B", None)
if lora_A is None or lora_B is None:
return
if init == "gaussian":
torch.nn.init.normal_(lora_A, std=1 / max(1, lora_rank))
torch.nn.init.zeros_(lora_B)
return
raise ValueError("Unsupported GenRLWanModel LoRA init_weights="
f"{init_weights!r}. Use 'gaussian' or 'default'.")
@contextmanager
def _disable_lora_adapters(transformer: Any):
"""Temporarily run a LoRA-wrapped transformer as its frozen base model."""
lora_layers = [module for module in transformer.modules() if hasattr(module, "disable_lora")]
previous = [bool(module.disable_lora) for module in lora_layers]
try:
for module in lora_layers:
module.disable_lora = True
yield
finally:
for module, was_disabled in zip(lora_layers, previous, strict=True):
module.disable_lora = was_disabled
def _attach_disable_adapter(transformer: Any) -> None:
"""Expose a PEFT-compatible disable_adapter context manager."""
def disable_adapter(self):
return _disable_lora_adapters(self)
transformer.disable_adapter = MethodType( # type: ignore[attr-defined]
disable_adapter,
transformer,
)
class _InfiniteDummyLoader:
"""Trivial iterable that yields empty dicts forever."""
def __iter__(self):
while True:
yield {}
class GenRLWanModel(WanModel):
"""Wan model with text encoder for RL training.
Compared to the base :class:`WanModel`, this variant:
* Loads the T5 text encoder and tokenizer from the
pretrained model path.
* Sets a dummy dataloader so the trainer can iterate
without blocking.
* Does **not** build the standard parquet video
dataloader.
"""
def __init__(
self,
*,
init_from: str,
training_config: TrainingConfig,
trainable: bool = True,
use_lora: bool = False,
lora_r: int = 32,
lora_alpha: int = 64,
lora_target_modules: list[str] | None = None,
lora_path: str | None = None,
lora_init_weights: str = "gaussian",
disable_custom_init_weights: bool = False,
flow_shift: float = 3.0,
enable_gradient_checkpointing_type: str
| None = None,
transformer_override_safetensor: str
| None = None,
) -> None:
super().__init__(
init_from=init_from,
training_config=training_config,
trainable=trainable,
disable_custom_init_weights=(disable_custom_init_weights),
flow_shift=flow_shift,
enable_gradient_checkpointing_type=(enable_gradient_checkpointing_type),
transformer_override_safetensor=(transformer_override_safetensor),
)
if use_lora:
if lora_target_modules is None:
raise ValueError("GenRLWanModel use_lora=True requires "
"lora_target_modules.")
if lora_path:
raise ValueError("GenRLWanModel lora_path is not supported for "
"FastVideo LoRA training yet.")
converted_count = _apply_fastvideo_lora(
self.transformer,
lora_rank=int(lora_r),
lora_alpha=int(lora_alpha),
target_modules=lora_target_modules,
init_weights=lora_init_weights,
)
if converted_count == 0:
raise ValueError("GenRLWanModel use_lora=True did not match any "
f"FastVideo linear layers: {lora_target_modules}")
logger.info(
"Converted %d GenRL Wan transformer layers to LoRA",
converted_count,
)
_attach_disable_adapter(self.transformer)
self.text_encoder: Any = None
self.tokenizer: Any = None
def disable_adapter(self):
"""PEFT-compatible context manager for reference KL with LoRA."""
return _disable_lora_adapters(self.transformer)
# ------------------------------------------------------------------
# Lifecycle
# ------------------------------------------------------------------
def init_preprocessors(
self,
training_config: TrainingConfig,
) -> None: # type: ignore[override]
"""Load VAE, text encoder, and tokenizer."""
# Load VAE.
self.vae = load_module_from_path(
model_path=str(training_config.model_path),
module_type="vae",
training_config=training_config,
)
self.world_group = get_world_group()
self.sp_group = get_sp_group()
self._init_timestep_mechanics()
# Load text encoder and tokenizer.
model_path = str(training_config.model_path)
self._load_text_encoder(model_path, training_config)
# Dummy dataloader for the trainer's outer loop.
self.dataloader = _InfiniteDummyLoader()
self.start_step = 0
def _load_text_encoder(
self,
model_path: str,
training_config: TrainingConfig,
) -> None:
from transformers import AutoTokenizer
logger.info("Loading tokenizer from %s", model_path)
self.tokenizer = AutoTokenizer.from_pretrained(model_path, subfolder="tokenizer")
logger.info("Loading text encoder from %s", model_path)
self.text_encoder = load_module_from_path(
model_path=model_path,
module_type="text_encoder",
training_config=training_config,
)
self.text_encoder.requires_grad_(False)
def on_train_start(self) -> None:
"""Skip negative conditioning (handled by method)."""
+1
View File
@@ -170,6 +170,7 @@ class Trainer:
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
outputs=outputs,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
+9
View File
@@ -209,6 +209,15 @@ disallow_untyped_calls = true
check_untyped_defs = true
follow_imports = "silent"
[[tool.mypy.overrides]]
module = [
"fastvideo.train.methods.rl.reward.HPSv3.*",
"fastvideo.train.methods.rl.reward.VideoAlign.*",
"*.fastvideo.train.methods.rl.reward.HPSv3.*",
"*.fastvideo.train.methods.rl.reward.VideoAlign.*",
]
ignore_errors = true
[tool.codespell]
# ``*/_vendored/*`` matches upstream-provenance files vendored under any
# ``_vendored/`` subdir (project-wide convention; mirrors the