Commit bd8d9d99 authored by jameskrw's avatar jameskrw
Browse files

sokoban vision tested

parent bb7e2d30
Loading
Loading
Loading
Loading
+10 −3
Original line number Diff line number Diff line
@@ -15,12 +15,19 @@ bash scripts/install.sh
```
bash vagen/examples/sokoban/debug_qwen0_5_1_gpu_grpo.sh
bash vagen/examples/sokoban/debug_qwen0_5_4_gpu_ppo.sh
bash vagen/examples/sokoban/debug_qwen2_5_vl_4gpu_grpo.sh

# Verified on 1 and 4 A100 GPUs
```
## Current Status
1. sokoban-text is runnbale: both single A100 and 4 A100s, grpo (fake), performance is bad
2. sokoban-vision is testing in 4 A100s, minor bugs to fix


## TODO
1. Implement real grpo: rollout.n>1
2. Make PPO runnable (modify ppo scripts and code)
3. Add more metrics and image visualization
1. DEBUG: 之前看validation好像有连续2个llm response中间没有user的情况,不知到现在的code fix 没有
1. Make sure loss mask is working correctly
2. Implement real grpo: rollout.n>1
3. Make PPO runnable (modify ppo scripts and code)
4. Add more metrics and image visualization
+51 −48
Original line number Diff line number Diff line
# set -x
set -x

# export VLLM_ATTENTION_BACKEND=XFORMERS
export VLLM_ATTENTION_BACKEND=XFORMERS

# python -m vagen.env.sokoban.create_dataset --visual_env --data_dir data/sokoban
python -m vagen.env.sokoban.create_dataset --visual_env --data_dir data/sokoban

# 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=2 \
#     data.max_prompt_length=512 \
#     data.max_response_length=1536 \
#     data.max_trajectory_length=2048 \
#     +data.max_response_per_turn=256 \
#     +actor_rollout_ref.rollout.max_response_per_turn=256 \
#     +actor_rollout_ref.rollout.max_trajectory_length=2048 \
#     data.image_key=images \
#     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=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 \
#     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=1 \
#     actor_rollout_ref.rollout.tensor_model_parallel_size=2 \
#     actor_rollout_ref.rollout.name=vllm \
#     actor_rollout_ref.rollout.gpu_memory_utilization=0.6 \
#     actor_rollout_ref.rollout.enable_chunked_prefill=False \
#     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=1 \
#     actor_rollout_ref.ref.fsdp_config.param_offload=True \
#     algorithm.kl_ctrl.kl_coef=0.001 \
#     trainer.critic_warmup=0 \
#     trainer.logger=['console','wandb'] \
#     trainer.project_name='vagen' \
#     trainer.experiment_name='qwen2_5_vl_3b_function_rm' \
#     trainer.n_gpus_per_node=2 \
#     trainer.nnodes=1 \
#     trainer.save_freq=50 \
#     trainer.test_freq=2 \
#     trainer.total_epochs=15 \
#     +max_turns=2 \
#     2>&1 | tee debug_qwen2_5_vl_4gpu_grpo.log
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=32 \
    data.max_prompt_length=1536 \
    data.max_response_length=128 \
    data.max_trajectory_length=3072 \
    data.image_key=images \
    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=32 \
    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=1 \
    actor_rollout_ref.rollout.tensor_model_parallel_size=2 \
    actor_rollout_ref.rollout.name=vllm \
    actor_rollout_ref.rollout.gpu_memory_utilization=0.6 \
    actor_rollout_ref.rollout.enable_chunked_prefill=False \
    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=1 \
    actor_rollout_ref.ref.fsdp_config.param_offload=True \
    algorithm.kl_ctrl.kl_coef=0.001 \
    trainer.critic_warmup=0 \
    trainer.logger=['console','wandb'] \
    trainer.project_name='vagen' \
    trainer.experiment_name='qwen2_5_vl_3b_function_rm' \
    trainer.n_gpus_per_node=4 \
    trainer.nnodes=1 \
    trainer.save_freq=50 \
    trainer.test_freq=2 \
    trainer.total_epochs=15 \
    rollout_manger.max_turns=2 \
    rollout_manger.window_size=5 \
    trainer.val_before_train=True \
    trainer.val_generations_to_log_to_wandb=5 \
    2>&1 | tee debug_qwen2_5_vl_4gpu_grpo.log


# Tested and get oom error
 No newline at end of file
+17 −30
Original line number Diff line number Diff line
@@ -45,13 +45,13 @@ class QwenVLRolloutManger():
        self.env_states = None # dict
        self.batch_idx_to_env_id = None # dict

    def _handle_special_tokens(self, llm_raw_response: str, compute_loss_mask: bool) -> str:
    def _handle_special_tokens(self, llm_raw_response: str, prep_for_loss_mask: bool) -> str:
        """
        1. Filter out special tokens: <image> and special tokens marking environment observation in the llm generated response
        2. Add special tokens to the beginning and end of the response if compute_loss_mask is True
        2. prep_for_loss_mask: if true, add special tokens to the beginning and end of the response if compute_loss_mask is True
        """
        llm_raw_response = re.sub(r'<image>', '', llm_raw_response)
        if compute_loss_mask:
        if prep_for_loss_mask:
            # filtering special tokens for llm_raw_response, then adding them to the beginning and end of the response for loss mask computation
            sptk_b = self.config.sptk_for_loss_mask[0]
            sptk_e = self.config.sptk_for_loss_mask[1]
@@ -81,6 +81,7 @@ class QwenVLRolloutManger():
            image_inputs = self.processor.image_processor(image_data, return_tensors='pt')
            image_grid_thw = image_inputs['image_grid_thw']
            row_dict['multi_modal_inputs'] = {key: val for key, val in image_inputs.items()}
            # print(f"[DEBUG] number of image_data in rollout: {len(image_data)}")
        if image_grid_thw is not None:
            merge_length = self.processor.image_processor.merge_size**2
            index = 0
@@ -95,7 +96,9 @@ class QwenVLRolloutManger():

            prompt_template = prompt_template.replace('<|placeholder|>',
                                                        self.processor.image_token)
        
            # print(f"[DEBUG] number of image_data in final trajectory: {len(image_data)}")
            # number_of_image_tokens=prompt_template.count(self.processor.image_token)
            # print(f"[DEBUG] number_of_image_tokens: {number_of_image_tokens}")
        return prompt_template, row_dict, image_grid_thw, raw_prompt
    
    def _compute_loss_mask(self, input_ids, attention_mask):
@@ -294,6 +297,7 @@ class QwenVLRolloutManger():
                            step: int, 
                            window_size: int = None,
                            is_final: bool = False,
                            prep_for_loss_mask: bool = False,
        ):
        """
        Given a recording, generate the prompt for MLLM
@@ -305,6 +309,7 @@ class QwenVLRolloutManger():
            window_size: Number of past steps to include in the context
            is_final: Whether the prompt is for the final step 
                - if True, the end of the chat is from the last assistant's response
            prep_for_loss_mask: whether to use special token to wrap llm response
        """
        
        assert step >= 0
@@ -318,38 +323,19 @@ 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:
                chat.append({"role": "user", "content": record['text_template']})
                if 'image_data' in record:
                    for img in record['image_data']:
                        image_data.append(img)
            else:
            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)
                filtered_llm_raw_response = self._handle_special_tokens(llm_raw_response, prep_for_loss_mask=prep_for_loss_mask)
                chat.append({"role": "assistant", "content": filtered_llm_raw_response})
                if not is_final:
            if i<len(history)-1 or not is_final:
                chat.append({"role": "user", "content": record['text_template']})
                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=not is_final, tokenize=False)
        prompt_with_chat_template = self.tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=False)
        return {
            "prompt": prompt_with_chat_template,
            "image_data": image_data,
@@ -378,7 +364,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, is_final=False)
        rst=self._single_recording_to_prompt(recording, step, window_size, is_final=False, prep_for_loss_mask=False)
        prompt_with_chat_template=rst['prompt']
        image_data=rst['image_data']        
        has_images = len(image_data) > 0        
@@ -436,7 +422,7 @@ class QwenVLRolloutManger():
        prompt_with_chat_template=self.tokenizer.pad_token 
        
        # handle response
        response_rst=self._single_recording_to_prompt(recording, step, window_size, is_final=True)
        response_rst=self._single_recording_to_prompt(recording, step, window_size, is_final=True, prep_for_loss_mask=True)
        response_with_chat_template=response_rst['prompt']
        image_data=response_rst['image_data']
       
@@ -532,7 +518,8 @@ class QwenVLRolloutManger():
        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:
                batch.append(batch[-1])
                # do we need to use copy or not here?
                batch.append(batch[-1].copy())
        return collate_fn(batch)