Commit eed12d53 authored by jameskrw's avatar jameskrw
Browse files

minor bug fixing added validation config

parent 4d7c8b6e
Loading
Loading
Loading
Loading
+3 −3
Original line number Diff line number Diff line
@@ -10,7 +10,7 @@ python3 -m vagen.trainer.main_ppo \
    data.val_files=data/sokoban/test.parquet \
    data.train_batch_size=16 \
    data.max_prompt_length=2048 \
    data.max_response_length=512 \
    data.max_response_length=256 \
    data.max_trajectory_length=3072 \
    data.image_key=images \
    actor_rollout_ref.model.path=Qwen/Qwen2.5-VL-3B-Instruct \
@@ -41,8 +41,8 @@ python3 -m vagen.trainer.main_ppo \
    trainer.experiment_name='qwen2_5_vl_3b_function_rm' \
    trainer.n_gpus_per_node=4 \
    trainer.nnodes=1 \
    trainer.save_freq=-1 \
    trainer.test_freq=-1 \
    trainer.save_freq=50 \
    trainer.test_freq=2 \
    trainer.total_epochs=15 \
    +max_turns=2 \
    2>&1 | tee debug.log
+2 −2
Original line number Diff line number Diff line
@@ -20,7 +20,7 @@ from vagen.env.base import EnvConfig,IMAGE_PLACEHOLDER
@dataclass
class QwenVLRolloutConifg:
    window_size: int = 5
    max_trajectories_length: int = 3072
    max_trajectory_length: int = 3072
    max_turns: int = 5
    n_gpu_per_node: int = 1 # used for multigpu batch balancing
    sptk_for_loss_mask: List[str] = field(default_factory=lambda: ['<|box_start|>', '<|box_end|>'])
@@ -420,7 +420,7 @@ class QwenVLRolloutManger():
        
        input_ids_response, attention_mask_response = verl_F.tokenize_and_postprocess_data(prompt=response_with_chat_template,
                                                                         tokenizer=self.tokenizer,
                                                                         max_length=self.config.max_trajectories_length,
                                                                         max_length=self.config.max_trajectory_length,
                                                                         pad_token_id=self.tokenizer.pad_token_id,
                                                                         left_pad=False,
                                                                         truncation=self.truncation)
+2 −0
Original line number Diff line number Diff line
@@ -184,3 +184,5 @@ trainer:
  remove_previous_ckpt_in_save: False
  del_local_ckpt_after_load: False
  default_local_dir: checkpoints/${trainer.project_name}/${trainer.experiment_name}
  val_before_train: True
  val_only: False
 No newline at end of file
+8 −8
Original line number Diff line number Diff line
@@ -707,7 +707,7 @@ class RayPPOTrainer(object):

        if self.test_rollout_config==None:
            self.test_rollout_config = QwenVLRolloutConifg(
                max_trajectories_length=self.config.data.max_trajectories_length,
                max_trajectory_length=self.config.data.max_trajectory_length,
                max_turns=self.config.max_turns,
                n_gpu_per_node=self.config.trainer.n_gpus_per_node,
            )
@@ -956,19 +956,19 @@ class RayPPOTrainer(object):
        # perform validation before training
        # currently, we only support validation using the reward_function.
        # TODO implement validation
        # if self.val_reward_fn is not None and self.config.trainer.get('val_before_train', True):
        #     val_metrics = self._validate()
        #     pprint(f'Initial validation metrics: {val_metrics}')
        #     logger.log(data=val_metrics, step=self.global_steps)
        #     if self.config.trainer.get('val_only', False):
        #         return
        if self.val_reward_fn is not None and self.config.trainer.get('val_before_train', True):
            val_metrics = self._validate()
            pprint(f'Initial validation metrics: {val_metrics}')
            logger.log(data=val_metrics, step=self.global_steps)
            if self.config.trainer.get('val_only', False):
                return

        # we start from step 1
        self.global_steps += 1


        rollout_config = QwenVLRolloutConifg(
            max_trajectories_length=self.config.data.max_trajectories_length,
            max_trajectory_length=self.config.data.max_trajectory_length,
            max_turns=self.config.max_turns,
            n_gpu_per_node=self.config.trainer.n_gpus_per_node,
        )