Skip to content
Closed
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion examples/split_placement/split_monkey_patch.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
"""
from pprint import pprint
from verl import DataProto
from verl.trainer.ppo.ray_trainer import compute_advantage, apply_kl_penalty, reduce_metrics, compute_data_metrics, _timer, compute_timing_metrics, AdvantageEstimator
from verl.trainer.ppo.ray_trainer import compute_advantage, apply_kl_penalty, reduce_metrics, compute_data_metrics, _timer, compute_timing_metrics, AdvantageEstimator, calc_mini_batch_loss_token_nums
from copy import deepcopy
import numpy as np
import torch
Expand Down Expand Up @@ -101,6 +101,12 @@ def fit(self):
# compute global_valid tokens
batch.meta_info['global_token_num'] = torch.sum(batch.batch['attention_mask'], dim=-1).tolist()

batch.meta_info["mini_batch_loss_token_nums"] = calc_mini_batch_loss_token_nums(
batch,
traj_mini_bsz=self.config.actor_rollout_ref.actor.ppo_mini_batch_size *
self.config.actor_rollout_ref.rollout.n,
num_dp_ranks=self.actor_rollout_wg.world_size)

# recompute old_log_probs
with _timer('old_log_prob', timing_raw):
old_log_prob = self.actor_rollout_wg.compute_log_prob(batch)
Expand Down
8 changes: 7 additions & 1 deletion recipe/dapo/src/dapo_ray_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
import torch

from verl import DataProto
from verl.trainer.ppo.ray_trainer import RayPPOTrainer, _timer, apply_kl_penalty, compute_advantage, AdvantageEstimator
from verl.trainer.ppo.ray_trainer import RayPPOTrainer, _timer, apply_kl_penalty, compute_advantage, AdvantageEstimator, calc_mini_batch_loss_token_nums
from verl.trainer.ppo.metric_utils import (compute_data_metrics, compute_throughout_metrics, compute_timing_metrics,
reduce_metrics)

Expand Down Expand Up @@ -223,6 +223,12 @@ def fit(self):
# compute global_valid tokens
batch.meta_info['global_token_num'] = torch.sum(batch.batch['attention_mask'], dim=-1).tolist()

batch.meta_info["mini_batch_loss_token_nums"] = calc_mini_batch_loss_token_nums(
batch,
traj_mini_bsz=self.config.actor_rollout_ref.actor.ppo_mini_batch_size *
self.config.actor_rollout_ref.rollout.n,
num_dp_ranks=self.actor_rollout_wg.world_size)

# recompute old_log_probs
with _timer('old_log_prob', timing_raw):
old_log_prob = self.actor_rollout_wg.compute_log_prob(batch)
Expand Down
8 changes: 7 additions & 1 deletion recipe/prime/prime_ray_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
from verl import DataProto
from verl.single_controller.ray import RayWorkerGroup
from verl.trainer.ppo.ray_trainer import RayPPOTrainer
from verl.trainer.ppo.ray_trainer import Role, WorkerType, ResourcePoolManager, reduce_metrics, _timer
from verl.trainer.ppo.ray_trainer import Role, WorkerType, ResourcePoolManager, reduce_metrics, _timer, calc_mini_batch_loss_token_nums
from verl.trainer.ppo.metric_utils import _compute_response_info
from verl.utils.checkpoint.checkpoint_manager import find_latest_ckpt_path
from verl.utils.dataset.rl_dataset import RLHFDataset, collate_fn
Expand Down Expand Up @@ -394,6 +394,12 @@ def fit(self):
# compute global_valid tokens
batch.meta_info['global_token_num'] = torch.sum(batch.batch['attention_mask'], dim=-1).tolist()

batch.meta_info["mini_batch_loss_token_nums"] = calc_mini_batch_loss_token_nums(
batch,
traj_mini_bsz=self.config.actor_rollout_ref.actor.ppo_mini_batch_size *
self.config.actor_rollout_ref.rollout.n,
num_dp_ranks=self.actor_rollout_wg.world_size)

# verify
with _timer('verify', timing_raw):
scores = self.reward_fn.verify(batch)
Expand Down
1 change: 1 addition & 0 deletions verl/trainer/config/ppo_trainer.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -147,6 +147,7 @@ critic:
shuffle: ${actor_rollout_ref.actor.shuffle}
grad_clip: 1.0
cliprange_value: 0.5
loss_agg_mode: ${actor_rollout_ref.actor.loss_agg_mode}
checkpoint:
contents: ['model', 'optimizer', 'extra'] # with 'hf_model' you can save whole model as hf format, now only use sharded model checkpoint to save space

Expand Down
39 changes: 34 additions & 5 deletions verl/trainer/ppo/ray_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,12 +137,12 @@ def _check_resource_available(self):


def apply_kl_penalty(data: DataProto, kl_ctrl: core_algos.AdaptiveKLController, kl_penalty='kl'):
responses = data.batch['responses']
response_length = responses.size(1)
# Back-compatible with trainers that do not compute response mask in fit
if "response_mask" not in data.batch.keys():
data.batch['response_mask'] = compute_response_mask(data)
response_mask = data.batch['response_mask']
token_level_scores = data.batch['token_level_scores']
batch_size = data.batch.batch_size[0]
attention_mask = data.batch['attention_mask']
response_mask = attention_mask[:, -response_length:]

# compute kl between ref_policy and current policy
# When apply_kl_penalty, algorithm.use_kl_in_reward=True, so the reference model has been enabled.
Expand Down Expand Up @@ -172,14 +172,37 @@ def compute_response_mask(data: DataProto):
return attention_mask[:, -response_length:]


def calc_mini_batch_loss_token_nums(batch: DataProto, traj_mini_bsz: int, num_dp_ranks: int) -> list[int]:
if "response_mask" not in batch.batch.keys():
batch.batch['response_mask'] = compute_response_mask(batch)
response_mask = batch.batch['response_mask']

traj_bsz = len(batch.batch)
num_mini_batches = (traj_bsz + traj_mini_bsz - 1) // traj_mini_bsz
traj_mini_bsz_per_rank = traj_mini_bsz // num_dp_ranks

mini_batch_loss_token_nums = []
for _ in range(num_mini_batches):
mini_batch_traj_idxs = []
for dp_rank in range(num_dp_ranks):
start_traj_idx = int(traj_bsz / num_dp_ranks * dp_rank)
next_start_traj_idx = int(traj_bsz / num_dp_ranks * (dp_rank + 1))
end_traj_idx = int(min(start_traj_idx + traj_mini_bsz_per_rank, next_start_traj_idx))
mini_batch_traj_idxs.extend(list(range(start_traj_idx, end_traj_idx)))
mini_batch_resp_mask = response_mask[mini_batch_traj_idxs]
mini_batch_loss_token_num = mini_batch_resp_mask.sum()
mini_batch_loss_token_nums.append(mini_batch_loss_token_num)

return mini_batch_loss_token_nums


def compute_advantage(data: DataProto, adv_estimator, gamma=1.0, lam=1.0, num_repeat=1):
# Back-compatible with trainers that do not compute response mask in fit
if "response_mask" not in data.batch.keys():
data.batch['response_mask'] = compute_response_mask(data)
# prepare response group
# TODO: add other ways to estimate advantages
if adv_estimator == AdvantageEstimator.GAE:
values = data.batch['values']
advantages, returns = core_algos.compute_gae_advantage_return(
token_level_rewards=data.batch['token_level_rewards'],
values=data.batch['values'],
Expand Down Expand Up @@ -871,6 +894,12 @@ def fit(self):
# compute global_valid tokens
batch.meta_info['global_token_num'] = torch.sum(batch.batch['attention_mask'], dim=-1).tolist()

batch.meta_info["mini_batch_loss_token_nums"] = self.calc_mini_batch_loss_token_nums(
batch,
traj_mini_bsz=self.config.actor_rollout_ref.actor.ppo_mini_batch_size *
self.config.actor_rollout_ref.rollout.n,
num_dp_ranks=self.actor_rollout_wg.world_size)

# recompute old_log_probs
with _timer('old_log_prob', timing_raw):
old_log_prob = self.actor_rollout_wg.compute_log_prob(batch)
Expand Down
47 changes: 27 additions & 20 deletions verl/workers/actor/dp_actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,13 +255,11 @@ def update_policy(self, data: DataProto):

metrics = {}
for epoch in range(self.config.ppo_epochs):
for batch_idx, data in enumerate(dataloader):
# split batch into micro_batches
mini_batch = data
for mini_idx, mini_batch in enumerate(dataloader):
if has_multi_modal_inputs:
self.gradient_accumulation = self.config.ppo_mini_batch_size // self.config.ppo_micro_batch_size_per_gpu
num_micro_batches = mini_batch.batch.batch_size[0] // self.config.ppo_micro_batch_size_per_gpu
micro_batches = data.select(select_keys, non_tensor_select_keys).chunk(num_micro_batches)
micro_batches = mini_batch.select(select_keys, non_tensor_select_keys).chunk(num_micro_batches)
elif self.config.use_dynamic_bsz:
max_token_len = self.config.ppo_max_token_len_per_gpu * self.ulysses_sequence_parallel_size
micro_batches, _ = rearrange_micro_batches(batch=mini_batch, max_token_len=max_token_len)
Expand All @@ -272,18 +270,22 @@ def update_policy(self, data: DataProto):

self.actor_optimizer.zero_grad()

for data in micro_batches:
for micro_batch in micro_batches:
# Support all hardwares
if isinstance(data, DataProto):
data = {**data.batch.to(torch.cuda.current_device()), **data.non_tensor_batch}
if isinstance(micro_batch, DataProto):
micro_batch = {
**micro_batch.batch.to(torch.cuda.current_device()),
**micro_batch.non_tensor_batch
}
else:
data = data.to(torch.cuda.current_device()) # actor device is cpu when using offload
responses = data['responses']
micro_batch = micro_batch.to(
torch.cuda.current_device()) # actor device is cpu when using offload
responses = micro_batch['responses']
response_length = responses.size(1)
attention_mask = data['attention_mask']
attention_mask = micro_batch['attention_mask']
response_mask = attention_mask[:, -response_length:]
old_log_prob = data['old_log_probs']
advantages = data['advantages']
old_log_prob = micro_batch['old_log_probs']
advantages = micro_batch['advantages']

clip_ratio = self.config.clip_ratio
clip_ratio_low = self.config.clip_ratio_low if self.config.clip_ratio_low is not None else clip_ratio
Expand All @@ -293,7 +295,7 @@ def update_policy(self, data: DataProto):
loss_agg_mode = self.config.loss_agg_mode

# all return: (bsz, response_length)
entropy, log_prob = self._forward_micro_batch(micro_batch=data, temperature=temperature)
entropy, log_prob = self._forward_micro_batch(micro_batch=micro_batch, temperature=temperature)

pg_loss, pg_clipfrac, ppo_kl, pg_clipfrac_lower = compute_policy_loss(
old_log_prob=old_log_prob,
Expand All @@ -311,7 +313,7 @@ def update_policy(self, data: DataProto):
policy_loss = pg_loss - entropy_loss * entropy_coeff

if self.config.use_kl_loss:
ref_log_prob = data['ref_log_prob']
ref_log_prob = micro_batch['ref_log_prob']
# compute kl loss
kld = kl_penalty(logprob=log_prob,
ref_logprob=ref_log_prob,
Expand All @@ -325,23 +327,28 @@ def update_policy(self, data: DataProto):
metrics['actor/kl_coef'] = self.config.kl_loss_coef

if self.config.use_dynamic_bsz:
# relative to the dynamic bsz
loss = policy_loss * (len(data) / self.config.ppo_mini_batch_size)
if self.config.loss_agg_mode == 'token-mean':
mini_batch_loss_token_nums = data.meta_info['mini_batch_loss_token_nums']
mini_batch_loss_token_num = mini_batch_loss_token_nums[mini_idx]
num_valid_toks = response_mask.sum()
loss = policy_loss * num_valid_toks / mini_batch_loss_token_num
else: # seq-mean
loss = policy_loss * (len(data) / self.config.ppo_mini_batch_size)
else:
loss = policy_loss / self.gradient_accumulation
loss.backward()

data = {
mini_metric_data = {
'actor/entropy': entropy_loss.detach().item(),
'actor/pg_loss': pg_loss.detach().item(),
'actor/pg_clipfrac': pg_clipfrac.detach().item(),
'actor/ppo_kl': ppo_kl.detach().item(),
'actor/pg_clipfrac_lower': pg_clipfrac_lower.detach().item(),
}
append_to_dict(metrics, data)
append_to_dict(metrics, mini_metric_data)

grad_norm = self._optimizer_step()
data = {'actor/grad_norm': grad_norm.detach().item()}
append_to_dict(metrics, data)
metric_data = {'actor/grad_norm': grad_norm.detach().item()}
append_to_dict(metrics, metric_data)
self.actor_optimizer.zero_grad()
return metrics
50 changes: 28 additions & 22 deletions verl/workers/critic/dp_critic.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,12 +188,10 @@ def update_critic(self, data: DataProto):
dataloader = batch.split(self.config.ppo_mini_batch_size)

for epoch in range(self.config.ppo_epochs):
for batch_idx, data in enumerate(dataloader):
# split batch into micro_batches
mini_batch = data
for mini_idx, mini_batch in enumerate(dataloader):
if has_multi_modal_inputs:
num_micro_batches = mini_batch.batch.batch_size[0] // self.config.ppo_micro_batch_size_per_gpu
micro_batches = data.select(select_keys, non_tensor_select_keys).chunk(num_micro_batches)
micro_batches = mini_batch.select(select_keys, non_tensor_select_keys).chunk(num_micro_batches)
elif self.config.use_dynamic_bsz:
max_token_len = self.config.ppo_max_token_len_per_gpu * self.ulysses_sequence_parallel_size
micro_batches, _ = rearrange_micro_batches(batch=mini_batch, max_token_len=max_token_len)
Expand All @@ -203,23 +201,27 @@ def update_critic(self, data: DataProto):

self.critic_optimizer.zero_grad()

for data in micro_batches:
for micro_batch in micro_batches:
#Support all devices
if isinstance(data, DataProto):
data = {**data.batch.to(torch.cuda.current_device()), **data.non_tensor_batch}
if isinstance(micro_batch, DataProto):
micro_batch = {
**micro_batch.batch.to(torch.cuda.current_device()),
**micro_batch.non_tensor_batch
}
else:
data = data.to(torch.cuda.current_device()) # critic device is cpu when using offload
input_ids = data['input_ids']
responses = data['responses']
attention_mask = data['attention_mask']
position_ids = data['position_ids']
values = data['values']
returns = data['returns']
micro_batch = micro_batch.to(
torch.cuda.current_device()) # critic device is cpu when using offload
input_ids = micro_batch['input_ids']
responses = micro_batch['responses']
attention_mask = micro_batch['attention_mask']
position_ids = micro_batch['position_ids']
values = micro_batch['values']
returns = micro_batch['returns']
response_length = responses.size(1)

response_mask = attention_mask[:, -response_length - 1:-1]

vpreds = self._forward_micro_batch(data)
vpreds = self._forward_micro_batch(micro_batch)

# assert not torch.any(torch.isnan(vpreds)).item()

Expand All @@ -229,23 +231,27 @@ def update_critic(self, data: DataProto):
response_mask=response_mask,
cliprange_value=self.config.cliprange_value)
if self.config.use_dynamic_bsz:
# relative to the dynamic bsz
loss = vf_loss * (len(data) / self.config.ppo_mini_batch_size)
if self.config.loss_agg_mode == 'token-mean':
mini_batch_loss_token_nums = data.meta_info['mini_batch_loss_token_nums']
mini_batch_loss_token_num = mini_batch_loss_token_nums[mini_idx]
num_valid_toks = response_mask.sum()
loss = vf_loss * num_valid_toks / mini_batch_loss_token_num
else: # seq-mean
loss = vf_loss * (len(data) / self.config.ppo_mini_batch_size)
else:
loss = vf_loss / self.gradient_accumulation

loss.backward()

data = {
mini_metric_data = {
'critic/vf_loss': vf_loss.detach().item(),
'critic/vf_clipfrac': vf_clipfrac.detach().item(),
'critic/vpred_mean': masked_mean(vpreds, response_mask).detach().item(),
}

append_to_dict(metrics, data)
append_to_dict(metrics, mini_metric_data)

grad_norm = self._optimizer_step()
data = {'critic/grad_norm': grad_norm.detach().item()}
append_to_dict(metrics, data)
metric_data = {'critic/grad_norm': grad_norm.detach().item()}
append_to_dict(metrics, metric_data)
self.critic_optimizer.zero_grad()
return metrics