from collections import defaultdict
import contextlib
import datetime
from concurrent import futures
import hashlib
import itertools
import json
import math
import os
import random
import sys
import tempfile
import time
from absl import app, flags
from accelerate import Accelerator
from accelerate.logging import get_logger
from accelerate.utils import ProjectConfiguration, set_seed
from huggingface_hub import snapshot_download
from ml_collections import config_flags
import numpy as np
from peft import LoraConfig, PeftModel, get_peft_model
from PIL import Image, ImageDraw
import torch
from torch.utils.data import DataLoader, Dataset, Sampler
import tqdm
import wandb
from diffusers.models import AutoencoderKL
from transformers.cache_utils import Cache, DynamicCache
from transformers.modeling_outputs import BaseModelOutputWithPast
import flow_grpo.rewards
from flow_grpo.dataset_paths import DatasetPathResolver, resolve_metadata_paths
from flow_grpo.ema import EMAModuleWrapper
from flow_grpo.omnigen_patch.omnigen_pipeline_with_logprob import (
build_group_ids,
compute_cfg_model_output,
compute_cfg_model_output_joint,
filter_branch_inputs,
omnigen_sde_step_with_logprob,
pipeline_with_logprob,
pipeline_with_logprob_joint_image_reward,
)
from flow_grpo.omnigen_patch.joint_model_loader import load_joint_omnigen_components_for_rl
from flow_grpo.stat_tracking import PerPromptStatTracker
from flow_grpo.utils.run_metadata import write_run_metadata
DEFAULT_OMNIGEN_CODE_ROOT = "/NAS_REMOTE/vicky/wt/codes/cxr_syn/gen_code"
OMNIGEN_CODE_ROOT = os.environ.get("OMNIGEN_CODE_ROOT", DEFAULT_OMNIGEN_CODE_ROOT)
if os.path.exists(OMNIGEN_CODE_ROOT) and OMNIGEN_CODE_ROOT not in sys.path:
sys.path.insert(0, OMNIGEN_CODE_ROOT)
import transformers.cache_utils as hf_cache_utils # noqa: E402
if not hasattr(hf_cache_utils, "OffloadedCache"):
hf_cache_utils.OffloadedCache = hf_cache_utils.DynamicCache
from OmniGen import OmniGen, OmniGenProcessor # noqa: E402
from OmniGen.transformer import Phi3Transformer # noqa: E402
tqdm = tqdm.tqdm
FLAGS = flags.FLAGS
config_flags.DEFINE_config_file("config", "config/base.py", "Training configuration.")
logger = get_logger(__name__)
def _debug_flag(name: str) -> bool:
return os.environ.get(name, "0").strip().lower() in {"1", "true", "yes", "on"}
def _to_rgb_pil(image):
if isinstance(image, Image.Image):
return image.convert("RGB")
if isinstance(image, torch.Tensor):
image = image.detach().cpu()
if image.ndim == 3 and image.shape[0] in (1, 3):
if image.shape[0] == 1:
image = image.repeat(3, 1, 1)
image = image.clamp(0, 1).permute(1, 2, 0).numpy()
elif image.ndim == 2:
image = image.numpy()
else:
image = image.numpy()
array = np.asarray(image)
if array.ndim == 3 and array.shape[0] in (1, 3) and array.shape[-1] not in (1, 3):
array = np.transpose(array, (1, 2, 0))
if array.dtype != np.uint8:
array = np.clip(array, 0, 1) if np.issubdtype(array.dtype, np.floating) else np.clip(array, 0, 255)
array = (array * 255).round().astype(np.uint8) if array.max(initial=0) <= 1.0 else array.astype(np.uint8)
if array.ndim == 2:
return Image.fromarray(array, mode="L").convert("RGB")
if array.ndim == 3 and array.shape[-1] == 1:
return Image.fromarray(array[..., 0], mode="L").convert("RGB")
return Image.fromarray(array).convert("RGB")
def _load_optional_image(image_path):
if not image_path:
return None
if not os.path.exists(image_path):
return None
return Image.open(image_path).convert("RGB")
def _truncate_text(text, max_length=220):
text = str(text).strip().replace("\n", " ")
if len(text) <= max_length:
return text
return text[: max_length - 3] + "..."
def _build_visual_panel(input_image, output_image, gt_image=None):
panels = [("Input", _to_rgb_pil(input_image)), ("Output", _to_rgb_pil(output_image))]
if gt_image is not None:
panels.append(("GT", _to_rgb_pil(gt_image)))
target_width = max(image.width for _, image in panels)
target_height = max(image.height for _, image in panels)
resized = [
(
label,
image.resize((target_width, target_height), resample=Image.Resampling.BILINEAR),
)
for label, image in panels
]
gap = 12
title_height = 24
canvas_width = len(resized) * target_width + (len(resized) - 1) * gap
canvas_height = title_height + target_height
canvas = Image.new("RGB", (canvas_width, canvas_height), color=(255, 255, 255))
draw = ImageDraw.Draw(canvas)
x = 0
for label, image in resized:
draw.text((x, 4), label, fill=(0, 0, 0))
canvas.paste(image, (x, title_height))
x += target_width + gap
return canvas
def _build_wandb_examples(
output_images,
prompts,
metadatas,
ref_images,
*,
instructions=None,
rewards=None,
max_samples=8,
note=None,
):
num_samples = min(max_samples, len(output_images))
examples = []
for idx in range(num_samples):
metadata = metadatas[idx]
gt_image = _load_optional_image(metadata.get("gt_image"))
panel = _build_visual_panel(ref_images[idx], output_images[idx], gt_image=gt_image)
caption_parts = []
sample_id = metadata.get("sample_id")
if sample_id:
caption_parts.append(str(sample_id))
caption_parts.append(f"prompt: {_truncate_text(prompts[idx])}")
if instructions is not None and instructions[idx] != prompts[idx]:
caption_parts.append(f"instruction: {_truncate_text(instructions[idx])}")
if rewards is not None:
reward_parts = []
for key, value in rewards.items():
current_value = float(value[idx])
if current_value != -10:
reward_parts.append(f"{key}: {current_value:.2f}")
if reward_parts:
caption_parts.append(" | ".join(reward_parts))
if note:
caption_parts.append(note)
examples.append(wandb.Image(panel, caption="\n".join(caption_parts)))
return examples
def patch_omnigen_phi3_forward():
if getattr(Phi3Transformer.forward, "_flow_grpo_patched", False):
return
def _patched_forward(
self,
input_ids=None,
attention_mask=None,
position_ids=None,
past_key_values=None,
inputs_embeds=None,
use_cache=None,
output_attentions=None,
output_hidden_states=None,
return_dict=None,
cache_position=None,
offload_model=False,
):
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
use_cache = use_cache if use_cache is not None else self.config.use_cache
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if (input_ids is None) ^ (inputs_embeds is not None):
raise ValueError("You must specify exactly one of input_ids or inputs_embeds")
if self.gradient_checkpointing and self.training and use_cache:
logger.warning_once(
"`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
)
use_cache = False
return_legacy_cache = False
if use_cache and not isinstance(past_key_values, Cache):
return_legacy_cache = True
if past_key_values is None:
past_key_values = DynamicCache(config=self.config)
else:
past_key_values = DynamicCache.from_legacy_cache(past_key_values)
logger.warning_once(
"We detected tuple-based `past_key_values`. This is deprecated and was converted to DynamicCache."
)
if cache_position is None:
past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
cache_position = torch.arange(
past_seen_tokens,
past_seen_tokens + inputs_embeds.shape[1],
device=inputs_embeds.device,
)
if position_ids is None:
position_ids = cache_position.unsqueeze(0)
if attention_mask is not None and attention_mask.dim() == 3:
dtype = inputs_embeds.dtype
min_dtype = torch.finfo(dtype).min
attention_mask = (1 - attention_mask) * min_dtype
attention_mask = attention_mask.unsqueeze(1).to(inputs_embeds.dtype)
else:
raise Exception("attention_mask parameter was unavailable or invalid")
hidden_states = inputs_embeds
has_external_position_embeddings = hasattr(self, "rotary_emb")
position_embeddings = (
self.rotary_emb(hidden_states, position_ids=position_ids)
if has_external_position_embeddings
else None
)
all_hidden_states = () if output_hidden_states else None
all_self_attns = () if output_attentions else None
next_decoder_cache = None
for layer_idx, decoder_layer in enumerate(self.layers):
if output_hidden_states:
all_hidden_states += (hidden_states,)
if self.gradient_checkpointing and self.training:
if has_external_position_embeddings:
layer_outputs = self._gradient_checkpointing_func(
decoder_layer.__call__,
hidden_states,
attention_mask,
position_ids,
past_key_values,
use_cache,
position_embeddings,
)
else:
layer_outputs = self._gradient_checkpointing_func(
decoder_layer.__call__,
hidden_states,
attention_mask,
position_ids,
past_key_values,
output_attentions,
use_cache,
cache_position,
)
else:
if offload_model and not self.training:
self.get_offlaod_layer(layer_idx, device=inputs_embeds.device)
if has_external_position_embeddings:
layer_outputs = decoder_layer(
hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
output_attentions=output_attentions,
use_cache=use_cache,
position_embeddings=position_embeddings,
cache_position=cache_position,
)
else:
layer_outputs = decoder_layer(
hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_value=past_key_values,
output_attentions=output_attentions,
use_cache=use_cache,
cache_position=cache_position,
)
if isinstance(layer_outputs, tuple):
hidden_states = layer_outputs[0]
if use_cache and not has_external_position_embeddings:
next_decoder_cache = layer_outputs[2 if output_attentions else 1]
if output_attentions and len(layer_outputs) > 1:
all_self_attns += (layer_outputs[1],)
else:
hidden_states = layer_outputs
hidden_states = self.norm(hidden_states)
if output_hidden_states:
all_hidden_states += (hidden_states,)
next_cache = (past_key_values if has_external_position_embeddings else next_decoder_cache) if use_cache else None
if return_legacy_cache and next_cache is not None:
next_cache = next_cache.to_legacy_cache()
if not return_dict:
return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
return BaseModelOutputWithPast(
last_hidden_state=hidden_states,
past_key_values=next_cache,
hidden_states=all_hidden_states,
attentions=all_self_attns,
)
_patched_forward._flow_grpo_patched = True
Phi3Transformer.forward = _patched_forward
patch_omnigen_phi3_forward()
class RadiomicsEditDataset(Dataset):
def __init__(self, dataset, split="train", condition_dropout_prob=0.0):
if isinstance(dataset, dict) or hasattr(dataset, "get"):
key = f"{split}_jsonl"
self.file_path = dataset.get(key) or dataset.get("jsonl")
if self.file_path is None and dataset.get("root"):
self.file_path = os.path.join(dataset.get("root"), f"{split}_metadata.jsonl")
if self.file_path is None:
raise ValueError(f"Dataset config is missing {key!r} for split={split}.")
else:
self.file_path = os.path.join(dataset, f"{split}_metadata.jsonl")
self.path_resolver = DatasetPathResolver.from_dataset_config(dataset)
self.condition_dropout_prob = condition_dropout_prob if split == "train" else 0.0
with open(self.file_path, "r", encoding="utf-8") as f:
self.metadatas = [json.loads(line) for line in f if line.strip()]
def __len__(self):
return len(self.metadatas)
def __getitem__(self, idx):
metadata = resolve_metadata_paths(
dict(self.metadatas[idx]),
self.path_resolver,
required_keys=("input_images", "output_image", "output_mask"),
)
input_images = metadata.get("input_images") or []
if len(input_images) == 0:
raise ValueError("OmniGen edit training requires at least one input image.")
instruction = metadata.get("instruction")
if not instruction:
instruction = f"
<|image_1|> {metadata['prompt']}"
primary_input_image = input_images[0]
prompt = metadata.get("prompt") or metadata.get("instruction", "")
output_image = metadata.get("output_image") or metadata.get("gt_image")
group_key = str(
metadata.get("sample_id")
or output_image
or (f"{primary_input_image}::{output_image}" if output_image else None)
or f"{prompt}::{primary_input_image}"
)
if random.random() < self.condition_dropout_prob:
instruction = ""
input_images = []
return {
"prompt": prompt,
"instruction": instruction,
"metadata": metadata,
"input_image_paths": input_images,
"ref_image": Image.open(primary_input_image).convert("RGB"),
"group_key": group_key,
}
@staticmethod
def collate_fn(examples):
prompts = [example["prompt"] for example in examples]
instructions = [example["instruction"] for example in examples]
metadatas = [example["metadata"] for example in examples]
input_image_paths = [example["input_image_paths"] for example in examples]
ref_images = [example["ref_image"] for example in examples]
group_keys = [example["group_key"] for example in examples]
return prompts, instructions, metadatas, input_image_paths, ref_images, group_keys
class DistributedKRepeatSampler(Sampler):
def __init__(self, dataset, batch_size, k, num_replicas, rank, num_batches_per_epoch, seed=0):
self.dataset = dataset
self.batch_size = batch_size
self.k = k
self.num_replicas = num_replicas
self.rank = rank
self.num_batches_per_epoch = num_batches_per_epoch
self.seed = seed
self.total_samples = self.num_replicas * self.batch_size
self.group_span = self.k // math.gcd(self.total_samples, self.k)
if self.num_batches_per_epoch % self.group_span != 0:
raise ValueError(
"num_batches_per_epoch must be divisible by the sampler group span. "
f"Got num_batches_per_epoch={num_batches_per_epoch}, group_span={self.group_span}, "
f"k={k}, num_replicas={num_replicas}, batch_size={batch_size}."
)
self.groups_per_epoch = self.num_batches_per_epoch // self.group_span
self.prompts_per_group_pack = self.total_samples * self.group_span // self.k
self.epoch = 0
self._scheduled_epoch = None
self._scheduled_batches = []
self._batch_ptr = 0
def _sample_unique_indices(self, generator, count):
indices = []
dataset_size = len(self.dataset)
while len(indices) < count:
indices.extend(torch.randperm(dataset_size, generator=generator).tolist())
return indices[:count]
def _build_epoch_batches(self):
generator = torch.Generator()
generator.manual_seed(self.seed + self.epoch)
prompts_needed = self.groups_per_epoch * self.prompts_per_group_pack
prompt_indices = self._sample_unique_indices(generator, prompts_needed)
scheduled_batches = []
for pack_index in range(self.groups_per_epoch):
start = pack_index * self.prompts_per_group_pack
end = start + self.prompts_per_group_pack
current_indices = prompt_indices[start:end]
repeated_indices = [idx for idx in current_indices for _ in range(self.k)]
shuffled_indices = torch.randperm(len(repeated_indices), generator=generator).tolist()
shuffled_samples = [repeated_indices[i] for i in shuffled_indices]
for batch_offset in range(self.group_span):
batch_start = batch_offset * self.total_samples
batch_end = batch_start + self.total_samples
global_batch = shuffled_samples[batch_start:batch_end]
replica_start = self.rank * self.batch_size
replica_end = replica_start + self.batch_size
scheduled_batches.append(global_batch[replica_start:replica_end])
return scheduled_batches
def __iter__(self):
while True:
if self._scheduled_epoch != self.epoch or self._batch_ptr >= len(self._scheduled_batches):
self._scheduled_batches = self._build_epoch_batches()
self._scheduled_epoch = self.epoch
self._batch_ptr = 0
batch_indices = self._scheduled_batches[self._batch_ptr]
self._batch_ptr += 1
yield batch_indices
def set_epoch(self, epoch):
self.epoch = epoch
self._scheduled_epoch = None
self._scheduled_batches = []
self._batch_ptr = 0
def create_generator(group_keys, base_seed, device):
generators = []
generator_device = device.type if isinstance(device, torch.device) else str(device)
for key in group_keys:
digest = hashlib.sha256(str(key).encode("utf-8")).digest()
key_hash = int.from_bytes(digest[:4], "big")
seed = (base_seed + key_hash) % (2**31)
generators.append(torch.Generator(device=generator_device).manual_seed(seed))
return generators
def calculate_zero_std_ratio(group_keys, gathered_rewards):
group_array = np.array(group_keys)
unique_groups, inverse_indices, counts = np.unique(
group_array, return_inverse=True, return_counts=True
)
del unique_groups
grouped_rewards = gathered_rewards["ori_avg"][np.argsort(inverse_indices)]
split_indices = np.cumsum(counts)[:-1]
reward_groups = np.split(grouped_rewards, split_indices)
group_std = np.array([np.std(group) for group in reward_groups])
zero_std_ratio = np.count_nonzero(group_std == 0) / len(group_std)
return zero_std_ratio, group_std.mean()
def requires_grad(model, flag=True):
for param in model.parameters():
param.requires_grad = flag
def count_parameters(model):
total = 0
trainable = 0
for param in model.parameters():
numel = param.numel()
total += numel
if param.requires_grad:
trainable += numel
return trainable, total
def unwrap_model(model, accelerator):
model = accelerator.unwrap_model(model)
return getattr(model, "_orig_mod", model)
def disable_adapter(model, accelerator):
return unwrap_model(model, accelerator).disable_adapter()
def resolve_model_root(model_name_or_path):
if os.path.exists(model_name_or_path):
return model_name_or_path
cache_dir = os.getenv("HF_HUB_CACHE", "/tmp/flow_grpo_hf_cache/hub")
os.makedirs(cache_dir, exist_ok=True)
return snapshot_download(
repo_id=model_name_or_path,
cache_dir=cache_dir,
ignore_patterns=["flax_model.msgpack", "rust_model.ot", "tf_model.h5"],
)
def patch_omnigen_config_compat(model_root):
config_path = os.path.join(model_root, "config.json")
if not os.path.exists(config_path):
return
with open(config_path, "r", encoding="utf-8") as f:
config = json.load(f)
rope_scaling = config.get("rope_scaling")
original_max_position_embeddings = config.get("original_max_position_embeddings")
updated = False
if isinstance(rope_scaling, dict):
if rope_scaling.get("type") in {"su", "yarn"}:
rope_scaling["type"] = "longrope"
updated = True
allowed_rope_keys = {"type", "short_factor", "long_factor"}
extra_rope_keys = set(rope_scaling) - allowed_rope_keys
for key in extra_rope_keys:
rope_scaling.pop(key)
updated = updated or bool(extra_rope_keys)
if updated:
with open(config_path, "w", encoding="utf-8") as f:
json.dump(config, f, ensure_ascii=False, indent=2)
f.write("\n")
def load_omnigen_model(config, device, weight_dtype, *, enable_gradient_checkpointing=False):
model_root = resolve_model_root(config.pretrained.model)
patch_omnigen_config_compat(model_root)
model = OmniGen.from_pretrained(model_root)
model.llm.config.use_cache = False
if hasattr(model.llm, "gradient_checkpointing_enable") and enable_gradient_checkpointing:
try:
model.llm.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
except TypeError:
model.llm.gradient_checkpointing_enable()
model.to(device=device, dtype=weight_dtype)
return model, model_root
def load_omnigen_components(config, device, weight_dtype):
model, model_root = load_omnigen_model(
config,
device,
weight_dtype,
enable_gradient_checkpointing=config.activation_checkpointing,
)
processor = OmniGenProcessor.from_pretrained(model_root)
vae_path = getattr(config.pretrained, "vae_path", None)
if vae_path:
vae = AutoencoderKL.from_pretrained(vae_path)
else:
local_vae_path = os.path.join(model_root, "vae")
if os.path.exists(local_vae_path):
vae = AutoencoderKL.from_pretrained(local_vae_path)
else:
vae = AutoencoderKL.from_pretrained("stabilityai/sdxl-vae")
vae.to(device=device, dtype=torch.float32)
return model, vae, processor
def merge_lora_into_base_model(model, lora_path, weight_dtype, *, trainable=True):
logger.info("Merging OmniGen LoRA into base weights from %s", lora_path)
merged_model = PeftModel.from_pretrained(
model,
lora_path,
is_trainable=False,
)
if hasattr(merged_model, "set_adapter"):
merged_model.set_adapter("default")
if not hasattr(merged_model, "merge_and_unload"):
raise AttributeError(
"Loaded OmniGen LoRA does not support merge_and_unload(); "
"cannot start full-parameter training from adapter weights."
)
model = merged_model.merge_and_unload()
model.to(dtype=weight_dtype)
requires_grad(model, trainable)
return model
def load_lora_adapter_config(lora_path):
adapter_config_path = os.path.join(lora_path, "adapter_config.json")
try:
with open(adapter_config_path, "r", encoding="utf-8") as handle:
return json.load(handle)
except OSError:
logger.warning("Could not read LoRA adapter config from %s", adapter_config_path)
except json.JSONDecodeError:
logger.warning("Invalid LoRA adapter config JSON at %s", adapter_config_path)
return None
def load_full_finetune_reference_model(config, device, weight_dtype):
# Keep a frozen copy of the initial full-finetuning policy for KL regularization.
reference_model, _ = load_omnigen_model(config, device, weight_dtype)
if config.train.lora_path:
reference_model = merge_lora_into_base_model(
reference_model,
config.train.lora_path,
weight_dtype,
trainable=False,
)
else:
requires_grad(reference_model, False)
reference_model.eval()
return reference_model
def compute_log_prob(model, sample, timestep_index, config):
current_latents = sample["latents"][:, timestep_index]
current_timestep = sample["timesteps"][:, timestep_index]
next_timestep = sample["next_timesteps"][:, timestep_index]
if sample.get("use_joint_mask", False):
current_mask_latents = sample["mask_latents"][:, timestep_index]
model_output, _ = compute_cfg_model_output_joint(
model,
current_latents,
current_mask_latents,
current_timestep,
sample["branch_inputs"],
cfg_scale=config.sample.guidance_scale,
img_cfg_scale=config.sample.img_guidance_scale,
use_img_guidance=config.sample.use_img_guidance,
)
else:
model_output = compute_cfg_model_output(
model,
current_latents,
current_timestep,
sample["branch_inputs"],
cfg_scale=config.sample.guidance_scale,
img_cfg_scale=config.sample.img_guidance_scale,
use_img_guidance=config.sample.use_img_guidance,
)
stochastic_step_mask = sample.get("stochastic_step_mask")
use_stochastic_step = True
if stochastic_step_mask is not None:
use_stochastic_step = bool(stochastic_step_mask[timestep_index].detach().cpu().item())
return omnigen_sde_step_with_logprob(
model_output,
current_timestep,
next_timestep,
current_latents,
prev_sample=sample["next_latents"][:, timestep_index],
noise_level=config.sample.noise_level if use_stochastic_step else 0.0,
sde_type=config.sample.sde_type if use_stochastic_step else "ode",
logprob_mode=getattr(config.sample, "omnigen_logprob_mode", "mse"),
)
def resolve_train_timestep_indices(config):
num_steps = int(config.sample.num_steps)
if bool(getattr(config.sample, "sde_window_enabled", False)) and bool(
getattr(config.sample, "sde_window_train_only", True)
):
mode = getattr(config.sample, "sde_window_mode", "fixed")
if mode != "fixed":
raise ValueError(f"Unsupported OmniGen SDE window mode: {mode!r}")
start = int(getattr(config.sample, "sde_window_start", 0))
end = int(getattr(config.sample, "sde_window_end", num_steps))
if start < 0 or end < 0 or start >= end or end > num_steps:
raise ValueError(
"Invalid SDE window train range. Expected 0 <= start < end <= num_steps, "
f"got start={start}, end={end}, num_steps={num_steps}."
)
return list(range(start, end))
num_train_timesteps = max(1, int(num_steps * config.train.timestep_fraction))
return list(range(num_train_timesteps))
def save_ckpt(save_dir, model, global_step, accelerator, ema, trainable_parameters, config):
save_root = os.path.join(save_dir, "checkpoints", f"checkpoint-{global_step}")
adapter_root = os.path.join(save_root, "lora")
os.makedirs(adapter_root, exist_ok=True)
if not accelerator.is_main_process:
return
if config.train.ema:
ema.copy_ema_to(trainable_parameters, store_temp=True)
unwrapped = unwrap_model(model, accelerator)
if config.use_lora:
unwrapped.save_pretrained(adapter_root)
else:
torch.save(unwrapped.state_dict(), os.path.join(save_root, "model.pt"))
if config.train.ema:
ema.copy_temp_to(trainable_parameters)
def eval(
model,
vae,
processor,
test_dataloader,
config,
accelerator,
global_step,
reward_fn,
executor,
autocast,
weight_dtype,
ema,
trainable_parameters,
):
eval_start_time = time.time()
if config.train.ema:
ema.copy_ema_to(trainable_parameters, store_temp=True)
all_rewards = defaultdict(list)
eval_gallery = []
max_logged_eval_samples = 15
max_eval_batches = getattr(config.sample, "max_eval_batches", None)
eval_dataloader_len = len(test_dataloader)
if max_eval_batches is None:
eval_iterable = test_dataloader
eval_total = eval_dataloader_len
else:
eval_iterable = itertools.islice(test_dataloader, max_eval_batches)
eval_total = min(max_eval_batches, eval_dataloader_len)
if accelerator.is_main_process:
logger.info(
"Starting train-time Eval: this is no-grad monitoring, not optimizer training. "
"max_eval_batches=%s, eval_num_steps=%s, eval_noise_level=0.0, eval_batches=%s, eval_dataset_size=%s",
max_eval_batches,
config.sample.eval_num_steps,
eval_total,
len(getattr(test_dataloader, "dataset", [])),
)
for batch_index, test_batch in enumerate(
tqdm(
eval_iterable,
desc="Eval",
total=eval_total,
disable=not accelerator.is_local_main_process,
dynamic_ncols=True,
)
):
prompts, instructions, prompt_metadata, input_image_paths, ref_images, _ = test_batch
with autocast():
with torch.no_grad():
eval_pipeline_fn = pipeline_with_logprob_joint_image_reward if bool(getattr(config, "use_joint_mask", False)) else pipeline_with_logprob
collected = eval_pipeline_fn(
model,
vae,
processor,
instructions,
input_image_paths,
height=config.resolution,
width=config.resolution,
num_inference_steps=config.sample.eval_num_steps,
guidance_scale=config.sample.eval_guidance_scale,
img_guidance_scale=config.sample.eval_img_guidance_scale,
max_input_image_size=config.sample.max_input_image_size,
use_img_guidance=config.sample.use_img_guidance,
use_input_image_size_as_output=config.sample.use_input_image_size_as_output,
dtype=weight_dtype,
output_type="pt",
noise_level=0.0,
sde_type=config.sample.sde_type,
mask_scale_factor=getattr(getattr(config, "joint", {}), "mask_scale_factor", 1.0),
sde_window_enabled=False,
)
processor = collected["processor"]
images = collected["images"]
rewards_future = executor.submit(
reward_fn,
images,
prompts,
prompt_metadata,
ref_images,
only_strict=False,
)
time.sleep(0)
rewards, _ = rewards_future.result()
for key, value in rewards.items():
rewards_gather = accelerator.gather(
torch.as_tensor(value, device=accelerator.device).float()
).cpu().numpy()
all_rewards[key].append(rewards_gather)
local_images = images.float().cpu().numpy()
for sample_index in range(len(local_images)):
if len(eval_gallery) >= max_logged_eval_samples:
break
eval_gallery.append(
{
"image": local_images[sample_index],
"prompt": prompts[sample_index],
"instruction": instructions[sample_index],
"metadata": prompt_metadata[sample_index],
"ref_image": ref_images[sample_index],
"rewards": {
key: float(np.asarray(value)[sample_index])
for key, value in rewards.items()
},
}
)
if not all_rewards:
if config.train.ema:
ema.copy_temp_to(trainable_parameters)
return processor
all_rewards = {key: np.concatenate(value) for key, value in all_rewards.items()}
if accelerator.is_main_process and eval_gallery:
wandb.log(
{
"eval_images": _build_wandb_examples(
[item["image"] for item in eval_gallery],
[item["prompt"] for item in eval_gallery],
[item["metadata"] for item in eval_gallery],
[item["ref_image"] for item in eval_gallery],
instructions=[item["instruction"] for item in eval_gallery],
rewards={
key: np.array([item["rewards"].get(key, -10.0) for item in eval_gallery], dtype=np.float32)
for key in eval_gallery[0]["rewards"].keys()
},
max_samples=max_logged_eval_samples,
note="eval output",
),
**{
f"eval_reward_{key}": np.mean(value[value != -10])
for key, value in all_rewards.items()
},
},
step=global_step,
)
if config.train.ema:
ema.copy_temp_to(trainable_parameters)
if accelerator.is_main_process:
logger.info("Train-time Eval finished in %.1fs", time.time() - eval_start_time)
return processor
def main(_):
config = FLAGS.config
unique_id = datetime.datetime.now().strftime("%Y.%m.%d_%H.%M.%S")
config.run_name = f"{config.run_name}_{unique_id}" if config.run_name else unique_id
train_timestep_indices = resolve_train_timestep_indices(config)
num_train_timesteps = len(train_timestep_indices)
accelerator_config = ProjectConfiguration(
project_dir=os.path.join(config.logdir, config.run_name),
automatic_checkpoint_naming=True,
total_limit=config.num_checkpoint_limit,
)
accelerator = Accelerator(
mixed_precision=config.mixed_precision,
project_config=accelerator_config,
gradient_accumulation_steps=config.train.gradient_accumulation_steps * num_train_timesteps,
)
if accelerator.is_main_process:
wandb.init(
project=os.environ.get("WANDB_PROJECT", "flow_grpo"),
name=os.environ.get("WANDB_NAME"),
mode=os.environ.get("WANDB_MODE", "online"),
)
write_run_metadata(
config.save_dir,
config,
config_entry=os.environ.get("CONFIG_ENTRY"),
extra={
"base_model": getattr(config.pretrained, "model", None),
"omnigen_code_root": getattr(config.pretrained, "local_code_root", None),
"sft_lora_path": getattr(getattr(config, "joint", {}), "sft_lora_path", None),
"mask_modules_path": getattr(getattr(config, "joint", {}), "mask_modules_path", None),
"reward_mask_channels": os.environ.get("REWARD_MASK_CHANNELS"),
"rl_lora_target_modules": getattr(config.train, "rl_lora_target_modules", None),
"train_jsonl": os.environ.get("TRAIN_JSONL"),
"test_jsonl": os.environ.get("TEST_JSONL"),
"save_dir": config.save_dir,
"sde_window_enabled": bool(getattr(config.sample, "sde_window_enabled", False)),
"sde_window_start": getattr(config.sample, "sde_window_start", None),
"sde_window_end": getattr(config.sample, "sde_window_end", None),
"sde_window_mode": getattr(config.sample, "sde_window_mode", None),
"omnigen_logprob_mode": getattr(config.sample, "omnigen_logprob_mode", "mse"),
"train_timestep_indices": train_timestep_indices,
},
)
if accelerator.is_main_process:
if _debug_flag("DEBUG_VERBOSE"):
logger.info(f"\n{config}")
else:
logger.info(
"Resolved joint config summary: max_eval_batches=%s, eval_num_steps=%s, num_steps=%s, "
"noise_level=%s, train_batch_size=%s, test_batch_size=%s, num_image_per_prompt=%s, beta=%s, "
"sde_window_enabled=%s, sde_window=[%s,%s), train_timesteps=%s, logprob_mode=%s",
getattr(config.sample, "max_eval_batches", None),
config.sample.eval_num_steps,
config.sample.num_steps,
config.sample.noise_level,
getattr(config.sample, "train_batch_size", None),
getattr(config.sample, "test_batch_size", None),
config.sample.num_image_per_prompt,
getattr(config.train, "beta", None),
bool(getattr(config.sample, "sde_window_enabled", False)),
getattr(config.sample, "sde_window_start", None),
getattr(config.sample, "sde_window_end", None),
train_timestep_indices,
getattr(config.sample, "omnigen_logprob_mode", "mse"),
)
set_seed(config.seed, device_specific=True)
weight_dtype = torch.float32
if accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16
model_load_start = time.time()
use_joint_mask = bool(getattr(config, "use_joint_mask", False))
joint_helpers = {}
if use_joint_mask:
model, vae, processor, joint_helpers = load_joint_omnigen_components_for_rl(
config,
device=accelerator.device,
weight_dtype=weight_dtype,
attach_rl_lora=True,
)
else:
model, vae, processor = load_omnigen_components(config, accelerator.device, weight_dtype)
requires_grad(vae, False)
if config.use_lora:
merge_lora_path = getattr(config.train, "merge_lora_path", None)
if merge_lora_path:
model = merge_lora_into_base_model(
model,
merge_lora_path,
weight_dtype,
trainable=False,
)
requires_grad(model, False)
lora_rank = getattr(config.train, "lora_rank", 8)
lora_alpha = getattr(config.train, "lora_alpha", lora_rank)
lora_config = LoraConfig(
r=lora_rank,
lora_alpha=lora_alpha,
init_lora_weights="gaussian",
target_modules=getattr(config.train, "rl_lora_target_modules", ["qkv_proj", "o_proj"]),
)
model.llm.enable_input_require_grads()
if config.train.lora_path:
logger.info("Loading pretrained OmniGen LoRA from %s", config.train.lora_path)
saved_lora_config = load_lora_adapter_config(config.train.lora_path)
if saved_lora_config is not None:
saved_rank = saved_lora_config.get("r")
saved_alpha = saved_lora_config.get("lora_alpha")
if saved_rank != lora_rank or saved_alpha != lora_alpha:
logger.warning(
"Loaded LoRA adapter was saved with r=%s, alpha=%s; "
"requested r=%s, alpha=%s only applies when creating a fresh adapter.",
saved_rank,
saved_alpha,
lora_rank,
lora_alpha,
)
model = PeftModel.from_pretrained(
model,
config.train.lora_path,
is_trainable=True,
)
model.set_adapter("default")
else:
model = get_peft_model(model, lora_config)
model.to(dtype=weight_dtype)
elif config.train.lora_path:
model = merge_lora_into_base_model(model, config.train.lora_path, weight_dtype)
if accelerator.is_main_process:
logger.info("Model/component load finished in %.1fs", time.time() - model_load_start)
trainable_parameters = [p for p in model.parameters() if p.requires_grad]
trainable_param_count, total_param_count = count_parameters(model)
ema = None
if config.train.ema:
ema = EMAModuleWrapper(
trainable_parameters,
decay=0.9,
update_step_interval=8,
device=accelerator.device,
)
if config.allow_tf32:
torch.backends.cuda.matmul.allow_tf32 = True
if config.train.use_8bit_adam:
try:
import bitsandbytes as bnb
except ImportError:
raise ImportError(
"Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`"
)
optimizer_cls = bnb.optim.AdamW8bit
else:
optimizer_cls = torch.optim.AdamW
optimizer = optimizer_cls(
trainable_parameters,
lr=config.train.learning_rate,
betas=(config.train.adam_beta1, config.train.adam_beta2),
weight_decay=config.train.adam_weight_decay,
eps=config.train.adam_epsilon,
)
train_dataset = RadiomicsEditDataset(
config.dataset,
"train",
condition_dropout_prob=getattr(config.train, "condition_dropout_prob", 0.0),
)
test_dataset = RadiomicsEditDataset(config.dataset, "test")
train_sampler = DistributedKRepeatSampler(
dataset=train_dataset,
batch_size=config.sample.train_batch_size,
k=config.sample.num_image_per_prompt,
num_replicas=accelerator.num_processes,
rank=accelerator.process_index,
num_batches_per_epoch=config.sample.num_batches_per_epoch,
seed=42,
)
if config.sample.same_latent and train_sampler.group_span > 1:
raise ValueError(
"same_latent=True currently requires each prompt group to fit inside one global sampling batch. "
f"Got group_span={train_sampler.group_span}, k={config.sample.num_image_per_prompt}, "
f"num_replicas={accelerator.num_processes}, train_batch_size={config.sample.train_batch_size}."
)
train_dataloader = DataLoader(
train_dataset,
batch_sampler=train_sampler,
num_workers=0,
collate_fn=RadiomicsEditDataset.collate_fn,
)
test_dataloader = DataLoader(
test_dataset,
batch_size=config.sample.test_batch_size,
shuffle=False,
num_workers=2,
collate_fn=RadiomicsEditDataset.collate_fn,
)
if config.sample.num_image_per_prompt == 1:
config.per_prompt_stat_tracking = False
if config.per_prompt_stat_tracking:
stat_tracker = PerPromptStatTracker(config.sample.global_std)
autocast = contextlib.nullcontext if config.use_lora else accelerator.autocast
reward_fn = getattr(flow_grpo.rewards, "multi_score")(accelerator.device, config.reward_fn)
eval_reward_fn = getattr(flow_grpo.rewards, "multi_score")(accelerator.device, config.reward_fn)
model, optimizer, train_dataloader, test_dataloader = accelerator.prepare(
model, optimizer, train_dataloader, test_dataloader
)
reference_model = None
kl_reference_mode = "disabled"
if config.train.beta > 0:
if config.use_lora:
kl_reference_mode = "base_model_without_adapter"
else:
logger.info("Loading frozen OmniGen reference model for full-finetune KL regularization.")
reference_model = load_full_finetune_reference_model(
config,
accelerator.device,
weight_dtype,
)
kl_reference_mode = "frozen_initial_policy"
trainable_parameters = [p for p in model.parameters() if p.requires_grad]
if config.train.ema:
ema = EMAModuleWrapper(
trainable_parameters,
decay=0.9,
update_step_interval=8,
device=accelerator.device,
)
executor = futures.ThreadPoolExecutor(max_workers=8)
samples_per_epoch = (
config.sample.train_batch_size
* accelerator.num_processes
* config.sample.num_batches_per_epoch
)
total_train_batch_size = (
config.train.batch_size
* accelerator.num_processes
* config.train.gradient_accumulation_steps
)
logger.info("***** Running OmniGen GRPO training *****")
logger.info(f" Finetuning mode = {'LoRA' if config.use_lora else 'full'}")
if config.use_lora:
logger.info(
" LoRA rank / alpha = %s / %s",
getattr(config.train, "lora_rank", 8),
getattr(config.train, "lora_alpha", getattr(config.train, "lora_rank", 8)),
)
logger.info(" Merged base LoRA path = %s", getattr(config.train, "merge_lora_path", None))
logger.info(" Trainable LoRA path = %s", getattr(config.train, "lora_path", None))
logger.info(" KL reference mode = %s", kl_reference_mode)
trainable_tensor_count = sum(1 for p in model.parameters() if p.requires_grad)
logger.info(
" Trainable parameters = %s / %s (%.2f%%); trainable tensors=%s; LoRA targets=%s",
f"{trainable_param_count:,}",
f"{total_param_count:,}",
100.0 * trainable_param_count / max(total_param_count, 1),
trainable_tensor_count,
getattr(config.train, "rl_lora_target_modules", None),
)
if _debug_flag("DEBUG_MODEL_PARAMS") and accelerator.is_main_process:
logger.info(" Trainable parameter names = %s", [name for name, param in model.named_parameters() if param.requires_grad])
logger.info(f" Sample batch size per device = {config.sample.train_batch_size}")
logger.info(f" Train batch size per device = {config.train.batch_size}")
logger.info(f" Gradient accumulation steps = {config.train.gradient_accumulation_steps}")
logger.info(f" Samples per epoch = {samples_per_epoch}")
logger.info(f" Effective train batch size = {total_train_batch_size}")
epoch = 0
global_step = 0
train_iter = iter(train_dataloader)
max_train_steps = int(getattr(config.train, "max_train_steps", 0) or 0)
while epoch < config.num_epochs:
if max_train_steps > 0 and global_step >= max_train_steps:
logger.info("Reached max_train_steps=%s at global_step=%s.", max_train_steps, global_step)
break
train_sampler.set_epoch(epoch)
model.eval()
# Disabled per request while we focus on train rollouts and throughput.
if epoch % config.eval_freq == 0:
processor = eval(
model,
vae,
processor,
test_dataloader,
config,
accelerator,
global_step,
eval_reward_fn,
executor,
autocast,
weight_dtype,
ema,
trainable_parameters,
)
if epoch % config.save_freq == 0 and epoch > 0:
save_ckpt(config.save_dir, model, global_step, accelerator, ema, trainable_parameters, config)
if accelerator.is_main_process:
logger.info("Starting no-grad rollout sampling for epoch %s", epoch)
rollout_start_time = time.time()
reward_wait_start_time = None
update_start_time = None
samples = []
last_images = None
last_prompts = None
last_instructions = None
last_rewards = None
last_ref_images = None
last_metadata = None
last_input_image_paths = None
for batch_index in tqdm(
range(config.sample.num_batches_per_epoch),
desc=f"Epoch {epoch}: sampling",
disable=not accelerator.is_local_main_process,
dynamic_ncols=True,
):
prompts, instructions, prompt_metadata, input_image_paths, ref_images, group_keys = next(train_iter)
group_ids = build_group_ids(group_keys, device=accelerator.device)
generator = (
create_generator(
group_keys,
base_seed=epoch * 10000 + batch_index,
device=accelerator.device,
)
if config.sample.same_latent
else None
)
with autocast():
with torch.no_grad():
pipeline_fn = pipeline_with_logprob_joint_image_reward if use_joint_mask else pipeline_with_logprob
collected = pipeline_fn(
model,
vae,
processor,
instructions,
input_image_paths,
height=config.resolution,
width=config.resolution,
num_inference_steps=config.sample.num_steps,
guidance_scale=config.sample.guidance_scale,
img_guidance_scale=config.sample.img_guidance_scale,
max_input_image_size=config.sample.max_input_image_size,
use_img_guidance=config.sample.use_img_guidance,
use_input_image_size_as_output=config.sample.use_input_image_size_as_output,
dtype=weight_dtype,
generator=generator,
output_type="pt",
noise_level=config.sample.noise_level,
sde_type=config.sample.sde_type,
mask_scale_factor=getattr(getattr(config, "joint", {}), "mask_scale_factor", 1.0),
sde_window_enabled=bool(getattr(config.sample, "sde_window_enabled", False)),
sde_window_start=int(getattr(config.sample, "sde_window_start", 0)),
sde_window_end=int(getattr(config.sample, "sde_window_end", config.sample.num_steps)),
sde_window_mode=getattr(config.sample, "sde_window_mode", "fixed"),
logprob_mode=getattr(config.sample, "omnigen_logprob_mode", "mse"),
)
processor = collected["processor"]
latents = torch.stack(collected["all_latents"], dim=1)
mask_latents = (
torch.stack(collected["all_mask_latents"], dim=1)
if use_joint_mask and "all_mask_latents" in collected
else None
)
log_probs = torch.stack(collected["all_log_probs"], dim=1)
timesteps = torch.stack(collected["all_timesteps"], dim=1)
next_timesteps = torch.stack(collected["all_next_timesteps"], dim=1)
std_dev_ts = torch.stack(collected["all_std_dev_ts"], dim=1)
stochastic_step_mask = collected["stochastic_step_mask"].to(accelerator.device)
images = collected["images"]
rewards_future = executor.submit(
reward_fn,
images,
prompts,
prompt_metadata,
ref_images,
only_strict=True,
)
time.sleep(0)
samples.append(
{
"group_ids": group_ids,
"branch_inputs": collected["branch_inputs"],
"timesteps": timesteps,
"next_timesteps": next_timesteps,
"latents": latents[:, :-1],
"next_latents": latents[:, 1:],
"mask_latents": mask_latents[:, :-1] if mask_latents is not None else None,
"next_mask_latents": mask_latents[:, 1:] if mask_latents is not None else None,
"use_joint_mask": use_joint_mask,
"log_probs": log_probs,
"std_dev_ts": std_dev_ts,
"stochastic_step_mask": stochastic_step_mask,
"rewards": rewards_future,
}
)
last_images = images
last_prompts = list(prompts)
last_instructions = list(instructions)
last_ref_images = list(ref_images)
last_metadata = list(prompt_metadata)
last_input_image_paths = [list(paths) for paths in input_image_paths]
if accelerator.is_main_process:
logger.info("No-grad rollout sampling finished in %.1fs; waiting for radiomics rewards", time.time() - rollout_start_time)
reward_wait_start_time = time.time()
for sample in tqdm(
samples,
desc="Waiting for rewards",
disable=not accelerator.is_local_main_process,
dynamic_ncols=True,
):
rewards, _ = sample["rewards"].result()
sample["rewards"] = {
key: torch.as_tensor(value, device=accelerator.device).float()
for key, value in rewards.items()
}
last_rewards = sample["rewards"]
if accelerator.is_main_process:
logger.info("Radiomics reward wait finished in %.1fs", time.time() - reward_wait_start_time)
reward_tensors = {
key: torch.cat([sample["rewards"][key] for sample in samples], dim=0)
for key in samples[0]["rewards"].keys()
}
reward_tensors["ori_avg"] = reward_tensors["avg"]
reward_tensors["avg"] = reward_tensors["avg"].unsqueeze(1).repeat(1, config.sample.num_steps)
gathered_rewards = {
key: accelerator.gather(value).float().cpu().numpy()
for key, value in reward_tensors.items()
}
gathered_group_ids = accelerator.gather(
torch.cat([sample["group_ids"] for sample in samples], dim=0)
).cpu().numpy()
if epoch % 10 == 0 and last_images is not None:
preview_images = None
with autocast():
with torch.no_grad():
preview_pipeline_fn = pipeline_with_logprob_joint_image_reward if use_joint_mask else pipeline_with_logprob
preview_collected = preview_pipeline_fn(
model,
vae,
processor,
last_instructions,
last_input_image_paths,
height=config.resolution,
width=config.resolution,
num_inference_steps=config.sample.eval_num_steps,
guidance_scale=config.sample.eval_guidance_scale,
img_guidance_scale=config.sample.eval_img_guidance_scale,
max_input_image_size=config.sample.max_input_image_size,
use_img_guidance=config.sample.use_img_guidance,
use_input_image_size_as_output=config.sample.use_input_image_size_as_output,
dtype=weight_dtype,
output_type="pt",
noise_level=0.0,
sde_type=config.sample.sde_type,
mask_scale_factor=getattr(getattr(config, "joint", {}), "mask_scale_factor", 1.0),
sde_window_enabled=False,
)
if accelerator.is_main_process:
preview_images = preview_collected["images"].float().cpu().numpy()
if accelerator.is_main_process:
reward_arrays = {
key: value.detach().float().cpu().numpy()
for key, value in last_rewards.items()
}
wandb.log(
{
"images": _build_wandb_examples(
preview_images,
last_prompts,
last_metadata,
last_ref_images,
instructions=last_instructions,
max_samples=8,
note="deterministic preview (noise_level=0.0)",
),
"train_rollout_images": _build_wandb_examples(
last_images.float().cpu().numpy(),
last_prompts,
last_metadata,
last_ref_images,
instructions=last_instructions,
rewards=reward_arrays,
max_samples=8,
note=f"stochastic rollout (noise_level={config.sample.noise_level})",
),
},
step=global_step,
)
accelerator.wait_for_everyone()
reward_summary = None
if accelerator.is_main_process:
reward_summary = {}
extra_logs = {}
if "radiomics_missing" in gathered_rewards:
extra_logs["reward/radiomics_missing_score_rate"] = float(np.mean(gathered_rewards["radiomics_missing"]))
if "radiomics_roi_pixels" in gathered_rewards:
extra_logs["reward/roi_pixels_mean"] = float(np.mean(gathered_rewards["radiomics_roi_pixels"]))
wandb.log(
{
"epoch": epoch,
"config/use_joint_mask": float(bool(getattr(config, "use_joint_mask", False))),
"config/sde_window_enabled": float(bool(getattr(config.sample, "sde_window_enabled", False))),
"config/sde_window_start": float(getattr(config.sample, "sde_window_start", 0)),
"config/sde_window_end": float(getattr(config.sample, "sde_window_end", config.sample.num_steps)),
"config/sde_window_count": float(num_train_timesteps),
"config/omnigen_logprob_mode": getattr(config.sample, "omnigen_logprob_mode", "mse"),
**extra_logs,
**{
f"reward_{key}": value.mean()
for key, value in gathered_rewards.items()
if "_strict_accuracy" not in key and "_accuracy" not in key
},
},
step=global_step,
)
for key, value in gathered_rewards.items():
if key == "ori_avg" or "_strict_accuracy" in key or "_accuracy" in key:
continue
valid = value != -10
reward_summary[key] = float(np.mean(value[valid])) if np.any(valid) else float(np.mean(value))
if config.per_prompt_stat_tracking:
group_key_strings = [str(group_id) for group_id in gathered_group_ids.tolist()]
advantages = stat_tracker.update(group_key_strings, gathered_rewards["avg"])
group_size, trained_prompt_num = stat_tracker.get_stats()
zero_std_ratio, reward_std_mean = calculate_zero_std_ratio(
group_key_strings,
gathered_rewards,
)
if accelerator.is_main_process:
wandb.log(
{
"group_size": group_size,
"trained_prompt_num": trained_prompt_num,
"zero_std_ratio": zero_std_ratio,
"reward_std_mean": reward_std_mean,
},
step=global_step,
)
stat_tracker.clear()
else:
group_size = None
trained_prompt_num = None
zero_std_ratio = None
reward_std_mean = None
advantages = (
gathered_rewards["avg"] - gathered_rewards["avg"].mean()
) / (gathered_rewards["avg"].std() + 1e-4)
if accelerator.is_main_process and reward_summary is not None:
summary_parts = [f"{key}={value:.4f}" for key, value in sorted(reward_summary.items())]
if zero_std_ratio is not None:
summary_parts.append(f"zero_std_ratio={zero_std_ratio:.4f}")
if reward_std_mean is not None:
summary_parts.append(f"reward_std_mean={reward_std_mean:.4f}")
if group_size is not None:
summary_parts.append(f"group_size={group_size:.2f}")
if trained_prompt_num is not None:
summary_parts.append(f"trained_prompt_num={trained_prompt_num}")
logger.info("Epoch %s reward summary: %s", epoch, ", ".join(summary_parts))
local_advantages = (
torch.as_tensor(advantages)
.reshape(accelerator.num_processes, -1, advantages.shape[-1])[accelerator.process_index]
.to(accelerator.device)
)
offset = 0
for sample in samples:
batch_size = sample["latents"].shape[0]
sample["advantages"] = local_advantages[offset : offset + batch_size]
offset += batch_size
sample["train_mask"] = sample["advantages"].abs().sum(dim=1) != 0
del sample["rewards"]
if len(samples) == 0:
epoch += 1
continue
if accelerator.is_main_process:
logger.info("Starting RL training updates")
update_start_time = time.time()
for inner_epoch in range(config.train.num_inner_epochs):
random.shuffle(samples)
model.train()
info = defaultdict(list)
for sample in tqdm(
samples,
desc=f"Epoch {epoch}.{inner_epoch}: training",
disable=not accelerator.is_local_main_process,
dynamic_ncols=True,
):
for timestep_index in tqdm(
train_timestep_indices,
desc="Timestep",
leave=False,
disable=not accelerator.is_local_main_process,
dynamic_ncols=True,
):
with accelerator.accumulate(model):
with autocast():
_, log_prob, prev_sample_mean, std_dev_t = compute_log_prob(
model, sample, timestep_index, config
)
if config.train.beta > 0:
with torch.no_grad():
if reference_model is None:
with disable_adapter(model, accelerator):
_, _, prev_sample_mean_ref, _ = compute_log_prob(
model, sample, timestep_index, config
)
else:
_, _, prev_sample_mean_ref, _ = compute_log_prob(
reference_model, sample, timestep_index, config
)
advantages_t = torch.clamp(
sample["advantages"][:, timestep_index],
-config.train.adv_clip_max,
config.train.adv_clip_max,
).to(device=log_prob.device, dtype=log_prob.dtype)
train_mask = sample["train_mask"].to(device=log_prob.device, dtype=log_prob.dtype)
train_weight_sum = train_mask.sum().clamp_min(1.0)
ratio = torch.exp(log_prob - sample["log_probs"][:, timestep_index])
log_prob_delta = log_prob - sample["log_probs"][:, timestep_index]
unclipped_loss = -advantages_t * ratio
clipped_loss = -advantages_t * torch.clamp(
ratio,
1.0 - config.train.clip_range,
1.0 + config.train.clip_range,
)
policy_loss = (
torch.maximum(unclipped_loss, clipped_loss) * train_mask
).sum() / train_weight_sum
if config.train.beta > 0:
kl_per_sample = (
((prev_sample_mean - prev_sample_mean_ref) ** 2)
.mean(dim=tuple(range(1, prev_sample_mean.ndim)), keepdim=False)
/ (2 * (std_dev_t.view(std_dev_t.shape[0], -1).mean(dim=1) ** 2 + 1e-6))
)
kl_loss = (kl_per_sample * train_mask).sum() / train_weight_sum
loss = policy_loss + config.train.beta * kl_loss
else:
kl_loss = None
loss = policy_loss
info["approx_kl"].append(
0.5
* (
((log_prob - sample["log_probs"][:, timestep_index]) ** 2) * train_mask
).sum()
/ train_weight_sum
)
info["clipfrac"].append(
(
(torch.abs(ratio - 1.0) > config.train.clip_range).float()
* train_mask
).sum()
/ train_weight_sum
)
info["clipfrac_gt_one"].append(
((ratio - 1.0 > config.train.clip_range).float() * train_mask).sum()
/ train_weight_sum
)
info["clipfrac_lt_one"].append(
((1.0 - ratio > config.train.clip_range).float() * train_mask).sum()
/ train_weight_sum
)
info["ratio_mean"].append((ratio * train_mask).sum() / train_weight_sum)
info["ratio_std"].append(
torch.sqrt(
((((ratio - ((ratio * train_mask).sum() / train_weight_sum)) ** 2) * train_mask).sum() / train_weight_sum)
+ 1e-12
)
)
info["log_prob_delta_mean"].append((log_prob_delta * train_mask).sum() / train_weight_sum)
info["std_dev_t_mean"].append(
(std_dev_t.view(std_dev_t.shape[0], -1).mean(dim=1) * train_mask).sum() / train_weight_sum
)
info[f"per_step/approx_kl_step_{timestep_index}"].append(info["approx_kl"][-1])
info[f"per_step/clipfrac_step_{timestep_index}"].append(info["clipfrac"][-1])
info[f"per_step/ratio_mean_step_{timestep_index}"].append(info["ratio_mean"][-1])
info[f"per_step/std_dev_t_step_{timestep_index}"].append(info["std_dev_t_mean"][-1])
info["policy_loss"].append(policy_loss)
info["loss"].append(loss)
if kl_loss is not None:
info["kl_loss"].append(kl_loss)
accelerator.backward(loss)
if accelerator.sync_gradients:
accelerator.clip_grad_norm_(model.parameters(), config.train.max_grad_norm)
optimizer.step()
optimizer.zero_grad()
if accelerator.sync_gradients:
reduced = {key: torch.mean(torch.stack(value)) for key, value in info.items()}
reduced = accelerator.reduce(reduced, reduction="mean")
reduced.update({"epoch": epoch, "inner_epoch": inner_epoch})
if accelerator.is_main_process:
wandb.log(reduced, step=global_step)
global_step += 1
if max_train_steps > 0 and global_step >= max_train_steps:
break
info = defaultdict(list)
if config.train.ema:
ema.step(trainable_parameters, global_step)
if max_train_steps > 0 and global_step >= max_train_steps:
break
if max_train_steps > 0 and global_step >= max_train_steps:
break
if accelerator.is_main_process:
logger.info("RL update phase finished in %.1fs", time.time() - update_start_time)
epoch += 1
logger.info("Training finished after %s epochs.", epoch)
if __name__ == "__main__":
app.run(main)