Commit 2b8a9347 authored by jameskrw's avatar jameskrw
Browse files

added support for real grpo and ppo

parent e60df3d0
Loading
Loading
Loading
Loading
+5 −3
Original line number Diff line number Diff line
@@ -31,14 +31,14 @@ python3 -m vagen.trainer.main_ppo \
    actor_rollout_ref.actor.optim.lr=1e-6 \
    actor_rollout_ref.model.use_remove_padding=False \
    actor_rollout_ref.actor.ppo_mini_batch_size=32 \
    actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=8 \
    actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=1 \
    actor_rollout_ref.actor.use_kl_loss=True \
    actor_rollout_ref.actor.kl_loss_coef=0.001 \
    actor_rollout_ref.actor.kl_loss_type=low_var_kl \
    actor_rollout_ref.model.enable_gradient_checkpointing=True \
    actor_rollout_ref.actor.fsdp_config.param_offload=False \
    actor_rollout_ref.actor.fsdp_config.optimizer_offload=False \
    actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=8 \
    actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=1 \
    actor_rollout_ref.rollout.tensor_model_parallel_size=1 \
    actor_rollout_ref.rollout.name=vllm \
    actor_rollout_ref.rollout.gpu_memory_utilization=0.4 \
@@ -46,7 +46,7 @@ python3 -m vagen.trainer.main_ppo \
    actor_rollout_ref.rollout.enforce_eager=False \
    actor_rollout_ref.rollout.free_cache_engine=False \
    actor_rollout_ref.rollout.n=1 \
    actor_rollout_ref.ref.log_prob_micro_batch_size_per_gpu=8 \
    actor_rollout_ref.ref.log_prob_micro_batch_size_per_gpu=1 \
    actor_rollout_ref.ref.fsdp_config.param_offload=True \
    algorithm.kl_ctrl.kl_coef=0.001 \
    trainer.critic_warmup=0 \
@@ -62,4 +62,6 @@ python3 -m vagen.trainer.main_ppo \
    rollout_manger.window_size=5 \
    trainer.val_before_train=True \
    trainer.val_generations_to_log_to_wandb=5 \
    # grpo sampling param
    rollout_manger.n_trajectory=8 \ 
    2>&1 | tee debug_qwen0_5_1_gpu_grpo.log
+2 −0
Original line number Diff line number Diff line
@@ -61,6 +61,8 @@ python3 -m vagen.trainer.main_ppo \
    rollout_manger.window_size=5 \
    trainer.val_before_train=True \
    trainer.val_generations_to_log_to_wandb=5 \
    # grpo sampling param
    rollout_manger.n_trajectory=8 \ 
    2>&1 | tee debug_qwen0_5_4_gpu_grpo.log

# NOTE change gpu_memory_utilization to a smaller value (0.4) to avoid oom error
 No newline at end of file
+2 −0
Original line number Diff line number Diff line
@@ -59,4 +59,6 @@ python3 -m vagen.trainer.main_ppo \
    rollout_manger.window_size=5 \
    trainer.val_before_train=True \
    trainer.val_generations_to_log_to_wandb=5 \
    # grpo sampling param
    rollout_manger.n_trajectory=8 \ 
    2>&1 | tee debug_qwen2_5_vl_4gpu_grpo.log
 No newline at end of file
+1 −0
Original line number Diff line number Diff line
@@ -193,3 +193,4 @@ rollout_manger:
  window_size: 5
  max_turns: 5
  n_gpus_per_node: ${trainer.n_gpus_per_node}
  n_trajectory: 1 # >1 for grpo, please set actor_rollout_ref.rollout.n = 1 since we don't use vllm n sample parameter, set it >1 here for grpo agentic setting
 No newline at end of file
+27 −89
Original line number Diff line number Diff line
@@ -142,6 +142,14 @@ def compute_advantage(data: DataProto, adv_estimator, gamma=1.0, lam=1.0, num_re
        attention_mask = data.batch['attention_mask']
        response_mask = attention_mask[:, -response_length:]
        token_level_rewards = data.batch['token_level_rewards']
        if "loss_mask" in data.batch.keys():
            loss_mask = data.batch['loss_mask']
            advantages, returns =core_algos.compute_gae_advantage_return_with_loss_mask(token_level_rewards=token_level_rewards,
                                                                    values=values,
                                                                    eos_mask=loss_mask,
                                                                    gamma=gamma,
                                                                    lam=lam)
        else:
            advantages, returns = core_algos.compute_gae_advantage_return(token_level_rewards=token_level_rewards,
                                                                        values=values,
                                                                        eos_mask=response_mask,
@@ -156,6 +164,13 @@ def compute_advantage(data: DataProto, adv_estimator, gamma=1.0, lam=1.0, num_re
        response_length = responses.size(-1)
        attention_mask = data.batch['attention_mask']
        response_mask = attention_mask[:, -response_length:]
        if "loss_mask" in data.batch.keys():
            loss_mask = data.batch['loss_mask']
            # seems here only need to replace eos_mask with loss_mask
            advantages, returns = core_algos.compute_grpo_outcome_advantage(token_level_rewards=token_level_rewards,
                                                                        eos_mask=loss_mask,
                                                                        index=index)
        else:
            advantages, returns = core_algos.compute_grpo_outcome_advantage(token_level_rewards=token_level_rewards,
                                                                            eos_mask=response_mask,
                                                                            index=index)
@@ -612,91 +627,6 @@ class RayPPOTrainer(object):
        wandb.log({"val/generations": new_table}, step=self.global_steps)
        self.validation_table = new_table
    
    
    # Original _validate
    # def _validate(self):
    #     reward_tensor_lst = []
    #     data_source_lst = []

    #     # Lists to collect samples for the table
    #     sample_inputs = []
    #     sample_outputs = []
    #     sample_scores = []

    #     for test_data in self.val_dataloader:
    #         test_batch = DataProto.from_single_dict(test_data)

    #         # we only do validation on rule-based rm
    #         if self.config.reward_model.enable and test_batch[0].non_tensor_batch['reward_model']['style'] == 'model':
    #             return {}

    #         # Store original inputs
    #         input_ids = test_batch.batch['input_ids']
    #         input_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in input_ids]
    #         sample_inputs.extend(input_texts)

    #         if 'multi_modal_inputs' in test_batch.non_tensor_batch.keys():
    #             test_gen_batch = test_batch.pop(
    #                 batch_keys=['input_ids', 'attention_mask', 'position_ids'],
    #                 non_tensor_batch_keys=['raw_prompt_ids', 'multi_modal_data', 'multi_modal_inputs'],
    #             )
    #         else:
    #             test_gen_batch = test_batch.pop(
    #                 batch_keys=['input_ids', 'attention_mask', 'position_ids'],
    #                 non_tensor_batch_keys=['raw_prompt_ids'],
    #             )

    #         test_gen_batch.meta_info = {
    #             'eos_token_id': self.tokenizer.eos_token_id,
    #             'pad_token_id': self.tokenizer.pad_token_id,
    #             'recompute_log_prob': False,
    #             'do_sample': False,
    #             'validate': True,
    #         }

    #         # pad to be divisible by dp_size
    #         test_gen_batch_padded, pad_size = pad_dataproto_to_divisor(test_gen_batch, self.actor_rollout_wg.world_size)
    #         test_output_gen_batch_padded = self.actor_rollout_wg.generate_sequences(test_gen_batch_padded)
    #         # unpad
    #         test_output_gen_batch = unpad_dataproto(test_output_gen_batch_padded, pad_size=pad_size)
    #         print('validation generation end')

    #         # Store generated outputs
    #         output_ids = test_output_gen_batch.batch['responses']
    #         output_texts = [self.tokenizer.decode(ids, skip_special_tokens=True) for ids in output_ids]
    #         sample_outputs.extend(output_texts)

    #         test_batch = test_batch.union(test_output_gen_batch)

    #         # evaluate using reward_function
    #         reward_tensor = self.val_reward_fn(test_batch)

    #         # Store scores
    #         scores = reward_tensor.sum(-1).cpu().tolist()
    #         sample_scores.extend(scores)

    #         reward_tensor_lst.append(reward_tensor)
    #         data_source_lst.append(test_batch.non_tensor_batch.get('data_source', ['unknown'] * reward_tensor.shape[0]))

    #     self._maybe_log_val_generations_to_wandb(inputs=sample_inputs, outputs=sample_outputs, scores=sample_scores)

    #     reward_tensor = torch.cat(reward_tensor_lst, dim=0).sum(-1).cpu()  # (batch_size,)
    #     data_sources = np.concatenate(data_source_lst, axis=0)

    #     # evaluate test_score based on data source
    #     data_source_reward = {}
    #     for i in range(reward_tensor.shape[0]):
    #         data_source = data_sources[i]
    #         if data_source not in data_source_reward:
    #             data_source_reward[data_source] = []
    #         data_source_reward[data_source].append(reward_tensor[i].item())

    #     metric_dict = {}
    #     for data_source, rewards in data_source_reward.items():
    #         metric_dict[f'val/test_score/{data_source}'] = np.mean(rewards)

    #     return metric_dict
    
    # Agentic Setting _validate
    def _validate(self):
        print(f"[DEBUG] validation at global step {self.global_steps} begins")
@@ -1034,6 +964,12 @@ class RayPPOTrainer(object):

                #             del gen_baseline_batch, gen_baseline_output

                
                # We control grpo sampling param here
                batch.non_tensor_batch['uid'] = np.array([str(uuid.uuid4()) for _ in range(len(batch.batch))],dtype=object)
                batch = batch.repeat(repeat_times=self.config.rollout_manger.n_trajectory, interleave=True)
                
                    
                env_configs = [
                    EnvConfig(env_name=batch.non_tensor_batch['extra_info'][i]['env_name'],
                              env_config=batch.non_tensor_batch['extra_info'][i]['env_config'],
@@ -1050,10 +986,12 @@ class RayPPOTrainer(object):
                        final_gen_batch_output = rollout_manager.get_final_trajectory()

                    print(f"[DEBUG] step {self.global_steps} rollout ends")
                    batch.non_tensor_batch['uid'] = np.array([str(uuid.uuid4()) for _ in range(len(batch.batch))],
                                                             dtype=object)
                    # repeat to align with repeated responses in rollout
                    batch = batch.repeat(repeat_times=self.config.actor_rollout_ref.rollout.n, interleave=True)
                    
                    # This is moved to before rollout because we don't use vllm n sample param for grpo due to multi-turn nature of our method
                    # batch.non_tensor_batch['uid'] = np.array([str(uuid.uuid4()) for _ in range(len(batch.batch))],
                    #                                          dtype=object)
                    # # repeat to align with repeated responses in rollout
                    # batch = batch.repeat(repeat_times=self.config.actor_rollout_ref.rollout.n, interleave=True)
                    batch = batch.union(final_gen_batch_output)

                    # balance the number of valid tokens on each dp rank.