-
Notifications
You must be signed in to change notification settings - Fork 41
feat: add standard PPO training with GAE and value critic (Transformers/FSDP) #256
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
xxyyrr598
wants to merge
7
commits into
modelscope:main
Choose a base branch
from
xxyyrr598:add-ppo
base: main
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
+911
−9
Open
Changes from all commits
Commits
Show all changes
7 commits
Select commit
Hold shift + click to select a range
d5fb82c
feat: add standard PPO training with GAE and value critic (Transforme…
xxyyrr598 46fb917
fix: force CPU in value model tests for GPU hosts
xxyyrr598 8eb12a4
fix: align value model test inputs with model device
xxyyrr598 aa3bf04
fix: wrap model before reading device in value model tests
xxyyrr598 13910f6
Add .gitattributes to normalize line endings
xxyyrr598 5edbde8
refactor: unify PPO epochs via training.num_train_epochs and extract …
xxyyrr598 1ce486d
feat(loss): add loss_agg_mode to PPOLoss, default to token-mean
xxyyrr598 File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1 @@ | ||
| * text=auto | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,240 @@ | ||
| """Standard PPO training on GSM8K with a LoRA policy and full-parameter critic. | ||
|
|
||
| The first implementation supports the Transformers/Accelerate-FSDP backend. Policy, | ||
| critic, and vLLM sampler use separate GPU groups. The frozen policy base model is | ||
| used as the reference policy. | ||
| """ | ||
| 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 | ||
| 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, PPOMetric, PPOValueMetric | ||
| from twinkle.model import TransformersModel, TransformersValueModel | ||
| from twinkle.processor import InputProcessor | ||
| from twinkle.preprocessor.llm import GSM8KProcessor | ||
| 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 | ||
| NUM_GENERATIONS = args.rl.num_generations or 4 | ||
| MAX_NEW_TOKENS = args.sampling.max_tokens or 1024 | ||
| POLICY_LR = args.optimizer.learning_rate or 1e-5 | ||
| 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 | ||
| # Number of policy/value updates over each rollout batch. Reuse the common | ||
| # training argument whaohile preserving PPO's historical default. | ||
| PPO_EPOCHS = args.training.num_train_epochs if args.training.num_train_epochs is not None else 4 | ||
| SAVE_STEPS = args.training.save_steps or 50 | ||
| ADAPTER_NAME = args.lora.adapter_name or 'default' | ||
|
|
||
|
|
||
| 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]]: | ||
| """Extract response-token rows from collected model outputs.""" | ||
| 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 | ||
| rows.append(value_row[:mask.numel()][mask].detach().float().cpu().tolist()) | ||
| return rows | ||
|
|
||
|
|
||
| def main(): | ||
| 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') | ||
| lora_config = 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, | ||
| ) | ||
| policy.add_adapter_to_model(ADAPTER_NAME, lora_config, 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( | ||
| 'PPOLoss', | ||
| epsilon=args.loss.epsilon, | ||
| entropy_coef=args.loss.entropy_coef, | ||
| loss_agg_mode='token-mean', | ||
| ) | ||
| policy.add_metric(PPOMetric, epsilon=args.loss.epsilon) | ||
| 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') | ||
| critic.set_optimizer('AdamW', lr=CRITIC_LR) | ||
| critic.set_lr_scheduler('CosineAnnealingLR', T_max=MAX_STEPS, eta_min=0) | ||
| critic.set_loss('PPOValueLoss', epsilon=args.loss.value_clip) | ||
| critic.add_metric(PPOValueMetric, epsilon=args.loss.value_clip) | ||
| 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', | ||
| ) | ||
| gae = GAEAdvantage(args.rl.gamma, args.rl.gae_lambda, args.rl.normalize_advantages) | ||
| reward_metric = CompletionRewardMetric() | ||
| sampling_params = SamplingParams(max_tokens=MAX_NEW_TOKENS, 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() | ||
| expanded = [prompt for prompt in prompts for _ in range(NUM_GENERATIONS)] | ||
| samples = sampler.sample(expanded, sampling_params) | ||
|
|
||
| trajectories, old_logps, lengths = [], [], [] | ||
| for response in samples: | ||
| for sequence in response.sequences: | ||
| trajectories.append(sequence.new_input_feature) | ||
| old_logps.append([entry[0][1] for entry in sequence.logprobs]) | ||
| lengths.append(len(sequence.tokens)) | ||
| rewards, format_rewards, accuracy_rewards = compute_rewards(trajectories) | ||
| reward_metric.accumulate( | ||
| completion_lengths=lengths, | ||
| rewards={'total': rewards, 'format': format_rewards, 'accuracy': accuracy_rewards}, | ||
| ) | ||
|
|
||
| reference = policy.forward_only(inputs=trajectories, disable_lora=True) | ||
| ref_logps = response_rows(reference['logps'], trajectories) | ||
| critic_outputs = critic.forward_only(inputs=trajectories) | ||
| old_values = response_rows(critic_outputs['values'], trajectories) | ||
| token_rewards = gae.build_token_rewards( | ||
| rewards, lengths, old_logps=old_logps, ref_logps=ref_logps, kl_coef=args.rl.kl_coef) | ||
| max_len = max(lengths) | ||
| padded_rewards = [row + [0.0] * (max_len - len(row)) for row in token_rewards] | ||
| padded_values = [row + [0.0] * (max_len - len(row)) for row in old_values] | ||
| masks = [[True] * length + [False] * (max_len - length) for length in lengths] | ||
| advantages, returns = gae(padded_rewards, padded_values, masks=masks) | ||
| advantages = [advantages[i, :length].tolist() for i, length in enumerate(lengths)] | ||
| returns = [returns[i, :length].tolist() for i, length in enumerate(lengths)] | ||
|
|
||
| indices = list(range(len(trajectories))) | ||
| for _ in range(PPO_EPOCHS): | ||
| 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_old_logps = [old_logps[i] for i in chosen] | ||
| mb_old_values = [old_values[i] for i in chosen] | ||
| mb_advantages = [advantages[i] for i in chosen] | ||
| mb_returns = [returns[i] for i in chosen] | ||
| policy.forward_backward( | ||
| inputs=mb_inputs, | ||
| old_logps=mb_old_logps, | ||
| advantages=mb_advantages, | ||
| micro_batch_size=MICRO_BATCH_SIZE, | ||
| ) | ||
| policy.clip_grad_and_step() | ||
| critic.forward_backward( | ||
| inputs=mb_inputs, | ||
| old_values=mb_old_values, | ||
| returns=mb_returns, | ||
| advantages=mb_advantages, | ||
| micro_batch_size=MICRO_BATCH_SIZE, | ||
| ) | ||
| critic.clip_grad_and_step() | ||
| optim_step += 1 | ||
| if optim_step % SAVE_STEPS == 0: | ||
| policy.save(f'ppo-policy-checkpoint-{optim_step}') | ||
| critic.save(f'ppo-critic-checkpoint-{optim_step}') | ||
| if optim_step >= MAX_STEPS: | ||
| break | ||
| if optim_step >= MAX_STEPS: | ||
| break | ||
|
|
||
| logs = reward_metric.calculate() | ||
| logs.update(policy.calculate_metric(is_training=True)) | ||
| logs.update(critic.calculate_metric(is_training=True)) | ||
| rollout_step += 1 | ||
| logger.info(f'[Rollout {rollout_step}, optim step {optim_step}/{MAX_STEPS}] {logs}') | ||
|
|
||
| policy.save('ppo-policy-final') | ||
| critic.save('ppo-critic-final') | ||
|
|
||
|
|
||
| if __name__ == '__main__': | ||
| main() |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,29 @@ | ||
| #!/bin/sh | ||
| set -eu | ||
|
|
||
| # Standard PPO on GSM8K via Ray. | ||
| # Transformers/Accelerate-FSDP: 4 policy + 4 full-parameter critic + 4 sampler GPUs. | ||
| # Override any option after the defaults, for example: | ||
| # sh ppo.sh --max-steps 20 --num-train-epochs 1 | ||
|
|
||
| python ppo.py \ | ||
| --model-id ms://Qwen/Qwen3.5-4B \ | ||
| --model-gpus 4 \ | ||
| --critic-model-gpus 4 \ | ||
| --sampler-gpus 4 \ | ||
| --num-generations 2 \ | ||
| --max-tokens 1024 \ | ||
| --batch-size 4 \ | ||
| --mini-batch-size 4 \ | ||
| --micro-batch-size 1 \ | ||
| --num-train-epochs 4 \ | ||
| --gamma 1.0 \ | ||
| --gae-lambda 0.95 \ | ||
| --kl-coef 0.01 \ | ||
| --value-clip 0.2 \ | ||
| --lr 1e-5 \ | ||
| --critic-learning-rate 1e-5 \ | ||
| --max-steps 200 \ | ||
| --save-steps 50 \ | ||
| --adapter-name default \ | ||
| "$@" |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1,10 +1,12 @@ | ||
| # Copyright (c) ModelScope Contributors. All rights reserved. | ||
| from .base import Advantage | ||
| from .gae import GAEAdvantage | ||
| from .grpo import GRPOAdvantage | ||
| from .rloo import RLOOAdvantage | ||
|
|
||
| __all__ = [ | ||
| 'Advantage', | ||
| 'GAEAdvantage', | ||
| 'GRPOAdvantage', | ||
| 'RLOOAdvantage', | ||
| ] |
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
这里是为什么
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
这个改动纯粹是行尾规范化,不影响任何运行逻辑,整个文件只有一行:
作用:
text=auto 让 Git 自动区分文本/二进制文件。文本文件入库时统一按 LF 存储(无论提交者是什么平台),检出时再按当前平台转回(Windows 用 CRLF,Linux/macOS 用 LF);二进制文件完全不碰。
为什么加:
我在 Windows/WSL 下开发,而 CI 跑在 Linux 上。不做规范化的话,只改一行代码就可能因为行尾(CRLF vs LF)变化导致整个文件显示成全量修改,污染 diff、干扰 review。
同时避免把 CRLF 意外提交进仓库,防止 Linux CI 上出现由 \r 引起的偶发失败(比如 shell/python 文件)。
它只影响 Git 对后续提交的行尾处理,不改变任何代码、依赖或行为。如果维护者认为这个 PR 不需要包含该规范化,我可以移除。