| """ |
| Gradient Ascent utilities for reward-guided diffusion generation. |
| |
| This module implements gradient ascent on the LRM reward score to guide |
| the diffusion process toward higher preference scores. |
| """ |
|
|
| import torch |
| import torch.nn.functional as F |
| from typing import Optional, Tuple, List, Literal |
| from tqdm import tqdm |
| from lr_scheduler import create_lr_scheduler, LRScheduler |
|
|
|
|
| class RewardGuidedDiffusion: |
| """ |
| Implements reward-guided generation using gradient ascent. |
| |
| During denoising, at specified timesteps, we: |
| 1. Compute the reward score for current latents |
| 2. Calculate gradients of reward w.r.t. latents |
| 3. Update latents in the direction that increases reward |
| |
| This guides generation toward higher preference scores. |
| """ |
| |
| def __init__( |
| self, |
| reward_model, |
| grad_scale: float = 1.0, |
| grad_timestep_range: Optional[Tuple[int, int]] = None, |
| num_grad_steps: int = 5, |
| grad_step_size: float = 0.1, |
| gradient_checkpoint: bool = False, |
| |
| lr_scheduler_type: Literal["constant", "linear", "cosine", "exponential", "step"] = "constant", |
| lr_scheduler_kwargs: Optional[dict] = None, |
| |
| use_momentum: bool = False, |
| momentum: float = 0.9, |
| use_nesterov: bool = False, |
| use_iso_projection: bool = False |
| ): |
| """ |
| Initialize reward-guided diffusion. |
| |
| Args: |
| reward_model: LRM reward model for computing preference scores |
| grad_scale: Scale factor for gradient updates (default: 1.0) |
| grad_timestep_range: Tuple of (min_t, max_t) for gradient ascent. |
| If None, applies to all timesteps. |
| num_grad_steps: Number of gradient ascent steps per timestep |
| grad_step_size: Step size for each gradient update (initial LR) |
| gradient_checkpoint: Whether to use gradient checkpointing |
| lr_scheduler_type: Type of LR scheduler ("constant", "linear", "cosine", "exponential", "step") |
| lr_scheduler_kwargs: Additional kwargs for LR scheduler (e.g., end_lr, min_lr, warmup_steps) |
| use_momentum: Whether to use momentum in gradient updates |
| momentum: Momentum coefficient (typically 0.9) |
| use_nesterov: Whether to use Nesterov momentum |
| use_iso_projection: Whether to use Iso Projection |
| """ |
| self.reward_model = reward_model |
| self.grad_scale = grad_scale |
| self.grad_timestep_range = grad_timestep_range |
| self.num_grad_steps = num_grad_steps |
| self.grad_step_size = grad_step_size |
| self.gradient_checkpoint = gradient_checkpoint |
| |
| |
| self.lr_scheduler_type = lr_scheduler_type |
| self.lr_scheduler_kwargs = lr_scheduler_kwargs or {} |
| self.lr_scheduler: Optional[LRScheduler] = None |
| self.global_lr_scheduler: Optional[LRScheduler] = None |
| |
| |
| self.use_momentum = use_momentum |
| self.momentum = momentum |
| self.use_nesterov = use_nesterov |
| self.velocity = None |
|
|
| self.use_iso_projection = use_iso_projection |
| |
| |
| self.grad_stats = [] |
| self.timestep_counter = 0 |
| |
| def should_apply_gradient(self, timestep: int) -> bool: |
| """Check if gradient ascent should be applied at this timestep.""" |
|
|
| if self.grad_timestep_range is None: |
| return False |
| |
| min_t, max_t = self.grad_timestep_range |
| return min_t <= timestep <= max_t |
| |
| @torch.enable_grad() |
| def compute_reward_gradient( |
| self, |
| latents: torch.Tensor, |
| prompt: str, |
| timestep: int, |
| ) -> Tuple[torch.Tensor, float]: |
| """ |
| Compute gradient of reward score w.r.t. latents in FP32 to prevent underflow. |
| """ |
| |
| latents_fp32 = latents.detach().to(torch.float32).clone() |
| latents_fp32.requires_grad_(True) |
| |
| |
| |
| |
| reward_score = self.reward_model.get_reward_score( |
| latents_fp32, |
| prompt, |
| timestep, |
| enable_grad=True |
| ) |
| |
| if reward_score.numel() > 1: |
| reward_score = reward_score.mean() |
| |
| |
| |
| |
| grad = torch.autograd.grad( |
| outputs=reward_score, |
| inputs=latents_fp32, |
| create_graph=False, |
| retain_graph=True, |
| allow_unused=True, |
| )[0] |
| |
| |
| if grad is None: |
| grad = torch.zeros_like(latents) |
| else: |
| grad = grad.to(latents.dtype) |
| |
| return grad, reward_score.item() |
| |
| def apply_gradient_ascent( |
| self, |
| latents: torch.Tensor, |
| prompt: str, |
| timestep: int, |
| base_noise: Optional[torch.Tensor] = None, |
| verbose: bool = True, |
| total_denoising_steps: Optional[int] = None, |
| ) -> Tuple[torch.Tensor, dict]: |
| |
| |
| original_latents = latents.detach().clone().to(torch.float32) |
| current_latents = torch.nn.Parameter(original_latents.clone()) |
| self.reward_model.unet.conv_in.weight.requires_grad_(True) |
| |
| |
| with torch.no_grad(): |
| initial_reward = self.reward_model.get_reward_score( |
| latents, |
| prompt, |
| timestep |
| ) |
| initial_reward_val = initial_reward.item() if initial_reward.numel() == 1 else initial_reward.mean().item() |
| |
| |
| grad_norms = [] |
| reward_history = [initial_reward_val] |
| lr_history = [] |
| |
| |
| reward = self.reward_model.get_reward_score( |
| current_latents.to(latents.dtype), |
| prompt, |
| timestep, |
| enable_grad=True |
| ) |
| |
| loss = -reward.mean() |
| loss.backward() |
|
|
| |
| raw_grad = current_latents.grad |
| reward_history.append(reward.mean().item()) |
| |
| |
| if raw_grad is not None and base_noise is not None and self.use_iso_projection: |
| gamma = 1e-8 |
| B = raw_grad.shape[0] |
|
|
| grad_flat = raw_grad.view(B, -1) |
| noise_flat = base_noise.view(B, -1).to(torch.float32) |
|
|
| |
| dot_product = (grad_flat * noise_flat).sum(dim=1, keepdim=True) |
| noise_norm_sq = (noise_flat * noise_flat).sum(dim=1, keepdim=True) |
|
|
| proj_scalar = dot_product / (noise_norm_sq + gamma) |
| proj_scalar = proj_scalar.view(B, 1, 1, 1) |
|
|
| |
| grad_parallel = proj_scalar * base_noise.to(torch.float32) |
| grad_perp = raw_grad - grad_parallel |
|
|
| |
| |
| |
| safe_proj_scalar = torch.clamp(proj_scalar, min=0.0) |
| |
| beta = 1.0 |
| safe_grad_parallel = beta * (safe_proj_scalar * base_noise.to(torch.float32)) |
|
|
| |
| |
| else: |
| grad_perp = raw_grad |
| if base_noise is None: |
| print("?? WARNING: base_noise missing. Skipping Iso-Marginal projection.") |
|
|
| |
| if grad_perp is not None: |
| max_grad = grad_perp.norm().item() |
| |
| if max_grad > 0: |
| kinetic_direction = grad_perp / (max_grad + 1e-8) |
| |
| |
| alpha = self.grad_step_size |
| |
| with torch.no_grad(): |
| rectified_latents = original_latents - (alpha * kinetic_direction) |
| else: |
| print("?? WARNING: Gradient exists but max value is 0.0") |
| rectified_latents = original_latents.clone() |
| alpha = 0.0 |
| else: |
| print("?? FATAL: PyTorch completely dropped the latent gradient!") |
| rectified_latents = original_latents.clone() |
| max_grad = 0.0 |
| alpha = 0.0 |
|
|
| if verbose: |
| print(f" Grad step | LR: {alpha:.6f} | Reward: {reward.mean().item():.4f} | Max Grad: {max_grad:.4f}") |
| |
| |
| final_latents = rectified_latents.detach().to(latents.dtype) |
| |
| with torch.no_grad(): |
| final_reward = self.reward_model.get_reward_score( |
| final_latents, prompt, timestep |
| ) |
| final_reward_val = final_reward.item() if final_reward.numel() == 1 else final_reward.mean().item() |
| |
| stats = { |
| 'timestep': timestep, |
| 'initial_reward': initial_reward_val, |
| 'final_reward': final_reward_val, |
| 'reward_improvement': final_reward_val - initial_reward_val, |
| 'grad_norms': [max_grad], |
| 'reward_history': reward_history, |
| 'lr_history': [alpha], |
| 'latent_change': (final_latents - original_latents.to(latents.dtype)).norm().item(), |
| } |
| |
| self.grad_stats.append(stats) |
| |
| return final_latents, stats |
|
|
| def get_statistics(self) -> dict: |
| """Get aggregated statistics across all gradient ascent applications.""" |
| if not self.grad_stats: |
| return {} |
| |
| total_improvement = sum(s['reward_improvement'] for s in self.grad_stats) |
| avg_improvement = total_improvement / len(self.grad_stats) |
| |
| all_grad_norms = [n for s in self.grad_stats for n in s['grad_norms']] |
| |
| return { |
| 'num_applications': len(self.grad_stats), |
| 'total_reward_improvement': total_improvement, |
| 'avg_reward_improvement': avg_improvement, |
| 'avg_grad_norm': sum(all_grad_norms) / len(all_grad_norms) if all_grad_norms else 0, |
| 'max_grad_norm': max(all_grad_norms) if all_grad_norms else 0, |
| 'detailed_stats': self.grad_stats, |
| } |
| |
| def reset_statistics(self): |
| """Reset statistics and global scheduler.""" |
| self.grad_stats = [] |
| self.global_lr_scheduler = None |
| self.timestep_counter = 0 |
|
|
|
|
| def create_reward_guided_generator( |
| reward_model, |
| grad_timestep_range: Tuple[int, int] = (500, 700), |
| grad_scale: float = 1.0, |
| num_grad_steps: int = 5, |
| grad_step_size: float = 0.1, |
| lr_scheduler_type: str = "constant", |
| lr_scheduler_kwargs: Optional[dict] = None, |
| use_momentum: bool = False, |
| momentum: float = 0.9, |
| use_nesterov: bool = False, |
| use_iso_projection: bool = False |
| ) -> RewardGuidedDiffusion: |
| """ |
| Convenience function to create a reward-guided diffusion generator. |
| |
| Args: |
| reward_model: LRM reward model |
| grad_timestep_range: Tuple of (min_t, max_t) for applying gradients |
| grad_scale: Scale factor for gradient magnitude |
| num_grad_steps: Number of gradient ascent iterations per timestep |
| grad_step_size: Step size for each gradient update (initial LR) |
| lr_scheduler_type: Type of LR scheduler |
| lr_scheduler_kwargs: Additional kwargs for LR scheduler |
| use_momentum: Whether to use momentum |
| momentum: Momentum coefficient |
| use_nesterov: Whether to use Nesterov momentum |
| use_iso_projection: Whether to use Iso Projection |
| |
| Returns: |
| RewardGuidedDiffusion instance |
| """ |
| return RewardGuidedDiffusion( |
| reward_model=reward_model, |
| grad_scale=grad_scale, |
| grad_timestep_range=grad_timestep_range, |
| num_grad_steps=num_grad_steps, |
| grad_step_size=grad_step_size, |
| lr_scheduler_type=lr_scheduler_type, |
| lr_scheduler_kwargs=lr_scheduler_kwargs, |
| use_momentum=use_momentum, |
| momentum=momentum, |
| use_nesterov=use_nesterov, |
| use_iso_projection= False |
| ) |
|
|