diff --git a/README.md b/README.md index 3d2b40ec..1a2ed03b 100644 --- a/README.md +++ b/README.md @@ -94,6 +94,7 @@ sh INSTALL_MEGATRON.sh | Multimodal FSDP finetuning | transformers | [Script](cookbook/mm/fsdp2.py) | | GRPO RL training | megatron | [Script](cookbook/rl/grpo/grpo.py) | | PPO RL training | transformers | [Script](cookbook/rl/ppo/ppo.py) | +| SAO synchronous correctness baseline | transformers | [Script](cookbook/rl/sao/sao_sync.py) / [Guide](cookbook/rl/sao/README.md) | | GRPO Multimodal RL training | megatron | [Script](cookbook/rl/grpo/grpo_mm.py) | | GRPO Math RL training | megatron | [Script](cookbook/rl/grpo/short_math_grpo.py) | | DPO full-parameter training | transformers | [Script](cookbook/rl/dpo/dpo_full.py) | diff --git a/README_ZH.md b/README_ZH.md index f2f214f4..62a0b0e3 100644 --- a/README_ZH.md +++ b/README_ZH.md @@ -88,6 +88,7 @@ sh INSTALL_MEGATRON.sh | 多模态 FSDP 微调 | transformers | [脚本](cookbook/mm/fsdp2.py) | | GRPO 强化学习训练 | megatron | [脚本](cookbook/rl/grpo/grpo.py) | | PPO 强化学习训练 | transformers | [脚本](cookbook/rl/ppo/ppo.py) | +| SAO 同步正确性基线 | transformers | [脚本](cookbook/rl/sao/sao_sync.py) / [说明](cookbook/rl/sao/README.md) | | GRPO 多模态强化学习训练 | megatron | [脚本](cookbook/rl/grpo/grpo_mm.py) | | GRPO 数学强化学习训练 | megatron | [脚本](cookbook/rl/grpo/short_math_grpo.py) | | DPO 全参数训练 | transformers | [脚本](cookbook/rl/dpo/dpo_full.py) | diff --git a/cookbook/rl/sao/sao_sync.py b/cookbook/rl/sao/sao_sync.py new file mode 100644 index 00000000..0a8d49eb --- /dev/null +++ b/cookbook/rl/sao/sao_sync.py @@ -0,0 +1,293 @@ +"""Synchronous algorithm-correctness baseline for SAO on GSM8K. + +This intentionally retains a rollout-batch barrier. It validates single rollout, +DIS, fixed critic targets, frozen-attention critic training, and a 2:1 +critic-to-actor update ratio. It is not the asynchronous SAO pipeline. +""" +import random +from typing import Any, Dict, List, Tuple + +from peft import LoraConfig + +import twinkle +from twinkle import DeviceGroup, DeviceMesh, get_device_placement, get_logger +from twinkle.advantage import GAEAdvantage, SAOGAEAdvantage +from twinkle.checkpoint_engine import CheckpointEngineManager +from twinkle.cli import CLI +from twinkle.data_format import SamplingParams +from twinkle.dataloader import DataLoader +from twinkle.dataset import Dataset, DatasetMeta +from twinkle.metric import CompletionRewardMetric, SAOMetric +from twinkle.model import TransformersModel, TransformersValueModel +from twinkle.preprocessor.llm import GSM8KProcessor +from twinkle.processor import InputProcessor +from twinkle.reward import GSM8KAccuracyReward, GSM8KFormatReward +from twinkle.sampler import vLLMSampler + +logger = get_logger() +args = CLI.from_args() + +MODEL_ID = args.model.model_id or 'ms://Qwen/Qwen3.5-4B' +POLICY_GPUS = args.infra.model_gpus or 4 +CRITIC_GPUS = args.infra.critic_model_gpus or 4 +SAMPLER_GPUS = args.infra.sampler_gpus or 4 +NUM_GPUS = POLICY_GPUS + CRITIC_GPUS + SAMPLER_GPUS +MAX_NEW_TOKENS = args.sampling.max_tokens or 1024 +POLICY_LR = args.optimizer.learning_rate +CRITIC_LR = args.rl.critic_learning_rate +MAX_STEPS = args.training.max_steps or 200 +BATCH_SIZE = args.training.batch_size or 4 +MINI_BATCH_SIZE = args.training.mini_batch_size or 4 +MICRO_BATCH_SIZE = args.training.micro_batch_size or 1 +SAVE_STEPS = args.training.save_steps or 50 +ADAPTER_NAME = args.lora.adapter_name or 'default' +CRITIC_UPDATES = args.rl.critic_updates_per_actor_update +EPSILON_HIGH = 5.0 if args.loss.epsilon_high is None else args.loss.epsilon_high + + +def create_gsm8k_dataset(): + dataset = Dataset(DatasetMeta('ms://modelscope/gsm8k', subset_name='main', split='train')) + dataset.set_template('Qwen3_5Template', model_id=MODEL_ID, max_length=400) + dataset.map(GSM8KProcessor()) + dataset.encode(add_generation_prompt=True) + return dataset + + +def compute_rewards(trajectories: List[Dict[str, Any]]) -> Tuple[List[float], List[float], List[float]]: + accuracy = GSM8KAccuracyReward()(trajectories) + formatting = GSM8KFormatReward()(trajectories) + return [a + f for a, f in zip(accuracy, formatting)], formatting, accuracy + + +def response_rows(full_values, trajectories) -> List[List[float]]: + import torch + value_rows = [] + tensors = full_values if isinstance(full_values, list) else [full_values] + for tensor in tensors: + if tensor is None: + continue + tensor = torch.as_tensor(tensor) + if tensor.dim() == 1: + tensor = tensor.unsqueeze(0) + value_rows.extend(tensor) + if len(value_rows) != len(trajectories): + raise ValueError(f'model output batch mismatch: {len(value_rows)} rows for {len(trajectories)} trajectories') + rows = [] + for value_row, trajectory in zip(value_rows, trajectories): + mask = torch.as_tensor(trajectory['labels'], device=value_row.device) != -100 + if value_row.numel() < mask.numel(): + raise ValueError( + f'value sequence is shorter than labels: values={value_row.numel()}, labels={mask.numel()}') + response_values = value_row[:mask.numel()][mask] + if response_values.numel() == 0: + raise ValueError('trajectory contains no action-token value predictions') + rows.append(response_values.detach().float().cpu().tolist()) + return rows + + +def pad_rows(rows, lengths, fill=0.0): + max_len = max(lengths) + return [list(row) + [fill] * (max_len - len(row)) for row in rows] + + +def main(): + if args.rl.num_generations != 1: + raise ValueError('SAO requires --num-generations 1') + if CRITIC_UPDATES <= 0: + raise ValueError('--critic-updates-per-actor-update must be positive') + if min(BATCH_SIZE, MINI_BATCH_SIZE, MICRO_BATCH_SIZE) <= 0: + raise ValueError('batch-size, mini-batch-size, and micro-batch-size must be positive') + if MINI_BATCH_SIZE > BATCH_SIZE: + raise ValueError('--mini-batch-size cannot exceed --batch-size') + if MINI_BATCH_SIZE % MICRO_BATCH_SIZE != 0: + raise ValueError('--mini-batch-size must be divisible by --micro-batch-size') + # DIS compares vLLM rollout log-probs with the learner's raw model log-probs. + # Temperature/top-k/top-p/repetition transforms would make those two + # distributions different and invalidate the importance ratio. + if args.sampling.temperature != 1.0: + raise ValueError('SAO requires --temperature 1.0 for an exact rollout/current log-prob ratio') + if args.sampling.top_p != 1.0 or args.sampling.top_k != -1: + raise ValueError('SAO requires --top-p 1.0 and --top-k -1 for an exact behavior-policy log-prob') + if args.sampling.repetition_penalty != 1.0: + raise ValueError('SAO requires --repetition-penalty 1.0 for an exact behavior-policy log-prob') + if BATCH_SIZE % MINI_BATCH_SIZE != 0: + logger.warning('batch-size is not divisible by mini-batch-size; the final actor/critic update ' + 'of each rollout batch will use a smaller mini-batch') + actor_updates_per_rollout = (BATCH_SIZE + MINI_BATCH_SIZE - 1) // MINI_BATCH_SIZE + if actor_updates_per_rollout != 1: + logger.warning( + f'Each rollout batch performs {actor_updates_per_rollout} actor optimizer steps. ' + 'For the paper-style one update over a global batch of 128, set ' + '--batch-size 128 --mini-batch-size 128.') + + critic_start = POLICY_GPUS + sampler_start = POLICY_GPUS + CRITIC_GPUS + groups = [ + DeviceGroup(name='policy', ranks=list(range(POLICY_GPUS)), device_type='GPU'), + DeviceGroup(name='critic', ranks=list(range(critic_start, sampler_start)), device_type='GPU'), + DeviceGroup(name='sampler', ranks=list(range(sampler_start, NUM_GPUS)), device_type='GPU'), + ] + policy_mesh = DeviceMesh.from_sizes(world_size=POLICY_GPUS, fsdp_size=POLICY_GPUS) + critic_mesh = DeviceMesh.from_sizes(world_size=CRITIC_GPUS, fsdp_size=CRITIC_GPUS) + sampler_mesh = DeviceMesh.from_sizes(world_size=SAMPLER_GPUS, dp_size=SAMPLER_GPUS) + twinkle.initialize(mode='ray', nproc_per_node=NUM_GPUS, groups=groups, lazy_collect=False) + + policy = TransformersModel(model_id=MODEL_ID, device_mesh=policy_mesh, remote_group='policy') + policy.add_adapter_to_model( + ADAPTER_NAME, + LoraConfig( + target_modules=['q_proj', 'k_proj', 'v_proj', 'o_proj', 'gate_proj', 'up_proj', 'down_proj'], + r=32, lora_alpha=64, lora_dropout=0.05), + gradient_accumulation_steps=1, + ) + policy.set_optimizer('AdamW', lr=POLICY_LR) + policy.set_lr_scheduler('CosineAnnealingLR', T_max=MAX_STEPS, eta_min=0) + policy.set_loss( + 'SAOLoss', epsilon_low=args.loss.epsilon_low, epsilon_high=EPSILON_HIGH, + detach_importance_weight=args.loss.detach_importance_weight, entropy_coef=args.loss.entropy_coef) + policy.add_metric(SAOMetric, epsilon=args.loss.epsilon_low, epsilon_high=EPSILON_HIGH) + policy.set_processor(InputProcessor) + policy.set_template('Qwen3_5Template', model_id=MODEL_ID) + + critic = TransformersValueModel(model_id=MODEL_ID, device_mesh=critic_mesh, remote_group='critic') + if args.rl.freeze_critic_attention: + frozen = critic.freeze_attention_for_value_training() + logger.info(f'Frozen-attention critic: {frozen}') + logger.info(f'Critic parameters: {critic.trainable_parameter_summary()}') + critic.set_optimizer('AdamW', lr=CRITIC_LR) + critic.set_lr_scheduler( + 'LinearWarmupScheduler', num_warmup_steps=10, + num_training_steps=MAX_STEPS * CRITIC_UPDATES) + critic.set_loss('SAOValueLoss') + critic.set_processor(InputProcessor) + critic.set_template('Qwen3_5Template', model_id=MODEL_ID) + + sampler = vLLMSampler( + model_id=MODEL_ID, + engine_args={ + 'gpu_memory_utilization': 0.8, + 'max_model_len': 400 + MAX_NEW_TOKENS, + 'max_lora_rank': 32, + 'enable_lora': True, + 'tensor_parallel_size': 1, + }, + device_mesh=sampler_mesh, + remote_group='sampler', + ) + sampler.set_template('Qwen3_5Template', model_id=MODEL_ID) + checkpoint_manager = CheckpointEngineManager(model=policy, sampler=sampler) + dataloader = DataLoader( + dataset=create_gsm8k_dataset, batch_size=BATCH_SIZE, min_batch_size=BATCH_SIZE, + device_mesh=policy_mesh, remote_group='policy') + critic_gae = SAOGAEAdvantage( + gamma=args.rl.gamma, gae_lambda=args.rl.sao_critic_lambda, normalize=False) + actor_gae = SAOGAEAdvantage( + gamma=args.rl.gamma, alpha=args.rl.sao_alpha, + gae_lambda=None if args.rl.sao_policy_lambda_adaptive else args.rl.gae_lambda, + normalize=args.rl.normalize_advantages) + reward_metric = CompletionRewardMetric() + sampling_params = SamplingParams( + max_tokens=MAX_NEW_TOKENS, + seed=args.sampling.seed, + stop=args.sampling.stop, + temperature=args.sampling.temperature, + top_k=args.sampling.top_k, + top_p=args.sampling.top_p, + repetition_penalty=args.sampling.repetition_penalty, + num_samples=1, + logprobs=1, + ) + + optim_step = 0 + rollout_step = 0 + logger.info(get_device_placement()) + while optim_step < MAX_STEPS: + for batch in dataloader: + if optim_step >= MAX_STEPS: + break + reward_metric.reset() + prompts = batch if isinstance(batch, list) else [batch] + checkpoint_manager.sync_weights(merge_and_sync=False) + sampler.reset_prefix_cache() + samples = sampler.sample(prompts, sampling_params) + + trajectories, rollout_logps, lengths = [], [], [] + for response in samples: + for sequence in response.sequences: + if not sequence.tokens or sequence.logprobs is None: + raise ValueError('SAO rollout must contain generated tokens and log-probabilities') + action_count = sum(label != -100 for label in sequence.new_input_feature['labels']) + if len(sequence.tokens) != len(sequence.logprobs) or action_count != len(sequence.tokens): + raise ValueError( + 'SAO rollout token alignment failed: ' + f'tokens={len(sequence.tokens)}, logprobs={len(sequence.logprobs)}, ' + f'action_labels={action_count}') + trajectories.append(sequence.new_input_feature) + rollout_logps.append([entry[0][1] for entry in sequence.logprobs]) + lengths.append(len(sequence.tokens)) + if not trajectories: + logger.warning('No trajectories in rollout batch; skipping learner update') + continue + rewards, format_rewards, accuracy_rewards = compute_rewards(trajectories) + reward_metric.accumulate( + completion_lengths=lengths, + rewards={'total': rewards, 'format': format_rewards, 'accuracy': accuracy_rewards}) + token_rewards = GAEAdvantage.build_token_rewards(rewards, lengths) + padded_rewards = pad_rows(token_rewards, lengths) + masks = [[True] * length + [False] * (max(lengths) - length) for length in lengths] + terminated = [True] * len(trajectories) + truncated = [False] * len(trajectories) + + initial = critic.forward_only(inputs=trajectories) + initial_values = response_rows(initial['values'], trajectories) + _, fixed_returns = critic_gae( + padded_rewards, pad_rows(initial_values, lengths), action_masks=masks, + terminated=terminated, truncated=truncated, effective_lengths=lengths) + fixed_returns = [fixed_returns[i, :length].tolist() for i, length in enumerate(lengths)] + + indices = list(range(len(trajectories))) + random.shuffle(indices) + for start in range(0, len(indices), MINI_BATCH_SIZE): + chosen = indices[start:start + MINI_BATCH_SIZE] + mb_inputs = [trajectories[i] for i in chosen] + mb_returns = [fixed_returns[i] for i in chosen] + for _ in range(CRITIC_UPDATES): + critic.forward_backward( + inputs=mb_inputs, returns=mb_returns, micro_batch_size=MICRO_BATCH_SIZE) + critic.clip_grad_and_step() + + updated = critic.forward_only(inputs=mb_inputs) + new_values = response_rows(updated['values'], mb_inputs) + mb_lengths = [lengths[i] for i in chosen] + mb_rewards = [token_rewards[i] for i in chosen] + mb_masks = [[True] * length + [False] * (max(mb_lengths) - length) for length in mb_lengths] + advantages, _ = actor_gae( + pad_rows(mb_rewards, mb_lengths), pad_rows(new_values, mb_lengths), action_masks=mb_masks, + terminated=[True] * len(chosen), truncated=[False] * len(chosen), + effective_lengths=mb_lengths) + mb_advantages = [advantages[i, :length].tolist() for i, length in enumerate(mb_lengths)] + policy.forward_backward( + inputs=mb_inputs, old_logps=[rollout_logps[i] for i in chosen], advantages=mb_advantages, + micro_batch_size=MICRO_BATCH_SIZE) + policy.clip_grad_and_step() + optim_step += 1 + if optim_step % SAVE_STEPS == 0: + policy.save(f'sao-policy-checkpoint-{optim_step}') + critic.save(f'sao-critic-checkpoint-{optim_step}') + if optim_step >= MAX_STEPS: + break + + logs = reward_metric.calculate() + logs.update(policy.calculate_metric(is_training=True)) + logs.update({f'critic/{key}': value for key, value in critic.calculate_metric(is_training=True).items()}) + logs['train/critic_updates_per_actor_update'] = CRITIC_UPDATES + logs['train/actor_updates_per_rollout_batch'] = actor_updates_per_rollout + rollout_step += 1 + logger.info(f'[SAO sync rollout {rollout_step}, actor step {optim_step}/{MAX_STEPS}] {logs}') + + policy.save('sao-policy-final') + critic.save('sao-critic-final') + + +if __name__ == '__main__': + main() diff --git a/cookbook/rl/sao/sao_sync.sh b/cookbook/rl/sao/sao_sync.sh new file mode 100644 index 00000000..f760aa9a --- /dev/null +++ b/cookbook/rl/sao/sao_sync.sh @@ -0,0 +1,29 @@ +#!/bin/sh +set -eu + +# Synchronous SAO correctness baseline: 4 policy + 4 critic + 4 sampler GPUs. +# This is deliberately not the asynchronous rollout/learner pipeline. +python sao_sync.py \ + --model-id ms://Qwen/Qwen3.5-4B \ + --model-gpus 4 \ + --critic-model-gpus 4 \ + --sampler-gpus 4 \ + --num-generations 1 \ + --max-tokens 1024 \ + --batch-size 4 \ + --mini-batch-size 4 \ + --micro-batch-size 1 \ + --gamma 1.0 \ + --sao-alpha 1.5 \ + --sao-critic-lambda 1.0 \ + --critic-updates-per-actor-update 2 \ + --epsilon-low 0.3 \ + --epsilon-high 5.0 \ + --detach-importance-weight \ + --freeze-critic-attention \ + --lr 1e-6 \ + --critic-learning-rate 5e-6 \ + --max-steps 200 \ + --save-steps 50 \ + --adapter-name default \ + "$@" diff --git a/src/twinkle/advantage/__init__.py b/src/twinkle/advantage/__init__.py index 519dd1fe..52e855b6 100644 --- a/src/twinkle/advantage/__init__.py +++ b/src/twinkle/advantage/__init__.py @@ -3,10 +3,12 @@ from .gae import GAEAdvantage from .grpo import GRPOAdvantage from .rloo import RLOOAdvantage +from .sao_gae import SAOGAEAdvantage __all__ = [ 'Advantage', 'GAEAdvantage', 'GRPOAdvantage', 'RLOOAdvantage', + 'SAOGAEAdvantage', ] diff --git a/src/twinkle/advantage/sao_gae.py b/src/twinkle/advantage/sao_gae.py new file mode 100644 index 00000000..b6662f2c --- /dev/null +++ b/src/twinkle/advantage/sao_gae.py @@ -0,0 +1,118 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from typing import TYPE_CHECKING, List, Optional, Tuple, Union + +from .base import Advantage + +if TYPE_CHECKING: + import torch + + +class SAOGAEAdvantage(Advantage): + """Batched skip-observation GAE with terminal/truncation semantics. + + True positions in ``action_masks`` form the Bellman chain. Consequently, + prompt, padding, and environment-observation tokens are skipped. + """ + + def __init__( + self, + gamma: float = 1.0, + alpha: float = 1.5, + gae_lambda: Optional[float] = None, + normalize: bool = True, + ): + if not 0.0 <= gamma <= 1.0: + raise ValueError('gamma must be in [0, 1]') + if alpha <= 0.0: + raise ValueError('alpha must be positive') + if gae_lambda is not None and not 0.0 <= gae_lambda <= 1.0: + raise ValueError('gae_lambda must be in [0, 1]') + self.gamma = gamma + self.alpha = alpha + self.gae_lambda = gae_lambda + self.normalize = normalize + + def __call__( + self, + rewards: Union['torch.Tensor', List[List[float]]], + values: Union['torch.Tensor', List[List[float]]], + *, + action_masks: Union['torch.Tensor', List[List[bool]]], + terminated: Union['torch.Tensor', List[bool]], + truncated: Union['torch.Tensor', List[bool]], + bootstrap_values: Optional[Union['torch.Tensor', List[Optional[float]]]] = None, + effective_lengths: Optional[Union['torch.Tensor', List[int]]] = None, + normalize: Optional[bool] = None, + **kwargs, + ) -> Tuple['torch.Tensor', 'torch.Tensor']: + import torch + + rewards = torch.as_tensor(rewards, dtype=torch.float32) + values = torch.as_tensor(values, dtype=torch.float32, device=rewards.device) + action_masks = torch.as_tensor(action_masks, dtype=torch.bool, device=rewards.device) + if rewards.dim() == 1: + rewards, values, action_masks = rewards.unsqueeze(0), values.unsqueeze(0), action_masks.unsqueeze(0) + if rewards.shape != values.shape or rewards.shape != action_masks.shape: + raise ValueError('rewards, values, and action_masks must have identical shapes') + + batch_size = rewards.shape[0] + terminated = torch.as_tensor(terminated, dtype=torch.bool, device=rewards.device).flatten() + truncated = torch.as_tensor(truncated, dtype=torch.bool, device=rewards.device).flatten() + if terminated.numel() != batch_size or truncated.numel() != batch_size: + raise ValueError('terminated and truncated must have one value per sequence') + if bool((terminated & truncated).any()): + raise ValueError('a trajectory cannot be both terminated and truncated') + + if bootstrap_values is None: + bootstrap = [None] * batch_size + else: + bootstrap = list(bootstrap_values) + if len(bootstrap) != batch_size: + raise ValueError('bootstrap_values must have one value per sequence') + + if effective_lengths is None: + lengths = action_masks.sum(dim=-1).tolist() + else: + lengths = torch.as_tensor(effective_lengths).flatten().tolist() + if len(lengths) != batch_size: + raise ValueError('effective_lengths must have one value per sequence') + + advantages = torch.zeros_like(rewards) + for batch_idx in range(batch_size): + positions = action_masks[batch_idx].nonzero(as_tuple=True)[0] + if positions.numel() == 0: + raise ValueError(f'trajectory {batch_idx} has no action tokens') + if lengths[batch_idx] <= 0: + raise ValueError('effective lengths must be positive') + if bool(truncated[batch_idx]) and bootstrap[batch_idx] is None: + raise ValueError(f'truncated trajectory {batch_idx} requires a bootstrap value') + if not bool(terminated[batch_idx]) and not bool(truncated[batch_idx]): + raise ValueError(f'trajectory {batch_idx} must be terminated or truncated') + + lambda_value = self.gae_lambda + if lambda_value is None: + lambda_value = 1.0 - 1.0 / (self.alpha * float(lengths[batch_idx])) + lambda_value = min(1.0, max(0.0, lambda_value)) + + next_advantage = rewards.new_zeros(()) + for index in range(positions.numel() - 1, -1, -1): + position = positions[index] + if index + 1 < positions.numel(): + next_value = values[batch_idx, positions[index + 1]].detach() + elif bool(truncated[batch_idx]): + next_value = torch.as_tensor(bootstrap[batch_idx], device=rewards.device, dtype=torch.float32) + else: + next_value = rewards.new_zeros(()) + delta = rewards[batch_idx, position] + self.gamma * next_value - values[batch_idx, position] + next_advantage = delta + self.gamma * lambda_value * next_advantage + advantages[batch_idx, position] = next_advantage + + returns = (advantages + values.detach()).masked_fill(~action_masks, 0.0) + should_normalize = self.normalize if normalize is None else normalize + if should_normalize: + valid = advantages[action_masks] + if valid.numel() > 1: + normalized = (advantages - valid.mean()) / valid.std(unbiased=False).clamp_min(1e-8) + advantages = torch.where(action_masks, normalized, advantages) + advantages = advantages.masked_fill(~action_masks, 0.0) + return advantages, returns diff --git a/src/twinkle/cli/cli.py b/src/twinkle/cli/cli.py index d88960a9..af9edb5c 100644 --- a/src/twinkle/cli/cli.py +++ b/src/twinkle/cli/cli.py @@ -136,6 +136,8 @@ class LossArgs: sft_weight: float = 1.0 entropy_coef: float = 0.0 value_clip: float = 0.2 + epsilon_low: float = 0.3 + detach_importance_weight: bool = True ignore_index: int = -100 @@ -213,6 +215,11 @@ class RLArgs: kl_coef: float = 0.0 normalize_advantages: bool = True critic_learning_rate: float = 1e-5 + critic_updates_per_actor_update: int = 2 + sao_alpha: float = 1.5 + sao_policy_lambda_adaptive: bool = True + sao_critic_lambda: float = 1.0 + freeze_critic_attention: bool = True @dataclass diff --git a/src/twinkle/loss/__init__.py b/src/twinkle/loss/__init__.py index 18b93f5a..fe97c20c 100644 --- a/src/twinkle/loss/__init__.py +++ b/src/twinkle/loss/__init__.py @@ -8,7 +8,8 @@ from .infonce import InfonceLoss from .liger_fused_linear_cross_entropy import LigerFusedLinearCrossEntropyLoss from .mse import MSELoss -from .value import PPOValueLoss +from .sao import SAOLoss +from .value import PPOValueLoss, SAOValueLoss torch_loss_mapping = { 'mse': MSELoss, @@ -21,6 +22,8 @@ 'grpo': GRPOLoss, 'ppo': PPOLoss, 'ppo_value': PPOValueLoss, + 'sao': SAOLoss, + 'sao_value': SAOValueLoss, 'gspo': GSPOLoss, 'sapo': SAPOLoss, 'cispo': CISPOLoss, diff --git a/src/twinkle/loss/policy_objective.py b/src/twinkle/loss/policy_objective.py new file mode 100644 index 00000000..4a7253f3 --- /dev/null +++ b/src/twinkle/loss/policy_objective.py @@ -0,0 +1,49 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + import torch + + +class PolicyObjective: + """Base interface for a per-token policy optimization objective.""" + + def __call__( + self, + ratio: 'torch.Tensor', + advantages: 'torch.Tensor', + per_token_logps: 'torch.Tensor', + ) -> 'torch.Tensor': + raise NotImplementedError + + +class DISPolicyObjective(PolicyObjective): + """Direct double-sided importance-sampling objective used by SAO.""" + + def __init__( + self, + epsilon_low: float = 0.3, + epsilon_high: float = 5.0, + detach_importance_weight: bool = True, + ): + if not 0.0 <= epsilon_low < 1.0: + raise ValueError('epsilon_low must be in [0, 1)') + if epsilon_high < 0.0: + raise ValueError('epsilon_high must be non-negative') + self.epsilon_low = epsilon_low + self.epsilon_high = epsilon_high + self.detach_importance_weight = detach_importance_weight + + def __call__( + self, + ratio: 'torch.Tensor', + advantages: 'torch.Tensor', + per_token_logps: 'torch.Tensor', + ) -> 'torch.Tensor': + import torch + + trusted = (ratio > 1.0 - self.epsilon_low) & (ratio < 1.0 + self.epsilon_high) + weight = torch.where(trusted, ratio, torch.zeros_like(ratio)) + if self.detach_importance_weight: + weight = weight.detach() + return -weight * advantages.detach() * per_token_logps.float() diff --git a/src/twinkle/loss/sao.py b/src/twinkle/loss/sao.py new file mode 100644 index 00000000..69a79865 --- /dev/null +++ b/src/twinkle/loss/sao.py @@ -0,0 +1,42 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from typing import TYPE_CHECKING + +from .grpo import GRPOLoss +from .policy_objective import DISPolicyObjective + +if TYPE_CHECKING: + import torch + + +class SAOLoss(GRPOLoss): + """SAO direct double-sided importance-sampling policy loss. + + Tokens whose current/rollout policy ratio lies outside the strict trust + interval are assigned zero weight instead of being clipped to a boundary. + """ + + def __init__( + self, + epsilon_low: float = 0.3, + epsilon_high: float = 5.0, + detach_importance_weight: bool = True, + **kwargs, + ): + self.policy_objective = DISPolicyObjective( + epsilon_low=epsilon_low, + epsilon_high=epsilon_high, + detach_importance_weight=detach_importance_weight, + ) + super().__init__(epsilon=epsilon_low, epsilon_high=epsilon_high, **kwargs) + + def _compute_per_token_loss( + self, + ratio: 'torch.Tensor', + advantages: 'torch.Tensor', + per_token_logps: 'torch.Tensor', + ) -> 'torch.Tensor': + return self.policy_objective(ratio, advantages, per_token_logps) + + def _aggregate_loss(self, per_token_loss, loss_mask, **kwargs): + mask = loss_mask.to(per_token_loss.dtype) + return (per_token_loss * mask).sum() / mask.sum().clamp(min=1.0) diff --git a/src/twinkle/loss/value.py b/src/twinkle/loss/value.py index 563cdaef..870fd945 100644 --- a/src/twinkle/loss/value.py +++ b/src/twinkle/loss/value.py @@ -60,3 +60,35 @@ def __call__( mask_f = mask.to(values.dtype) loss = (per_token_loss * mask_f).sum() / mask_f.sum().clamp(min=1.0) return LossOutput(loss=loss, num_tokens=0) + + +class SAOValueLoss(Loss): + """Masked MSE critic loss used by the SAO baseline.""" + + require_logps = False + require_values = True + + def __init__(self, ignore_index: int = -100, **kwargs): + self.ignore_index = ignore_index + self._aligner = GRPOLoss(ignore_index=ignore_index) + + def __call__(self, inputs: Dict, outputs: Dict, *, returns=None, **kwargs) -> LossOutput: + import torch + + labels = torch.as_tensor(inputs.get('labels')) + if labels.dim() == 1: + labels = labels.unsqueeze(0) + mask = labels != self.ignore_index + values = outputs.get('values') + assert values is not None, "outputs must contain 'values'" + if values.dim() == 3 and values.shape[-1] == 1: + values = values.squeeze(-1) + if values.dim() == 1: + values = values.unsqueeze(0) + if values.shape != mask.shape: + raise AssertionError(f'values/mask shape mismatch: values={tuple(values.shape)} mask={tuple(mask.shape)}') + assert returns is not None, 'returns are required for SAO value loss' + returns = self._aligner._pad_and_align_to_batch(returns, mask, values.device, values.dtype) + mask_f = mask.to(values.dtype) + loss = ((values.float() - returns.detach().float()).square() * mask_f).sum() / mask_f.sum().clamp(min=1.0) + return LossOutput(loss=loss, num_tokens=0) diff --git a/src/twinkle/metric/__init__.py b/src/twinkle/metric/__init__.py index f6ac5120..1cd4fc25 100644 --- a/src/twinkle/metric/__init__.py +++ b/src/twinkle/metric/__init__.py @@ -4,7 +4,7 @@ from .completion_and_reward import CompletionRewardMetric from .dpo import DPOMetric from .embedding import EmbeddingMetric -from .grpo import CISPOMetric, GRPOMetric, GSPOMetric, PPOMetric +from .grpo import CISPOMetric, GRPOMetric, GSPOMetric, PPOMetric, SAOMetric from .loss import LossMetric from .ppo import PPOValueMetric from .train_metric import TrainMetric diff --git a/src/twinkle/metric/grpo.py b/src/twinkle/metric/grpo.py index e3eaacd2..d51caf3d 100644 --- a/src/twinkle/metric/grpo.py +++ b/src/twinkle/metric/grpo.py @@ -377,6 +377,19 @@ class PPOMetric(GRPOMetric): """PPO policy metric; shares token-level ratio and clipping statistics with GRPO.""" +class SAOMetric(GRPOMetric): + """SAO metric; trust rejection is unconditional on advantage sign.""" + + def _accumulate_clip(self, log_ratio, advantages, mask, mask_f): + import torch + ratio = torch.exp(log_ratio.clamp(min=-20.0, max=20.0)) + is_low = ratio <= 1 - self.epsilon + is_high = ratio >= 1 + self.epsilon_high + self.sum_clip_low += float((is_low.float() * mask_f).sum().item()) + self.sum_clip_high += float((is_high.float() * mask_f).sum().item()) + self.clip_n_total += float(mask_f.sum().item()) + + class GSPOMetric(GRPOMetric): """GRPOMetric variant for GSPO: clip applies to per-sequence geometric-mean ratio.""" diff --git a/src/twinkle/model/transformers/value_model.py b/src/twinkle/model/transformers/value_model.py index 110650f6..b82b87b1 100644 --- a/src/twinkle/model/transformers/value_model.py +++ b/src/twinkle/model/transformers/value_model.py @@ -3,13 +3,22 @@ from torch import nn from typing import Optional -from twinkle import DeviceMesh, remote_class +from twinkle import DeviceMesh, remote_class, remote_function from .transformers import TransformersModel @remote_class() class TransformersValueModel(TransformersModel): - """Transformers causal-LM backbone with a scalar value head.""" + """Transformers causal-LM backbone with a scalar value head. + + The implementation supports causal language models with hybrid attention. + In particular, Qwen3.5 alternates full-attention decoder layers + (``self_attn``) with GatedDeltaNet linear-attention layers + (``linear_attn``). Both token mixers are treated as attention modules + when applying the frozen-attention critic training strategy. + """ + + _ATTENTION_ATTRIBUTES = ('self_attn', 'linear_attn', 'attn') def __init__(self, *args, device_mesh: Optional[DeviceMesh] = None, **kwargs): super().__init__(*args, device_mesh=device_mesh, **kwargs) @@ -23,6 +32,40 @@ def __init__(self, *args, device_mesh: Optional[DeviceMesh] = None, **kwargs): self.model.set_output_embeddings(value_head) self.model.config.tie_word_embeddings = False + @remote_function(dispatch='all', collect='first', lazy_collect=False) + def freeze_attention_for_value_training(self): + """Freeze attention/token-mixer modules while leaving feed-forward/value-head parameters trainable. + + ``self_attn`` covers conventional and Qwen3.5 full-attention decoder + layers, ``linear_attn`` covers Qwen3.5 GatedDeltaNet layers, and + ``attn`` preserves compatibility with models + such as GPT-2 and vision backbones. + """ + model = self.strategy.unwrap_model(self.model) + attention_modules = [] + seen_attention_ids = set() + for module in model.modules(): + for attribute in self._ATTENTION_ATTRIBUTES: + attention = getattr(module, attribute, None) + if isinstance(attention, nn.Module) and id(attention) not in seen_attention_ids: + attention_modules.append(attention) + seen_attention_ids.add(id(attention)) + if not attention_modules: + raise ValueError(f'No decoder attention modules found for {type(model).__name__}') + frozen = 0 + for attention in attention_modules: + for parameter in attention.parameters(): + parameter.requires_grad = False + frozen += parameter.numel() + return {'attention_modules': len(attention_modules), 'frozen_parameters': frozen} + + @remote_function(dispatch='all', collect='first', lazy_collect=False) + def trainable_parameter_summary(self): + model = self.strategy.unwrap_model(self.model) + total = sum(parameter.numel() for parameter in model.parameters()) + trainable = sum(parameter.numel() for parameter in model.parameters() if parameter.requires_grad) + return {'total_parameters': total, 'trainable_parameters': trainable, 'frozen_parameters': total - trainable} + def add_adapter_to_model(self, *args, **kwargs): raise NotImplementedError('PPO critic LoRA is not supported; train the critic as a full-parameter model') diff --git a/tests/advantage/test_sao_gae.py b/tests/advantage/test_sao_gae.py new file mode 100644 index 00000000..a28fab9f --- /dev/null +++ b/tests/advantage/test_sao_gae.py @@ -0,0 +1,46 @@ +import pytest +import torch + +from twinkle.advantage import SAOGAEAdvantage + + +def test_skip_observation_gae_crosses_observation(): + gae = SAOGAEAdvantage(gamma=1.0, gae_lambda=1.0, normalize=False) + advantages, returns = gae( + [[0.0, 99.0, 1.0]], [[0.2, 42.0, 0.4]], + action_masks=[[True, False, True]], terminated=[True], truncated=[False]) + torch.testing.assert_close(advantages, torch.tensor([[0.8, 0.0, 0.6]])) + torch.testing.assert_close(returns, torch.tensor([[1.0, 0.0, 1.0]])) + + +def test_terminal_has_zero_bootstrap(): + gae = SAOGAEAdvantage(gamma=1.0, gae_lambda=1.0, normalize=False) + advantages, _ = gae([[2.0]], [[0.5]], action_masks=[[True]], terminated=[True], truncated=[False]) + assert advantages.item() == pytest.approx(1.5) + + +def test_truncated_requires_and_uses_bootstrap(): + gae = SAOGAEAdvantage(gamma=1.0, gae_lambda=1.0, normalize=False) + with pytest.raises(ValueError, match='requires a bootstrap'): + gae([[0.0]], [[0.5]], action_masks=[[True]], terminated=[False], truncated=[True]) + advantages, _ = gae( + [[0.0]], [[0.5]], action_masks=[[True]], terminated=[False], truncated=[True], bootstrap_values=[2.0]) + assert advantages.item() == pytest.approx(1.5) + + +def test_batch_sequences_do_not_link(): + gae = SAOGAEAdvantage(gamma=1.0, gae_lambda=1.0, normalize=False) + advantages, _ = gae( + [[1.0, 0.0], [3.0, 0.0]], torch.zeros(2, 2), + action_masks=[[True, False], [True, False]], terminated=[True, True], truncated=[False, False]) + torch.testing.assert_close(advantages[:, 0], torch.tensor([1.0, 3.0])) + + +def test_length_adaptive_lambda_is_per_sequence(): + gae = SAOGAEAdvantage(gamma=1.0, alpha=1.0, normalize=False) + advantages, _ = gae( + [[0.0, 1.0], [0.0, 1.0]], torch.zeros(2, 2), + action_masks=[[True, True], [True, True]], terminated=[True, True], truncated=[False, False], + effective_lengths=[2, 4]) + assert advantages[0, 0].item() == pytest.approx(0.5) + assert advantages[1, 0].item() == pytest.approx(0.75) diff --git a/tests/cli/test_cli.py b/tests/cli/test_cli.py index 3b8847d3..22b18dc5 100644 --- a/tests/cli/test_cli.py +++ b/tests/cli/test_cli.py @@ -334,6 +334,26 @@ def test_cli_ppo_fields(self): assert args.rl.critic_learning_rate == pytest.approx(2e-5) assert args.loss.value_clip == pytest.approx(0.1) + def test_cli_sao_fields(self): + args = CLI.from_args(argv=[ + '--num-generations', '1', + '--epsilon-low', '0.3', + '--epsilon-high', '5.0', + '--no-detach-importance-weight', + '--critic-updates-per-actor-update', '2', + '--sao-alpha', '1.5', + '--sao-critic-lambda', '1.0', + '--no-freeze-critic-attention', + ]) + assert args.rl.num_generations == 1 + assert args.loss.epsilon_low == pytest.approx(0.3) + assert args.loss.epsilon_high == pytest.approx(5.0) + assert args.loss.detach_importance_weight is False + assert args.rl.critic_updates_per_actor_update == 2 + assert args.rl.sao_alpha == pytest.approx(1.5) + assert args.rl.sao_critic_lambda == pytest.approx(1.0) + assert args.rl.freeze_critic_attention is False + def test_use_megatron_true_flips_strategy(self): args = CLI.from_args(argv=['--use_megatron', 'true']) assert args.model.strategy == 'native_fsdp' diff --git a/tests/loss/test_sao.py b/tests/loss/test_sao.py new file mode 100644 index 00000000..9be3e31e --- /dev/null +++ b/tests/loss/test_sao.py @@ -0,0 +1,86 @@ +import math + +import pytest +import torch + +from twinkle.loss import SAOLoss, SAOValueLoss +from twinkle.loss.policy_objective import DISPolicyObjective, PolicyObjective + + +def _loss(ratios, advantages=None): + logps = torch.tensor([[math.log(value) for value in ratios]], requires_grad=True) + result = SAOLoss(epsilon_low=0.3, epsilon_high=5.0)( + {'labels': torch.ones_like(logps, dtype=torch.long)}, + {'logps': logps}, + old_logps=torch.zeros_like(logps), + advantages=advantages or [[1.0] * len(ratios)], + )['loss'] + return logps, result + + +def test_sao_ratio_boundaries_are_strict(): + logps, loss = _loss([0.7, 0.7001, 5.999, 6.0]) + loss.backward() + assert logps.grad[0, 0] == 0 + assert logps.grad[0, 1] != 0 + assert logps.grad[0, 2] != 0 + assert logps.grad[0, 3] == 0 + + +def test_sao_outside_trust_has_zero_gradient(): + logps, loss = _loss([0.1, 7.0]) + loss.backward() + torch.testing.assert_close(logps.grad, torch.zeros_like(logps)) + + +def test_sao_inside_gradient_uses_detached_ratio(): + logps, loss = _loss([2.0]) + loss.backward() + assert logps.grad.item() == pytest.approx(-2.0) + + +def test_sao_uses_reusable_dis_policy_objective(): + loss = SAOLoss(epsilon_low=0.3, epsilon_high=5.0) + assert isinstance(loss.policy_objective, PolicyObjective) + assert isinstance(loss.policy_objective, DISPolicyObjective) + + +def test_dis_policy_objective_matches_original_sao_formula(): + logps = torch.tensor([[math.log(0.7), math.log(2.0), math.log(6.0)]], requires_grad=True) + ratio = torch.exp(logps) + advantages = torch.tensor([[1.0, -0.5, 1.0]]) + objective = DISPolicyObjective(epsilon_low=0.3, epsilon_high=5.0) + + actual = objective(ratio, advantages, logps) + + trusted = (ratio > 0.7) & (ratio < 6.0) + weight = torch.where(trusted, ratio, torch.zeros_like(ratio)).detach() + expected = -weight * advantages.detach() * logps.float() + torch.testing.assert_close(actual, expected) + + actual.sum().backward() + torch.testing.assert_close(logps.grad, torch.tensor([[0.0, 1.0, 0.0]])) + + +def test_sao_ragged_alignment_and_token_mean_denominator(): + logps = torch.tensor([[0.0, math.log(7.0), 0.0], [math.log(2.0), 0.0, 0.0]], requires_grad=True) + result = SAOLoss()( + {'labels': torch.tensor([[1, 2, -100], [3, -100, -100]])}, + {'logps': logps}, + old_logps=[[0.0, 0.0], [0.0]], + advantages=[[1.0, 1.0], [1.0]], + )['loss'] + result.backward() + # Three action tokens form the denominator; the rejected ratio=7 token contributes zero. + assert logps.grad[0, 0].item() == pytest.approx(-1 / 3) + assert logps.grad[0, 1].item() == 0 + assert logps.grad[1, 0].item() == pytest.approx(-2 / 3) + + +def test_sao_value_loss_is_masked_mse(): + values = torch.tensor([[0.0, 2.0, 99.0]], requires_grad=True) + loss = SAOValueLoss()( + {'labels': torch.tensor([[1, 2, -100]])}, {'values': values}, returns=[[1.0, 0.0]])['loss'] + assert loss.item() == pytest.approx(2.5) + loss.backward() + assert values.grad[0, 2] == 0 diff --git a/tests/model/test_value_model.py b/tests/model/test_value_model.py index c4ee7c14..68a1a3f1 100644 --- a/tests/model/test_value_model.py +++ b/tests/model/test_value_model.py @@ -74,3 +74,51 @@ def test_value_model_forward_only_returns_values(): assert outputs['values'].shape == (1, 3) assert outputs.get('logps') is None assert not outputs['values'].requires_grad + + +def test_freeze_attention_keeps_mlp_and_value_head_trainable(): + model = TransformersValueModel(model_id=_tiny_model_dir(), mixed_precision='no') + summary = model.freeze_attention_for_value_training() + assert summary['attention_modules'] == 1 + backbone = model.model.transformer.h[0] + assert all(not parameter.requires_grad for parameter in backbone.attn.parameters()) + assert any(parameter.requires_grad for parameter in backbone.mlp.parameters()) + assert all(parameter.requires_grad for parameter in model.model.get_output_embeddings().parameters()) + counts = model.trainable_parameter_summary() + assert counts['frozen_parameters'] == summary['frozen_parameters'] + + +def test_freeze_attention_supports_qwen35_hybrid_token_mixers(): + """Qwen3.5 exposes full and linear attention under different names.""" + + class FullAttentionLayer(torch.nn.Module): + + def __init__(self): + super().__init__() + self.self_attn = torch.nn.Linear(8, 8) + self.mlp = torch.nn.Linear(8, 8) + + class LinearAttentionLayer(torch.nn.Module): + + def __init__(self): + super().__init__() + self.linear_attn = torch.nn.Linear(8, 8) + self.mlp = torch.nn.Linear(8, 8) + + model = TransformersValueModel(model_id=_tiny_model_dir(), mixed_precision='no') + full_attention_layer = FullAttentionLayer() + linear_attention_layer = LinearAttentionLayer() + model.model.qwen35_hybrid_layers = torch.nn.ModuleList([ + full_attention_layer, + linear_attention_layer, + ]) + + summary = model.freeze_attention_for_value_training() + + # One GPT-2 ``attn`` module plus the two Qwen3.5-style token mixers. + assert summary['attention_modules'] == 3 + assert all(not parameter.requires_grad for parameter in full_attention_layer.self_attn.parameters()) + assert all(not parameter.requires_grad for parameter in linear_attention_layer.linear_attn.parameters()) + assert all(parameter.requires_grad for parameter in full_attention_layer.mlp.parameters()) + assert all(parameter.requires_grad for parameter in linear_attention_layer.mlp.parameters()) + assert all(parameter.requires_grad for parameter in model.model.get_output_embeddings().parameters())