From e664e0bc384d73fdd500e4da494915332e16b60d Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Tue, 28 Jul 2026 20:10:46 +0800 Subject: [PATCH 1/5] add mask replay code and the cookbook Signed-off-by: vx120 <893600387@qq.com> --- cookbook/rl/grpo/grpo_sampling_replay.py | 330 ++++++++++++++++++ cookbook/rl/grpo/grpo_sampling_replay.sh | 19 + src/twinkle/data_format/__init__.py | 2 +- src/twinkle/data_format/sampling.py | 8 + src/twinkle/loss/grpo.py | 16 + src/twinkle/metric/__init__.py | 1 + src/twinkle/metric/grpo.py | 21 +- src/twinkle/metric/rollout.py | 88 +++++ .../model/transformers/transformers.py | 48 ++- .../sampler/vllm_sampler/vllm_engine.py | 68 +++- .../sampler/vllm_sampler/vllm_sampler.py | 1 + src/twinkle/utils/__init__.py | 5 +- src/twinkle/utils/nccl_safe.py | 1 + src/twinkle/utils/torch_utils.py | 116 ++++++ 14 files changed, 708 insertions(+), 16 deletions(-) create mode 100644 cookbook/rl/grpo/grpo_sampling_replay.py create mode 100644 cookbook/rl/grpo/grpo_sampling_replay.sh create mode 100644 src/twinkle/metric/rollout.py diff --git a/cookbook/rl/grpo/grpo_sampling_replay.py b/cookbook/rl/grpo/grpo_sampling_replay.py new file mode 100644 index 000000000..2f8a4a41f --- /dev/null +++ b/cookbook/rl/grpo/grpo_sampling_replay.py @@ -0,0 +1,330 @@ +import os +import time +from typing import List, Tuple, Dict, Any + +from peft import LoraConfig + +import twinkle +from twinkle import DeviceMesh, DeviceGroup, get_device_placement, get_logger +from twinkle.advantage import GRPOAdvantage +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.model import TransformersModel +from twinkle.processor import InputProcessor +from twinkle.reward import GSM8KAccuracyReward, GSM8KFormatReward +from twinkle.sampler import vLLMSampler +from twinkle.metric import ( + CompletionRewardMetric, + compute_grpo_rollout_metrics, +) +from twinkle.preprocessor.llm import GSM8KProcessor + +logger = get_logger() +args = CLI.from_args() + +MODEL_ID = args.model.model_id or 'ms://Qwen/Qwen3.5-4B' +USE_MEGATRON = args.model.strategy != 'native_fsdp' +# This entry point is exclusively for sampling-distribution replay. +ENABLE_SAMPLING_REPLAY = True + +MODEL_GPUS = args.infra.model_gpus or 4 +SAMPLER_GPUS = args.infra.sampler_gpus or 4 +NUM_GPUS = MODEL_GPUS + SAMPLER_GPUS + +NUM_GENERATIONS = args.rl.num_generations or 8 +MAX_NEW_TOKENS = args.sampling.max_tokens or 4096 +LEARNING_RATE = args.optimizer.learning_rate or 1e-5 +MAX_STEPS = args.training.max_steps or 200 +BATCH_SIZE = args.training.batch_size or 8 +MINI_BATCH_SIZE = args.training.mini_batch_size or 8 +MICRO_BATCH_SIZE = args.training.micro_batch_size or 2 +GRADIENT_ACCUMULATION_STEPS = args.training.gradient_accumulation_steps or 1 +ADAPTER_NAME = args.lora.adapter_name or 'default' +SAVE_STEPS = args.training.save_steps or 50 +LOGPROBS_MODE = ( + 'processed_logprobs' + if ENABLE_SAMPLING_REPLAY + else os.getenv('TWINKLE_LOGPROBS_MODE', 'processed_logprobs') +) + +if ENABLE_SAMPLING_REPLAY and USE_MEGATRON: + raise ValueError('Sampling replay currently requires --strategy native_fsdp') + +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_reward_fn = GSM8KAccuracyReward() + format_reward_fn = GSM8KFormatReward() + + accuracy_rewards = accuracy_reward_fn(trajectories) + format_rewards = format_reward_fn(trajectories) + total_rewards = [a + f for a, f in zip(accuracy_rewards, format_rewards)] + return total_rewards, format_rewards, accuracy_rewards + + +def extract_rollout_batch(sample_responses, *, require_sampling_masks: bool): + """Flatten sampler responses into aligned lists used by reward and training.""" + rollout_batch = { + 'input_data': [], + 'old_logps': [], + 'sampling_masks': [], + 'completion_lengths': [], + 'stop_reasons': [], + } + for sample_response in sample_responses: + for sequence in sample_response.sequences: + if sequence.logprobs is None: + raise RuntimeError('A sampled sequence is missing token log probabilities') + if require_sampling_masks and sequence.sampling_mask is None: + raise RuntimeError( + 'Sampling replay is enabled but a sampled sequence has no sampling mask') + rollout_batch['input_data'].append(sequence.new_input_feature) + rollout_batch['old_logps'].append( + [logprob[0][1] for logprob in sequence.logprobs]) + rollout_batch['sampling_masks'].append(sequence.sampling_mask) + rollout_batch['completion_lengths'].append(len(sequence.tokens)) + rollout_batch['stop_reasons'].append(sequence.stop_reason) + return rollout_batch + + +def main(): + # set sampler and model separate to use different gpus + device_groups = [ + DeviceGroup(name='model',ranks=list(range(MODEL_GPUS)),device_type='GPU'), + DeviceGroup(name='sampler',ranks=list(range(MODEL_GPUS, NUM_GPUS)),device_type='GPU'), + ] + if USE_MEGATRON: + model_mesh = DeviceMesh.from_sizes(world_size=MODEL_GPUS, dp_size=MODEL_GPUS) + else: + model_mesh = DeviceMesh.from_sizes(world_size=MODEL_GPUS, dp_size=MODEL_GPUS) + sampler_mesh = DeviceMesh.from_sizes(world_size=SAMPLER_GPUS, dp_size=SAMPLER_GPUS) + twinkle.initialize(mode='ray', nproc_per_node=NUM_GPUS, groups=device_groups, lazy_collect=False) + + # lora_config = LoraConfig(target_modules='all-linear', r=32, lora_alpha=64, lora_dropout=0.05) + # Since we are training on text-only data, we avoid using 'all-linear' which would include the ViT layers. + lora_config = LoraConfig( + target_modules=[ + 'q_proj', 'k_proj', 'v_proj', 'o_proj', + 'gate_proj', 'up_proj', 'down_proj', + 'in_proj_qkv', 'in_proj_z', 'in_proj_a', 'in_proj_b', 'out_proj', + ], + r=32, lora_alpha=64, lora_dropout=0.0, + ) + if USE_MEGATRON: + from twinkle.model.megatron import MegatronModel + model = MegatronModel(model_id=MODEL_ID, device_mesh=model_mesh, remote_group='model', mixed_precision='bf16') + else: + from transformers import Qwen3_5ForConditionalGeneration + model = TransformersModel( + model_id=MODEL_ID, + model_cls=Qwen3_5ForConditionalGeneration, + device_mesh=model_mesh, + remote_group='model', + ) + + model.add_adapter_to_model(ADAPTER_NAME, lora_config, gradient_accumulation_steps=1) + if USE_MEGATRON: + model.set_optimizer('default', lr=LEARNING_RATE) + model.set_lr_scheduler('default', lr_decay_steps=MAX_STEPS, max_lr=LEARNING_RATE) + else: + model.set_optimizer('AdamW', lr=LEARNING_RATE) + model.set_lr_scheduler('CosineAnnealingLR', T_max=MAX_STEPS, eta_min=0) + model.set_loss( + 'GRPOLoss', + epsilon=0.2, + beta=0.0, + entropy_coef=0.0, + enable_sampling_replay=ENABLE_SAMPLING_REPLAY, + ) + model.set_processor(InputProcessor) + model.set_template('Qwen3_5Template', model_id=MODEL_ID) + + sampler = vLLMSampler( + model_id=MODEL_ID, + engine_args={ + 'gpu_memory_utilization': 0.8, + 'max_model_len': 4496, + 'max_lora_rank': 32, # save as lora_config + # NOTE: To use enable_lora with qwen3.5, ensure vLLM includes + # PR https://github.com/vllm-project/vllm/pull/36976 + # enable_lora=True used with ckpt_manager.sync_weights(merge_and_sync=False) + # meaning only sync lora weights, if merge_and_sync=True, + # lora will be merged into the base model and sync all weights to vLLM + 'enable_lora': True, + 'enable_sampling_replay': ENABLE_SAMPLING_REPLAY, + 'logprobs_mode': LOGPROBS_MODE, + }, + device_mesh=sampler_mesh, + remote_group='sampler', + ) + sampler.set_template('Qwen3_5Template', model_id=MODEL_ID) + + ckpt_manager = CheckpointEngineManager(model=model, sampler=sampler) + + GLOBAL_BATCH_SIZE = BATCH_SIZE * GRADIENT_ACCUMULATION_STEPS + dataloader = DataLoader( + dataset=create_gsm8k_dataset, + batch_size=GLOBAL_BATCH_SIZE, + min_batch_size=GLOBAL_BATCH_SIZE, + device_mesh=model_mesh, + remote_group='model', + ) + advantage_fn = GRPOAdvantage() + metrics = CompletionRewardMetric() + + sampling_params = SamplingParams( + max_tokens=MAX_NEW_TOKENS, + num_samples=1, + logprobs=1, + temperature=1.0, + top_p=0.95 if ENABLE_SAMPLING_REPLAY else 1.0, + top_k=-1, + repetition_penalty=1.0, + ) + if ENABLE_SAMPLING_REPLAY: + model.add_metric( + 'GRPOMetric', + is_training=True, + temperature=sampling_params.temperature, + epsilon=0.2, + ) + logger.info( + 'Sampling replay enabled: model_runner=v2, logprobs_mode=processed_logprobs, ' + 'temperature=%s, top_p=%s, top_k=%s', + sampling_params.temperature, + sampling_params.top_p, + sampling_params.top_k, + ) + + optim_step = 0 + sampling_replay_stats_logged = False + logger.info(get_device_placement()) + + for batch in dataloader: + if optim_step >= MAX_STEPS: + break + metrics.reset() + global_prompts = batch if isinstance(batch, list) else [batch] + # enable_lora=True used with ckpt_manager.sync_weights(merge_and_sync=False) + # meaning only sync lora weights, if merge_and_sync=True, + # lora will be merged into the base model and sync all weights to vLLM + weight_sync_started = time.perf_counter() + ckpt_manager.sync_weights(merge_and_sync=False) + weight_sync_seconds = time.perf_counter() - weight_sync_started + sampler.reset_prefix_cache() + def sample_prompt_groups(prompts): + expand_prompts = [] + for prompt in prompts: + expand_prompts.extend([prompt] * NUM_GENERATIONS) + started = time.perf_counter() + responses = sampler.sample(expand_prompts, sampling_params) + elapsed = time.perf_counter() - started + return extract_rollout_batch( + responses, + require_sampling_masks=ENABLE_SAMPLING_REPLAY, + ), elapsed + + rollout_batch, sampling_seconds = sample_prompt_groups(global_prompts) + sampled_tokens_total = sum(rollout_batch['completion_lengths']) + # Match the original GRPO control flow: every sampled rollout is scored, + # logged, and trained. Zero-variance groups keep their zero advantages; + # they are diagnosed below but never resampled, dropped, or skipped. + total_rewards, format_rewards, accuracy_rewards = compute_rewards( + rollout_batch['input_data']) + + all_input_data: List[Dict[str, Any]] = rollout_batch['input_data'] + all_old_logps: List[List[float]] = rollout_batch['old_logps'] + all_sampling_masks = rollout_batch['sampling_masks'] + all_completion_lengths: List[int] = rollout_batch['completion_lengths'] + all_stop_reasons = rollout_batch['stop_reasons'] + metrics.accumulate( + completion_lengths=all_completion_lengths, + rewards={ + 'total': total_rewards, + 'format': format_rewards, + 'accuracy': accuracy_rewards, + }, + ) + rollout_reward_log_dict = metrics.calculate() + + advantages = advantage_fn(total_rewards, num_generations=NUM_GENERATIONS, scale='group').tolist() + rollout_log_dict = compute_grpo_rollout_metrics( + completion_lengths=all_completion_lengths, + stop_reasons=all_stop_reasons, + rewards=total_rewards, + advantages=advantages, + num_generations=NUM_GENERATIONS, + sampling_masks=all_sampling_masks if ENABLE_SAMPLING_REPLAY else None, + ) + num_rollout_tokens = sum(all_completion_lengths) + rollout_log_dict['profiling/weight_sync_seconds'] = weight_sync_seconds + rollout_log_dict['profiling/sampling_seconds'] = sampling_seconds + rollout_log_dict['profiling/sampling_tokens_per_second'] = ( + sampled_tokens_total / sampling_seconds if sampling_seconds else 0.0) + rollout_log_dict['profiling/sampling_generated_tokens'] = sampled_tokens_total + + # Split completions into mini-batches and run one optim step per mini-batch. + total_completions = len(all_input_data) + for mb_start in range(0, total_completions, MINI_BATCH_SIZE): + mb_end = min(mb_start + MINI_BATCH_SIZE, total_completions) + mb_inputs = all_input_data[mb_start:mb_end] + mb_old_logps = all_old_logps[mb_start:mb_end] + mb_advantages = advantages[mb_start:mb_end] + replay_kwargs = {} + if ENABLE_SAMPLING_REPLAY: + replay_kwargs = { + 'sampling_masks': all_sampling_masks[mb_start:mb_end], + 'temperature': sampling_params.temperature, + } + + training_started = time.perf_counter() + model.forward_backward( + inputs=mb_inputs, + old_logps=mb_old_logps, + advantages=mb_advantages, + micro_batch_size=MICRO_BATCH_SIZE, + **replay_kwargs, + ) + model.clip_grad_and_step() + training_seconds = time.perf_counter() - training_started + if ENABLE_SAMPLING_REPLAY and not sampling_replay_stats_logged: + logger.info( + 'Sampling replay active: sequences=%d, tokens=%d, mean_kept_tokens=%.2f', + len(all_sampling_masks), + num_rollout_tokens, + rollout_log_dict['replay/support_size_mean'], + ) + sampling_replay_stats_logged = True + optim_step += 1 + + if optim_step % SAVE_STEPS == 0: + model.save(f'grpo-gsm8k-checkpoint-{optim_step}') + # Copy the rollout reward into every optimizer-step log line. A + # rollout can span multiple mini-batches, but no Step lacks reward. + log_dict = dict(rollout_reward_log_dict) + log_dict.update(model.calculate_metric(is_training=True)) + if mb_start == 0: + log_dict.update(rollout_log_dict) + num_training_tokens = sum(all_completion_lengths[mb_start:mb_end]) + log_dict['profiling/training_seconds'] = training_seconds + log_dict['profiling/training_completion_tokens_per_second'] = ( + num_training_tokens / training_seconds if training_seconds else 0.0) + logger.info(f'[Step {optim_step}/{MAX_STEPS}] {log_dict}') + if optim_step >= MAX_STEPS: + break + + logger.info(f'Training completed. optim_steps={optim_step}') + model.save('grpo-gsm8k-checkpoint') + +if __name__ == '__main__': + main() diff --git a/cookbook/rl/grpo/grpo_sampling_replay.sh b/cookbook/rl/grpo/grpo_sampling_replay.sh new file mode 100644 index 000000000..f1120ecb0 --- /dev/null +++ b/cookbook/rl/grpo/grpo_sampling_replay.sh @@ -0,0 +1,19 @@ +#!/bin/sh +set -eu + +# Sampling-distribution replay example. +python grpo_sampling_replay.py \ + --model-id ms://Qwen/Qwen3.5-4B \ + --strategy native_fsdp \ + --model-gpus 4 \ + --sampler-gpus 4 \ + --num-generations 8 \ + --max-tokens 4096 \ + --batch-size 8 \ + --mini-batch-size 8 \ + --micro-batch-size 2 \ + --max-steps 200 \ + --lr 1e-5 \ + --save-steps 50 \ + --adapter-name default \ + "$@" diff --git a/src/twinkle/data_format/__init__.py b/src/twinkle/data_format/__init__.py index c93bebd2d..1dff273c7 100644 --- a/src/twinkle/data_format/__init__.py +++ b/src/twinkle/data_format/__init__.py @@ -2,5 +2,5 @@ from .input_feature import InputFeature from .message import Message, Tool, ToolCall from .output import LossOutput, ModelOutput -from .sampling import SampledSequence, SampleResponse, SamplingParams +from .sampling import SampledSequence, SampleResponse, SamplingMask, SamplingParams from .trajectory import Trajectory, pack_value, user_data_get diff --git a/src/twinkle/data_format/sampling.py b/src/twinkle/data_format/sampling.py index 05ecdd641..cdd2233a8 100644 --- a/src/twinkle/data_format/sampling.py +++ b/src/twinkle/data_format/sampling.py @@ -166,6 +166,13 @@ def from_dict(cls, d: Dict[str, Any]) -> 'SamplingParams': return cls(**filtered) +@dataclass +class SamplingMask: + """CSR token support sets aligned with sampled sequence tokens.""" + token_ids: List[int] + offsets: List[int] + + @dataclass class SampledSequence: """A single sampled sequence with tokens and logprobs.""" @@ -175,6 +182,7 @@ class SampledSequence: decoded: str = None new_input_feature: InputFeature = None routed_experts: Optional[Any] = None + sampling_mask: Optional[SamplingMask] = None @dataclass diff --git a/src/twinkle/loss/grpo.py b/src/twinkle/loss/grpo.py index 781b22060..36970636d 100644 --- a/src/twinkle/loss/grpo.py +++ b/src/twinkle/loss/grpo.py @@ -32,12 +32,20 @@ def __init__( beta: float = 0.0, entropy_coef: float = 0.0, ignore_index: int = -100, + enable_sampling_replay: bool = False, **kwargs, ): self.epsilon = epsilon self.epsilon_high = epsilon_high if epsilon_high is not None else epsilon self.beta = beta self.entropy_coef = entropy_coef + self.enable_sampling_replay = enable_sampling_replay + if enable_sampling_replay and self.__class__ is not GRPOLoss: + raise ValueError('sampling replay is only supported by GRPOLoss') + if enable_sampling_replay and beta != 0.0: + raise ValueError('sampling replay does not support a GRPO KL penalty (beta must be 0)') + if enable_sampling_replay and entropy_coef != 0.0: + raise ValueError('sampling replay does not support a GRPO entropy bonus') # Gate the expensive entropy compute path in the model forward. self.require_entropy = entropy_coef > 0.0 self.ignore_index = ignore_index @@ -201,6 +209,7 @@ def __call__( old_logps: Optional[Union['torch.Tensor', List[List[float]]]] = None, ref_logps: Optional['torch.Tensor'] = None, advantages: Optional[Union['torch.Tensor', List[float], np.ndarray]] = None, + sampling_masks=None, **kwargs, ): """ @@ -222,6 +231,11 @@ def __call__( **kwargs: Additional arguments """ import torch + if self.enable_sampling_replay: + if sampling_masks is None: + raise ValueError('sampling_masks are required when sampling replay is enabled') + if old_logps is None: + raise ValueError('old_logps are required when sampling replay is enabled') labels = inputs.get('labels') assert labels is not None, "inputs must contain 'labels'" if not torch.is_tensor(labels): @@ -230,6 +244,8 @@ def __call__( labels = labels.unsqueeze(0) logps = outputs.get('logps') + if self.enable_sampling_replay and logps is None: + raise RuntimeError('sampling replay logps must be computed by the model forward') loss_mask = (labels != self.ignore_index).bool() if logps is None: logits = outputs.get('logits') diff --git a/src/twinkle/metric/__init__.py b/src/twinkle/metric/__init__.py index baeb6c1c9..cd7d8c99d 100644 --- a/src/twinkle/metric/__init__.py +++ b/src/twinkle/metric/__init__.py @@ -6,4 +6,5 @@ from .embedding import EmbeddingMetric from .grpo import CISPOMetric, GRPOMetric, GSPOMetric from .loss import LossMetric +from .rollout import compute_grpo_rollout_metrics, zero_variance_reward_group_indices from .train_metric import TrainMetric diff --git a/src/twinkle/metric/grpo.py b/src/twinkle/metric/grpo.py index bd85aab67..71fde1de4 100644 --- a/src/twinkle/metric/grpo.py +++ b/src/twinkle/metric/grpo.py @@ -1,6 +1,6 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import math -from typing import Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from twinkle.data_format import InputFeature, ModelOutput from twinkle.utils import get_logger @@ -9,6 +9,9 @@ logger = get_logger() +if TYPE_CHECKING: + import torch + class GRPOMetric(Metric): @@ -41,6 +44,8 @@ def reset(self): self.sum_new: float = 0.0 self.sum_old: float = 0.0 self.sum_diff: float = 0.0 + self.sum_diff_sq: float = 0.0 + self.sum_ratio: float = 0.0 self.sum_approx_kl: float = 0.0 self.max_token_kl: float = 0.0 self.max_token_ratio: float = 0.0 @@ -185,13 +190,16 @@ def _accumulate_mb( old_f = old_f * scale d = logps_f - old_f # new - old + ratio = torch.exp(d) self.sum_old += float((old_f * mask_f).sum().item()) self.sum_diff += float((d * mask_f).sum().item()) + self.sum_diff_sq += float((d.square() * mask_f).sum().item()) + self.sum_ratio += float((ratio * mask_f).sum().item()) # Schulman K3 estimator of KL(old || new): # samples x ~ old, r(x) = new(x) / old(x), # k3 = r - 1 - log(r) = exp(new - old) - (new - old) - 1. - kl = torch.exp(d) - d - 1.0 + kl = ratio - d - 1.0 kl_masked = kl * mask_f self.sum_approx_kl += float(kl_masked.sum().item()) # Per-token extremes for collapse detection @@ -200,7 +208,7 @@ def _accumulate_mb( if cur_max_kl > self.max_token_kl: self.max_token_kl = cur_max_kl # Track ratio extremes - ratio_masked = torch.exp(d) * mask_f + ratio_masked = ratio * mask_f cur_max_ratio = float(ratio_masked.max().item()) if cur_max_ratio > self.max_token_ratio: self.max_token_ratio = cur_max_ratio @@ -311,11 +319,12 @@ def accumulate( cursor += advanced def calculate(self) -> Dict[str, Any]: - import torch local = [{ 'sum_new': self.sum_new, 'sum_old': self.sum_old, 'sum_diff': self.sum_diff, + 'sum_diff_sq': self.sum_diff_sq, + 'sum_ratio': self.sum_ratio, 'sum_kl': self.sum_approx_kl, 'max_token_kl': self.max_token_kl, 'max_token_ratio': self.max_token_ratio, @@ -344,11 +353,15 @@ def calculate(self) -> Dict[str, Any]: if any(r['has_old'] for r in all_results): mean_old = sum(r['sum_old'] for r in all_results) / n_total mean_diff = sum(r['sum_diff'] for r in all_results) / n_total + mean_diff_sq = sum(r['sum_diff_sq'] for r in all_results) / n_total + mean_ratio = sum(r['sum_ratio'] for r in all_results) / n_total mean_kl = sum(r['sum_kl'] for r in all_results) / n_total global_max_kl = max(r['max_token_kl'] for r in all_results) global_max_ratio = max(r['max_token_ratio'] for r in all_results) results['train/mean_old_logp'] = mean_old results['train/logp_diff_mean'] = mean_diff + results['train/logp_diff_std'] = math.sqrt(max(mean_diff_sq - mean_diff**2, 0.0)) + results['train/importance_ratio_mean'] = mean_ratio results['train/approx_kl'] = mean_kl results['train/token_kl_max'] = global_max_kl results['train/token_ratio_max'] = global_max_ratio diff --git a/src/twinkle/metric/rollout.py b/src/twinkle/metric/rollout.py new file mode 100644 index 000000000..cc261ccd9 --- /dev/null +++ b/src/twinkle/metric/rollout.py @@ -0,0 +1,88 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from typing import Any, Dict, Optional, Sequence + +import numpy as np + + +def zero_variance_reward_group_indices( + rewards: Sequence[float], + num_generations: int, +) -> list[int]: + """Return GRPO group indices that cannot produce a relative advantage.""" + if num_generations <= 0: + raise ValueError('num_generations must be positive') + if len(rewards) % num_generations != 0: + raise ValueError('rewards must form complete num_generations groups') + if len(rewards) == 0: + return [] + + grouped_rewards = np.asarray(rewards, dtype=np.float64).reshape(-1, num_generations) + group_ranges = np.ptp(grouped_rewards, axis=1) + return np.flatnonzero(np.isclose(group_ranges, 0.0)).astype(int).tolist() + + +def compute_grpo_rollout_metrics( + *, + completion_lengths: Sequence[int], + stop_reasons: Sequence[str], + rewards: Sequence[float], + advantages: Sequence[float], + num_generations: int, + sampling_masks: Optional[Sequence[Any]] = None, +) -> Dict[str, float]: + """Reduce one GRPO rollout batch into scalar diagnostics.""" + if len(stop_reasons) != len(completion_lengths): + raise ValueError('stop_reasons must align with completion_lengths') + if len(rewards) != len(completion_lengths): + raise ValueError('rewards must align with completion_lengths') + if num_generations <= 0 or len(rewards) % num_generations != 0: + raise ValueError('rewards must form complete num_generations groups') + if len(advantages) != len(rewards): + raise ValueError('advantages must align with rewards') + + metrics: Dict[str, float] = {} + if len(completion_lengths) > 0: + lengths = np.asarray(completion_lengths, dtype=np.float64) + metrics['rollout/completion_length_p95'] = float(np.percentile(lengths, 95)) + + if len(stop_reasons) > 0: + num_sequences = len(stop_reasons) + metrics['rollout/stop_rate'] = sum(reason == 'stop' for reason in stop_reasons) / num_sequences + metrics['rollout/length_stop_rate'] = ( + sum(reason == 'length' for reason in stop_reasons) / num_sequences) + + if len(rewards) > 0: + grouped_rewards = np.asarray(rewards, dtype=np.float64).reshape(-1, num_generations) + if num_generations > 1: + group_stds = grouped_rewards.std(axis=1, ddof=1) + else: + group_stds = np.zeros(grouped_rewards.shape[0], dtype=np.float64) + metrics['grpo/group_reward_std_mean'] = float(group_stds.mean()) + zero_variance_groups = zero_variance_reward_group_indices(rewards, num_generations) + metrics['grpo/zero_variance_group_fraction'] = ( + len(zero_variance_groups) / grouped_rewards.shape[0]) + metrics['grpo/nonzero_advantage_fraction'] = float( + (~np.isclose(np.asarray(advantages, dtype=np.float64), 0.0)).mean()) + + if sampling_masks is not None: + if len(sampling_masks) != len(completion_lengths): + raise ValueError('sampling_masks must align with completion_lengths') + support_sizes = [] + for sequence_idx, (sampling_mask, completion_length) in enumerate( + zip(sampling_masks, completion_lengths)): + offsets = sampling_mask.offsets + if len(offsets) - 1 != completion_length: + raise ValueError( + f'sampling mask {sequence_idx} has {len(offsets) - 1} rows, ' + f'expected {completion_length}') + support_sizes.extend(end - start for start, end in zip(offsets, offsets[1:])) + + if support_sizes: + sizes = np.asarray(support_sizes, dtype=np.float64) + metrics['replay/support_size_mean'] = float(sizes.mean()) + metrics['replay/support_size_p50'] = float(np.percentile(sizes, 50)) + metrics['replay/support_size_p95'] = float(np.percentile(sizes, 95)) + metrics['replay/support_size_max'] = float(sizes.max()) + metrics['replay/singleton_fraction'] = float((sizes == 1).mean()) + + return metrics diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index 017a515b3..aef56f1ff 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -39,7 +39,7 @@ from twinkle.patch import Patch, apply_context, apply_patch from twinkle.processor import InputProcessor from twinkle.template import Template -from twinkle.utils import construct_class, get_logger, selective_log_softmax, torch_util +from twinkle.utils import construct_class, get_logger, replayed_selective_log_softmax, selective_log_softmax, torch_util from twinkle.utils.framework import Torch from twinkle.utils.grad_clip import normalize_and_clip_grad_norm from twinkle.utils.transformers_utils import filter_from_config_kwargs @@ -446,6 +446,7 @@ def forward(self, *, inputs: Union[InputFeature, List[InputFeature], List[Trajec """ adapter_name = kwargs.pop('adapter_name', self._get_default_group()) temperature = float(kwargs.pop('temperature', 1.0)) + sampling_masks = kwargs.pop('sampling_masks', None) return_logits = kwargs.pop('return_logits', False) task = kwargs.pop('task', 'causal_lm') optimizer_config = self.optimizer_group[adapter_name] @@ -466,6 +467,15 @@ def forward(self, *, inputs: Union[InputFeature, List[InputFeature], List[Trajec loss_require_logits = getattr(loss_instance, 'require_logits', False) loss_require_entropy = getattr(loss_instance, 'require_entropy', False) loss_require_logps = getattr(loss_instance, 'require_logps', True) + enable_sampling_replay = getattr(loss_instance, 'enable_sampling_replay', False) + if enable_sampling_replay: + if sampling_masks is None: + raise ValueError('sampling_masks are required when sampling replay is enabled') + if kwargs.get('old_logps') is None: + raise ValueError('old_logps are required when sampling replay is enabled') + cp_world_size = self.device_mesh.cp_world_size if self.device_mesh is not None else 1 + if getattr(self, '_enable_sp', False) or cp_world_size > 1: + raise ValueError('sampling replay does not support sequence or context parallelism') assert isinstance(processor, InputProcessor), 'Set a correct `InputProcessor` before forwarding' inputs: Dict[str, Any] = processor( inputs, @@ -497,11 +507,20 @@ def forward(self, *, inputs: Union[InputFeature, List[InputFeature], List[Trajec masked_labels = labels.clone() masked_labels[~loss_mask] = 0 logits = outputs['logits'] - logits.div_(temperature) - if loss_require_entropy: + if enable_sampling_replay: + outputs['logps'] = replayed_selective_log_softmax( + logits=logits, + labels=masked_labels, + loss_mask=loss_mask, + sampling_masks=sampling_masks, + temperature=temperature, + ) + elif loss_require_entropy: + logits.div_(temperature) outputs['logps'], outputs['entropies'] = selective_log_softmax( logits, masked_labels, return_entropy=True) else: + logits.div_(temperature) outputs['logps'] = selective_log_softmax(logits, masked_labels) del logits outputs['past_key_values'] = None @@ -535,6 +554,7 @@ def forward_only(self, *, inputs: Union[InputFeature, List[InputFeature], List[T adapter_name = kwargs.pop('adapter_name', self._get_default_group()) disable_lora = kwargs.pop('disable_lora', False) temperature = float(kwargs.pop('temperature', 1.0)) + sampling_masks = kwargs.pop('sampling_masks', None) return_logits = kwargs.pop('return_logits', False) task = kwargs.pop('task', 'causal_lm') optimizer_config = self.optimizer_group[adapter_name] @@ -557,6 +577,15 @@ def forward_only(self, *, inputs: Union[InputFeature, List[InputFeature], List[T loss_require_logits = getattr(loss_instance, 'require_logits', False) loss_require_entropy = getattr(loss_instance, 'require_entropy', False) loss_require_logps = getattr(loss_instance, 'require_logps', True) + enable_sampling_replay = getattr(loss_instance, 'enable_sampling_replay', False) + if enable_sampling_replay: + if sampling_masks is None: + raise ValueError('sampling_masks are required when sampling replay is enabled') + if kwargs.get('old_logps') is None: + raise ValueError('old_logps are required when sampling replay is enabled') + cp_world_size = self.device_mesh.cp_world_size if self.device_mesh is not None else 1 + if getattr(self, '_enable_sp', False) or cp_world_size > 1: + raise ValueError('sampling replay does not support sequence or context parallelism') inputs: Dict[str, Any] = processor( inputs, sp_strategy=self.sp_strategy, @@ -591,11 +620,20 @@ def forward_only(self, *, inputs: Union[InputFeature, List[InputFeature], List[T masked_labels = labels.clone() masked_labels[~loss_mask] = 0 logits = outputs['logits'] - logits.div_(temperature) - if loss_require_entropy: + if enable_sampling_replay: + outputs['logps'] = replayed_selective_log_softmax( + logits=logits, + labels=masked_labels, + loss_mask=loss_mask, + sampling_masks=sampling_masks, + temperature=temperature, + ) + elif loss_require_entropy: + logits.div_(temperature) outputs['logps'], outputs['entropies'] = selective_log_softmax( logits, masked_labels, return_entropy=True) else: + logits.div_(temperature) outputs['logps'] = selective_log_softmax(logits, masked_labels) del logits outputs['past_key_values'] = None diff --git a/src/twinkle/sampler/vllm_sampler/vllm_engine.py b/src/twinkle/sampler/vllm_sampler/vllm_engine.py index b1e1790de..5fcf1d844 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_engine.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_engine.py @@ -8,7 +8,7 @@ from typing import Any, Dict, List, Optional, Union from twinkle import get_logger -from twinkle.data_format.sampling import SampledSequence, SampleResponse, SamplingParams, StopReason +from twinkle.data_format.sampling import SampledSequence, SampleResponse, SamplingMask, SamplingParams, StopReason from twinkle.sampler.base_engine import BaseSamplerEngine from twinkle.utils import Platform from twinkle.utils.framework import Torch @@ -29,6 +29,48 @@ def _map_finish_reason(reason: str | None) -> StopReason: return _FINISH_REASON_MAP.get(str(reason), 'length') +def _filter_engine_config( + engine_config: Dict[str, Any], + valid_args, + enable_sampling_replay: bool, +): + valid_args = set(valid_args) + invalid_args = set(engine_config) - valid_args + if enable_sampling_replay and 'enable_return_sampling_mask' in invalid_args: + raise RuntimeError( + 'Sampling replay requires a vLLM build whose AsyncEngineArgs accepts ' + 'enable_return_sampling_mask') + filtered_engine_config = {key: value for key, value in engine_config.items() if key in valid_args} + return filtered_engine_config, invalid_args + + +def _copy_sampling_mask(mask, num_tokens: int, required: bool) -> Optional[SamplingMask]: + if mask is None: + if required: + raise RuntimeError('vLLM output is missing sampling mask while sampling replay is enabled') + return None + + token_ids = [int(token_id) for token_id in mask.token_ids] + offsets = [int(offset) for offset in mask.offsets] + num_rows = len(offsets) - 1 + if num_rows != num_tokens: + raise RuntimeError( + f'vLLM sampling mask has {num_rows} rows for {num_tokens} sampled tokens') + if not offsets or offsets[0] != 0 or offsets[-1] != len(token_ids): + raise RuntimeError('vLLM sampling mask has invalid CSR endpoints') + if any(start >= end for start, end in zip(offsets, offsets[1:])): + raise RuntimeError('vLLM sampling mask contains an empty or invalid CSR row') + return SamplingMask(token_ids=token_ids, offsets=offsets) + + +def _set_sampling_replay_output_kind(vllm_params, enable_sampling_replay: bool) -> None: + """Use the only vLLM output mode that carries the full sampling mask.""" + if not enable_sampling_replay: + return + from vllm.sampling_params import RequestOutputKind + vllm_params.output_kind = RequestOutputKind.FINAL_ONLY + + def get_vllm_max_lora_rank(lora_rank: int) -> int: """Get the nearest allowed vLLM LoRA rank.""" from typing import get_args @@ -78,6 +120,7 @@ def __init__( quantization: Optional[str] = None, load_format: str = 'auto', logprobs_mode: Optional[str] = None, + enable_sampling_replay: bool = False, **kwargs, ): from twinkle.hub import HubOperation @@ -97,7 +140,9 @@ def __init__( self.dtype = dtype self.quantization = quantization self.load_format = load_format - self.logprobs_mode = logprobs_mode or 'processed_logprobs' + self.enable_sampling_replay = enable_sampling_replay + self.logprobs_mode = 'processed_logprobs' if enable_sampling_replay else ( + logprobs_mode or 'processed_logprobs') self.engine_kwargs = kwargs or {} self._lora_request_cache: Dict[str, Any] = {} @@ -130,6 +175,8 @@ def __init__( def _create_engine(self): """Create and return the vLLM engine.""" os.environ['VLLM_USE_V1'] = '1' + if self.enable_sampling_replay: + os.environ['VLLM_USE_V2_MODEL_RUNNER'] = '1' from vllm.engine.arg_utils import AsyncEngineArgs from vllm.usage.usage_lib import UsageContext from vllm.v1.engine.async_llm import AsyncLLM @@ -175,9 +222,15 @@ def _create_engine(self): 'twinkle.sampler.vllm_sampler.vllm_worker_extension.TwinkleWorkerExtension') engine_config.update(self.engine_kwargs) + if self.enable_sampling_replay: + engine_config['enable_return_sampling_mask'] = True + engine_config['logprobs_mode'] = 'processed_logprobs' valid_args = inspect.signature(AsyncEngineArgs).parameters.keys() - filtered_engine_config = {k: v for k, v in engine_config.items() if k in valid_args} - invalid_args = set(engine_config.keys()) - set(valid_args) + filtered_engine_config, invalid_args = _filter_engine_config( + engine_config, + valid_args, + self.enable_sampling_replay, + ) if invalid_args: logger.warning(f'VLLMEngine: Filtered out invalid arguments: {invalid_args}') # Create engine using vLLM v1 API @@ -244,6 +297,7 @@ async def sample(self, prompt_logprobs_k = sampling_params.prompt_logprobs or 0 logprobs = sampling_params.logprobs or 0 vllm_params = sampling_params.to_vllm(**kwargs) + _set_sampling_replay_output_kind(vllm_params, self.enable_sampling_replay) # Build request if request_id is None: @@ -291,6 +345,11 @@ async def sample(self, sequences = [] for output in result.outputs: token_ids = list(output.token_ids) + sampling_mask = _copy_sampling_mask( + getattr(output, 'sampling_mask', None), + num_tokens=len(token_ids), + required=self.enable_sampling_replay, + ) # Extract logprobs seq_logprobs = None @@ -319,6 +378,7 @@ async def sample(self, tokens=token_ids, logprobs=seq_logprobs, routed_experts=routed_experts, + sampling_mask=sampling_mask, )) # Extract prompt logprobs if requested diff --git a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py index 3c7b2f686..def44c793 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_sampler.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_sampler.py @@ -269,6 +269,7 @@ async def _sample_single( logprobs=seq.logprobs, decoded=self.template.decode(seq.tokens), new_input_feature=new_input_feature, + sampling_mask=seq.sampling_mask, ) sequences.append(sampled_seq) return SampleResponse( diff --git a/src/twinkle/utils/__init__.py b/src/twinkle/utils/__init__.py index d5d1b698b..53829fa2b 100644 --- a/src/twinkle/utils/__init__.py +++ b/src/twinkle/utils/__init__.py @@ -10,8 +10,9 @@ from .parallel import processing_lock from .platforms import GPU, NPU, Platform, ensure_hccl_socket_env, ensure_npu_backend from .safetensors import LazyTensor, SafetensorLazyLoader, StreamingSafetensorSaver -from .torch_utils import (clone_state_dict_to_cpu, pad_and_stack_tensors, pad_sequence_to_length, selective_log_softmax, - split_cp_inputs, stateless_init_process_group, to_device) +from .torch_utils import (clone_state_dict_to_cpu, pad_and_stack_tensors, pad_sequence_to_length, + replayed_selective_log_softmax, selective_log_softmax, split_cp_inputs, + stateless_init_process_group, to_device) from .transformers_utils import find_all_linears, find_layers, get_modules_to_not_convert from .unsafe import check_unsafe, trust_remote_code from .utils import copy_files_by_pattern, deep_getattr, get_runtime_meta diff --git a/src/twinkle/utils/nccl_safe.py b/src/twinkle/utils/nccl_safe.py index b22b10137..311590b14 100644 --- a/src/twinkle/utils/nccl_safe.py +++ b/src/twinkle/utils/nccl_safe.py @@ -78,6 +78,7 @@ def __init__(self, loss_instance): self.require_logps = getattr(loss_instance, 'require_logps', True) self.require_entropy = getattr(loss_instance, 'require_entropy', False) self.require_logits = getattr(loss_instance, 'require_logits', False) + self.enable_sampling_replay = getattr(loss_instance, 'enable_sampling_replay', False) self.reduction = getattr(loss_instance, 'reduction', 'mean') self._nccl_safe_wrapped = True diff --git a/src/twinkle/utils/torch_utils.py b/src/twinkle/utils/torch_utils.py index 84a335852..487289e5c 100644 --- a/src/twinkle/utils/torch_utils.py +++ b/src/twinkle/utils/torch_utils.py @@ -136,6 +136,122 @@ def selective_log_softmax(logits, index, return_entropy: bool = False): return per_token_logps +def replayed_selective_log_softmax( + logits: 'torch.Tensor', + labels: 'torch.Tensor', + loss_mask: 'torch.Tensor', + sampling_masks, + temperature: float, +) -> 'torch.Tensor': + """Compute selected log probabilities on rollout-time CSR support sets.""" + import math + import torch + + if not math.isfinite(temperature) or temperature <= 0: + raise ValueError('temperature must be greater than 0 for sampling replay') + if logits.dim() != 3: + raise ValueError(f'logits must have shape [batch, seq_len, vocab], got {tuple(logits.shape)}') + if labels.shape != logits.shape[:2] or loss_mask.shape != labels.shape: + raise ValueError('labels and loss_mask must match the first two logits dimensions') + if len(sampling_masks) != labels.shape[0]: + raise ValueError( + f'sampling mask batch has {len(sampling_masks)} samples, expected {labels.shape[0]}') + + vocab_size = logits.shape[-1] + flat_token_ids = [] + global_offsets = [0] + for batch_idx, sampling_mask in enumerate(sampling_masks): + if sampling_mask is None: + raise ValueError(f'sampling mask is missing for sample {batch_idx}') + token_ids = [int(token_id) for token_id in sampling_mask.token_ids] + offsets = [int(offset) for offset in sampling_mask.offsets] + if not offsets or offsets[0] != 0: + raise ValueError(f'sampling mask offsets for sample {batch_idx} must start at 0') + if offsets[-1] != len(token_ids): + raise ValueError( + f'sampling mask offsets for sample {batch_idx} must end at {len(token_ids)}') + for row_idx, (start, end) in enumerate(zip(offsets, offsets[1:])): + if start > end: + raise ValueError( + f'sampling mask offsets are not monotonic at sample {batch_idx}, row {row_idx}') + if start == end: + raise ValueError(f'sampling mask contains an empty row at sample {batch_idx}, row {row_idx}') + row_token_ids = token_ids[start:end] + if len(set(row_token_ids)) != len(row_token_ids): + raise ValueError( + f'sampling mask contains duplicate token IDs at sample {batch_idx}, row {row_idx}') + + num_rows = len(offsets) - 1 + num_train_tokens = int(loss_mask[batch_idx].sum().item()) + if num_rows != num_train_tokens: + raise ValueError( + f'sampling mask for sample {batch_idx} has {num_rows} rows but ' + f'{num_train_tokens} training tokens') + invalid_token_id = next( + (token_id for token_id in token_ids if token_id < 0 or token_id >= vocab_size), + None, + ) + if invalid_token_id is not None: + raise ValueError( + f'sampling mask token ID {invalid_token_id} is outside vocabulary [0, {vocab_size})') + + base_offset = global_offsets[-1] + flat_token_ids.extend(token_ids) + global_offsets.extend(base_offset + offset for offset in offsets[1:]) + + positions = loss_mask.nonzero(as_tuple=False) + num_rows = positions.shape[0] + if len(global_offsets) != num_rows + 1: + raise ValueError( + f'sampling masks contain {len(global_offsets) - 1} rows for {num_rows} training tokens') + result = torch.zeros(labels.shape, dtype=torch.float32, device=logits.device) + if num_rows == 0: + return result + + offsets_tensor = torch.tensor(global_offsets, dtype=torch.long, device=logits.device) + lengths = offsets_tensor[1:] - offsets_tensor[:-1] + row_ids = torch.repeat_interleave( + torch.arange(num_rows, device=logits.device), + lengths, + ) + kept_token_ids = torch.tensor(flat_token_ids, dtype=torch.long, device=logits.device) + sampled_labels = labels[positions[:, 0], positions[:, 1]].long() + + matches = kept_token_ids == sampled_labels[row_ids] + match_counts = torch.zeros(num_rows, dtype=torch.int32, device=logits.device) + match_counts.scatter_add_(0, row_ids, matches.to(torch.int32)) + missing_rows = (match_counts == 0).nonzero(as_tuple=False) + if missing_rows.numel(): + row_idx = int(missing_rows[0].item()) + raise ValueError( + f'sampled label {int(sampled_labels[row_idx].item())} is absent from ' + f'sampling mask row {row_idx}') + + kept_logits = logits[ + positions[row_ids, 0], + positions[row_ids, 1], + kept_token_ids, + ].float() / temperature + selected_logits = logits[ + positions[:, 0], + positions[:, 1], + sampled_labels, + ].float() / temperature + + row_max = torch.full( + (num_rows,), + -torch.inf, + dtype=torch.float32, + device=logits.device, + ) + row_max.scatter_reduce_(0, row_ids, kept_logits, reduce='amax', include_self=True) + row_exp_sums = torch.zeros(num_rows, dtype=torch.float32, device=logits.device) + row_exp_sums.scatter_add_(0, row_ids, torch.exp(kept_logits - row_max[row_ids])) + flat_logps = selected_logits - (row_max + torch.log(row_exp_sums)) + result[positions[:, 0], positions[:, 1]] = flat_logps + return result + + def _vocab_parallel_selective_log_softmax( logits: 'torch.Tensor', index: 'torch.Tensor', From ef2d5210e443668b490c7d6410147aef7df7d88e Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Tue, 28 Jul 2026 22:02:24 +0800 Subject: [PATCH 2/5] remove additional metric Signed-off-by: vx120 <893600387@qq.com> --- cookbook/rl/grpo/grpo_sampling_replay.py | 66 ++---------------- src/twinkle/metric/__init__.py | 1 - src/twinkle/metric/grpo.py | 21 ++---- src/twinkle/metric/rollout.py | 88 ------------------------ 4 files changed, 8 insertions(+), 168 deletions(-) delete mode 100644 src/twinkle/metric/rollout.py diff --git a/cookbook/rl/grpo/grpo_sampling_replay.py b/cookbook/rl/grpo/grpo_sampling_replay.py index 2f8a4a41f..56e215266 100644 --- a/cookbook/rl/grpo/grpo_sampling_replay.py +++ b/cookbook/rl/grpo/grpo_sampling_replay.py @@ -1,5 +1,4 @@ import os -import time from typing import List, Tuple, Dict, Any from peft import LoraConfig @@ -16,10 +15,7 @@ from twinkle.processor import InputProcessor from twinkle.reward import GSM8KAccuracyReward, GSM8KFormatReward from twinkle.sampler import vLLMSampler -from twinkle.metric import ( - CompletionRewardMetric, - compute_grpo_rollout_metrics, -) +from twinkle.metric import CompletionRewardMetric from twinkle.preprocessor.llm import GSM8KProcessor logger = get_logger() @@ -79,7 +75,6 @@ def extract_rollout_batch(sample_responses, *, require_sampling_masks: bool): 'old_logps': [], 'sampling_masks': [], 'completion_lengths': [], - 'stop_reasons': [], } for sample_response in sample_responses: for sequence in sample_response.sequences: @@ -93,7 +88,6 @@ def extract_rollout_batch(sample_responses, *, require_sampling_masks: bool): [logprob[0][1] for logprob in sequence.logprobs]) rollout_batch['sampling_masks'].append(sequence.sampling_mask) rollout_batch['completion_lengths'].append(len(sequence.tokens)) - rollout_batch['stop_reasons'].append(sequence.stop_reason) return rollout_batch @@ -191,23 +185,7 @@ def main(): top_k=-1, repetition_penalty=1.0, ) - if ENABLE_SAMPLING_REPLAY: - model.add_metric( - 'GRPOMetric', - is_training=True, - temperature=sampling_params.temperature, - epsilon=0.2, - ) - logger.info( - 'Sampling replay enabled: model_runner=v2, logprobs_mode=processed_logprobs, ' - 'temperature=%s, top_p=%s, top_k=%s', - sampling_params.temperature, - sampling_params.top_p, - sampling_params.top_k, - ) - optim_step = 0 - sampling_replay_stats_logged = False logger.info(get_device_placement()) for batch in dataloader: @@ -218,27 +196,22 @@ def main(): # enable_lora=True used with ckpt_manager.sync_weights(merge_and_sync=False) # meaning only sync lora weights, if merge_and_sync=True, # lora will be merged into the base model and sync all weights to vLLM - weight_sync_started = time.perf_counter() ckpt_manager.sync_weights(merge_and_sync=False) - weight_sync_seconds = time.perf_counter() - weight_sync_started sampler.reset_prefix_cache() def sample_prompt_groups(prompts): expand_prompts = [] for prompt in prompts: expand_prompts.extend([prompt] * NUM_GENERATIONS) - started = time.perf_counter() responses = sampler.sample(expand_prompts, sampling_params) - elapsed = time.perf_counter() - started return extract_rollout_batch( responses, require_sampling_masks=ENABLE_SAMPLING_REPLAY, - ), elapsed + ) - rollout_batch, sampling_seconds = sample_prompt_groups(global_prompts) - sampled_tokens_total = sum(rollout_batch['completion_lengths']) + rollout_batch = sample_prompt_groups(global_prompts) # Match the original GRPO control flow: every sampled rollout is scored, # logged, and trained. Zero-variance groups keep their zero advantages; - # they are diagnosed below but never resampled, dropped, or skipped. + # they are never resampled, dropped, or skipped. total_rewards, format_rewards, accuracy_rewards = compute_rewards( rollout_batch['input_data']) @@ -246,7 +219,6 @@ def sample_prompt_groups(prompts): all_old_logps: List[List[float]] = rollout_batch['old_logps'] all_sampling_masks = rollout_batch['sampling_masks'] all_completion_lengths: List[int] = rollout_batch['completion_lengths'] - all_stop_reasons = rollout_batch['stop_reasons'] metrics.accumulate( completion_lengths=all_completion_lengths, rewards={ @@ -258,20 +230,6 @@ def sample_prompt_groups(prompts): rollout_reward_log_dict = metrics.calculate() advantages = advantage_fn(total_rewards, num_generations=NUM_GENERATIONS, scale='group').tolist() - rollout_log_dict = compute_grpo_rollout_metrics( - completion_lengths=all_completion_lengths, - stop_reasons=all_stop_reasons, - rewards=total_rewards, - advantages=advantages, - num_generations=NUM_GENERATIONS, - sampling_masks=all_sampling_masks if ENABLE_SAMPLING_REPLAY else None, - ) - num_rollout_tokens = sum(all_completion_lengths) - rollout_log_dict['profiling/weight_sync_seconds'] = weight_sync_seconds - rollout_log_dict['profiling/sampling_seconds'] = sampling_seconds - rollout_log_dict['profiling/sampling_tokens_per_second'] = ( - sampled_tokens_total / sampling_seconds if sampling_seconds else 0.0) - rollout_log_dict['profiling/sampling_generated_tokens'] = sampled_tokens_total # Split completions into mini-batches and run one optim step per mini-batch. total_completions = len(all_input_data) @@ -287,7 +245,6 @@ def sample_prompt_groups(prompts): 'temperature': sampling_params.temperature, } - training_started = time.perf_counter() model.forward_backward( inputs=mb_inputs, old_logps=mb_old_logps, @@ -296,15 +253,6 @@ def sample_prompt_groups(prompts): **replay_kwargs, ) model.clip_grad_and_step() - training_seconds = time.perf_counter() - training_started - if ENABLE_SAMPLING_REPLAY and not sampling_replay_stats_logged: - logger.info( - 'Sampling replay active: sequences=%d, tokens=%d, mean_kept_tokens=%.2f', - len(all_sampling_masks), - num_rollout_tokens, - rollout_log_dict['replay/support_size_mean'], - ) - sampling_replay_stats_logged = True optim_step += 1 if optim_step % SAVE_STEPS == 0: @@ -313,12 +261,6 @@ def sample_prompt_groups(prompts): # rollout can span multiple mini-batches, but no Step lacks reward. log_dict = dict(rollout_reward_log_dict) log_dict.update(model.calculate_metric(is_training=True)) - if mb_start == 0: - log_dict.update(rollout_log_dict) - num_training_tokens = sum(all_completion_lengths[mb_start:mb_end]) - log_dict['profiling/training_seconds'] = training_seconds - log_dict['profiling/training_completion_tokens_per_second'] = ( - num_training_tokens / training_seconds if training_seconds else 0.0) logger.info(f'[Step {optim_step}/{MAX_STEPS}] {log_dict}') if optim_step >= MAX_STEPS: break diff --git a/src/twinkle/metric/__init__.py b/src/twinkle/metric/__init__.py index cd7d8c99d..baeb6c1c9 100644 --- a/src/twinkle/metric/__init__.py +++ b/src/twinkle/metric/__init__.py @@ -6,5 +6,4 @@ from .embedding import EmbeddingMetric from .grpo import CISPOMetric, GRPOMetric, GSPOMetric from .loss import LossMetric -from .rollout import compute_grpo_rollout_metrics, zero_variance_reward_group_indices from .train_metric import TrainMetric diff --git a/src/twinkle/metric/grpo.py b/src/twinkle/metric/grpo.py index 71fde1de4..bd85aab67 100644 --- a/src/twinkle/metric/grpo.py +++ b/src/twinkle/metric/grpo.py @@ -1,6 +1,6 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import math -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Union from twinkle.data_format import InputFeature, ModelOutput from twinkle.utils import get_logger @@ -9,9 +9,6 @@ logger = get_logger() -if TYPE_CHECKING: - import torch - class GRPOMetric(Metric): @@ -44,8 +41,6 @@ def reset(self): self.sum_new: float = 0.0 self.sum_old: float = 0.0 self.sum_diff: float = 0.0 - self.sum_diff_sq: float = 0.0 - self.sum_ratio: float = 0.0 self.sum_approx_kl: float = 0.0 self.max_token_kl: float = 0.0 self.max_token_ratio: float = 0.0 @@ -190,16 +185,13 @@ def _accumulate_mb( old_f = old_f * scale d = logps_f - old_f # new - old - ratio = torch.exp(d) self.sum_old += float((old_f * mask_f).sum().item()) self.sum_diff += float((d * mask_f).sum().item()) - self.sum_diff_sq += float((d.square() * mask_f).sum().item()) - self.sum_ratio += float((ratio * mask_f).sum().item()) # Schulman K3 estimator of KL(old || new): # samples x ~ old, r(x) = new(x) / old(x), # k3 = r - 1 - log(r) = exp(new - old) - (new - old) - 1. - kl = ratio - d - 1.0 + kl = torch.exp(d) - d - 1.0 kl_masked = kl * mask_f self.sum_approx_kl += float(kl_masked.sum().item()) # Per-token extremes for collapse detection @@ -208,7 +200,7 @@ def _accumulate_mb( if cur_max_kl > self.max_token_kl: self.max_token_kl = cur_max_kl # Track ratio extremes - ratio_masked = ratio * mask_f + ratio_masked = torch.exp(d) * mask_f cur_max_ratio = float(ratio_masked.max().item()) if cur_max_ratio > self.max_token_ratio: self.max_token_ratio = cur_max_ratio @@ -319,12 +311,11 @@ def accumulate( cursor += advanced def calculate(self) -> Dict[str, Any]: + import torch local = [{ 'sum_new': self.sum_new, 'sum_old': self.sum_old, 'sum_diff': self.sum_diff, - 'sum_diff_sq': self.sum_diff_sq, - 'sum_ratio': self.sum_ratio, 'sum_kl': self.sum_approx_kl, 'max_token_kl': self.max_token_kl, 'max_token_ratio': self.max_token_ratio, @@ -353,15 +344,11 @@ def calculate(self) -> Dict[str, Any]: if any(r['has_old'] for r in all_results): mean_old = sum(r['sum_old'] for r in all_results) / n_total mean_diff = sum(r['sum_diff'] for r in all_results) / n_total - mean_diff_sq = sum(r['sum_diff_sq'] for r in all_results) / n_total - mean_ratio = sum(r['sum_ratio'] for r in all_results) / n_total mean_kl = sum(r['sum_kl'] for r in all_results) / n_total global_max_kl = max(r['max_token_kl'] for r in all_results) global_max_ratio = max(r['max_token_ratio'] for r in all_results) results['train/mean_old_logp'] = mean_old results['train/logp_diff_mean'] = mean_diff - results['train/logp_diff_std'] = math.sqrt(max(mean_diff_sq - mean_diff**2, 0.0)) - results['train/importance_ratio_mean'] = mean_ratio results['train/approx_kl'] = mean_kl results['train/token_kl_max'] = global_max_kl results['train/token_ratio_max'] = global_max_ratio diff --git a/src/twinkle/metric/rollout.py b/src/twinkle/metric/rollout.py deleted file mode 100644 index cc261ccd9..000000000 --- a/src/twinkle/metric/rollout.py +++ /dev/null @@ -1,88 +0,0 @@ -# Copyright (c) ModelScope Contributors. All rights reserved. -from typing import Any, Dict, Optional, Sequence - -import numpy as np - - -def zero_variance_reward_group_indices( - rewards: Sequence[float], - num_generations: int, -) -> list[int]: - """Return GRPO group indices that cannot produce a relative advantage.""" - if num_generations <= 0: - raise ValueError('num_generations must be positive') - if len(rewards) % num_generations != 0: - raise ValueError('rewards must form complete num_generations groups') - if len(rewards) == 0: - return [] - - grouped_rewards = np.asarray(rewards, dtype=np.float64).reshape(-1, num_generations) - group_ranges = np.ptp(grouped_rewards, axis=1) - return np.flatnonzero(np.isclose(group_ranges, 0.0)).astype(int).tolist() - - -def compute_grpo_rollout_metrics( - *, - completion_lengths: Sequence[int], - stop_reasons: Sequence[str], - rewards: Sequence[float], - advantages: Sequence[float], - num_generations: int, - sampling_masks: Optional[Sequence[Any]] = None, -) -> Dict[str, float]: - """Reduce one GRPO rollout batch into scalar diagnostics.""" - if len(stop_reasons) != len(completion_lengths): - raise ValueError('stop_reasons must align with completion_lengths') - if len(rewards) != len(completion_lengths): - raise ValueError('rewards must align with completion_lengths') - if num_generations <= 0 or len(rewards) % num_generations != 0: - raise ValueError('rewards must form complete num_generations groups') - if len(advantages) != len(rewards): - raise ValueError('advantages must align with rewards') - - metrics: Dict[str, float] = {} - if len(completion_lengths) > 0: - lengths = np.asarray(completion_lengths, dtype=np.float64) - metrics['rollout/completion_length_p95'] = float(np.percentile(lengths, 95)) - - if len(stop_reasons) > 0: - num_sequences = len(stop_reasons) - metrics['rollout/stop_rate'] = sum(reason == 'stop' for reason in stop_reasons) / num_sequences - metrics['rollout/length_stop_rate'] = ( - sum(reason == 'length' for reason in stop_reasons) / num_sequences) - - if len(rewards) > 0: - grouped_rewards = np.asarray(rewards, dtype=np.float64).reshape(-1, num_generations) - if num_generations > 1: - group_stds = grouped_rewards.std(axis=1, ddof=1) - else: - group_stds = np.zeros(grouped_rewards.shape[0], dtype=np.float64) - metrics['grpo/group_reward_std_mean'] = float(group_stds.mean()) - zero_variance_groups = zero_variance_reward_group_indices(rewards, num_generations) - metrics['grpo/zero_variance_group_fraction'] = ( - len(zero_variance_groups) / grouped_rewards.shape[0]) - metrics['grpo/nonzero_advantage_fraction'] = float( - (~np.isclose(np.asarray(advantages, dtype=np.float64), 0.0)).mean()) - - if sampling_masks is not None: - if len(sampling_masks) != len(completion_lengths): - raise ValueError('sampling_masks must align with completion_lengths') - support_sizes = [] - for sequence_idx, (sampling_mask, completion_length) in enumerate( - zip(sampling_masks, completion_lengths)): - offsets = sampling_mask.offsets - if len(offsets) - 1 != completion_length: - raise ValueError( - f'sampling mask {sequence_idx} has {len(offsets) - 1} rows, ' - f'expected {completion_length}') - support_sizes.extend(end - start for start, end in zip(offsets, offsets[1:])) - - if support_sizes: - sizes = np.asarray(support_sizes, dtype=np.float64) - metrics['replay/support_size_mean'] = float(sizes.mean()) - metrics['replay/support_size_p50'] = float(np.percentile(sizes, 50)) - metrics['replay/support_size_p95'] = float(np.percentile(sizes, 95)) - metrics['replay/support_size_max'] = float(sizes.max()) - metrics['replay/singleton_fraction'] = float((sizes == 1).mean()) - - return metrics From cf6ff2fce3cd137b19bdae70b0d631e7517970a2 Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Wed, 29 Jul 2026 17:11:30 +0800 Subject: [PATCH 3/5] Added some code comments. Signed-off-by: vx120 <893600387@qq.com> --- src/twinkle/utils/torch_utils.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/src/twinkle/utils/torch_utils.py b/src/twinkle/utils/torch_utils.py index 487289e5c..34c45c1ec 100644 --- a/src/twinkle/utils/torch_utils.py +++ b/src/twinkle/utils/torch_utils.py @@ -157,6 +157,7 @@ def replayed_selective_log_softmax( raise ValueError( f'sampling mask batch has {len(sampling_masks)} samples, expected {labels.shape[0]}') + # Flatten per-sample CSR rows into one global CSR layout. vocab_size = logits.shape[-1] flat_token_ids = [] global_offsets = [0] @@ -199,6 +200,7 @@ def replayed_selective_log_softmax( flat_token_ids.extend(token_ids) global_offsets.extend(base_offset + offset for offset in offsets[1:]) + # CSR rows are ordered exactly like the masked training-token positions. positions = loss_mask.nonzero(as_tuple=False) num_rows = positions.shape[0] if len(global_offsets) != num_rows + 1: @@ -227,6 +229,7 @@ def replayed_selective_log_softmax( f'sampled label {int(sampled_labels[row_idx].item())} is absent from ' f'sampling mask row {row_idx}') + # Gather only logits retained by the rollout sampler, then normalize per CSR row. kept_logits = logits[ positions[row_ids, 0], positions[row_ids, 1], @@ -238,6 +241,7 @@ def replayed_selective_log_softmax( sampled_labels, ].float() / temperature + # Use max-shifted log-sum-exp for numerically stable restricted softmax. row_max = torch.full( (num_rows,), -torch.inf, From 710021a77e1bd967dbca9702ef0598a962c7a673 Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Wed, 29 Jul 2026 17:17:56 +0800 Subject: [PATCH 4/5] add pre-commit code Signed-off-by: vx120 <893600387@qq.com> --- .../sampler/vllm_sampler/vllm_engine.py | 11 +++---- src/twinkle/utils/torch_utils.py | 30 +++++++------------ 2 files changed, 15 insertions(+), 26 deletions(-) diff --git a/src/twinkle/sampler/vllm_sampler/vllm_engine.py b/src/twinkle/sampler/vllm_sampler/vllm_engine.py index 5fcf1d844..d4f3448a6 100644 --- a/src/twinkle/sampler/vllm_sampler/vllm_engine.py +++ b/src/twinkle/sampler/vllm_sampler/vllm_engine.py @@ -37,9 +37,8 @@ def _filter_engine_config( valid_args = set(valid_args) invalid_args = set(engine_config) - valid_args if enable_sampling_replay and 'enable_return_sampling_mask' in invalid_args: - raise RuntimeError( - 'Sampling replay requires a vLLM build whose AsyncEngineArgs accepts ' - 'enable_return_sampling_mask') + raise RuntimeError('Sampling replay requires a vLLM build whose AsyncEngineArgs accepts ' + 'enable_return_sampling_mask') filtered_engine_config = {key: value for key, value in engine_config.items() if key in valid_args} return filtered_engine_config, invalid_args @@ -54,8 +53,7 @@ def _copy_sampling_mask(mask, num_tokens: int, required: bool) -> Optional[Sampl offsets = [int(offset) for offset in mask.offsets] num_rows = len(offsets) - 1 if num_rows != num_tokens: - raise RuntimeError( - f'vLLM sampling mask has {num_rows} rows for {num_tokens} sampled tokens') + raise RuntimeError(f'vLLM sampling mask has {num_rows} rows for {num_tokens} sampled tokens') if not offsets or offsets[0] != 0 or offsets[-1] != len(token_ids): raise RuntimeError('vLLM sampling mask has invalid CSR endpoints') if any(start >= end for start, end in zip(offsets, offsets[1:])): @@ -141,8 +139,7 @@ def __init__( self.quantization = quantization self.load_format = load_format self.enable_sampling_replay = enable_sampling_replay - self.logprobs_mode = 'processed_logprobs' if enable_sampling_replay else ( - logprobs_mode or 'processed_logprobs') + self.logprobs_mode = 'processed_logprobs' if enable_sampling_replay else (logprobs_mode or 'processed_logprobs') self.engine_kwargs = kwargs or {} self._lora_request_cache: Dict[str, Any] = {} diff --git a/src/twinkle/utils/torch_utils.py b/src/twinkle/utils/torch_utils.py index 34c45c1ec..9892b79db 100644 --- a/src/twinkle/utils/torch_utils.py +++ b/src/twinkle/utils/torch_utils.py @@ -154,8 +154,7 @@ def replayed_selective_log_softmax( if labels.shape != logits.shape[:2] or loss_mask.shape != labels.shape: raise ValueError('labels and loss_mask must match the first two logits dimensions') if len(sampling_masks) != labels.shape[0]: - raise ValueError( - f'sampling mask batch has {len(sampling_masks)} samples, expected {labels.shape[0]}') + raise ValueError(f'sampling mask batch has {len(sampling_masks)} samples, expected {labels.shape[0]}') # Flatten per-sample CSR rows into one global CSR layout. vocab_size = logits.shape[-1] @@ -169,32 +168,27 @@ def replayed_selective_log_softmax( if not offsets or offsets[0] != 0: raise ValueError(f'sampling mask offsets for sample {batch_idx} must start at 0') if offsets[-1] != len(token_ids): - raise ValueError( - f'sampling mask offsets for sample {batch_idx} must end at {len(token_ids)}') + raise ValueError(f'sampling mask offsets for sample {batch_idx} must end at {len(token_ids)}') for row_idx, (start, end) in enumerate(zip(offsets, offsets[1:])): if start > end: - raise ValueError( - f'sampling mask offsets are not monotonic at sample {batch_idx}, row {row_idx}') + raise ValueError(f'sampling mask offsets are not monotonic at sample {batch_idx}, row {row_idx}') if start == end: raise ValueError(f'sampling mask contains an empty row at sample {batch_idx}, row {row_idx}') row_token_ids = token_ids[start:end] if len(set(row_token_ids)) != len(row_token_ids): - raise ValueError( - f'sampling mask contains duplicate token IDs at sample {batch_idx}, row {row_idx}') + raise ValueError(f'sampling mask contains duplicate token IDs at sample {batch_idx}, row {row_idx}') num_rows = len(offsets) - 1 num_train_tokens = int(loss_mask[batch_idx].sum().item()) if num_rows != num_train_tokens: - raise ValueError( - f'sampling mask for sample {batch_idx} has {num_rows} rows but ' - f'{num_train_tokens} training tokens') + raise ValueError(f'sampling mask for sample {batch_idx} has {num_rows} rows but ' + f'{num_train_tokens} training tokens') invalid_token_id = next( (token_id for token_id in token_ids if token_id < 0 or token_id >= vocab_size), None, ) if invalid_token_id is not None: - raise ValueError( - f'sampling mask token ID {invalid_token_id} is outside vocabulary [0, {vocab_size})') + raise ValueError(f'sampling mask token ID {invalid_token_id} is outside vocabulary [0, {vocab_size})') base_offset = global_offsets[-1] flat_token_ids.extend(token_ids) @@ -204,8 +198,7 @@ def replayed_selective_log_softmax( positions = loss_mask.nonzero(as_tuple=False) num_rows = positions.shape[0] if len(global_offsets) != num_rows + 1: - raise ValueError( - f'sampling masks contain {len(global_offsets) - 1} rows for {num_rows} training tokens') + raise ValueError(f'sampling masks contain {len(global_offsets) - 1} rows for {num_rows} training tokens') result = torch.zeros(labels.shape, dtype=torch.float32, device=logits.device) if num_rows == 0: return result @@ -225,9 +218,8 @@ def replayed_selective_log_softmax( missing_rows = (match_counts == 0).nonzero(as_tuple=False) if missing_rows.numel(): row_idx = int(missing_rows[0].item()) - raise ValueError( - f'sampled label {int(sampled_labels[row_idx].item())} is absent from ' - f'sampling mask row {row_idx}') + raise ValueError(f'sampled label {int(sampled_labels[row_idx].item())} is absent from ' + f'sampling mask row {row_idx}') # Gather only logits retained by the rollout sampler, then normalize per CSR row. kept_logits = logits[ @@ -243,7 +235,7 @@ def replayed_selective_log_softmax( # Use max-shifted log-sum-exp for numerically stable restricted softmax. row_max = torch.full( - (num_rows,), + (num_rows, ), -torch.inf, dtype=torch.float32, device=logits.device, From 8678bd150c24c5d61b5ec05f9c4e2e2dfa17e3ed Mon Sep 17 00:00:00 2001 From: vx120 <893600387@qq.com> Date: Wed, 29 Jul 2026 19:48:40 +0800 Subject: [PATCH 5/5] reduce the judge code Signed-off-by: vx120 <893600387@qq.com> --- cookbook/rl/grpo/grpo_sampling_replay.py | 10 ++-------- src/twinkle/loss/grpo.py | 5 ----- src/twinkle/model/transformers/transformers.py | 4 ---- src/twinkle/utils/torch_utils.py | 17 +++-------------- 4 files changed, 5 insertions(+), 31 deletions(-) diff --git a/cookbook/rl/grpo/grpo_sampling_replay.py b/cookbook/rl/grpo/grpo_sampling_replay.py index 56e215266..3fba27562 100644 --- a/cookbook/rl/grpo/grpo_sampling_replay.py +++ b/cookbook/rl/grpo/grpo_sampling_replay.py @@ -68,7 +68,7 @@ def compute_rewards( return total_rewards, format_rewards, accuracy_rewards -def extract_rollout_batch(sample_responses, *, require_sampling_masks: bool): +def extract_rollout_batch(sample_responses): """Flatten sampler responses into aligned lists used by reward and training.""" rollout_batch = { 'input_data': [], @@ -80,9 +80,6 @@ def extract_rollout_batch(sample_responses, *, require_sampling_masks: bool): for sequence in sample_response.sequences: if sequence.logprobs is None: raise RuntimeError('A sampled sequence is missing token log probabilities') - if require_sampling_masks and sequence.sampling_mask is None: - raise RuntimeError( - 'Sampling replay is enabled but a sampled sequence has no sampling mask') rollout_batch['input_data'].append(sequence.new_input_feature) rollout_batch['old_logps'].append( [logprob[0][1] for logprob in sequence.logprobs]) @@ -203,10 +200,7 @@ def sample_prompt_groups(prompts): for prompt in prompts: expand_prompts.extend([prompt] * NUM_GENERATIONS) responses = sampler.sample(expand_prompts, sampling_params) - return extract_rollout_batch( - responses, - require_sampling_masks=ENABLE_SAMPLING_REPLAY, - ) + return extract_rollout_batch(responses) rollout_batch = sample_prompt_groups(global_prompts) # Match the original GRPO control flow: every sampled rollout is scored, diff --git a/src/twinkle/loss/grpo.py b/src/twinkle/loss/grpo.py index 36970636d..52479b92b 100644 --- a/src/twinkle/loss/grpo.py +++ b/src/twinkle/loss/grpo.py @@ -40,8 +40,6 @@ def __init__( self.beta = beta self.entropy_coef = entropy_coef self.enable_sampling_replay = enable_sampling_replay - if enable_sampling_replay and self.__class__ is not GRPOLoss: - raise ValueError('sampling replay is only supported by GRPOLoss') if enable_sampling_replay and beta != 0.0: raise ValueError('sampling replay does not support a GRPO KL penalty (beta must be 0)') if enable_sampling_replay and entropy_coef != 0.0: @@ -209,7 +207,6 @@ def __call__( old_logps: Optional[Union['torch.Tensor', List[List[float]]]] = None, ref_logps: Optional['torch.Tensor'] = None, advantages: Optional[Union['torch.Tensor', List[float], np.ndarray]] = None, - sampling_masks=None, **kwargs, ): """ @@ -232,8 +229,6 @@ def __call__( """ import torch if self.enable_sampling_replay: - if sampling_masks is None: - raise ValueError('sampling_masks are required when sampling replay is enabled') if old_logps is None: raise ValueError('old_logps are required when sampling replay is enabled') labels = inputs.get('labels') diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index aef56f1ff..5afc46c0a 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -471,8 +471,6 @@ def forward(self, *, inputs: Union[InputFeature, List[InputFeature], List[Trajec if enable_sampling_replay: if sampling_masks is None: raise ValueError('sampling_masks are required when sampling replay is enabled') - if kwargs.get('old_logps') is None: - raise ValueError('old_logps are required when sampling replay is enabled') cp_world_size = self.device_mesh.cp_world_size if self.device_mesh is not None else 1 if getattr(self, '_enable_sp', False) or cp_world_size > 1: raise ValueError('sampling replay does not support sequence or context parallelism') @@ -581,8 +579,6 @@ def forward_only(self, *, inputs: Union[InputFeature, List[InputFeature], List[T if enable_sampling_replay: if sampling_masks is None: raise ValueError('sampling_masks are required when sampling replay is enabled') - if kwargs.get('old_logps') is None: - raise ValueError('old_logps are required when sampling replay is enabled') cp_world_size = self.device_mesh.cp_world_size if self.device_mesh is not None else 1 if getattr(self, '_enable_sp', False) or cp_world_size > 1: raise ValueError('sampling replay does not support sequence or context parallelism') diff --git a/src/twinkle/utils/torch_utils.py b/src/twinkle/utils/torch_utils.py index 9892b79db..f6c7a008e 100644 --- a/src/twinkle/utils/torch_utils.py +++ b/src/twinkle/utils/torch_utils.py @@ -136,6 +136,9 @@ def selective_log_softmax(logits, index, return_entropy: bool = False): return per_token_logps +# Re-normalize trainer logits over each rollout-time top-p/top-k support set +# before reading the sampled token's log probability. Replaying the sampler's +# action space removes the sampling/training distribution mismatch in GRPO. def replayed_selective_log_softmax( logits: 'torch.Tensor', labels: 'torch.Tensor', @@ -165,18 +168,6 @@ def replayed_selective_log_softmax( raise ValueError(f'sampling mask is missing for sample {batch_idx}') token_ids = [int(token_id) for token_id in sampling_mask.token_ids] offsets = [int(offset) for offset in sampling_mask.offsets] - if not offsets or offsets[0] != 0: - raise ValueError(f'sampling mask offsets for sample {batch_idx} must start at 0') - if offsets[-1] != len(token_ids): - raise ValueError(f'sampling mask offsets for sample {batch_idx} must end at {len(token_ids)}') - for row_idx, (start, end) in enumerate(zip(offsets, offsets[1:])): - if start > end: - raise ValueError(f'sampling mask offsets are not monotonic at sample {batch_idx}, row {row_idx}') - if start == end: - raise ValueError(f'sampling mask contains an empty row at sample {batch_idx}, row {row_idx}') - row_token_ids = token_ids[start:end] - if len(set(row_token_ids)) != len(row_token_ids): - raise ValueError(f'sampling mask contains duplicate token IDs at sample {batch_idx}, row {row_idx}') num_rows = len(offsets) - 1 num_train_tokens = int(loss_mask[batch_idx].sum().item()) @@ -197,8 +188,6 @@ def replayed_selective_log_softmax( # CSR rows are ordered exactly like the masked training-token positions. positions = loss_mask.nonzero(as_tuple=False) num_rows = positions.shape[0] - if len(global_offsets) != num_rows + 1: - raise ValueError(f'sampling masks contain {len(global_offsets) - 1} rows for {num_rows} training tokens') result = torch.zeros(labels.shape, dtype=torch.float32, device=logits.device) if num_rows == 0: return result