Compare commits
22
Commits
will/design
...
py/add_rl
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a520eceed7 | ||
|
|
b2c75534ac | ||
|
|
0b9c325984 | ||
|
|
2afa837d96 | ||
|
|
83bb602836 | ||
|
|
606b42b6f9 | ||
|
|
e9a7a128c2 | ||
|
|
78b6995e13 | ||
|
|
d858270708 | ||
|
|
298845ccd9 | ||
|
|
9c2fa5aa9d | ||
|
|
73e0fa727f | ||
|
|
1fbe5a868c | ||
|
|
e1be068c46 | ||
|
|
cd69575095 | ||
|
|
2c19cbdebd | ||
|
|
5ae633a551 | ||
|
|
6e663214a7 | ||
|
|
8b0ba24deb | ||
|
|
ef83574e44 | ||
|
|
9f2fc2f303 | ||
|
|
c52de7e9df |
@@ -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
|
||||
@@ -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[*]:-}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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 "
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
@@ -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."""
|
||||
+185
@@ -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
@@ -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."""
|
||||
+615
@@ -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
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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 = {}
|
||||
@@ -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, )
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)."""
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user