Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
141a1140f6 | ||
|
|
6294015389 | ||
|
|
e31b6c9e90 | ||
|
|
bf0ff21eeb | ||
|
|
6f937102ad | ||
|
|
0164e93019 | ||
|
|
67e457aa92 | ||
|
|
d795f0c443 | ||
|
|
02452dd6e7 | ||
|
|
91ef24bc14 | ||
|
|
e76e9fda15 | ||
|
|
3b17f5a621 | ||
|
|
d758878705 | ||
|
|
689e629420 | ||
|
|
873dc9695f | ||
|
|
bfc0f46d61 | ||
|
|
39907dbe4d | ||
|
|
abdd0c9b6a | ||
|
|
f32a12200d | ||
|
|
f1d2c9e6b7 | ||
|
|
450579cb42 | ||
|
|
26d7d6cc08 | ||
|
|
44f0124eaa | ||
|
|
d3ace51394 | ||
|
|
58954c660b |
Executable
+129
@@ -0,0 +1,129 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Change to FastVideo root directory (3 levels up from this script)
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
FASTVIDEO_ROOT="$(cd "$SCRIPT_DIR/../../.." && pwd)"
|
||||
cd "$FASTVIDEO_ROOT"
|
||||
|
||||
# Add FastVideo root to PYTHONPATH so Python can find the fastvideo package
|
||||
export PYTHONPATH="$FASTVIDEO_ROOT${PYTHONPATH:+:$PYTHONPATH}"
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
RL_DATASET_DIR="data/ocr/" # Path to RL prompt dataset directory (should contain train.txt and test.txt)
|
||||
VALIDATION_DATASET_FILE="$SCRIPT_DIR/validation.json"
|
||||
NUM_GPUS=1
|
||||
|
||||
# use GPU 3
|
||||
export CUDA_VISIBLE_DEVICES=3
|
||||
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_grpo"
|
||||
--output_dir "checkpoints/wan_t2v_grpo"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 4
|
||||
# --train_sp_batch_size 4
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 5
|
||||
--num_height 240
|
||||
--num_width 416
|
||||
--num_frames 33
|
||||
--lora_rank 32
|
||||
--lora_training True
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
# --use-fsdp-inference False
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments (for RL prompt dataset)
|
||||
dataset_args=(
|
||||
--data_path $RL_DATASET_DIR # Used as fallback if rl_dataset_path not set
|
||||
--rl_dataset_path $RL_DATASET_DIR # RL prompt dataset directory
|
||||
--rl_dataset_type "text" # "text" or "geneval"
|
||||
--rl_num_image_per_prompt 4 # k parameter (number of samples per prompt)
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation True
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 5
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 10
|
||||
--training_state_checkpointing_steps 10
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# RL-specific arguments
|
||||
rl_args=(
|
||||
--inference_mode False
|
||||
--rl_mode True
|
||||
--rl_algorithm "grpo"
|
||||
--rl_kl_beta 0.004 # KL regularization coefficient
|
||||
--rl_policy_clip_range 0.2 # Policy clipping range for GRPO
|
||||
--rl_kl_reward 0.0 # KL reward coefficient (typically 0)
|
||||
--rl_global_std False # Use per-prompt std (recommended for GRPO)
|
||||
--rl_per_prompt_stat_tracking True # Enable per-prompt stat tracking
|
||||
--rl_warmup_steps 0 # Number of warmup steps (SFT before RL)
|
||||
--reward-models "{\"paddle_ocr\": 1.0}" # use video_ocr reward function
|
||||
)
|
||||
|
||||
# CFG arguments
|
||||
cfg_args=(
|
||||
--guidance_scale 1.0 # use guidance_scale > 1.0 to enable CFG
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0 # No CFG during training (CFG used in sampling)
|
||||
--dit_precision "fp32"
|
||||
# --dit_precision "bf16"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
# --resume_from_checkpoint "checkpoints/wan_t2v_grpo/checkpoint-XXX"
|
||||
--enable-gradient-checkpointing-type "full"
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--master_port 29501 \
|
||||
"$FASTVIDEO_ROOT/fastvideo/training/wan_rl_training_pipeline.py" \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${rl_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -8,6 +8,7 @@ from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset,
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
from fastvideo.dataset.rl_prompt_dataset import build_rl_prompt_dataloader
|
||||
|
||||
|
||||
def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
@@ -47,5 +48,6 @@ def gettextdataset(args) -> TextDataset:
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset", "TextDataset"
|
||||
"VideoCaptionMergedDataset", "TextDataset",
|
||||
"build_rl_prompt_dataloader"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import torch
|
||||
from torch.utils.data import Dataset, DataLoader, Sampler
|
||||
import json
|
||||
import os
|
||||
|
||||
|
||||
class TextPromptDataset(Dataset):
|
||||
"""Dataset for loading text prompts from a simple text file (one prompt per line)."""
|
||||
|
||||
def __init__(self, dataset, split='train'):
|
||||
self.file_path = os.path.join(dataset, f'{split}.txt')
|
||||
with open(self.file_path, 'r') as f:
|
||||
self.prompts = [line.strip() for line in f.readlines()]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.prompts)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {"prompt": self.prompts[idx], "metadata": {}}
|
||||
|
||||
@staticmethod
|
||||
def collate_fn(examples):
|
||||
prompts = [example["prompt"] for example in examples]
|
||||
metadatas = [example["metadata"] for example in examples]
|
||||
return prompts, metadatas
|
||||
|
||||
|
||||
class GenevalPromptDataset(Dataset):
|
||||
"""Dataset for loading prompts with metadata from JSONL files (e.g., GenEval format)."""
|
||||
|
||||
def __init__(self, dataset, split='train'):
|
||||
self.file_path = os.path.join(dataset, f'{split}_metadata.jsonl')
|
||||
with open(self.file_path, 'r', encoding='utf-8') as f:
|
||||
self.metadatas = [json.loads(line) for line in f]
|
||||
self.prompts = [item['prompt'] for item in self.metadatas]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.prompts)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
return {"prompt": self.prompts[idx], "metadata": self.metadatas[idx]}
|
||||
|
||||
@staticmethod
|
||||
def collate_fn(examples):
|
||||
prompts = [example["prompt"] for example in examples]
|
||||
metadatas = [example["metadata"] for example in examples]
|
||||
return prompts, metadatas
|
||||
|
||||
|
||||
class KRepeatSampler(Sampler):
|
||||
"""Sampler that repeats each sample k times, ensuring synchronized random selection. For single-node training, set num_replicas=1 and rank=0."""
|
||||
|
||||
def __init__(self, dataset, batch_size, k, num_replicas, rank, seed=0):
|
||||
self.dataset = dataset
|
||||
self.batch_size = batch_size # Batch size per GPU/card
|
||||
self.k = k # Number of repetitions per sample
|
||||
self.num_replicas = num_replicas # Total number of GPUs/cards
|
||||
self.rank = rank # Current GPU/card rank
|
||||
self.seed = seed # Random seed for synchronization
|
||||
|
||||
# Calculate the number of unique samples needed for each iteration
|
||||
self.total_samples = self.num_replicas * self.batch_size
|
||||
assert self.total_samples % self.k == 0, f"k can not div n*b, k{k}-num_replicas{num_replicas}-batch_size{batch_size}"
|
||||
self.m = self.total_samples // self.k # different number of samples
|
||||
self.step = 0
|
||||
|
||||
def __iter__(self):
|
||||
while True:
|
||||
# Generate a deterministic random sequence to ensure all cards are synchronized
|
||||
g = torch.Generator()
|
||||
g.manual_seed(self.seed + self.step)
|
||||
|
||||
# Randomly select m unique samples
|
||||
indices = torch.randperm(len(self.dataset), generator=g)[:self.m].tolist()
|
||||
|
||||
# Repeat each sample k times to generate a total of n*b samples
|
||||
repeated_indices = [idx for idx in indices for _ in range(self.k)]
|
||||
|
||||
# Shuffle the order to ensure even distribution
|
||||
shuffled_indices = torch.randperm(len(repeated_indices), generator=g).tolist()
|
||||
shuffled_samples = [repeated_indices[i] for i in shuffled_indices]
|
||||
|
||||
# Split samples among all cards
|
||||
per_card_samples = []
|
||||
for i in range(self.num_replicas):
|
||||
start = i * self.batch_size
|
||||
end = start + self.batch_size
|
||||
per_card_samples.append(shuffled_samples[start:end])
|
||||
|
||||
# Return the sample indices for the current card
|
||||
yield per_card_samples[self.rank]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dataset) // self.batch_size
|
||||
|
||||
def set_step(self, step):
|
||||
"""Used to synchronize the random state for different epochs."""
|
||||
self.step = step
|
||||
|
||||
|
||||
def build_rl_prompt_dataloader(
|
||||
dataset_path: str,
|
||||
dataset_type: str = "text",
|
||||
split: str = "train",
|
||||
train_batch_size: int = 8,
|
||||
test_batch_size: int = 8,
|
||||
k: int = 1,
|
||||
seed: int = 42,
|
||||
train_num_workers: int = 1,
|
||||
test_num_workers: int = 8,
|
||||
num_replicas: int = 1,
|
||||
rank: int = 0,
|
||||
) -> tuple[DataLoader, DataLoader]:
|
||||
"""
|
||||
Factory function to create train and test dataloaders for RL prompt datasets.
|
||||
|
||||
Args:
|
||||
dataset_path: Path to dataset directory
|
||||
dataset_type: "text" for TextPromptDataset or "geneval" for GenevalPromptDataset
|
||||
split: Dataset split ("train" or "test")
|
||||
train_batch_size: Batch size per GPU for training
|
||||
test_batch_size: Batch size for testing
|
||||
k: Number of times to repeat each sample (num_image_per_prompt)
|
||||
seed: Random seed for sampler synchronization
|
||||
train_num_workers: Number of workers for training dataloader
|
||||
test_num_workers: Number of workers for test dataloader
|
||||
num_replicas: Number of replicas (default 1 for single-node)
|
||||
rank: Rank of current process (default 0 for single-node)
|
||||
|
||||
Returns:
|
||||
Tuple of (train_dataloader, test_dataloader)
|
||||
"""
|
||||
# Create datasets based on type
|
||||
if dataset_type == "text":
|
||||
train_dataset = TextPromptDataset(dataset_path, 'train')
|
||||
test_dataset = TextPromptDataset(dataset_path, 'test')
|
||||
collate_fn = TextPromptDataset.collate_fn
|
||||
elif dataset_type == "geneval":
|
||||
train_dataset = GenevalPromptDataset(dataset_path, 'train')
|
||||
test_dataset = GenevalPromptDataset(dataset_path, 'test')
|
||||
collate_fn = GenevalPromptDataset.collate_fn
|
||||
else:
|
||||
raise ValueError(f"Unknown dataset_type: {dataset_type}. Must be 'text' or 'geneval'")
|
||||
|
||||
# Create infinite-loop training sampler
|
||||
train_sampler = KRepeatSampler(
|
||||
dataset=train_dataset,
|
||||
batch_size=train_batch_size,
|
||||
k=k,
|
||||
num_replicas=num_replicas,
|
||||
rank=rank,
|
||||
seed=seed
|
||||
)
|
||||
|
||||
# Create training dataloader with batch_sampler (infinite loop)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_sampler=train_sampler,
|
||||
num_workers=train_num_workers,
|
||||
collate_fn=collate_fn,
|
||||
)
|
||||
|
||||
# Create standard test dataloader
|
||||
test_dataloader = DataLoader(
|
||||
test_dataset,
|
||||
batch_size=test_batch_size,
|
||||
collate_fn=collate_fn,
|
||||
shuffle=False,
|
||||
num_workers=test_num_workers,
|
||||
)
|
||||
|
||||
return train_dataloader, test_dataloader, train_dataset, test_dataset
|
||||
|
||||
+311
-1
@@ -740,6 +740,271 @@ def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
return _current_fastvideo_args
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class RLArgs:
|
||||
"""
|
||||
Reinforcement Learning (RL) specific arguments
|
||||
"""
|
||||
# ============================================================================
|
||||
# SHARED RL CONFIGURATION
|
||||
rl_mode: bool = False # Enable RL training mode
|
||||
rl_algorithm: str = "grpo" # RL algorithm to use: "grpo", "ppo", "dpo"
|
||||
|
||||
# Trajectory collection
|
||||
num_rollouts: int = 4 # Number of rollouts to collect per training step
|
||||
rollout_steps: str = "20,30" # Random intermediate steps for sampling (comma-separated)
|
||||
noise_injection_min: int = 10 # Minimum timestep for noise injection
|
||||
noise_injection_max: int = 40 # Maximum timestep for noise injection
|
||||
use_sde_sampling: bool = True # Use SDE sampling (Flow-GRPO-Fast)
|
||||
num_denoising_steps: int = 2 # Number of denoising steps per trajectory (1-2 for fast)
|
||||
|
||||
# Advantage estimation
|
||||
gamma: float = 0.99 # Discount factor for returns
|
||||
lambda_param: float = 0.95 # GAE lambda parameter
|
||||
use_gae: bool = True # Use Generalized Advantage Estimation
|
||||
normalize_advantages: bool = True # Normalize advantages before policy update
|
||||
|
||||
# Reward models
|
||||
reward_models: dict[str, float] = field(default_factory=lambda: {"dummy": 1.0}) # reward models (names, weight)
|
||||
value_model_path: str = "" # Path to value model (can be empty to train from scratch)
|
||||
value_model_share_backbone: bool = False # Share transformer backbone between policy and value
|
||||
|
||||
# Training schedule
|
||||
warmup_steps: int = 1000 # Collect SFT-style data before starting RL
|
||||
collect_on_policy: bool = True # Collect fresh rollouts each step (on-policy)
|
||||
timestep_fraction: float = 0.99 # Fraction of timesteps to train on
|
||||
num_inner_epochs: int = 1 # Number of inner epochs per outer epoch
|
||||
|
||||
# KL regularization
|
||||
kl_beta: float = 0.004 # KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)
|
||||
kl_reward: float = 0.0 # KL reward coefficient (alternative to KL loss, typically 0)
|
||||
|
||||
# SFT integration
|
||||
sft_weight: float = 0.0 # SFT loss weight for supervised learning in RL training
|
||||
sft_batch_size: int = 3 # Batch size for SFT data
|
||||
|
||||
# CFG
|
||||
guidance_scale = 1.0 # use guidance_scale > 1.0 to enable CFG
|
||||
|
||||
# Statistics tracking
|
||||
global_std: bool = False # Use global std across all samples vs per-group std
|
||||
per_prompt_stat_tracking: bool = True # Track statistics per prompt
|
||||
|
||||
# Training options
|
||||
use_diffusion_loss: bool = True # Use diffusion loss in training
|
||||
|
||||
# ============================================================================
|
||||
# GRPO-SPECIFIC CONFIGURATION
|
||||
|
||||
# Policy optimization
|
||||
grpo_policy_clip_range: float = 0.001 # PPO-style clipping range for policy ratio
|
||||
grpo_value_clip_range: float = 0.2 # Value function clipping range
|
||||
grpo_num_policy_epochs: int = 1 # Number of policy update epochs (GRPO typically uses 1)
|
||||
grpo_num_value_epochs: int = 1 # Number of value function update epochs
|
||||
grpo_target_kl: float = 0.01 # Target KL divergence for early stopping
|
||||
grpo_entropy_coef: float = 0.0 # Entropy coefficient for exploration
|
||||
grpo_value_loss_coef: float = 0.5 # Value loss coefficient
|
||||
|
||||
# GRPO-Guard safety mechanisms
|
||||
grpo_use_grpo_guard: bool = True # Enable GRPO-Guard safety mechanisms
|
||||
grpo_ratio_norm_correction: bool = True # RatioNorm: correct importance ratio bias
|
||||
grpo_gradient_reweighting: bool = True # Reweight gradients across denoising steps
|
||||
grpo_max_importance_ratio: float = 10.0 # Clip importance ratios above this value
|
||||
|
||||
# ============================================================================
|
||||
# DPO-SPECIFIC CONFIGURATION
|
||||
|
||||
dpo_beta: float = 100.0 # DPO regularization parameter (typically much larger than GRPO beta)
|
||||
dpo_ref_update_step: int = 10000000 # Reference model update frequency for OnlineDPO
|
||||
dpo_label_smoothing: float = 0.0 # Label smoothing for DPO loss
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
"""Add RL-specific CLI arguments to the parser."""
|
||||
# RL (Reinforcement Learning) arguments
|
||||
parser.add_argument("--rl-mode",
|
||||
action=StoreBoolean,
|
||||
help="Enable RL training mode")
|
||||
parser.add_argument("--rl-algorithm",
|
||||
type=str,
|
||||
default=RLArgs.rl_algorithm,
|
||||
choices=["grpo", "ppo", "dpo"],
|
||||
help="RL algorithm to use (grpo, ppo, dpo)")
|
||||
|
||||
# Trajectory collection (Flow-GRPO-Fast)
|
||||
parser.add_argument("--rl-num-rollouts",
|
||||
type=int,
|
||||
default=RLArgs.num_rollouts,
|
||||
help="Number of rollouts to collect per training step")
|
||||
parser.add_argument("--rl-rollout-steps",
|
||||
type=str,
|
||||
default=RLArgs.rollout_steps,
|
||||
help="Random intermediate steps for sampling (comma-separated)")
|
||||
parser.add_argument("--rl-noise-injection-min",
|
||||
type=int,
|
||||
default=RLArgs.noise_injection_min,
|
||||
help="Minimum timestep for noise injection")
|
||||
parser.add_argument("--rl-noise-injection-max",
|
||||
type=int,
|
||||
default=RLArgs.noise_injection_max,
|
||||
help="Maximum timestep for noise injection")
|
||||
parser.add_argument("--rl-use-sde-sampling",
|
||||
action=StoreBoolean,
|
||||
help="Use SDE sampling (Flow-GRPO-Fast)")
|
||||
parser.add_argument("--rl-num-denoising-steps",
|
||||
type=int,
|
||||
default=RLArgs.num_denoising_steps,
|
||||
help="Number of denoising steps per trajectory (1-2 for fast)")
|
||||
|
||||
# Advantage estimation
|
||||
parser.add_argument("--rl-gamma",
|
||||
type=float,
|
||||
default=RLArgs.gamma,
|
||||
help="Discount factor for returns")
|
||||
parser.add_argument("--rl-lambda",
|
||||
type=float,
|
||||
default=RLArgs.lambda_param,
|
||||
help="GAE lambda parameter")
|
||||
parser.add_argument("--rl-use-gae",
|
||||
action=StoreBoolean,
|
||||
help="Use Generalized Advantage Estimation")
|
||||
parser.add_argument("--rl-normalize-advantages",
|
||||
action=StoreBoolean,
|
||||
help="Normalize advantages before policy update")
|
||||
|
||||
# Policy optimization (GRPO/PPO)
|
||||
parser.add_argument("--rl-policy-clip-range",
|
||||
type=float,
|
||||
default=RLArgs.grpo_policy_clip_range,
|
||||
dest="grpo_policy_clip_range", # Map to RLArgs field name
|
||||
help="PPO-style clipping range for policy ratio")
|
||||
parser.add_argument("--rl-value-clip-range",
|
||||
type=float,
|
||||
default=RLArgs.grpo_value_clip_range,
|
||||
help="Value function clipping range")
|
||||
parser.add_argument("--rl-num-policy-epochs",
|
||||
type=int,
|
||||
default=RLArgs.grpo_num_policy_epochs,
|
||||
help="Number of policy update epochs (GRPO typically uses 1)")
|
||||
parser.add_argument("--rl-num-value-epochs",
|
||||
type=int,
|
||||
default=RLArgs.grpo_num_value_epochs,
|
||||
help="Number of value function update epochs")
|
||||
parser.add_argument("--rl-target-kl",
|
||||
type=float,
|
||||
default=RLArgs.grpo_target_kl,
|
||||
help="Target KL divergence for early stopping")
|
||||
parser.add_argument("--rl-entropy-coef",
|
||||
type=float,
|
||||
default=RLArgs.grpo_entropy_coef,
|
||||
help="Entropy coefficient for exploration")
|
||||
parser.add_argument("--rl-value-loss-coef",
|
||||
type=float,
|
||||
default=RLArgs.grpo_value_loss_coef,
|
||||
help="Value loss coefficient")
|
||||
|
||||
# GRPO-Guard (safety mechanisms)
|
||||
parser.add_argument("--rl-use-grpo-guard",
|
||||
action=StoreBoolean,
|
||||
help="Enable GRPO-Guard safety mechanisms")
|
||||
parser.add_argument("--rl-ratio-norm-correction",
|
||||
action=StoreBoolean,
|
||||
help="RatioNorm: correct importance ratio bias")
|
||||
parser.add_argument("--rl-gradient-reweighting",
|
||||
action=StoreBoolean,
|
||||
help="Reweight gradients across denoising steps")
|
||||
parser.add_argument("--rl-max-importance-ratio",
|
||||
type=float,
|
||||
default=RLArgs.grpo_max_importance_ratio,
|
||||
help="Clip importance ratios above this value")
|
||||
|
||||
# Reward models
|
||||
parser.add_argument("--reward-models",
|
||||
type=str,
|
||||
default='{"dummy": 1.0}',
|
||||
help="Reward models as JSON dict (e.g., '{\"video_ocr\": 1.0, \"pickscore\": 0.5}')")
|
||||
parser.add_argument("--value-model-path",
|
||||
type=str,
|
||||
default=RLArgs.value_model_path,
|
||||
help="Path to value model (can be empty to train from scratch)")
|
||||
parser.add_argument("--value-model-share-backbone",
|
||||
action=StoreBoolean,
|
||||
help="Share transformer backbone between policy and value")
|
||||
|
||||
# Training schedule
|
||||
parser.add_argument("--rl-warmup-steps",
|
||||
type=int,
|
||||
default=RLArgs.warmup_steps,
|
||||
help="Collect SFT-style data before starting RL")
|
||||
parser.add_argument("--rl-collect-on-policy",
|
||||
action=StoreBoolean,
|
||||
help="Collect fresh rollouts each step (on-policy)")
|
||||
parser.add_argument("--rl-timestep-fraction",
|
||||
type=float,
|
||||
default=RLArgs.timestep_fraction,
|
||||
help="Fraction of timesteps to train on")
|
||||
parser.add_argument("--rl-num-inner-epochs",
|
||||
type=int,
|
||||
default=RLArgs.num_inner_epochs,
|
||||
help="Number of inner epochs per outer epoch")
|
||||
|
||||
# KL regularization
|
||||
parser.add_argument("--rl-kl-beta",
|
||||
type=float,
|
||||
default=RLArgs.kl_beta,
|
||||
dest="kl_beta", # Map CLI arg to RLArgs field name
|
||||
help="KL loss coefficient (GRPO uses KL loss, DPO uses larger beta)")
|
||||
parser.add_argument("--rl-kl-reward",
|
||||
type=float,
|
||||
default=RLArgs.kl_reward,
|
||||
help="KL reward coefficient (alternative to KL loss, typically 0)")
|
||||
|
||||
# SFT integration
|
||||
parser.add_argument("--rl-sft-weight",
|
||||
type=float,
|
||||
default=RLArgs.sft_weight,
|
||||
help="SFT loss weight for supervised learning in RL training")
|
||||
parser.add_argument("--rl-sft-batch-size",
|
||||
type=int,
|
||||
default=RLArgs.sft_batch_size,
|
||||
help="Batch size for SFT data")
|
||||
|
||||
# CFG settings
|
||||
parser.add_argument("--guidance-scale",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Guidance scale for CFG")
|
||||
|
||||
# Statistics tracking
|
||||
parser.add_argument("--rl-global-std",
|
||||
action=StoreBoolean,
|
||||
help="Use global std across all samples vs per-group std")
|
||||
parser.add_argument("--rl-per-prompt-stat-tracking",
|
||||
action=StoreBoolean,
|
||||
help="Track statistics per prompt")
|
||||
|
||||
# Training options
|
||||
parser.add_argument("--rl-use-diffusion-loss",
|
||||
action=StoreBoolean,
|
||||
help="Use diffusion loss in training")
|
||||
|
||||
# DPO-specific
|
||||
parser.add_argument("--dpo-beta",
|
||||
type=float,
|
||||
default=RLArgs.dpo_beta,
|
||||
help="DPO regularization parameter (typically much larger than GRPO beta)")
|
||||
parser.add_argument("--dpo-ref-update-step",
|
||||
type=int,
|
||||
default=RLArgs.dpo_ref_update_step,
|
||||
help="Reference model update frequency for OnlineDPO")
|
||||
parser.add_argument("--dpo-label-smoothing",
|
||||
type=float,
|
||||
default=RLArgs.dpo_label_smoothing,
|
||||
help="Label smoothing for DPO loss")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class TrainingArgs(FastVideoArgs):
|
||||
"""
|
||||
@@ -752,6 +1017,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
num_height: int = 0
|
||||
num_width: int = 0
|
||||
num_frames: int = 0
|
||||
|
||||
# RL dataset configuration (for RL prompt datasets)
|
||||
rl_dataset_path: str = "" # Path to RL prompt dataset directory (defaults to data_path if not set)
|
||||
rl_dataset_type: str = "text" # "text" or "geneval"
|
||||
rl_num_image_per_prompt: int = 4 # k parameter for KRepeatSampler (num_image_per_prompt)
|
||||
|
||||
train_batch_size: int = 0
|
||||
num_latent_t: int = 0
|
||||
@@ -862,6 +1132,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
last_step_only: bool = False # Only use the last timestep for training
|
||||
context_noise: int = 0 # Context noise level for cache updates
|
||||
|
||||
# Nested RL configuration
|
||||
rl_args: RLArgs = dataclasses.field(default_factory=RLArgs)
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
provided_args = clean_cli_args(args)
|
||||
@@ -886,6 +1159,25 @@ class TrainingArgs(FastVideoArgs):
|
||||
kwargs[attr] = WorkloadType.from_string(
|
||||
workload_type_value) if isinstance(
|
||||
workload_type_value, str) else workload_type_value
|
||||
elif attr == 'rl_args':
|
||||
# Construct nested RLArgs from CLI arguments
|
||||
rl_kwargs = {}
|
||||
for rl_field in dataclasses.fields(RLArgs):
|
||||
rl_attr = rl_field.name
|
||||
if hasattr(args, rl_attr):
|
||||
value = getattr(args, rl_attr)
|
||||
# Special handling for reward_models: parse JSON string to dict
|
||||
if rl_attr == 'reward_models' and isinstance(value, str):
|
||||
rl_kwargs[rl_attr] = json.loads(value) if value else {}
|
||||
else:
|
||||
rl_kwargs[rl_attr] = value
|
||||
else:
|
||||
# Use default value from RLArgs
|
||||
if rl_field.default_factory is not dataclasses.MISSING:
|
||||
rl_kwargs[rl_attr] = rl_field.default_factory()
|
||||
elif rl_field.default is not dataclasses.MISSING:
|
||||
rl_kwargs[rl_attr] = rl_field.default
|
||||
kwargs[attr] = RLArgs(**rl_kwargs)
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
# Get the field to check its default value
|
||||
@@ -915,11 +1207,26 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--data-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to parquet files")
|
||||
help="Path to parquet files (or RL prompt dataset directory for RL training)")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
help="Number of workers for dataloader")
|
||||
|
||||
# RL dataset arguments (optional, defaults to data_path)
|
||||
parser.add_argument("--rl-dataset-path",
|
||||
type=str,
|
||||
default="",
|
||||
help="Path to RL prompt dataset directory (defaults to --data-path if not set)")
|
||||
parser.add_argument("--rl-dataset-type",
|
||||
type=str,
|
||||
default="text",
|
||||
choices=["text", "geneval"],
|
||||
help="RL dataset type: 'text' for TextPromptDataset or 'geneval' for GenevalPromptDataset")
|
||||
parser.add_argument("--rl-num-image-per-prompt",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Number of times to repeat each prompt (k parameter for KRepeatSampler)")
|
||||
parser.add_argument("--num-height",
|
||||
type=int,
|
||||
required=True,
|
||||
@@ -1284,6 +1591,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
|
||||
# RL (Reinforcement Learning) arguments
|
||||
RLArgs.add_cli_args(parser)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -153,6 +153,7 @@ def maybe_load_fsdp_model(
|
||||
f"Unexpected param or buffer {n} on meta device.")
|
||||
# Avoid unintended computation graph accumulation during inference
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
|
||||
p.requires_grad = False
|
||||
|
||||
compile_in_loader = enable_torch_compile and training_mode
|
||||
|
||||
@@ -201,8 +201,6 @@ class ComposedPipelineBase(ABC):
|
||||
# fwd, bwd, and other operations' precision.
|
||||
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
|
||||
pipe = cls(model_path,
|
||||
fastvideo_args,
|
||||
required_config_modules=required_config_modules,
|
||||
|
||||
@@ -67,6 +67,21 @@ class ForwardBatch:
|
||||
execution, allowing methods to update specific components without needing
|
||||
to manage numerous individual parameters.
|
||||
"""
|
||||
|
||||
@dataclass
|
||||
class RLData:
|
||||
"""RL-specific data collection options and outputs."""
|
||||
enabled: bool = False
|
||||
collect_log_probs: bool = True
|
||||
collect_kl: bool = False
|
||||
kl_reward: float = 0.0
|
||||
store_trajectory: bool = True
|
||||
keep_trajectory_on_cpu: bool = False
|
||||
log_probs: torch.Tensor | None = None
|
||||
kl: torch.Tensor | None = None
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_timesteps: torch.Tensor | None = None
|
||||
|
||||
# TODO(will): double check that args are separate from fastvideo_args
|
||||
# properly. Also maybe think about providing an abstraction for pipeline
|
||||
# specific arguments.
|
||||
@@ -197,6 +212,9 @@ class ForwardBatch:
|
||||
logging_info: PipelineLoggingInfo = field(
|
||||
default_factory=PipelineLoggingInfo)
|
||||
|
||||
# RL data collection
|
||||
rl_data: "ForwardBatch.RLData" = field(default_factory=RLData)
|
||||
|
||||
def __post_init__(self):
|
||||
"""Initialize dependent fields after dataclass initialization."""
|
||||
|
||||
@@ -267,6 +285,36 @@ class TrainingBatch:
|
||||
latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# RL/GRPO-specific attributes
|
||||
reward_scores: torch.Tensor | None = None # Computed rewards from reward models
|
||||
log_probs: torch.Tensor | None = None # Current policy log probabilities [B, num_steps] or [B]
|
||||
old_log_probs: torch.Tensor | None = None # Old policy log probs (for importance ratio) [B, num_steps] or [B]
|
||||
advantages: torch.Tensor | None = None # GAE advantages [B, num_steps] or [B]
|
||||
returns: torch.Tensor | None = None # TD returns (advantages + values) [B, num_steps] or [B]
|
||||
values: torch.Tensor | None = None # Value function predictions [B]
|
||||
old_values: torch.Tensor | None = None # Old value predictions (for clipping) [B]
|
||||
|
||||
# GRPO sampling-specific attributes
|
||||
kl: torch.Tensor | None = None # KL divergences from sampling [B, num_steps] (if kl_reward > 0)
|
||||
prompt_ids: torch.Tensor | None = None # Prompt token IDs for stat tracking [B, seq_len]
|
||||
prompt_embeds: torch.Tensor | None = None # Prompt embeddings used in sampling [B, seq_len, hidden_dim]
|
||||
negative_prompt_embeds: torch.Tensor | None = None # Negative prompt embeddings for CFG [B, seq_len, hidden_dim]
|
||||
|
||||
# RL loss components
|
||||
policy_loss: float = 0.0 # GRPO/PPO policy loss
|
||||
value_loss: float = 0.0 # Value function loss
|
||||
kl_divergence: float = 0.0 # KL(new_policy || old_policy)
|
||||
importance_ratio: float = 1.0 # exp(log_prob - old_log_prob)
|
||||
clip_fraction: float = 0.0 # Fraction of ratios that were clipped
|
||||
|
||||
# RL metrics
|
||||
advantage_mean: float = 0.0 # Mean advantage (should be ~0 after normalization)
|
||||
advantage_std: float = 1.0 # Std of advantages
|
||||
reward_mean: float = 0.0 # Mean reward across batch
|
||||
reward_std: float = 0.0 # Std of rewards
|
||||
value_mean: float = 0.0 # Mean value prediction
|
||||
entropy: float = 0.0 # Policy entropy (for exploration)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PreprocessBatch(ForwardBatch):
|
||||
|
||||
@@ -4,11 +4,14 @@ Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import math
|
||||
import weakref
|
||||
from collections.abc import Iterable
|
||||
from contextlib import nullcontext
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.attention import get_attn_backend
|
||||
@@ -52,6 +55,84 @@ except ImportError:
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def sde_step_with_logprob(
|
||||
scheduler,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: float | torch.FloatTensor,
|
||||
sample: torch.FloatTensor,
|
||||
prev_sample: torch.FloatTensor | None = None,
|
||||
generator: torch.Generator | None = None,
|
||||
deterministic: bool = False,
|
||||
return_pixel_log_prob: bool = False,
|
||||
return_dt_and_std_dev_t: bool = False
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, ...]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE and
|
||||
compute log probabilities for the transition.
|
||||
"""
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
if timestep.ndim == 0:
|
||||
timestep = timestep.unsqueeze(0)
|
||||
step_indices = [
|
||||
scheduler.index_for_timestep(t.item()) for t in timestep
|
||||
]
|
||||
else:
|
||||
step_indices = [scheduler.index_for_timestep(timestep)]
|
||||
|
||||
prev_step_indices = [step + 1 for step in step_indices]
|
||||
|
||||
sigmas = scheduler.sigmas.to(sample.device, sample.dtype)
|
||||
sigma = sigmas[step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_prev = sigmas[prev_step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_max = sigmas[0].item()
|
||||
sigma_min = sigmas[-1].item()
|
||||
|
||||
dt = sigma_prev - sigma
|
||||
|
||||
std_dev_t = sigma_min + (sigma_max - sigma_min) * sigma
|
||||
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 not None and generator is not None:
|
||||
raise ValueError(
|
||||
"Cannot pass both generator and prev_sample. Please make sure that either `generator` or"
|
||||
" `prev_sample` stays `None`.")
|
||||
|
||||
if prev_sample is None:
|
||||
variance_noise = randn_tensor(
|
||||
model_output.shape,
|
||||
generator=generator,
|
||||
device=model_output.device,
|
||||
dtype=model_output.dtype,
|
||||
)
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
prev_sample = prev_sample_mean + std_dev_t * sqrt_dt * variance_noise
|
||||
else:
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
if deterministic:
|
||||
prev_sample = sample + dt * model_output
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
if return_pixel_log_prob:
|
||||
raise NotImplementedError(
|
||||
"Pixel-level log prob is not supported in this helper.")
|
||||
|
||||
std_dev_sqrt_dt = std_dev_t * sqrt_dt
|
||||
log_prob = (
|
||||
-((prev_sample.detach() - prev_sample_mean)**2) /
|
||||
(2 *
|
||||
(std_dev_sqrt_dt**2)) - torch.log(std_dev_sqrt_dt + 1e-8) - torch.log(
|
||||
torch.sqrt(2 * torch.as_tensor(math.pi, device=sample.device))))
|
||||
|
||||
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
|
||||
|
||||
if return_dt_and_std_dev_t:
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt
|
||||
|
||||
|
||||
class DenoisingStage(PipelineStage):
|
||||
"""
|
||||
Stage for running the denoising loop in diffusion pipelines.
|
||||
@@ -203,9 +284,11 @@ class DenoisingStage(PipelineStage):
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
assert latent_model_input.shape[0] == 1, "only support batch size 1"
|
||||
rl_data = batch.rl_data if batch.rl_data and batch.rl_data.enabled else None
|
||||
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
assert latent_model_input.shape[
|
||||
0] == 1, "TI2V task only supports batch size 1"
|
||||
# TI2V directly replaces the first frame of the latent with
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
@@ -243,6 +326,12 @@ class DenoisingStage(PipelineStage):
|
||||
# Initialize lists for ODE trajectory
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
rl_timesteps: list[torch.Tensor] = []
|
||||
rl_latents: list[torch.Tensor] = []
|
||||
rl_log_probs: list[torch.Tensor] = []
|
||||
rl_kl: list[torch.Tensor] = []
|
||||
if rl_data is not None and rl_data.store_trajectory:
|
||||
rl_latents.append(latents)
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
@@ -329,6 +418,24 @@ class DenoisingStage(PipelineStage):
|
||||
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
def run_transformer(model, encoder_hidden_states, cond_kwargs,
|
||||
is_cfg_negative: bool):
|
||||
batch.is_cfg_negative = is_cfg_negative
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
):
|
||||
return model(
|
||||
latent_model_input,
|
||||
encoder_hidden_states,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
@@ -390,40 +497,13 @@ class DenoisingStage(PipelineStage):
|
||||
# support torch dynamo compilation. They pass in
|
||||
# attn_metadata, vllm_config, and num_tokens. We can pass in
|
||||
# fastvideo_args or training_args, and attn_metadata.
|
||||
batch.is_cfg_negative = False
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred = current_model(
|
||||
latent_model_input,
|
||||
prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
noise_pred = run_transformer(current_model, prompt_embeds,
|
||||
pos_cond_kwargs, False)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
batch.is_cfg_negative = True
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
forward_batch=batch,
|
||||
):
|
||||
noise_pred_uncond = current_model(
|
||||
latent_model_input,
|
||||
neg_prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
**neg_cond_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
noise_pred_uncond = run_transformer(
|
||||
current_model, neg_prompt_embeds, neg_cond_kwargs,
|
||||
True)
|
||||
|
||||
noise_pred_text = noise_pred
|
||||
noise_pred = noise_pred_uncond + current_guidance_scale * (
|
||||
@@ -438,11 +518,58 @@ class DenoisingStage(PipelineStage):
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
)
|
||||
# Compute the previous noisy sample
|
||||
prev_latents = latents
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
if rl_data is not None:
|
||||
if rl_data.collect_log_probs:
|
||||
_, log_prob, prev_latents_mean, std_dev_t, _ = sde_step_with_logprob(
|
||||
self.scheduler,
|
||||
noise_pred.float(),
|
||||
t,
|
||||
prev_latents.float(),
|
||||
prev_sample=latents.float(),
|
||||
deterministic=False,
|
||||
return_dt_and_std_dev_t=True,
|
||||
)
|
||||
rl_log_probs.append(log_prob)
|
||||
|
||||
if rl_data.collect_kl and rl_data.kl_reward > 0:
|
||||
adapter_ctx = nullcontext()
|
||||
if hasattr(current_model, "disable_adapter"):
|
||||
adapter_ctx = current_model.disable_adapter()
|
||||
with adapter_ctx:
|
||||
noise_pred_ref = run_transformer(
|
||||
current_model, prompt_embeds,
|
||||
pos_cond_kwargs, False)
|
||||
if batch.do_classifier_free_guidance:
|
||||
noise_pred_uncond_ref = run_transformer(
|
||||
current_model, neg_prompt_embeds,
|
||||
neg_cond_kwargs, True)
|
||||
noise_pred_text_ref = noise_pred_ref
|
||||
noise_pred_ref = noise_pred_uncond_ref + current_guidance_scale * (
|
||||
noise_pred_text_ref -
|
||||
noise_pred_uncond_ref)
|
||||
_, _, prev_latents_mean_ref, std_dev_t_ref, _ = sde_step_with_logprob(
|
||||
self.scheduler,
|
||||
noise_pred_ref.float(),
|
||||
t,
|
||||
prev_latents.float(),
|
||||
prev_sample=latents.float(),
|
||||
deterministic=False,
|
||||
return_dt_and_std_dev_t=True,
|
||||
)
|
||||
if not torch.allclose(std_dev_t, std_dev_t_ref):
|
||||
logger.warning(
|
||||
"std_dev_t mismatch in RL KL computation at step %s",
|
||||
i)
|
||||
kl = (prev_latents_mean -
|
||||
prev_latents_mean_ref)**2 / (2 * std_dev_t**2)
|
||||
kl = kl.mean(dim=tuple(range(1, kl.ndim)))
|
||||
rl_kl.append(kl)
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
latents = latents.squeeze(0)
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
@@ -452,6 +579,15 @@ class DenoisingStage(PipelineStage):
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(t)
|
||||
trajectory_latents.append(latents)
|
||||
if rl_data is not None:
|
||||
rl_timesteps.append(t)
|
||||
if rl_data.store_trajectory:
|
||||
rl_latents.append(latents)
|
||||
if rl_data.collect_kl and rl_data.kl_reward <= 0:
|
||||
rl_kl.append(
|
||||
torch.zeros(latents.shape[0],
|
||||
device=latents.device,
|
||||
dtype=latents.dtype))
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
@@ -472,6 +608,25 @@ class DenoisingStage(PipelineStage):
|
||||
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
|
||||
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
|
||||
batch.trajectory_latents = trajectory_tensor.cpu()
|
||||
if rl_data is not None:
|
||||
if rl_timesteps:
|
||||
rl_data.trajectory_timesteps = torch.stack(rl_timesteps, dim=0)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.trajectory_timesteps = rl_data.trajectory_timesteps.cpu(
|
||||
)
|
||||
if rl_data.store_trajectory and rl_latents:
|
||||
rl_data.trajectory_latents = torch.stack(rl_latents, dim=1)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.trajectory_latents = rl_data.trajectory_latents.cpu(
|
||||
)
|
||||
if rl_log_probs:
|
||||
rl_data.log_probs = torch.stack(rl_log_probs, dim=1)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.log_probs = rl_data.log_probs.cpu()
|
||||
if rl_kl:
|
||||
rl_data.kl = torch.stack(rl_kl, dim=1)
|
||||
if rl_data.keep_trajectory_on_cpu:
|
||||
rl_data.kl = rl_data.kl.cpu()
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
@@ -17,10 +17,11 @@ from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
os.environ["MASTER_PORT"] = "29701"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
@@ -121,4 +122,24 @@ def test_wan_transformer():
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert_close(output1, output2, atol=1e-1, rtol=1e-2)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from fastvideo.distributed import (
|
||||
cleanup_dist_env_and_memory,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
)
|
||||
# Allow running this test file directly without pytest.
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
try:
|
||||
test_wan_transformer()
|
||||
logger.info("test_wan_transformer finished successfully.")
|
||||
finally:
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -1,5 +1,12 @@
|
||||
from .distillation_pipeline import DistillationPipeline
|
||||
from .training_pipeline import TrainingPipeline
|
||||
from .wan_training_pipeline import WanTrainingPipeline
|
||||
from fastvideo.training.rl import RLPipeline, create_rl_pipeline
|
||||
|
||||
__all__ = ["TrainingPipeline", "WanTrainingPipeline", "DistillationPipeline"]
|
||||
__all__ = [
|
||||
"TrainingPipeline",
|
||||
"WanTrainingPipeline",
|
||||
"DistillationPipeline",
|
||||
"RLPipeline",
|
||||
"create_rl_pipeline",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
from .rl_pipeline import RLPipeline, create_rl_pipeline
|
||||
|
||||
__all__ = [
|
||||
"RLPipeline",
|
||||
"create_rl_pipeline",
|
||||
]
|
||||
@@ -0,0 +1,11 @@
|
||||
from .rewards import (
|
||||
create_reward_models,
|
||||
MultiRewardAggregator,
|
||||
ValueModel
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"create_reward_models",
|
||||
"MultiRewardAggregator",
|
||||
"ValueModel",
|
||||
]
|
||||
@@ -0,0 +1,63 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Abstract base class for VIDEO reward models.
|
||||
|
||||
All VIDEO reward models should inherit from this class and implement
|
||||
the compute_reward() method.
|
||||
|
||||
IMPORTANT: Reward models must process FULL VIDEO SEQUENCES, not individual frames.
|
||||
Input shape is [B, T, C, H, W] where T is the temporal (frame) dimension.
|
||||
|
||||
For video-specific rewards, consider:
|
||||
- Temporal coherence across frames
|
||||
- Motion quality and smoothness
|
||||
- Video-text alignment (not just frame-text)
|
||||
- Multi-frame aesthetic quality
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
class BaseRewardModel(ABC, nn.Module):
|
||||
def __init__(self, model_path: str | None = None, device: str = "cuda"):
|
||||
super().__init__()
|
||||
self.model_path = model_path
|
||||
self.device = device
|
||||
|
||||
@abstractmethod
|
||||
def compute_reward(
|
||||
self,
|
||||
videos: torch.Tensor, # [B, T, C, H, W] decoded video sequences
|
||||
prompts: list[str] | None, # Text prompts
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute rewards for generated VIDEO sequences.
|
||||
|
||||
IMPORTANT: This method must process the FULL temporal sequence [B, T, C, H, W].
|
||||
Do NOT evaluate individual frames independently and average.
|
||||
|
||||
Args:
|
||||
videos: Decoded video tensors [B, T, C, H, W] in range [0, 1]
|
||||
B = batch size
|
||||
T = number of frames (temporal dimension)
|
||||
C = channels (typically 3 for RGB)
|
||||
H, W = height, width
|
||||
prompts: List of text prompts (length B) describing each video
|
||||
**kwargs: Additional model-specific arguments
|
||||
|
||||
Returns:
|
||||
rewards: Tensor of shape [B] with reward scores for each video sequence
|
||||
|
||||
Example:
|
||||
>>> videos = torch.rand(4, 17, 3, 256, 256) # 4 videos, 17 frames each
|
||||
>>> prompts = ["A cat jumping", "A dog running", ...]
|
||||
>>> rewards = model.compute_reward(videos, prompts)
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement compute_reward()")
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(model_path={self.model_path})"
|
||||
@@ -0,0 +1,206 @@
|
||||
from paddleocr import PaddleOCR
|
||||
import torch
|
||||
import numpy as np
|
||||
from Levenshtein import distance
|
||||
from typing import Any
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.training.rl.rewards.base import BaseRewardModel
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class OcrScorerVideo(BaseRewardModel):
|
||||
"""
|
||||
OCR reward model for multi-frame video OCR evaluation.
|
||||
|
||||
This model evaluates multiple frames across the video sequence,
|
||||
sampling frames at a specified interval and averaging the OCR scores.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
model_path: str | None = None,
|
||||
device: str = "cpu",
|
||||
frame_interval: int = 4):
|
||||
"""
|
||||
OCR reward calculator for videos
|
||||
|
||||
Args:
|
||||
model_path: Not used for PaddleOCR (kept for BaseRewardModel compatibility)
|
||||
device: Device string (used to determine use_gpu if not explicitly set)
|
||||
frame_interval: Sample every Nth frame (default: 4)
|
||||
"""
|
||||
super().__init__(model_path=model_path, device=device)
|
||||
|
||||
self.frame_interval = frame_interval
|
||||
self.ocr = PaddleOCR(
|
||||
use_angle_cls=False,
|
||||
lang="en",
|
||||
use_gpu=False,
|
||||
show_log=False # Disable unnecessary log output
|
||||
)
|
||||
|
||||
logger.info("Initialized OcrScorerVideo (device=%s, frame_interval=%d)",
|
||||
device, frame_interval)
|
||||
|
||||
def _process_single_video(self, video_tensor: torch.Tensor,
|
||||
prompt: str) -> float:
|
||||
"""
|
||||
Process a single video tensor and return its OCR reward.
|
||||
|
||||
Args:
|
||||
video_tensor: Video tensor of shape [C, T, H, W]
|
||||
prompt: Text prompt containing target OCR text in quotes
|
||||
|
||||
Returns:
|
||||
Average reward across positive-scoring frames
|
||||
"""
|
||||
# Extract target text from prompt
|
||||
try:
|
||||
target_text = prompt.split('"')[1].replace(' ', '').lower()
|
||||
except IndexError:
|
||||
logger.warning("Failed to extract quoted text from prompt: %s",
|
||||
prompt)
|
||||
target_text = prompt.replace(' ', '').lower()
|
||||
|
||||
if not target_text:
|
||||
return 0.0
|
||||
|
||||
# video_tensor is [C, T, H, W]
|
||||
C, T, H, W = video_tensor.shape
|
||||
|
||||
# Convert to numpy and move to CPU if needed
|
||||
video_np = video_tensor.detach().cpu().numpy()
|
||||
|
||||
# Convert from [C, T, H, W] to [T, H, W, C] for easier frame extraction
|
||||
video_np = np.transpose(video_np, (1, 2, 3, 0)) # [T, H, W, C]
|
||||
logger.info(f"in ocr 1.5, video_np[0][0]: {video_np[0][0]}")
|
||||
|
||||
# Normalize to [0, 255] uint8 if needed
|
||||
if video_np.max() <= 1.0:
|
||||
video_np = (video_np * 255).astype(np.uint8)
|
||||
else:
|
||||
video_np = video_np.astype(np.uint8)
|
||||
|
||||
frame_rewards = []
|
||||
|
||||
# Sample frames at specified interval
|
||||
for frame_idx in range(0, T, self.frame_interval):
|
||||
frame = video_np[frame_idx] # [H, W, C]
|
||||
logger.info(f"in ocr 2, frame.shape: {frame.shape}")
|
||||
# Run OCR
|
||||
try:
|
||||
result = self.ocr.ocr(frame, cls=False)
|
||||
logger.info(f"in ocr 3, result: {result}")
|
||||
if result and result[0]:
|
||||
recognized_text = "".join(
|
||||
[line[1][0] for line in result[0] if line[1][1] > 0])
|
||||
else:
|
||||
recognized_text = ""
|
||||
except Exception as e:
|
||||
logger.info("OCR failed on frame %d: %s", frame_idx, str(e))
|
||||
recognized_text = ''
|
||||
|
||||
logger.info(f"in ocr 4, recognized_text: {recognized_text}")
|
||||
|
||||
recognized_text = recognized_text.replace(' ', '').lower()
|
||||
if target_text in recognized_text:
|
||||
dist = 0
|
||||
else:
|
||||
dist = distance(recognized_text, target_text)
|
||||
dist = min(dist, len(target_text))
|
||||
reward = 1.0 - dist / len(target_text)
|
||||
|
||||
logger.info(f"in ocr 5, reward: {reward}")
|
||||
if reward > 0:
|
||||
frame_rewards.append(reward)
|
||||
|
||||
logger.info(f"in ocr 6, frame_rewards: {frame_rewards}")
|
||||
|
||||
return sum([reward / len(frame_rewards)
|
||||
for reward in frame_rewards]) if frame_rewards else 0.0
|
||||
|
||||
@torch.no_grad()
|
||||
def compute_reward(self, videos: torch.Tensor, prompts: list[str],
|
||||
**kwargs: Any) -> torch.Tensor:
|
||||
"""
|
||||
Calculate OCR reward by evaluating sampled frames across the video.
|
||||
|
||||
Args:
|
||||
videos: Video tensor of shape [B, C, T, H, W]
|
||||
B = batch size
|
||||
C = channels (typically 3 for RGB)
|
||||
T = number of frames (temporal dimension)
|
||||
H, W = height, width
|
||||
prompts: List of text prompts containing target OCR text in quotes (length B)
|
||||
**kwargs: Additional arguments
|
||||
|
||||
Returns:
|
||||
Reward tensor [B] with averaged OCR similarity scores across frames
|
||||
"""
|
||||
# Ensure videos is a torch tensor with correct shape
|
||||
assert isinstance(
|
||||
videos,
|
||||
torch.Tensor), f"videos must be torch.Tensor, got {type(videos)}"
|
||||
assert videos.ndim == 5, f"videos must have 5 dimensions [B, C, T, H, W], got shape {videos.shape}"
|
||||
|
||||
logger.info(f"in ocr 1, videos.shape: {videos.shape}")
|
||||
|
||||
B, C, T, H, W = videos.shape
|
||||
assert len(
|
||||
prompts
|
||||
) == B, f"Number of prompts ({len(prompts)}) must match batch size ({B})"
|
||||
|
||||
rewards = []
|
||||
for b in range(B):
|
||||
# Extract single video: [C, T, H, W]
|
||||
video = videos[b]
|
||||
reward = self._process_single_video(video, prompts[b])
|
||||
rewards.append(reward)
|
||||
|
||||
logger.info(f"in ocr 7, rewards: {rewards}")
|
||||
|
||||
rewards = torch.tensor(rewards, dtype=torch.float32, device=self.device)
|
||||
|
||||
logger.info(f"in ocr 8, rewards: {rewards}")
|
||||
|
||||
# Check for NaN or Inf values
|
||||
if torch.isnan(rewards).any() or torch.isinf(rewards).any():
|
||||
logger.warning(
|
||||
"NaN or Inf detected in OCR rewards, returning zero tensor")
|
||||
return torch.zeros_like(rewards)
|
||||
|
||||
return rewards
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
example_image_path = "flowgrpo_cmd.png"
|
||||
example_image = Image.open(example_image_path)
|
||||
example_prompt = '/f1ow_grpo$'
|
||||
|
||||
# Convert image to RGB if needed
|
||||
if example_image.mode != 'RGB':
|
||||
example_image = example_image.convert('RGB')
|
||||
|
||||
# Convert PIL Image to numpy array [H, W, C]
|
||||
image_np = np.array(example_image)
|
||||
|
||||
# Normalize to [0, 1] range and convert to float32
|
||||
image_np = image_np.astype(np.float32) / 255.0
|
||||
|
||||
# Convert to torch tensor and reshape: [H, W, C] -> [C, H, W]
|
||||
image_tensor = torch.from_numpy(image_np).permute(2, 0, 1)
|
||||
|
||||
# Add temporal dimension: [C, H, W] -> [C, T, H, W] where T=1
|
||||
video_tensor = image_tensor.unsqueeze(1) # [C, 1, H, W]
|
||||
|
||||
# Add batch dimension: [C, T, H, W] -> [B, C, T, H, W] where B=1
|
||||
video_tensor = video_tensor.unsqueeze(0) # [1, C, 1, H, W]
|
||||
|
||||
# Instantiate scorer
|
||||
scorer = OcrScorerVideo(device="cpu")
|
||||
|
||||
# Call compute_reward method with video tensor
|
||||
reward = scorer.compute_reward(video_tensor, [example_prompt])
|
||||
print(f"OCR Reward: {reward.item()}")
|
||||
@@ -0,0 +1,338 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Base infrastructure for VIDEO reward models in RL/GRPO training.
|
||||
|
||||
IMPORTANT: This module is designed exclusively for VIDEO generation models.
|
||||
All reward models must operate on video sequences [B, T, C, H, W], not single frames.
|
||||
|
||||
This module provides:
|
||||
1. Multi-reward aggregation for video
|
||||
2. Value model wrapper
|
||||
3. Integration with FastVideo video generation infrastructure
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.training.rl.rewards.ocr import OcrScorerVideo
|
||||
from fastvideo.training.rl.rewards.base import BaseRewardModel
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MultiRewardAggregator(nn.Module):
|
||||
"""
|
||||
Aggregates multiple reward models with configurable weights.
|
||||
|
||||
This implements the multi-reward aggregation strategy from flow_grpo,
|
||||
allowing combination of different reward signals (aesthetic quality,
|
||||
text-video alignment, compositional understanding, etc.)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reward_models: list[BaseRewardModel],
|
||||
reward_weights: list[float] | None = None,
|
||||
normalize_rewards: bool = True
|
||||
):
|
||||
"""
|
||||
Initialize multi-reward aggregator.
|
||||
|
||||
Args:
|
||||
reward_models: List of reward model instances
|
||||
reward_weights: Weights for each reward model (default: uniform)
|
||||
normalize_rewards: Whether to normalize rewards before aggregation
|
||||
"""
|
||||
super().__init__()
|
||||
self.reward_models = nn.ModuleList(reward_models)
|
||||
|
||||
if reward_weights is None:
|
||||
reward_weights = [1.0 / len(reward_models)] * len(reward_models)
|
||||
|
||||
assert len(reward_weights) == len(reward_models), \
|
||||
f"Number of weights ({len(reward_weights)}) must match number of models ({len(reward_models)})"
|
||||
|
||||
assert abs(sum(reward_weights) - 1.0) < 1e-6, \
|
||||
f"Reward weights must sum to 1.0, got {sum(reward_weights)}"
|
||||
|
||||
self.reward_weights = reward_weights
|
||||
self.normalize_rewards = normalize_rewards
|
||||
|
||||
logger.info(
|
||||
"Initialized MultiRewardAggregator with %d models: %s",
|
||||
len(reward_models),
|
||||
[(type(m).__name__, w) for m, w in zip(reward_models, reward_weights, strict=False)]
|
||||
)
|
||||
|
||||
def compute_reward(
|
||||
self,
|
||||
videos: torch.Tensor,
|
||||
prompts: list[str],
|
||||
return_individual: bool = False,
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor | dict[str, torch.Tensor]:
|
||||
"""
|
||||
Compute aggregated reward from multiple models.
|
||||
|
||||
Args:
|
||||
videos: Decoded video tensors [B, C, T, H, W]
|
||||
prompts: List of text prompts
|
||||
return_individual: If True, return dict with individual rewards
|
||||
**kwargs: Additional arguments passed to reward models
|
||||
|
||||
Returns:
|
||||
If return_individual=False: aggregated_rewards [B]
|
||||
If return_individual=True: dict with "aggregated" and individual model rewards
|
||||
"""
|
||||
batch_size = videos.shape[0]
|
||||
individual_rewards: dict[str, torch.Tensor] = {}
|
||||
|
||||
# Collect rewards from all models
|
||||
all_rewards = []
|
||||
for i, (model, weight) in enumerate(zip(self.reward_models, self.reward_weights, strict=False)):
|
||||
reward = model.compute_reward(videos, prompts, **kwargs)
|
||||
assert reward.shape == (batch_size,), \
|
||||
f"Reward model {i} returned shape {reward.shape}, expected ({batch_size},)"
|
||||
|
||||
# Optionally normalize individual rewards
|
||||
if self.normalize_rewards:
|
||||
reward = (reward - reward.mean()) / (reward.std() + 1e-8)
|
||||
|
||||
individual_rewards[f"reward_{type(model).__name__}"] = reward
|
||||
all_rewards.append(weight * reward)
|
||||
|
||||
# Aggregate with weights
|
||||
aggregated = sum(all_rewards)
|
||||
|
||||
if return_individual:
|
||||
individual_rewards["aggregated"] = aggregated
|
||||
return individual_rewards
|
||||
|
||||
return aggregated
|
||||
|
||||
def __repr__(self) -> str:
|
||||
models_str = ", ".join([
|
||||
f"{type(m).__name__}(w={w:.3f})"
|
||||
for m, w in zip(self.reward_models, self.reward_weights, strict=False)
|
||||
])
|
||||
return f"MultiRewardAggregator({models_str})"
|
||||
|
||||
|
||||
class ValueModel(nn.Module):
|
||||
"""
|
||||
Value function model wrapper for RL training.
|
||||
|
||||
The value model can either:
|
||||
1. Share the transformer backbone with the policy (memory efficient)
|
||||
2. Use a separate transformer (more flexible)
|
||||
|
||||
For now, this is a placeholder that will be expanded based on
|
||||
the chosen architecture strategy.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer: nn.Module,
|
||||
share_backbone: bool = False,
|
||||
hidden_size: int | None = None
|
||||
):
|
||||
"""
|
||||
Initialize value model.
|
||||
|
||||
Args:
|
||||
transformer: Transformer model (policy or separate)
|
||||
share_backbone: Whether to share backbone with policy
|
||||
hidden_size: Hidden size for value head (inferred if None)
|
||||
"""
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.share_backbone = share_backbone
|
||||
|
||||
# Value head will be added later based on transformer architecture
|
||||
# For now, just store the transformer reference
|
||||
logger.info(
|
||||
"Initialized ValueModel (share_backbone=%s)",
|
||||
share_backbone
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Forward pass to compute value predictions.
|
||||
|
||||
Args:
|
||||
hidden_states: Latent states [B, C, T, H, W]
|
||||
encoder_hidden_states: Text embeddings [B, L, D]
|
||||
timestep: Timesteps [B]
|
||||
**kwargs: Additional transformer arguments
|
||||
|
||||
Returns:
|
||||
values: Value predictions [B]
|
||||
"""
|
||||
# TODO: Implement value prediction
|
||||
# For now, return dummy values
|
||||
batch_size = hidden_states.shape[0]
|
||||
return torch.zeros(batch_size, device=hidden_states.device)
|
||||
|
||||
|
||||
class DummyRewardModel(BaseRewardModel):
|
||||
"""
|
||||
Dummy VIDEO reward model for testing and development.
|
||||
|
||||
Returns random rewards in the range [0, 1] for VIDEO inputs.
|
||||
This is a placeholder for testing the RL pipeline before real video reward models
|
||||
are implemented.
|
||||
|
||||
NOTE: This does NOT actually evaluate video quality - it's just for testing!
|
||||
"""
|
||||
|
||||
def __init__(self, mean: float = 0.5, std: float = 0.1):
|
||||
super().__init__(model_path=None)
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
logger.info("Initialized DummyRewardModel (VIDEO) - mean=%.2f, std=%.2f", mean, std)
|
||||
logger.warning(
|
||||
"DummyRewardModel is for TESTING ONLY - does not evaluate actual video quality!"
|
||||
)
|
||||
|
||||
def compute_reward(
|
||||
self,
|
||||
videos: torch.Tensor, # [B, T, C, H, W]
|
||||
prompts: list[str],
|
||||
**kwargs: Any
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Return random rewards for testing.
|
||||
|
||||
Args:
|
||||
videos: Video sequences [B, T, C, H, W]
|
||||
prompts: Text prompts
|
||||
|
||||
Returns:
|
||||
Random rewards [B] in range [0, 1]
|
||||
"""
|
||||
batch_size = videos.shape[0]
|
||||
num_frames = videos.shape[1]
|
||||
|
||||
logger.debug(
|
||||
"DummyRewardModel processing %d videos with %d frames each",
|
||||
batch_size,
|
||||
num_frames
|
||||
)
|
||||
|
||||
# Generate random rewards (not based on actual video content!)
|
||||
rewards = torch.randn(batch_size, device=videos.device) * self.std + self.mean
|
||||
return rewards.clamp(0.0, 1.0)
|
||||
|
||||
def load_model(self) -> None:
|
||||
"""No model to load for dummy."""
|
||||
pass
|
||||
|
||||
|
||||
def create_reward_models(
|
||||
reward_models: dict,
|
||||
device: str = "cuda"
|
||||
) -> MultiRewardAggregator:
|
||||
"""
|
||||
Factory function to create VIDEO reward models from configuration strings.
|
||||
|
||||
IMPORTANT: Only creates VIDEO reward models. Image-only reward models
|
||||
(PickScore, ImageReward, GenEval, etc.) are NOT supported.
|
||||
|
||||
Args:
|
||||
reward_models: dictionary of reward model names to weights
|
||||
Example: {"paddle_ocr": 0.5, "video_score": 0.5}
|
||||
device: Device to load models on
|
||||
|
||||
Returns:
|
||||
MultiRewardAggregator with loaded VIDEO reward models
|
||||
|
||||
Supported VIDEO Reward Types:
|
||||
- "paddle_ocr": PaddleOCR multi-frame video text recognition
|
||||
- "video_score": Video aesthetic quality (multi-frame) - TODO
|
||||
- "video_text_alignment": CLIP-based video-text similarity - TODO
|
||||
- "temporal_coherence": Frame-to-frame consistency - TODO
|
||||
- "motion_quality": Motion smoothness and realism - TODO
|
||||
- "dummy": Random rewards for testing (VIDEO-aware)
|
||||
|
||||
NOT Supported (Image-Only):
|
||||
- "pickscore": Image aesthetic (use "video_score" instead)
|
||||
- "imagereward": Image quality (use "video_score" instead)
|
||||
- "geneval": Image compositional (no video equivalent yet)
|
||||
- Any single-frame reward models
|
||||
|
||||
Example:
|
||||
>>> models = create_reward_models(
|
||||
... reward_models={
|
||||
... "paddle_ocr": 0.5,
|
||||
... "video_text_alignment": 0.5
|
||||
... },
|
||||
... device="cuda"
|
||||
... )
|
||||
"""
|
||||
|
||||
|
||||
assert reward_models, "No reward models specified. Please select at least 1 reward model"
|
||||
|
||||
types = [t.strip() for t in reward_models.keys()]
|
||||
weights = list(reward_models.values())
|
||||
|
||||
assert len(types) == len(weights), \
|
||||
f"Number of models ({len(types)}) must match number of weights ({len(weights)})"
|
||||
|
||||
# Create reward models based on types
|
||||
models_list: list[BaseRewardModel] = []
|
||||
for reward_type in types:
|
||||
if reward_type == "dummy":
|
||||
model = DummyRewardModel()
|
||||
|
||||
elif reward_type == "paddle_ocr":
|
||||
logger.info("Creating PaddleOCR reward model")
|
||||
model = OcrScorerVideo(device=device)
|
||||
|
||||
elif reward_type == "video_score":
|
||||
# TODO: Implement VideoScore reward model (Phase 2)
|
||||
logger.warning(
|
||||
"VideoScore reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
elif reward_type == "video_text_alignment":
|
||||
# TODO: Implement VideoTextAlignment reward model (Phase 2)
|
||||
logger.warning(
|
||||
"VideoTextAlignment reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
elif reward_type == "temporal_coherence":
|
||||
# TODO: Implement TemporalCoherence reward model (Phase 2)
|
||||
logger.warning(
|
||||
"TemporalCoherence reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
elif reward_type == "motion_quality":
|
||||
# TODO: Implement MotionQuality reward model (Phase 2)
|
||||
logger.warning(
|
||||
"MotionQuality reward not implemented yet, using DummyRewardModel"
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
else:
|
||||
logger.warning(
|
||||
"Unknown VIDEO reward type '%s', using DummyRewardModel",
|
||||
reward_type
|
||||
)
|
||||
model = DummyRewardModel()
|
||||
|
||||
models_list.append(model)
|
||||
|
||||
logger.info(
|
||||
"Created MultiRewardAggregator with %d VIDEO reward models",
|
||||
len(models_list)
|
||||
)
|
||||
|
||||
return MultiRewardAggregator(models_list, weights, normalize_rewards=True)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,385 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Utility functions for RL/GRPO training.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def compute_gae(
|
||||
rewards: torch.Tensor,
|
||||
values: torch.Tensor,
|
||||
next_values: torch.Tensor,
|
||||
dones: torch.Tensor | None = None,
|
||||
gamma: float = 0.99,
|
||||
lambda_: float = 0.95
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Compute Generalized Advantage Estimation (GAE-lambda).
|
||||
|
||||
GAE reduces variance in advantage estimation while allowing some bias.
|
||||
This is a key component of modern policy gradient methods like PPO and GRPO.
|
||||
|
||||
Args:
|
||||
rewards: Rewards at each step [B, T] or [B]
|
||||
values: Value predictions at each step [B, T] or [B]
|
||||
next_values: Value predictions at next step [B, T] or [B]
|
||||
dones: Episode termination flags [B, T] or [B] (1 if done, 0 otherwise)
|
||||
gamma: Discount factor
|
||||
lambda_: GAE lambda parameter (0=TD(0), 1=Monte Carlo)
|
||||
|
||||
Returns:
|
||||
advantages: GAE advantages [B, T] or [B]
|
||||
returns: TD(lambda) returns [B, T] or [B]
|
||||
|
||||
Reference:
|
||||
Schulman et al. "High-Dimensional Continuous Control Using Generalized Advantage Estimation"
|
||||
https://arxiv.org/abs/1506.02438
|
||||
"""
|
||||
if dones is None:
|
||||
dones = torch.zeros_like(rewards)
|
||||
|
||||
# Compute TD residuals: delta_t = r_t + gamma * V(s_{t+1}) - V(s_t)
|
||||
deltas = rewards + gamma * next_values * (1.0 - dones) - values
|
||||
|
||||
# If single step (no time dimension), return directly
|
||||
if deltas.dim() == 1:
|
||||
advantages = deltas
|
||||
returns = advantages + values
|
||||
return advantages, returns
|
||||
|
||||
# Multi-step: compute GAE recursively
|
||||
batch_size, num_steps = deltas.shape
|
||||
advantages = torch.zeros_like(deltas)
|
||||
gae = torch.zeros(batch_size, device=deltas.device)
|
||||
|
||||
# Backward pass to compute GAE
|
||||
for t in reversed(range(num_steps)):
|
||||
gae = deltas[:, t] + gamma * lambda_ * (1.0 - dones[:, t]) * gae
|
||||
advantages[:, t] = gae
|
||||
|
||||
# Returns are advantages + values
|
||||
returns = advantages + values
|
||||
|
||||
return advantages, returns
|
||||
|
||||
|
||||
def normalize_advantages(
|
||||
advantages: torch.Tensor,
|
||||
epsilon: float = 1e-8
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Normalize advantages to have zero mean and unit variance.
|
||||
|
||||
This is a common practice in PPO and GRPO to stabilize training.
|
||||
|
||||
Args:
|
||||
advantages: Raw advantages [B, ...]
|
||||
epsilon: Small constant for numerical stability
|
||||
|
||||
Returns:
|
||||
normalized_advantages: Normalized advantages [B, ...]
|
||||
"""
|
||||
mean = advantages.mean()
|
||||
std = advantages.std()
|
||||
return (advantages - mean) / (std + epsilon)
|
||||
|
||||
#TODO(jiali): refactor into algorithm
|
||||
def compute_grpo_policy_loss(
|
||||
log_probs: torch.Tensor,
|
||||
old_log_probs: torch.Tensor,
|
||||
advantages: torch.Tensor,
|
||||
clip_range: float = 0.2,
|
||||
use_ratio_norm: bool = True,
|
||||
max_importance_ratio: float = 10.0
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""
|
||||
Compute GRPO policy loss with importance sampling and clipping.
|
||||
|
||||
This implements the core GRPO objective with safety mechanisms from GRPO-Guard:
|
||||
- Importance ratio clipping (PPO-style)
|
||||
- RatioNorm correction (GRPO-Guard)
|
||||
- Ratio clamping for extreme values
|
||||
|
||||
Args:
|
||||
log_probs: Log probabilities from current policy [B]
|
||||
old_log_probs: Log probabilities from old policy [B]
|
||||
advantages: Advantages [B]
|
||||
clip_range: Clipping range for importance ratios
|
||||
use_ratio_norm: Apply RatioNorm correction (GRPO-Guard)
|
||||
max_importance_ratio: Maximum importance ratio before clamping
|
||||
|
||||
Returns:
|
||||
loss: Policy loss (scalar)
|
||||
info: Dictionary with diagnostic information
|
||||
|
||||
Reference:
|
||||
- PPO: Schulman et al. "Proximal Policy Optimization Algorithms"
|
||||
- GRPO-Guard: RatioNorm and gradient reweighting
|
||||
"""
|
||||
# Compute importance ratio: r_t = pi_new(a|s) / pi_old(a|s)
|
||||
log_ratio = log_probs - old_log_probs
|
||||
ratio = torch.exp(log_ratio)
|
||||
|
||||
# Clamp extreme ratios for numerical stability
|
||||
ratio = torch.clamp(ratio, 1.0 / max_importance_ratio, max_importance_ratio)
|
||||
|
||||
# RatioNorm correction (GRPO-Guard)
|
||||
# Corrects bias in importance sampling when ratio >> 1
|
||||
if use_ratio_norm:
|
||||
ratio_mean = ratio.mean()
|
||||
ratio = ratio / (ratio_mean + 1e-8)
|
||||
|
||||
# Clipped surrogate objective
|
||||
ratio_clipped = torch.clamp(ratio, 1.0 - clip_range, 1.0 + clip_range)
|
||||
surrogate1 = ratio * advantages
|
||||
surrogate2 = ratio_clipped * advantages
|
||||
policy_loss = -torch.min(surrogate1, surrogate2).mean()
|
||||
|
||||
# Compute diagnostics
|
||||
with torch.no_grad():
|
||||
# Clip fraction: how often ratios were clipped
|
||||
clip_fraction = ((ratio < 1.0 - clip_range) | (ratio > 1.0 + clip_range)).float().mean()
|
||||
|
||||
# KL divergence (approximate)
|
||||
kl_div = log_ratio.mean()
|
||||
|
||||
# Importance ratio stats
|
||||
importance_ratio_mean = ratio.mean()
|
||||
importance_ratio_std = ratio.std()
|
||||
|
||||
info = {
|
||||
"policy_loss": policy_loss.item(),
|
||||
"clip_fraction": clip_fraction.item(),
|
||||
"kl_divergence": kl_div.item(),
|
||||
"importance_ratio_mean": importance_ratio_mean.item(),
|
||||
"importance_ratio_std": importance_ratio_std.item(),
|
||||
}
|
||||
|
||||
return policy_loss, info
|
||||
|
||||
|
||||
def compute_value_loss(
|
||||
values: torch.Tensor,
|
||||
returns: torch.Tensor,
|
||||
old_values: torch.Tensor | None = None,
|
||||
clip_range: float = 0.2,
|
||||
use_clipping: bool = True
|
||||
) -> tuple[torch.Tensor, dict[str, Any]]:
|
||||
"""
|
||||
Compute value function loss with optional clipping.
|
||||
|
||||
Args:
|
||||
values: Value predictions from current model [B]
|
||||
returns: Target returns (from GAE) [B]
|
||||
old_values: Value predictions from old model [B] (for clipping)
|
||||
clip_range: Clipping range for value updates
|
||||
use_clipping: Whether to use clipped value loss (PPO-style)
|
||||
|
||||
Returns:
|
||||
loss: Value loss (scalar)
|
||||
info: Dictionary with diagnostic information
|
||||
"""
|
||||
# Standard MSE loss
|
||||
value_loss_unclipped = F.mse_loss(values, returns, reduction="none")
|
||||
|
||||
# Clipped value loss (PPO-style)
|
||||
if use_clipping and old_values is not None:
|
||||
values_clipped = old_values + torch.clamp(
|
||||
values - old_values,
|
||||
-clip_range,
|
||||
clip_range
|
||||
)
|
||||
value_loss_clipped = F.mse_loss(values_clipped, returns, reduction="none")
|
||||
value_loss = torch.max(value_loss_unclipped, value_loss_clipped).mean()
|
||||
else:
|
||||
value_loss = value_loss_unclipped.mean()
|
||||
|
||||
# Compute diagnostics
|
||||
with torch.no_grad():
|
||||
explained_variance = 1.0 - (returns - values).var() / (returns.var() + 1e-8)
|
||||
|
||||
info = {
|
||||
"value_loss": value_loss.item(),
|
||||
"explained_variance": explained_variance.item(),
|
||||
"value_mean": values.mean().item(),
|
||||
"value_std": values.std().item(),
|
||||
}
|
||||
|
||||
return value_loss, info
|
||||
|
||||
|
||||
def compute_policy_entropy(log_probs: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Compute policy entropy for exploration bonus.
|
||||
|
||||
Args:
|
||||
log_probs: Log probabilities [B]
|
||||
|
||||
Returns:
|
||||
entropy: Mean entropy across batch (scalar)
|
||||
"""
|
||||
# For continuous actions: H = -log_prob (assuming Gaussian)
|
||||
# For discrete: H = -sum(p * log(p))
|
||||
# Here we use a simple approximation
|
||||
entropy = -log_probs.mean()
|
||||
return entropy
|
||||
|
||||
|
||||
def apply_gradient_reweighting(
|
||||
gradients: torch.Tensor,
|
||||
timesteps: torch.Tensor,
|
||||
num_train_timesteps: int = 1000
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply GRPO-Guard gradient reweighting across denoising steps.
|
||||
|
||||
This reweights gradients based on the timestep to balance learning
|
||||
across different noise levels.
|
||||
|
||||
Args:
|
||||
gradients: Gradients to reweight [B, ...]
|
||||
timesteps: Timesteps at which gradients were computed [B]
|
||||
num_train_timesteps: Total number of training timesteps
|
||||
|
||||
Returns:
|
||||
reweighted_gradients: Reweighted gradients [B, ...]
|
||||
"""
|
||||
# Compute timestep weights (higher weight for later timesteps)
|
||||
# This is a simple linear weighting, can be made more sophisticated
|
||||
timestep_weights = 1.0 + (timesteps.float() / num_train_timesteps)
|
||||
timestep_weights = timestep_weights.view(-1, *([1] * (gradients.dim() - 1)))
|
||||
|
||||
return gradients * timestep_weights
|
||||
|
||||
|
||||
def sample_random_timesteps(
|
||||
batch_size: int,
|
||||
min_timestep: int,
|
||||
max_timestep: int,
|
||||
device: torch.device,
|
||||
generator: torch.Generator | None = None
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Sample random timesteps for noise injection (Flow-GRPO-Fast).
|
||||
|
||||
Args:
|
||||
batch_size: Number of samples
|
||||
min_timestep: Minimum timestep
|
||||
max_timestep: Maximum timestep
|
||||
device: Device for tensor
|
||||
generator: Random generator for reproducibility
|
||||
|
||||
Returns:
|
||||
timesteps: Random timesteps [B]
|
||||
"""
|
||||
if generator is not None:
|
||||
timesteps = torch.randint(
|
||||
min_timestep,
|
||||
max_timestep + 1,
|
||||
(batch_size,),
|
||||
device=device,
|
||||
generator=generator
|
||||
)
|
||||
else:
|
||||
timesteps = torch.randint(
|
||||
min_timestep,
|
||||
max_timestep + 1,
|
||||
(batch_size,),
|
||||
device=device
|
||||
)
|
||||
|
||||
return timesteps
|
||||
|
||||
|
||||
def compute_reward_statistics(
|
||||
rewards: torch.Tensor
|
||||
) -> dict[str, float]:
|
||||
"""
|
||||
Compute statistics for reward distribution.
|
||||
|
||||
Args:
|
||||
rewards: Reward values [B]
|
||||
|
||||
Returns:
|
||||
stats: Dictionary with mean, std, min, max
|
||||
"""
|
||||
return {
|
||||
"reward_mean": rewards.mean().item(),
|
||||
"reward_std": rewards.std().item(),
|
||||
"reward_min": rewards.min().item(),
|
||||
"reward_max": rewards.max().item(),
|
||||
}
|
||||
|
||||
|
||||
def check_early_stopping(
|
||||
kl_divergence: float,
|
||||
target_kl: float
|
||||
) -> bool:
|
||||
"""
|
||||
Check if training should stop early based on KL divergence.
|
||||
|
||||
Args:
|
||||
kl_divergence: Current KL divergence
|
||||
target_kl: Target KL threshold
|
||||
|
||||
Returns:
|
||||
should_stop: True if KL exceeds target
|
||||
"""
|
||||
if kl_divergence > target_kl:
|
||||
logger.warning(
|
||||
"Early stopping triggered: KL divergence %.4f > target %.4f",
|
||||
kl_divergence,
|
||||
target_kl
|
||||
)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def compute_log_probs_from_model_output(
|
||||
model_output: torch.Tensor,
|
||||
target: torch.Tensor,
|
||||
noise_level: float = 0.1
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute log probabilities from model predictions.
|
||||
|
||||
For diffusion models, we approximate log probabilities using the
|
||||
negative squared error (assuming Gaussian likelihood).
|
||||
|
||||
Args:
|
||||
model_output: Model predictions [B, C, T, H, W]
|
||||
target: Target values [B, C, T, H, W]
|
||||
noise_level: Assumed noise level (std) for Gaussian likelihood
|
||||
|
||||
Returns:
|
||||
log_probs: Log probabilities [B]
|
||||
"""
|
||||
# Compute mean squared error per sample
|
||||
mse = ((model_output - target) ** 2).flatten(1).mean(dim=1)
|
||||
|
||||
# Log probability under Gaussian: log p(x) = -0.5 * (x - mu)^2 / sigma^2 + const
|
||||
log_probs = -0.5 * mse / (noise_level ** 2)
|
||||
|
||||
return log_probs
|
||||
|
||||
|
||||
def check_for_nan_inf(tensor: torch.Tensor, name: str) -> None:
|
||||
"""
|
||||
Check tensor for NaN or Inf values and raise error if found.
|
||||
|
||||
Args:
|
||||
tensor: Tensor to check
|
||||
name: Name for error message
|
||||
"""
|
||||
if torch.isnan(tensor).any():
|
||||
raise ValueError(f"{name} contains NaN values")
|
||||
if torch.isinf(tensor).any():
|
||||
raise ValueError(f"{name} contains Inf values")
|
||||
@@ -0,0 +1,189 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Per-prompt statistics tracking for GRPO training.
|
||||
|
||||
This module ports the PerPromptStatTracker from FlowGRPO to FastVideo.
|
||||
It tracks reward statistics per unique prompt and computes normalized advantages.
|
||||
|
||||
Ported from:
|
||||
- flow_grpo/flow_grpo/stat_tracking.py
|
||||
|
||||
Key adaptations:
|
||||
1. Uses FastVideo's logging instead of print statements
|
||||
2. Works with single GPU (no distributed logic)
|
||||
3. Supports numpy arrays and torch tensors
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PerPromptStatTracker:
|
||||
"""
|
||||
Tracks reward statistics per unique prompt for advantage normalization.
|
||||
|
||||
This class maintains running statistics (mean, std) for each unique prompt
|
||||
and computes normalized advantages using either per-prompt or global statistics.
|
||||
|
||||
Used in GRPO training to normalize advantages within groups of samples
|
||||
generated from the same prompt, which helps stabilize training when different
|
||||
prompts have different reward scales.
|
||||
"""
|
||||
|
||||
def __init__(self, global_std: bool = False):
|
||||
"""
|
||||
Initialize the per-prompt stat tracker.
|
||||
|
||||
Args:
|
||||
global_std: If True, use global std across all rewards for normalization.
|
||||
If False, use per-prompt std (default, recommended for GRPO).
|
||||
"""
|
||||
self.global_std = global_std
|
||||
self.stats: dict[str, list] = {} # Maps prompt -> list of rewards
|
||||
self.history_prompts: set[int] = set() # Set of hashed prompts seen
|
||||
|
||||
def update(
|
||||
self,
|
||||
prompts: Union[list[str], np.ndarray],
|
||||
rewards: Union[list[float], np.ndarray, torch.Tensor],
|
||||
type: str = 'grpo'
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Update statistics and compute normalized advantages.
|
||||
|
||||
Args:
|
||||
prompts: List or array of prompt strings (one per sample)
|
||||
rewards: Array or tensor of reward values (one per sample)
|
||||
type: Advantage computation type:
|
||||
- 'grpo': Normalize by (reward - mean) / std (default)
|
||||
- 'rwr': Return rewards as-is (reward-weighted regression)
|
||||
- 'sft': Binary advantages (1 for max, 0 otherwise)
|
||||
- 'dpo': DPO-style advantages (1 for max, -1 for min)
|
||||
|
||||
Returns:
|
||||
advantages: Normalized advantages array [num_samples] or [num_samples, ...]
|
||||
Shape matches rewards shape
|
||||
"""
|
||||
# Convert to numpy arrays
|
||||
prompts = np.array(prompts)
|
||||
if isinstance(rewards, torch.Tensor):
|
||||
rewards = rewards.detach().cpu().numpy()
|
||||
rewards = np.array(rewards, dtype=np.float64)
|
||||
|
||||
# Ensure rewards are 1D (one reward per sample)
|
||||
# FlowGRPO expects rewards to be aggregated per sample
|
||||
if rewards.ndim > 1:
|
||||
# If multi-dimensional, flatten or take mean
|
||||
# For [B, num_steps] shape, we typically want one reward per sample
|
||||
# So we take the mean across timesteps
|
||||
if rewards.ndim == 2:
|
||||
# Assume shape is [B, num_steps] - take mean across timesteps
|
||||
rewards = rewards.mean(axis=1)
|
||||
else:
|
||||
# Flatten and take mean for higher dimensions
|
||||
rewards = rewards.reshape(len(prompts), -1).mean(axis=1)
|
||||
|
||||
# Ensure prompts and rewards have matching lengths
|
||||
assert len(prompts) == len(rewards), \
|
||||
f"Prompts ({len(prompts)}) and rewards ({len(rewards)}) must have same length"
|
||||
|
||||
unique_prompts = np.unique(prompts)
|
||||
advantages = np.zeros_like(rewards, dtype=np.float64)
|
||||
|
||||
# First pass: collect rewards for each prompt
|
||||
for prompt in unique_prompts:
|
||||
prompt_mask = prompts == prompt
|
||||
prompt_rewards = rewards[prompt_mask]
|
||||
|
||||
# Store rewards in stats
|
||||
if prompt not in self.stats:
|
||||
self.stats[prompt] = []
|
||||
self.stats[prompt].extend(prompt_rewards.tolist())
|
||||
self.history_prompts.add(hash(prompt))
|
||||
|
||||
# Second pass: compute statistics and advantages
|
||||
for prompt in unique_prompts:
|
||||
prompt_mask = prompts == prompt
|
||||
prompt_rewards = rewards[prompt_mask]
|
||||
|
||||
# Stack all historical rewards for this prompt
|
||||
if len(self.stats[prompt]) > 0:
|
||||
all_prompt_rewards = np.array(self.stats[prompt])
|
||||
else:
|
||||
all_prompt_rewards = prompt_rewards
|
||||
|
||||
# Compute mean and std
|
||||
mean = np.mean(all_prompt_rewards, axis=0, keepdims=True)
|
||||
|
||||
if self.global_std:
|
||||
# Use global std across all rewards
|
||||
std = np.std(rewards, axis=0, keepdims=True) + 1e-4
|
||||
else:
|
||||
# Use per-prompt std
|
||||
std = np.std(all_prompt_rewards, axis=0, keepdims=True) + 1e-4
|
||||
|
||||
# Compute advantages based on type
|
||||
if type == 'grpo':
|
||||
# GRPO: normalize by (reward - mean) / std
|
||||
advantages[prompt_mask] = (prompt_rewards - mean) / std
|
||||
elif type == 'rwr':
|
||||
# Reward-weighted regression: use rewards as-is
|
||||
advantages[prompt_mask] = prompt_rewards
|
||||
elif type == 'sft':
|
||||
# Supervised fine-tuning: binary (1 for max, 0 otherwise)
|
||||
max_reward = np.max(prompt_rewards)
|
||||
advantages[prompt_mask] = (prompt_rewards == max_reward).astype(np.float64)
|
||||
elif type == 'dpo':
|
||||
# DPO-style: 1 for max, -1 for min
|
||||
prompt_rewards_tensor = torch.tensor(prompt_rewards)
|
||||
max_idx = torch.argmax(prompt_rewards_tensor)
|
||||
min_idx = torch.argmin(prompt_rewards_tensor)
|
||||
|
||||
# If all rewards are the same, use first two indices
|
||||
if max_idx == min_idx:
|
||||
min_idx = torch.tensor(0)
|
||||
max_idx = torch.tensor(1) if len(prompt_rewards_tensor) > 1 else torch.tensor(0)
|
||||
|
||||
result = torch.zeros_like(prompt_rewards_tensor, dtype=torch.float64)
|
||||
result[max_idx] = 1.0
|
||||
result[min_idx] = -1.0
|
||||
advantages[prompt_mask] = result.numpy()
|
||||
else:
|
||||
raise ValueError(f"Unknown advantage type: {type}. Must be one of: 'grpo', 'rwr', 'sft', 'dpo'")
|
||||
|
||||
return advantages
|
||||
|
||||
def get_stats(self) -> tuple[float, int]:
|
||||
"""
|
||||
Get statistics about tracked prompts.
|
||||
|
||||
Returns:
|
||||
avg_group_size: Average number of samples per unique prompt
|
||||
history_prompts: Number of unique prompts seen (across all updates)
|
||||
"""
|
||||
if not self.stats:
|
||||
avg_group_size = 0.0
|
||||
else:
|
||||
total_samples = sum(len(v) for v in self.stats.values())
|
||||
avg_group_size = total_samples / len(self.stats)
|
||||
|
||||
history_prompts = len(self.history_prompts)
|
||||
|
||||
return avg_group_size, history_prompts
|
||||
|
||||
def clear(self) -> None:
|
||||
"""
|
||||
Clear all statistics (but keep history_prompts for tracking).
|
||||
|
||||
This is typically called after each epoch to reset per-epoch statistics
|
||||
while maintaining a record of all prompts seen during training.
|
||||
"""
|
||||
self.stats = {}
|
||||
logger.debug("Cleared per-prompt statistics (kept %d unique prompts in history)",
|
||||
len(self.history_prompts))
|
||||
@@ -0,0 +1,877 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
GRPO utilities for Wan model in FastVideo.
|
||||
|
||||
This module ports the SDE step and pipeline functions from FlowGRPO to work with
|
||||
FastVideo's scheduler and pipeline interfaces.
|
||||
|
||||
Ported from:
|
||||
- flow_grpo/flow_grpo/diffusers_patch/wan_pipeline_with_logprob.py
|
||||
|
||||
Key adaptations:
|
||||
1. Uses FastVideo's FlowUniPCMultistepScheduler instead of diffusers' UniPCMultistepScheduler
|
||||
2. Works with FastVideo's WanPipeline (ComposedPipelineBase) instead of diffusers' WanPipeline
|
||||
3. Direct module access via pipeline.get_module() instead of pipeline attributes
|
||||
4. Simplified prompt encoding (direct text encoder usage instead of pipeline stages)
|
||||
"""
|
||||
|
||||
import math
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.utils import get_compute_dtype
|
||||
|
||||
# for test_wan_transformer2
|
||||
import os
|
||||
from diffusers import WanTransformer3DModel
|
||||
from fastvideo.utils import maybe_download_model
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def test_wan_transformer():
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
dit_cpu_offload=True,
|
||||
pipeline_config=PipelineConfig(
|
||||
dit_config=WanVideoConfig(),
|
||||
dit_precision=precision_str))
|
||||
args.device = device
|
||||
|
||||
loader = TransformerLoader()
|
||||
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
|
||||
|
||||
model1 = WanTransformer3DModel.from_pretrained(
|
||||
TRANSFORMER_PATH, device=device,
|
||||
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
forward_batch = ForwardBatch(data_type="dummy", )
|
||||
|
||||
with torch.amp.autocast('cuda', dtype=precision):
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
'''
|
||||
INFO 01-19 22:53:46 [wan_grpo_utils.py:74] Model 1 weight sum: 395834.3506456231████ | 1/2 [00:00<00:00, 7.84it/s]
|
||||
INFO 01-19 22:53:46 [wan_grpo_utils.py:75] Model 1 weight mean: 0.0002789536598289884
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:83] Model 2 weight sum: 395834.3506456231
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:84] Model 2 weight mean: 0.0002789536598289884
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:87] Weight sum difference: 0.0
|
||||
INFO 01-19 22:53:47 [wan_grpo_utils.py:89] Weight mean difference: 0.0
|
||||
INFO 01-19 22:53:54 [wan_grpo_utils.py:145] Max Diff: 0.08203125
|
||||
INFO 01-19 22:53:54 [wan_grpo_utils.py:146] Mean Diff: 0.01129150390625
|
||||
'''
|
||||
|
||||
|
||||
def test_wan_transformer2(model2):
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
|
||||
logger.info("loading model1 transformer weight")
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
|
||||
model1 = WanTransformer3DModel.from_pretrained(
|
||||
TRANSFORMER_PATH,
|
||||
device=device,
|
||||
torch_dtype=precision,
|
||||
).to(device, dtype=precision).requires_grad_(False)
|
||||
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
logger.info("Model 1 weight sum: %s", weight_sum_model1)
|
||||
logger.info("Model 1 weight mean: %s", weight_mean_model1)
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info("Model 2 weight sum: %s", weight_sum_model2)
|
||||
logger.info("Model 2 weight mean: %s", weight_mean_model2)
|
||||
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info("Weight sum difference: %s", weight_sum_diff)
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info("Weight mean difference: %s", weight_mean_diff)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1 = model1.eval()
|
||||
model2 = model2.eval()
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
seq_len = 30
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(
|
||||
batch_size,
|
||||
16,
|
||||
21,
|
||||
160,
|
||||
90,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(
|
||||
batch_size,
|
||||
seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
forward_batch = ForwardBatch(data_type="dummy", )
|
||||
|
||||
with torch.amp.autocast("cuda", dtype=precision):
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
with set_forward_context(
|
||||
current_timestep=0,
|
||||
attn_metadata=None,
|
||||
forward_batch=forward_batch,
|
||||
):
|
||||
output2 = model2(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
|
||||
# Print basic stats for debugging (cast to float32 for stability)
|
||||
out1 = output1.detach().float()
|
||||
out2 = output2.detach().float()
|
||||
logger.info(
|
||||
"output1 stats: min=%s max=%s mean=%s std=%s",
|
||||
out1.min().item(),
|
||||
out1.max().item(),
|
||||
out1.mean().item(),
|
||||
out1.std(unbiased=False).item(),
|
||||
)
|
||||
logger.info(
|
||||
"output2 stats: min=%s max=%s mean=%s std=%s",
|
||||
out2.min().item(),
|
||||
out2.max().item(),
|
||||
out2.mean().item(),
|
||||
out2.std(unbiased=False).item(),
|
||||
)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert (output1.shape == output2.shape
|
||||
), f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
assert (output1.dtype == output2.dtype
|
||||
), f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info("Max Diff: %s", max_diff.item())
|
||||
logger.info("Mean Diff: %s", mean_diff.item())
|
||||
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
|
||||
# mean diff
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
'''
|
||||
when --dit_precision "bf16", use_fsdp hardcoded to False:
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:65] Model 1 weight sum: 395834.3506456231████████████████████████████████████████████████████████| 2/2 [00:00<00:00, 3.25it/s]
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:66] Model 1 weight mean: 0.0002789536598289884
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:75] Model 2 weight sum: 395125.463677882
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:76] Model 2 weight mean: 0.0002739000890162162
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:79] Weight sum difference: 708.8869677411276
|
||||
INFO 01-19 22:01:24 [wan_grpo_utils.py:81] Weight mean difference: 5.053570812772192e-06
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:139] output1 stats: min=-2.28125 max=1.921875 mean=-0.16638492047786713 std=0.458170622587204
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:146] output2 stats: min=-2.296875 max=1.90625 mean=-0.166452556848526 std=0.4579130709171295
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:165] Max Diff: 0.08984375
|
||||
INFO 01-19 22:01:32 [wan_grpo_utils.py:166] Mean Diff: 0.0120849609375
|
||||
|
||||
when --dit_precision "fp32", use_fsdp not changed:
|
||||
|
||||
'''
|
||||
|
||||
|
||||
def sde_step_with_logprob(
|
||||
scheduler: FlowUniPCMultistepScheduler,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: float | torch.FloatTensor,
|
||||
sample: torch.FloatTensor,
|
||||
prev_sample: torch.FloatTensor | None = None,
|
||||
generator: torch.Generator | None = None,
|
||||
deterministic: bool = False,
|
||||
return_pixel_log_prob: bool = False,
|
||||
return_dt_and_std_dev_t: bool = False
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, ...]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE.
|
||||
This function propagates the flow process from the learned model outputs
|
||||
(most often the predicted velocity) and computes log probabilities.
|
||||
|
||||
Ported from FlowGRPO's sde_step_with_logprob to work with FastVideo's
|
||||
FlowUniPCMultistepScheduler.
|
||||
|
||||
Args:
|
||||
scheduler: FastVideo FlowUniPCMultistepScheduler instance
|
||||
model_output: The direct output from learned flow model
|
||||
timestep: The current discrete timestep in the diffusion chain
|
||||
sample: A current instance of a sample created by the diffusion process
|
||||
prev_sample: Optional previous sample (if provided, used instead of sampling)
|
||||
generator: Optional random number generator
|
||||
deterministic: If True, no noise is added (deterministic sampling)
|
||||
return_pixel_log_prob: If True, return pixel-level log probabilities (not used)
|
||||
return_dt_and_std_dev_t: If True, return dt and std_dev_t separately
|
||||
|
||||
Returns:
|
||||
If return_dt_and_std_dev_t=True:
|
||||
(prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt)
|
||||
Otherwise:
|
||||
(prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt)
|
||||
"""
|
||||
|
||||
# # Convert all variables to fp32 for numerical stability
|
||||
# model_output = model_output.float()
|
||||
# sample = sample.float()
|
||||
# if prev_sample is not None:
|
||||
# prev_sample = prev_sample.float()
|
||||
|
||||
# Get step indices for current and previous timesteps
|
||||
# Handle both single timestep and batch of timesteps
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
if timestep.ndim == 0:
|
||||
timestep = timestep.unsqueeze(0)
|
||||
step_indices = [
|
||||
scheduler.index_for_timestep(t.item()) for t in timestep
|
||||
]
|
||||
else:
|
||||
step_indices = [scheduler.index_for_timestep(timestep)]
|
||||
|
||||
prev_step_indices = [step + 1 for step in step_indices]
|
||||
|
||||
# Move sigmas to sample device
|
||||
sigmas = scheduler.sigmas.to(sample.device)
|
||||
# myregion debug: hardcode sigmas to flow_grpo's
|
||||
sigmas = torch.Tensor([
|
||||
0.9997, 0.9824, 0.9639, 0.9441, 0.9227, 0.8996, 0.8746, 0.8475, 0.8178,
|
||||
0.7853, 0.7496, 0.7102, 0.6663, 0.6173, 0.5621, 0.4997, 0.4283, 0.3459,
|
||||
0.2498, 0.1362, 0.0000
|
||||
]).to(sample.device, sample.dtype)
|
||||
# end region
|
||||
|
||||
# Get sigma values for current and previous steps
|
||||
sigma = sigmas[step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_prev = sigmas[prev_step_indices].view(-1, 1, 1, 1, 1)
|
||||
sigma_max = sigmas[0].item() # First sigma (highest)
|
||||
sigma_min = sigmas[-1].item() # Last sigma (lowest)
|
||||
|
||||
dt = sigma_prev - sigma
|
||||
|
||||
# myregion debug
|
||||
print(f"[DEBUG]: sigma_max: {sigma_max}, sigma_min: {sigma_min}, dt: {dt}")
|
||||
print(f"[DEBUG]: in sde_step_with_logprob(), timestep: {timestep}")
|
||||
print(f"[DEBUG]: in sde_step_with_logprob(), sigmas: {sigmas}")
|
||||
print(f"[DEBUG]: in sde_step_with_logprob(), step_indices: {step_indices}")
|
||||
print(
|
||||
f"[DEBUG]: in sde_step_with_logprob(), prev_step_indices: {prev_step_indices}"
|
||||
)
|
||||
'''
|
||||
[DEBUG]: in sde_step_with_logprob(), timestep: tensor([428, 428, 428, 428], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
|
||||
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
|
||||
0.2499, 0.1363, 0.0000], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), step_indices: [16, 16, 16, 16]
|
||||
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [17, 17, 17, 17]
|
||||
|
||||
DEBUG]: in sde_step_with_logprob(), timestep: tensor([249], device='cuda:0')███████████████▎ | 18/20 [00:04<00:00, 3.87step/s, step_time=0.26s, timestep=346.0]
|
||||
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
|
||||
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
|
||||
0.2499, 0.1363, 0.0000], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), step_indices: [18]
|
||||
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [19]
|
||||
|
||||
[DEBUG]: in sde_step_with_logprob(), timestep: tensor([617, 617, 617, 617], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), sigmas: tensor([0.9999, 0.9826, 0.9642, 0.9443, 0.9230, 0.8999, 0.8749, 0.8477, 0.8181,
|
||||
0.7856, 0.7499, 0.7104, 0.6665, 0.6175, 0.5624, 0.4999, 0.4285, 0.3461,
|
||||
0.2499, 0.1363, 0.0000], device='cuda:0')
|
||||
[DEBUG]: in sde_step_with_logprob(), step_indices: [13, 13, 13, 13]
|
||||
[DEBUG]: in sde_step_with_logprob(), prev_step_indices: [14, 14, 14, 14]
|
||||
'''
|
||||
# endregion
|
||||
|
||||
# Compute std_dev_t and prev_sample_mean using SDE formulation
|
||||
std_dev_t = sigma_min + (sigma_max - sigma_min) * sigma
|
||||
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 not None and generator is not None:
|
||||
raise ValueError(
|
||||
"Cannot pass both generator and prev_sample. Please make sure that either `generator` or"
|
||||
" `prev_sample` stays `None`.")
|
||||
|
||||
# Sample prev_sample if not provided
|
||||
if prev_sample is None:
|
||||
variance_noise = randn_tensor(
|
||||
model_output.shape,
|
||||
generator=generator,
|
||||
device=model_output.device,
|
||||
dtype=model_output.dtype,
|
||||
)
|
||||
sqrt_dt = torch.sqrt(-1 * dt) # dt is negative (going backwards)
|
||||
prev_sample = prev_sample_mean + std_dev_t * sqrt_dt * variance_noise
|
||||
else:
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
# No noise is added during evaluation (deterministic)
|
||||
if deterministic:
|
||||
prev_sample = sample + dt * model_output
|
||||
sqrt_dt = torch.sqrt(-1 * dt)
|
||||
|
||||
# Compute log probability: log p(prev_sample | sample, model_output)
|
||||
# Assuming Gaussian distribution: N(prev_sample_mean, (std_dev_t * sqrt_dt)^2)
|
||||
std_dev_sqrt_dt = std_dev_t * sqrt_dt
|
||||
log_prob = (
|
||||
-((prev_sample.detach() - prev_sample_mean)**2) /
|
||||
(2 * (std_dev_sqrt_dt**2)) - torch.log(
|
||||
std_dev_sqrt_dt + 1e-8) # Add small epsilon for numerical stability
|
||||
- torch.log(
|
||||
torch.sqrt(2 * torch.as_tensor(math.pi, device=sample.device))))
|
||||
|
||||
# Mean along all but batch dimension
|
||||
log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
|
||||
|
||||
if return_dt_and_std_dev_t:
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t, sqrt_dt
|
||||
return prev_sample, log_prob, prev_sample_mean, std_dev_t * sqrt_dt
|
||||
|
||||
|
||||
def wan_pipeline_with_logprob(
|
||||
pipeline,
|
||||
prompt: str | list[str] = None,
|
||||
negative_prompt: str | list[str] = None,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 5.0,
|
||||
num_videos_per_prompt: int | None = 1,
|
||||
generator: torch.Generator | list[torch.Generator] | None = None,
|
||||
latents: torch.Tensor | None = None,
|
||||
prompt_embeds: torch.Tensor | None = None,
|
||||
negative_prompt_embeds: torch.Tensor | None = None,
|
||||
output_type: str | None = "pt",
|
||||
return_dict: bool = False,
|
||||
attention_kwargs: dict[str, Any] | None = None,
|
||||
max_sequence_length: int = 512,
|
||||
deterministic: bool = False,
|
||||
kl_reward: float = 0.0,
|
||||
return_pixel_log_prob: bool = False,
|
||||
) -> tuple[torch.Tensor, list[torch.Tensor], list[torch.Tensor],
|
||||
list[torch.Tensor], torch.Tensor | None]:
|
||||
"""
|
||||
Wan pipeline with log probability computation for GRPO training.
|
||||
|
||||
Ported from FlowGRPO's wan_pipeline_with_logprob to work with FastVideo's WanPipeline.
|
||||
This function generates videos and computes log probabilities at each denoising step.
|
||||
|
||||
Args:
|
||||
pipeline: FastVideo WanPipeline instance
|
||||
prompt: Text prompt(s) for generation
|
||||
negative_prompt: Negative prompt(s) for classifier-free guidance
|
||||
height: Height of generated video
|
||||
width: Width of generated video
|
||||
num_frames: Number of frames in generated video
|
||||
num_inference_steps: Number of denoising steps
|
||||
guidance_scale: Classifier-free guidance scale
|
||||
num_videos_per_prompt: Number of videos to generate per prompt
|
||||
generator: Random generator for reproducibility
|
||||
latents: Optional initial latents
|
||||
prompt_embeds: Optional pre-computed prompt embeddings
|
||||
negative_prompt_embeds: Optional pre-computed negative prompt embeddings
|
||||
output_type: Output type ("pt" for PyTorch tensor, "np" for numpy, "latent" for latents only)
|
||||
return_dict: Whether to return dict (not used, always returns tuple)
|
||||
attention_kwargs: Optional attention kwargs
|
||||
max_sequence_length: Maximum sequence length for text encoding
|
||||
deterministic: If True, use deterministic sampling (no noise)
|
||||
kl_reward: KL reward coefficient (if > 0, computes KL divergence)
|
||||
return_pixel_log_prob: If True, return pixel-level log probabilities (not used)
|
||||
|
||||
Returns:
|
||||
Tuple of:
|
||||
- video: Generated video tensor [B, C, T, H, W] or latents if output_type="latent"
|
||||
- all_latents: List of latents at each step [num_steps+1] of shape [B, C, T, H, W]
|
||||
- all_log_probs: List of log probabilities at each step [num_steps] of shape [B]
|
||||
- all_kl: List of KL divergences at each step [num_steps] of shape [B] (if kl_reward > 0)
|
||||
- prompt_ids: Tokenized prompt IDs [B, seq_len] (None if prompt_embeds were provided)
|
||||
"""
|
||||
# Get device from transformer
|
||||
transformer = pipeline.get_module("transformer")
|
||||
|
||||
# myregion debug: test transformer output
|
||||
logger.info("testing transformer, running test_wan_transformer2")
|
||||
test_wan_transformer()
|
||||
# test_wan_transformer2(transformer)
|
||||
# endregion
|
||||
|
||||
# hardcode dtype for debug
|
||||
# transformer_dtype = torch.float32
|
||||
# use get_compute_dtype() to get dtype based on mixed precision
|
||||
transformer_dtype = get_compute_dtype()
|
||||
logger.info(f"[DEBUG]: transformer_dtype: {transformer_dtype}")
|
||||
|
||||
# Get scheduler and other modules
|
||||
scheduler = pipeline.get_module("scheduler")
|
||||
vae = pipeline.get_module("vae")
|
||||
|
||||
# Determine batch size
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
elif prompt_embeds is not None:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
else:
|
||||
raise ValueError("Either prompt or prompt_embeds must be provided")
|
||||
|
||||
# Encode prompts if not provided
|
||||
prompt_ids = None
|
||||
if prompt_embeds is None:
|
||||
# Encode prompts directly using text encoder and tokenizer
|
||||
# This is a simplified encoding - for full pipeline encoding, use TextEncodingStage
|
||||
text_encoder = pipeline.get_module("text_encoder")
|
||||
tokenizer = pipeline.get_module("tokenizer")
|
||||
|
||||
# Normalize to list
|
||||
if isinstance(prompt, str):
|
||||
prompts_list = [prompt]
|
||||
else:
|
||||
prompts_list = prompt
|
||||
|
||||
# Tokenize prompts
|
||||
text_inputs = tokenizer(prompts_list,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
return_tensors="pt").to(pipeline.device)
|
||||
|
||||
# Store prompt_ids for return
|
||||
prompt_ids = text_inputs["input_ids"]
|
||||
|
||||
# Encode with text encoder
|
||||
with torch.no_grad():
|
||||
outputs = text_encoder(
|
||||
text_inputs["input_ids"],
|
||||
attention_mask=text_inputs["attention_mask"],
|
||||
output_hidden_states=True,
|
||||
)
|
||||
# Get last hidden state (Wan typically uses last hidden state)
|
||||
prompt_embeds = outputs.last_hidden_state
|
||||
|
||||
# Encode negative prompts if CFG is enabled
|
||||
if guidance_scale > 1.0:
|
||||
if negative_prompt is None:
|
||||
negative_prompt = [""] * len(prompts_list)
|
||||
elif isinstance(negative_prompt, str):
|
||||
negative_prompt = [negative_prompt]
|
||||
|
||||
neg_text_inputs = tokenizer(negative_prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
return_tensors="pt").to(pipeline.device)
|
||||
|
||||
with torch.no_grad():
|
||||
neg_outputs = text_encoder(
|
||||
neg_text_inputs["input_ids"],
|
||||
attention_mask=neg_text_inputs["attention_mask"],
|
||||
output_hidden_states=True,
|
||||
)
|
||||
negative_prompt_embeds = neg_outputs.last_hidden_state
|
||||
else:
|
||||
negative_prompt_embeds = None
|
||||
|
||||
# myregion Debug: Print shapes of prompt embeddings
|
||||
logger.info(
|
||||
f"After encoding - prompt_embeds shape: {prompt_embeds.shape if prompt_embeds is not None else None}"
|
||||
)
|
||||
logger.info(
|
||||
f"After encoding - negative_prompt_embeds shape: {negative_prompt_embeds.shape if negative_prompt_embeds is not None else None}"
|
||||
)
|
||||
logger.info(
|
||||
f"After encoding - prompt_embeds dtype: {prompt_embeds.dtype if prompt_embeds is not None else None}"
|
||||
)
|
||||
logger.info(
|
||||
f"After encoding - negative_prompt_embeds dtype: {negative_prompt_embeds.dtype if negative_prompt_embeds is not None else None}"
|
||||
)
|
||||
'''
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:290] After encoding - prompt_embeds shape: torch.Size([4, 512, 4096])
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:291] After encoding - negative_prompt_embeds shape: None
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:292] After encoding - prompt_embeds dtype: torch.float32
|
||||
INFO 01-17 05:31:13 [wan_grpo_utils.py:293] After encoding - negative_prompt_embeds dtype: None
|
||||
'''
|
||||
# endregion
|
||||
# logger.info("wan_pipeline_with_logprob's transformer class type: %s", type(transformer))
|
||||
# logger.info("Variables in transformer: %s", str(dir(transformer)))
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
if negative_prompt_embeds is not None:
|
||||
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
|
||||
|
||||
# Prepare timesteps
|
||||
scheduler.set_timesteps(num_inference_steps, device=pipeline.device)
|
||||
timesteps = scheduler.timesteps
|
||||
|
||||
# Prepare latent variables
|
||||
num_channels_latents = transformer.config.in_channels
|
||||
vae = pipeline.get_module("vae")
|
||||
# Get VAE scale factors
|
||||
vae_scale_factor_spatial = vae.spatial_compression_ratio
|
||||
vae_scale_factor_temporal = vae.temporal_compression_ratio
|
||||
|
||||
if latents is None:
|
||||
# Generate random latents
|
||||
# Note: num_frames in latents accounts for temporal compression
|
||||
num_latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
|
||||
latents_shape = (
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
num_latent_frames,
|
||||
height // vae_scale_factor_spatial,
|
||||
width // vae_scale_factor_spatial,
|
||||
)
|
||||
if generator is not None:
|
||||
if isinstance(generator, list):
|
||||
latents = [
|
||||
torch.randn(
|
||||
latents_shape[1:],
|
||||
generator=gen,
|
||||
device=pipeline.device,
|
||||
dtype=transformer_dtype,
|
||||
) for gen in generator
|
||||
]
|
||||
latents = torch.stack(latents, dim=0)
|
||||
else:
|
||||
latents = torch.randn(
|
||||
latents_shape,
|
||||
generator=generator,
|
||||
device=pipeline.device,
|
||||
dtype=transformer_dtype,
|
||||
)
|
||||
else:
|
||||
latents = torch.randn(latents_shape,
|
||||
device=pipeline.device,
|
||||
dtype=transformer_dtype)
|
||||
else:
|
||||
latents = latents.to(device=pipeline.device, dtype=transformer_dtype)
|
||||
|
||||
|
||||
# myregion Debug: Print latents shape, dtype, and value range
|
||||
logger.info("=" * 80)
|
||||
logger.info("Latents Debug Information:")
|
||||
logger.info(f" Shape: {latents.shape}")
|
||||
logger.info(f" Dtype: {latents.dtype}")
|
||||
logger.info(f" Min value: {latents.min().item():.6f}")
|
||||
logger.info(f" Max value: {latents.max().item():.6f}")
|
||||
logger.info(f" Mean value: {latents.mean().item():.6f}")
|
||||
logger.info(f" Std value: {latents.std().item():.6f}")
|
||||
logger.info(f" Device: {latents.device}")
|
||||
logger.info("=" * 80)
|
||||
'''
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:355] ================================================================================
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:356] Latents Debug Information:
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:357] Shape: torch.Size([4, 16, 9, 30, 52])
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:358] Dtype: torch.bfloat16
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:359] Min value: -4.500000
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:360] Max value: 4.656250
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:361] Mean value: 0.000111
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:362] Std value: 1.000000
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:363] Device: cuda:0
|
||||
INFO 01-17 07:41:33 [wan_grpo_utils.py:364] ================================================================================
|
||||
'''
|
||||
# endregion
|
||||
|
||||
all_latents = [latents]
|
||||
all_log_probs = []
|
||||
all_kl = []
|
||||
|
||||
# myregion Debug
|
||||
logger.info("Tensor type issue debugging:")
|
||||
logger.info(f"latents: {type(latents)}")
|
||||
logger.info(f"prompt_embeds: {type(prompt_embeds)}")
|
||||
logger.info(
|
||||
f"[DEBUG]: before denoising loop: type(timesteps): {type(timesteps)}")
|
||||
logger.info(
|
||||
f"[DEBUG]: before denoising loop: timesteps.shape: {timesteps.shape}")
|
||||
# endregion
|
||||
|
||||
# Progress bar for denoising loop
|
||||
progress_bar = tqdm(enumerate(timesteps),
|
||||
total=len(timesteps),
|
||||
desc="Denoising steps",
|
||||
unit="step")
|
||||
|
||||
for i, t in progress_bar:
|
||||
step_start_time = time.time()
|
||||
latents_ori = latents.clone()
|
||||
timestep = t.expand(latents.shape[0]) if isinstance(
|
||||
t, torch.Tensor) else torch.tensor([t] * latents.shape[0],
|
||||
device=pipeline.device)
|
||||
|
||||
logger.info(
|
||||
f"[DEBUG]: before set_forward_context: current_timestep=i:{i}")
|
||||
# Predict noise with transformer
|
||||
with set_forward_context(
|
||||
current_timestep=t.item(),
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
noise_pred = transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_pred.to(prompt_embeds.dtype)
|
||||
|
||||
# Classifier-free guidance
|
||||
if guidance_scale > 1.0:
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
noise_uncond = transformer(
|
||||
hidden_states=latents,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_uncond + guidance_scale * (noise_pred -
|
||||
noise_uncond)
|
||||
|
||||
# SDE step with log probability
|
||||
latents, log_prob, prev_latents_mean, std_dev_t = sde_step_with_logprob(
|
||||
scheduler,
|
||||
noise_pred, #.float(),
|
||||
t.unsqueeze(0) if isinstance(t, torch.Tensor) else t,
|
||||
latents, #.float(),
|
||||
deterministic=deterministic,
|
||||
return_pixel_log_prob=return_pixel_log_prob)
|
||||
# sde_step_with_logprob returns fp32
|
||||
# latents = latents.to(transformer_dtype)
|
||||
prev_latents = latents.clone()
|
||||
|
||||
all_latents.append(latents)
|
||||
all_log_probs.append(log_prob)
|
||||
|
||||
# Compute KL divergence if kl_reward > 0 (for KL reward in sampling)
|
||||
if kl_reward > 0 and not deterministic:
|
||||
# Use reference model (disable adapter if using LoRA)
|
||||
latent_model_input_ref = torch.cat(
|
||||
[latents_ori] * 2) if guidance_scale > 1.0 else latents_ori
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=None,
|
||||
forward_batch=None,
|
||||
):
|
||||
with transformer.disable_adapter() if hasattr(
|
||||
transformer, 'disable_adapter') else torch.no_grad():
|
||||
noise_pred_ref = transformer(
|
||||
hidden_states=latent_model_input_ref,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred_ref = noise_pred_ref.to(prompt_embeds.dtype)
|
||||
|
||||
# Perform guidance for reference model
|
||||
if guidance_scale > 1.0:
|
||||
noise_pred_uncond_ref, noise_pred_text_ref = noise_pred_ref.chunk(
|
||||
2)
|
||||
noise_pred_ref = noise_pred_uncond_ref + guidance_scale * (
|
||||
noise_pred_text_ref - noise_pred_uncond_ref)
|
||||
|
||||
# Compute reference log prob
|
||||
_, ref_log_prob, ref_prev_latents_mean, ref_std_dev_t = sde_step_with_logprob(
|
||||
scheduler,
|
||||
noise_pred_ref.float(),
|
||||
t.unsqueeze(0) if isinstance(t, torch.Tensor) else t,
|
||||
latents_ori.float(),
|
||||
prev_sample=prev_latents.float(),
|
||||
deterministic=deterministic,
|
||||
)
|
||||
|
||||
# Compute KL divergence: KL = (mean_diff)^2 / (2 * std^2)
|
||||
assert torch.allclose(
|
||||
std_dev_t, ref_std_dev_t
|
||||
), "std_dev_t should match between current and reference"
|
||||
kl = (prev_latents_mean - ref_prev_latents_mean)**2 / (2 *
|
||||
std_dev_t**2)
|
||||
kl = kl.mean(dim=tuple(range(1, kl.ndim)))
|
||||
all_kl.append(kl)
|
||||
else:
|
||||
# No KL reward, set to zero
|
||||
all_kl.append(torch.zeros(len(latents), device=latents.device))
|
||||
|
||||
# Update progress bar with timing information
|
||||
step_time = time.time() - step_start_time
|
||||
progress_bar.set_postfix({
|
||||
"step_time":
|
||||
f"{step_time:.2f}s",
|
||||
"timestep":
|
||||
f"{t.item() if isinstance(t, torch.Tensor) else t:.1f}"
|
||||
})
|
||||
|
||||
# Decode latents to video if needed
|
||||
if output_type != "latent":
|
||||
latents = latents.to(vae.dtype)
|
||||
|
||||
# Apply VAE normalization (Wan VAE specific)
|
||||
# Wan VAE requires denormalization before decoding
|
||||
if hasattr(vae, 'config') and hasattr(vae.config,
|
||||
'latents_mean') and hasattr(
|
||||
vae.config, 'latents_std'):
|
||||
# Get z_dim from config or VAE
|
||||
z_dim = getattr(vae.config, 'z_dim', latents.shape[1])
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype).view(
|
||||
1, z_dim, 1, 1, 1))
|
||||
latents_std = (
|
||||
1.0 / torch.tensor(vae.config.latents_std,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype).view(1, z_dim, 1, 1, 1))
|
||||
latents = latents / latents_std + latents_mean
|
||||
elif hasattr(vae, 'latents_mean') and hasattr(vae, 'latents_std'):
|
||||
# Alternative: check if latents_mean/std are direct attributes
|
||||
z_dim = latents.shape[1]
|
||||
latents_mean = (torch.tensor(vae.latents_mean,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype).view(
|
||||
1, z_dim, 1, 1, 1))
|
||||
latents_std = (1.0 / torch.tensor(
|
||||
vae.latents_std, device=latents.device,
|
||||
dtype=latents.dtype).view(1, z_dim, 1, 1, 1))
|
||||
latents = latents / latents_std + latents_mean
|
||||
|
||||
# Decode using VAE
|
||||
with torch.no_grad():
|
||||
video = vae.decode(latents.float(), return_dict=False)[0]
|
||||
# VAE.decode returns tensor directly (not tuple)
|
||||
|
||||
# Postprocess video: convert from [-1, 1] to [0, 1]
|
||||
# FastVideo VAE typically outputs in [-1, 1] range
|
||||
video = (video / 2 + 0.5).clamp(0, 1)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
return video, all_latents, all_log_probs, all_kl, prompt_ids
|
||||
@@ -174,17 +174,18 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
training_args.train_batch_size,
|
||||
parquet_schema=self.train_dataset_schema,
|
||||
num_data_workers=training_args.dataloader_num_workers,
|
||||
cfg_rate=training_args.training_cfg_rate,
|
||||
drop_last=True,
|
||||
text_padding_length=training_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.
|
||||
text_len, # type: ignore[attr-defined]
|
||||
seed=self.seed)
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
training_args.train_batch_size,
|
||||
parquet_schema=self.train_dataset_schema,
|
||||
num_data_workers=training_args.dataloader_num_workers,
|
||||
cfg_rate=training_args.training_cfg_rate,
|
||||
drop_last=True,
|
||||
text_padding_length=training_args.pipeline_config.
|
||||
text_encoder_configs[0].arch_config.
|
||||
text_len, # type: ignore[attr-defined]
|
||||
seed=self.seed)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
if self.training_args.boundary_ratio is not None:
|
||||
@@ -192,19 +193,21 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
else:
|
||||
self.boundary_timestep = None
|
||||
|
||||
logger.info("train_dataloader length: %s", len(self.train_dataloader))
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
logger.info("train_dataloader length: %s", len(self.train_dataloader))
|
||||
logger.info("train_sp_batch_size: %s",
|
||||
training_args.train_sp_batch_size)
|
||||
logger.info("gradient_accumulation_steps: %s",
|
||||
training_args.gradient_accumulation_steps)
|
||||
logger.info("sp_size: %s", training_args.sp_size)
|
||||
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
training_args.train_sp_batch_size)
|
||||
self.num_train_epochs = math.ceil(training_args.max_train_steps /
|
||||
self.num_update_steps_per_epoch)
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
training_args.train_sp_batch_size)
|
||||
self.num_train_epochs = math.ceil(training_args.max_train_steps /
|
||||
self.num_update_steps_per_epoch)
|
||||
|
||||
# TODO(will): is there a cleaner way to track epochs?
|
||||
self.current_epoch = 0
|
||||
@@ -575,8 +578,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
if not self.training_args.rl_args.rl_mode:
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
else:
|
||||
self.noise_random_generator = torch.Generator(device=self.device).manual_seed(
|
||||
self.seed)
|
||||
self.noise_gen_cuda = torch.Generator(
|
||||
device=current_platform.device_name).manual_seed(self.seed)
|
||||
self.validation_random_generator = torch.Generator(
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
|
||||
from fastvideo.training.rl.rl_pipeline import RLPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanRLTrainingPipeline(RLPipeline):
|
||||
"""
|
||||
A training pipeline for Wan with RL/GRPO support.
|
||||
|
||||
This pipeline extends RLPipeline with Wan-specific initialization.
|
||||
"""
|
||||
_required_config_modules = [
|
||||
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _create_inference_pipeline(self, training_args: TrainingArgs,
|
||||
dit_cpu_offload: bool):
|
||||
args_copy = deepcopy(training_args)
|
||||
args_copy.inference_mode = True
|
||||
loaded_modules = {
|
||||
"transformer": self.get_module("transformer"),
|
||||
}
|
||||
transformer_2 = self.get_module("transformer_2", None)
|
||||
if transformer_2 is not None:
|
||||
loaded_modules["transformer_2"] = transformer_2
|
||||
text_encoder = self.get_module("text_encoder", None)
|
||||
if text_encoder is not None:
|
||||
loaded_modules["text_encoder"] = text_encoder
|
||||
tokenizer = self.get_module("tokenizer", None)
|
||||
if tokenizer is not None:
|
||||
loaded_modules["tokenizer"] = tokenizer
|
||||
vae = self.get_module("vae", None)
|
||||
if vae is not None:
|
||||
loaded_modules["vae"] = vae
|
||||
|
||||
return WanPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules=loaded_modules,
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=dit_cpu_offload)
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
self.validation_pipeline = self._create_inference_pipeline(
|
||||
training_args, dit_cpu_offload=True)
|
||||
|
||||
def _build_sampling_pipeline(self, training_args: TrainingArgs):
|
||||
return self._create_inference_pipeline(training_args,
|
||||
dit_cpu_offload=False)
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting RL training pipeline...")
|
||||
|
||||
pipeline = WanRLTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("RL training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
# Enable RL mode
|
||||
args.rl_mode = True
|
||||
main(args)
|
||||
Reference in New Issue
Block a user