Commit f435f565 authored by williamzhangNU's avatar williamzhangNU
Browse files

modify rollout

parent be81cb27
Loading
Loading
Loading
Loading
+8 −4
Original line number Diff line number Diff line
@@ -46,6 +46,9 @@ Rewards:
Move: -0.1
Box on target: +1.0
All boxes placed: +10.0

Include your thought in <think> </think> tags and your final answer in <answer> </answer> tags.
Your response should be like: <think> [Your thought] </think> <answer> [Your answer] </answer>
"""


@@ -359,7 +362,7 @@ class SokobanInterface(BaseInterface):
        """

        assert not self.env.finished(), "Environment finished before step"
        reward, done, info = 0, False, {}
        reward, done, final_info = 0, False, {}


        preprocess_result = self._preprocess(raw_text)
@@ -367,13 +370,13 @@ class SokobanInterface(BaseInterface):
        action_list = preprocess_result.action_list
        valid_list = preprocess_result.valid_list
        answer = preprocess_result.answer
        info['llm_raw_response'] = preprocess_result.llm_raw_response
        final_info['llm_raw_response'] = preprocess_result.llm_raw_response

        # deal with format
        if think and answer: # format is correct
            reward += self.FORMAT_REWARD


        info = {}
        for action, valid in zip(action_list, valid_list):
            if done or self.env.finished():
                break
@@ -383,6 +386,7 @@ class SokobanInterface(BaseInterface):
            else: # termiante at the first invalid action
                break
        self.traj_reward += reward
        final_info.update(info) # NOTE currently only use the last step info

        env_state = self.env._render(mode='text' if not self.visual_env else 'rgb_array') # NOTE currently called after step

@@ -390,7 +394,7 @@ class SokobanInterface(BaseInterface):
            env_state=env_state,
            reward=reward,
            done=done,
            info=info,
            info=final_info,
            preprocess_result=preprocess_result,
        )
    
+2 −2
Original line number Diff line number Diff line
@@ -8,7 +8,7 @@ python3 -m vagen.trainer.main_ppo \
    algorithm.adv_estimator=grpo \
    data.train_files=data/sokoban/train.parquet \
    data.val_files=data/sokoban/test.parquet \
    data.train_batch_size=16 \
    data.train_batch_size=2 \
    data.max_prompt_length=512 \
    data.max_response_length=1536 \
    data.max_trajectory_length=2048 \
@@ -19,7 +19,7 @@ python3 -m vagen.trainer.main_ppo \
    actor_rollout_ref.model.path=Qwen/Qwen2.5-VL-3B-Instruct \
    actor_rollout_ref.actor.optim.lr=1e-6 \
    actor_rollout_ref.model.use_remove_padding=True \
    actor_rollout_ref.actor.ppo_mini_batch_size=4 \
    actor_rollout_ref.actor.ppo_mini_batch_size=2 \
    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 \
+44 −10
Original line number Diff line number Diff line
@@ -294,7 +294,19 @@ class QwenVLRolloutManger():
                            recording: List[Dict], 
                            step: int, 
                            window_size: int = None,
                            last_question: bool = False,):
                            is_final: bool = False,
        ):
        """
        Given a recording, generate the prompt for MLLM
        Chat: Sys -> |InitUser| -> |Assistant, User| -> |Assistant, User| -> ... -> |Assistant, User Final|

        Args:
            recording: List of dictionaries containing recorded environment interactions
            step: Current step to generate prompt for
            window_size: Number of past steps to include in the context
            is_final: Whether the prompt is for the final step 
                - if True, the last one should be from assistant
        """
        
        assert step >= 0
        start_step = max(0, step - window_size) if window_size is not None else 0
@@ -307,21 +319,38 @@ class QwenVLRolloutManger():
        env_id = history[0]['env_id']
        chat.append({"role": "system", "content": self.envs[env_id].get_task_instruction()})

        # for i, record in enumerate(history):
        #     if i>0:
        #         llm_raw_response = record['info']['llm_raw_response']
        #         filtered_llm_raw_response = self._handle_special_tokens(llm_raw_response, compute_loss_mask=False)
        #         chat.append({"role": "assistant", "content": filtered_llm_raw_response})
        #     if i<len(history)-1 or last_question:
        #         chat.append({"role": "user", "content": record['text_template']})

        # image_data=[]
        # for record in history:
        #     if 'image_data' in record:
        #         for img in record['image_data']:
        #             image_data.append(img)

        image_data=[]
        for i, record in enumerate(history):
            if i>0:
            if i == 0:
                chat.append({"role": "user", "content": record['text_template']})
                if 'image_data' in record:
                    for img in record['image_data']:
                        image_data.append(img)
            else:
                llm_raw_response = record['info']['llm_raw_response']
                filtered_llm_raw_response = self._handle_special_tokens(llm_raw_response, compute_loss_mask=False)
                chat.append({"role": "assistant", "content": filtered_llm_raw_response})
            if i<len(history)-1 or last_question:
                if not is_final:
                    chat.append({"role": "user", "content": record['text_template']})

        image_data=[]
        for record in history:
                    if 'image_data' in record:
                        for img in record['image_data']:
                            image_data.append(img)
            
        prompt_with_chat_template = self.tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=False)
        prompt_with_chat_template = self.tokenizer.apply_chat_template(chat, add_generation_prompt=not is_final, tokenize=False)
        return {
            "prompt": prompt_with_chat_template,
            "image_data": image_data,
@@ -350,7 +379,7 @@ class QwenVLRolloutManger():
                - position_ids for prompts: rope
                - rest postion_ids: refer to vllm_rollout_spmd.py to check how to compute
        """
        rst=self._single_recording_to_prompt(recording, step, window_size,last_question=True)
        rst=self._single_recording_to_prompt(recording, step, window_size, is_final=False)
        prompt_with_chat_template=rst['prompt']
        image_data=rst['image_data']        
        has_images = len(image_data) > 0        
@@ -408,7 +437,7 @@ class QwenVLRolloutManger():
        prompt_with_chat_template=self.tokenizer.pad_token 
        
        # handle response
        response_rst=self._single_recording_to_prompt(recording, step, window_size, last_question=False)
        response_rst=self._single_recording_to_prompt(recording, step, window_size, is_final=True)
        response_with_chat_template=response_rst['prompt']
        image_data=response_rst['image_data']
       
@@ -486,7 +515,8 @@ class QwenVLRolloutManger():
            window_size: Number of past steps to include in the context
        
        Returns:
            Dictionary containing properly formatted inputs for the MLLM›
            Dictionary containing properly formatted inputs for the MLLM
            - None if no data is available (all environments are done)
        """
        batch = []
        self.batch_idx_to_env_id = {}
@@ -498,6 +528,8 @@ class QwenVLRolloutManger():
            batch.append(self._generate_input_item(self.recorder[env_id], step, window_size))
            self.batch_idx_to_env_id[batch_idx] = env_id
            batch_idx += 1
        if not batch:
            return None
        if len(batch) % self.config.n_gpu_per_node != 0:
            # Pad the batch to make it divisible by n_gpu_per_node
            while len(batch) % self.config.n_gpu_per_node != 0:
@@ -515,6 +547,8 @@ class QwenVLRolloutManger():
        """
        for step in range(self.config.max_turns):
            input_batch_dict = self.gen_batch(step, self.config.window_size)
            if input_batch_dict is None:
                break
            input_batch = DataProto.from_single_dict(input_batch_dict)
            if 'multi_modal_data' in input_batch.non_tensor_batch.keys():
                gen_batch = input_batch.pop(
+10 −10
Original line number Diff line number Diff line
@@ -708,6 +708,7 @@ class RayPPOTrainer(object):
        if self.test_rollout_config==None:
            self.test_rollout_config = QwenVLRolloutConifg(
                max_trajectory_length=self.config.data.max_trajectory_length,
                max_response_per_turn=self.config.data.max_response_per_turn,
                max_turns=self.config.max_turns,
                n_gpu_per_node=self.config.trainer.n_gpus_per_node,
            )
@@ -743,11 +744,11 @@ class RayPPOTrainer(object):
                    for i in range(len(batch))
                ]
            
            self.rollout_manager.reset(env_configs)
            self.test_rollout_manager.reset(env_configs)
            print('validation generation start')
            self.rollout_manager.rollout_loop()
            self.test_rollout_manager.rollout_loop()
            print('validation generation end')
            inputs, outputs, scores = self.rollout_manager.recording_to_log() # data source == inputs in our current setting, outputs=whole trjecotry
            inputs, outputs, scores = self.test_rollout_manager.recording_to_log() # data source == inputs in our current setting, outputs=whole trjecotry
            sample_inputs.extend(inputs)
            sample_outputs.extend(outputs)
            sample_scores.extend(scores)
@@ -955,13 +956,12 @@ 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