Commit 4f662578 authored by williamzhangNU's avatar williamzhangNU
Browse files

minor update

parent 40c2cd19
Loading
Loading
Loading
Loading
+21 −19
Original line number Diff line number Diff line
@@ -188,7 +188,6 @@ class SokobanEnv(BaseEnv, GymSokobanEnv):
class PreprocessResult:
    action_list: List[int]
    answer_list: List[str] # string of extracted answer (may be invalid action)
    valid_list: List[bool]
    think: str
    answer: str
    llm_raw_response: str
@@ -197,7 +196,6 @@ class PreprocessResult:
        return {
            'action_list': self.action_list,
            'answer_list': self.answer_list,
            'valid_list': self.valid_list,
            'think': self.think,
            'answer': self.answer,
            'llm_raw_response': self.llm_raw_response,
@@ -208,7 +206,9 @@ class PreprocessResult:
class SokobanInterface(BaseInterface):

    INVALID_ACTION = 0
    FORMAT_REWARD = 1
    FORMAT_REWARD = 0.1
    FORMAT_PENALTY = -0.1
    VALID_ACTION_REWARD = 0.2
    ACTION_LOOKUP = {
        0: "None",
        1: "Up",
@@ -266,32 +266,39 @@ class SokobanInterface(BaseInterface):
    @classmethod
    def _preprocess(cls, text: str) -> PreprocessResult:
        """Preprocess the raw text from LLM into a list of actions.
        Ensure at least one action (may be invalid).
        NOTE Only keep valid actions.

        Args:
            text: raw text from LLM

        Returns:
            PreprocessResult containing parsed actions and validity flags
            PreprocessResult containing parsed valid actions
        """
        first_step_preprocess = preprocess_text(text)
        preprocess_result = PreprocessResult(
            action_list=[],
            valid_list=[],
            answer_list=first_step_preprocess['answer_list'],
            think=first_step_preprocess['think'],
            answer=first_step_preprocess['answer'],
            llm_raw_response=text,
        )
        
        # for answer in preprocess_result.answer_list:
        #     action = cls._extract_one_action(answer)
        #     if action != cls.INVALID_ACTION:
        #         preprocess_result.action_list.append(action)
        #         preprocess_result.valid_list.append(True)
        #     else:
        #         preprocess_result.action_list.append(cls.INVALID_ACTION)
        #         preprocess_result.valid_list.append(False)

        # ensure there are only valid actions
        for answer in preprocess_result.answer_list:
            action = cls._extract_one_action(answer)
            if action != cls.INVALID_ACTION:
                preprocess_result.action_list.append(action)
                preprocess_result.valid_list.append(True)
            else:
                preprocess_result.action_list.append(cls.INVALID_ACTION)
                preprocess_result.valid_list.append(False)
                break
        
        return preprocess_result
        
@@ -323,11 +330,8 @@ class SokobanInterface(BaseInterface):

        answer = preprocess_result.answer
        valid_action = []
        for action, valid in zip(preprocess_result.action_list, preprocess_result.valid_list):
            if valid:
        for action in preprocess_result.action_list:
            valid_action.append(cls.ACTION_LOOKUP[action])
            else:
                break

        observation = IMAGE_PLACEHOLDER if not isinstance(env_state, str) else env_state
        text_template = action_template.format(
@@ -374,25 +378,23 @@ class SokobanInterface(BaseInterface):
        preprocess_result = self._preprocess(raw_text)
        think = preprocess_result.think
        action_list = preprocess_result.action_list
        valid_list = preprocess_result.valid_list
        answer = preprocess_result.answer
        final_info['llm_raw_response'] = preprocess_result.llm_raw_response

        # deal with format
        if think and answer: # format is correct
            reward += self.FORMAT_REWARD
            if action_list:
                reward += self.VALID_ACTION_REWARD
        else:
            reward -= self.FORMAT_REWARD*0.1
            reward += self.FORMAT_PENALTY

        info = {}
        for action, valid in zip(action_list, valid_list):
        for action in action_list:
            if done or self.env.finished():
                break
            if valid:
            _, env_reward, done, info = self.env.step(action)
            reward += env_reward
            else: # termiante at the first invalid action
                break
        self.traj_reward += reward
        final_info.update(info) # NOTE currently only use the last step info

+67 −48
Original line number Diff line number Diff line
# set -x
set -x

# export VLLM_ATTENTION_BACKEND=XFORMERS
export VLLM_ATTENTION_BACKEND=XFORMERS
export PYTHONHASHSEED=0

# python -m vagen.env.sokoban.create_dataset --data_dir data/sokoban-text
python -m vagen.env.sokoban.create_dataset \
    --data_dir data/sokoban-text-1-step \
    --max_action_length 1 \
    --dim_room 6 6 \
    --num_boxes 1 \
    --max_steps 100 \
    --search_depth 30 \
    --start_seed 0 \
    --train_ratio 0.8 \
    --n_candidate 20000

# # max_trajectory_length = max_prompt_length + max_response_length
# max_trajectory_length = max_prompt_length + max_response_length

# python3 -m vagen.trainer.main_ppo \
#     algorithm.adv_estimator=ppo \
#     data.train_files=data/sokoban-text/train.parquet \
#     data.val_files=data/sokoban-text/test.parquet \
#     data.train_batch_size=32 \
#     data.max_prompt_length=2048 \
#     data.max_response_length=128 \
#     data.max_trajectory_length=3072 \
#     data.image_key=images \
#     actor_rollout_ref.model.path=Qwen/Qwen2.5-0.5B-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=2 \
#     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=2 \
#     actor_rollout_ref.rollout.tensor_model_parallel_size=1 \
#     actor_rollout_ref.rollout.name=vllm \
#     actor_rollout_ref.rollout.gpu_memory_utilization=0.7 \
#     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=2 \
#     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_05b_function_rm' \
#     trainer.n_gpus_per_node=1 \
#     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 \
#     2>&1 | tee debug_qwen0_5_1_gpu_ppo.log
python3 -m vagen.trainer.main_ppo \
    algorithm.adv_estimator=gae \
    data.train_files=data/sokoban-text-1-step/train.parquet \
    data.val_files=data/sokoban-text-1-step/test.parquet \
    data.train_batch_size=64 \
    data.max_prompt_length=768 \
    data.max_response_length=128 \
    data.max_trajectory_length=1024 \
    data.image_key=images \
    actor_rollout_ref.model.path=Qwen/Qwen2.5-0.5B-Instruct \
    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=2 \
    actor_rollout_ref.actor.use_kl_loss=False \
    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=2 \
    actor_rollout_ref.rollout.tensor_model_parallel_size=1 \
    actor_rollout_ref.rollout.name=vllm \
    actor_rollout_ref.rollout.gpu_memory_utilization=0.4 \
    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=2 \
    actor_rollout_ref.ref.fsdp_config.param_offload=True \
    critic.optim.lr=1e-5 \
    critic.model.use_remove_padding=False \
    critic.model.path=Qwen/Qwen2.5-0.5B-Instruct \
    critic.model.enable_gradient_checkpointing=True \
    critic.ppo_micro_batch_size_per_gpu=2 \
    critic.model.fsdp_config.param_offload=False \
    critic.model.fsdp_config.optimizer_offload=False \
    algorithm.kl_ctrl.kl_coef=0.001 \
    trainer.critic_warmup=0 \
    trainer.logger=['console','wandb'] \
    trainer.project_name='vagen' \
    trainer.experiment_name='debug_qwen0_5_1_gpu_ppo' \
    trainer.n_gpus_per_node=1 \
    trainer.nnodes=1 \
    trainer.save_freq=100 \
    trainer.test_freq=5 \
    trainer.total_epochs=15 \
    rollout_manger.max_turns=1 \
    rollout_manger.window_size=5 \
    trainer.val_before_train=True \
    trainer.val_generations_to_log_to_wandb=4 \
    rollout_manger.n_trajectory=2 \
    2>&1 | tee debug_qwen0_5_1_gpu_ppo.log
+20 −16
Original line number Diff line number Diff line
@@ -23,7 +23,8 @@ class QwenVLRolloutConifg:
    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|>'])
    sptk_for_loss_mask: List[str] = field(init=False, default_factory=lambda: ['<|box_start|>', '<|box_end|>'])
    end_turn_token: str = field(init=False, default='<|im_end|>') # specially for Qwen2.5
    
class QwenVLRolloutManger():
    def __init__(self,
@@ -45,6 +46,7 @@ class QwenVLRolloutManger():
        self.env_states = None # dict
        self.batch_idx_to_env_id = None # dict

    @torch.no_grad()
    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
@@ -60,6 +62,7 @@ class QwenVLRolloutManger():
            llm_raw_response = sptk_b + llm_raw_response + sptk_e
        return llm_raw_response
    
    @torch.no_grad()
    def _handle_multi_modal_data(
            self, 
            prompt_template: str, 
@@ -101,6 +104,7 @@ class QwenVLRolloutManger():
            # print(f"[DEBUG] number_of_image_tokens: {number_of_image_tokens}")
        return prompt_template, row_dict, image_grid_thw, raw_prompt
    
    @torch.no_grad()
    def _compute_loss_mask(self, input_ids, attention_mask):
        """
        Compute loss mask for the input ids and attention mask
@@ -130,8 +134,8 @@ class QwenVLRolloutManger():
        """
        
        # Get token IDs for special tokens and pad token
        sptk_b = self.tokenizer.convert_tokens_to_ids('<|box_start|>')
        sptk_e = self.tokenizer.convert_tokens_to_ids('<|box_end|>')
        sptk_b = self.tokenizer.convert_tokens_to_ids(self.config.sptk_for_loss_mask[0])
        sptk_e = self.tokenizer.convert_tokens_to_ids(self.config.sptk_for_loss_mask[1])
        pad_token_id = self.tokenizer.pad_token_id

        batch_size = input_ids.shape[0]
@@ -184,7 +188,7 @@ class QwenVLRolloutManger():
        
        return new_input_ids, new_attention_mask, new_loss_mask, new_token_level_reward_mask
    
        
    @torch.no_grad()
    def reset(self, env_configs: List[EnvConfig]):
        """
        Reset environments based on provided configurations, reusing environments when possible.
@@ -274,7 +278,7 @@ class QwenVLRolloutManger():
        
        return initial_obs, initial_info
    
    
    @torch.no_grad()
    def record(self, env_id, obs, reward, done, info):
        """
        Record each step's obs, info, done, reward,
@@ -296,7 +300,7 @@ class QwenVLRolloutManger():
            record_entry['image_data'] = [process_image(image) for image in obs['multi_modal_data'][IMAGE_PLACEHOLDER]]
        self.recorder[env_id].append(record_entry)


    @torch.no_grad()
    def _single_recording_to_prompt(self,
                            recording: List[Dict], 
                            step: int, 
@@ -345,12 +349,17 @@ class QwenVLRolloutManger():
                        image_data.append(img)
            
        prompt_with_chat_template = self.tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=False)
        # switch box_end and im_end so that the model can learn to generate <|im_end|>
        prompt_with_chat_template = prompt_with_chat_template.replace(
            f'{self.config.sptk_for_loss_mask[1]}{self.config.end_turn_token}',
            f'{self.config.end_turn_token}{self.config.sptk_for_loss_mask[1]}')
        return {
            "prompt": prompt_with_chat_template,
            "image_data": image_data,
            "rewards": rewards,
        }
    
    @torch.no_grad()
    def _generate_input_item(
            self, 
            recording: List[Dict], 
@@ -400,7 +409,7 @@ class QwenVLRolloutManger():
        return row_dict



    @torch.no_grad()
    def _generate_input_final_item(
            self, 
            recording: List[Dict], 
@@ -514,7 +523,7 @@ class QwenVLRolloutManger():
        row_dict["multi_turn_token_level_reward"] = multi_turn_token_level_reward # (seq_len) later need to convert to (response_len)
        return row_dict


    @torch.no_grad()
    def gen_batch(self, step, window_size):
        """
        Generate a batch of data for the current step
@@ -546,8 +555,7 @@ class QwenVLRolloutManger():
                batch.append(batch[-1].copy())
        return collate_fn(batch)
    
    
    
    @torch.no_grad()
    def rollout_loop(self):
        """
        Step the environment and record the results
@@ -581,9 +589,6 @@ class QwenVLRolloutManger():
                    raw_prompt_ids_array[i] = raw_prompt_ids[i]
                else:
                    raw_prompt_ids_array[i] = raw_prompt_ids[i].tolist()
                # if not i % 4:
                #     print(f"[DEBUG] raw_prompt_ids_array({i}) length: {len(raw_prompt_ids_array[i])}")
                #     print(f"[DEBUG] raw_prompt_ids_array({i}) content: {self.tokenizer.decode(raw_prompt_ids_array[i])}")
            gen_batch.non_tensor_batch['raw_prompt_ids'] = raw_prompt_ids_array
            
            output_batch = self.actor_rollout_wg.generate_sequences(gen_batch)
@@ -608,8 +613,7 @@ class QwenVLRolloutManger():
                self.env_states[env_id]['done'] = done
                self.record(env_id, obs, reward, done, info)
        
        
        
    @torch.no_grad()
    def get_final_trajectory(self) -> DataProto:
        """
        Get the final trajectory of all environments
@@ -638,7 +642,7 @@ class QwenVLRolloutManger():
        # print(f"[DEBUG] --------------------------------------------")
        return batch
    
    
    @torch.no_grad()
    def recording_to_log(self):
        """
        Get the recording of all environments
+12 −2
Original line number Diff line number Diff line
@@ -16,6 +16,7 @@ Note that we don't combine the main with ray_trainer as ray_trainer is used by o
"""
from vagen.trainer.ppo.ray_trainer import RayPPOTrainer
from vagen.utils.compute_score import compute_score
from vagen.trainer.ppo.ray_trainer import AdvantageEstimator

import ray
import hydra
@@ -72,9 +73,17 @@ def main_task(config, compute_score=None):
    role_worker_mapping = {
        Role.ActorRollout: ray.remote(ActorRolloutRefWorker),
        Role.Critic: ray.remote(CriticWorker),
        Role.RefPolicy: ray.remote(ActorRolloutRefWorker)
    }

    use_ref = config.algorithm.adv_estimator != AdvantageEstimator.GAE.value and \
        config.actor_rollout_ref.ref.get('use_ref', True)
    print(f"[DEBUG] use_ref={use_ref}")
    if use_ref:
        role_worker_mapping[Role.RefPolicy] = ray.remote(ActorRolloutRefWorker)
    else:
        config.actor_rollout_ref.actor.use_kl_loss = False
        print("[WARNING] Ref policy is disabled, use_kl_loss is set to False")

    global_pool_id = 'global_pool'
    resource_pool_spec = {
        global_pool_id: [config.trainer.n_gpus_per_node] * config.trainer.nnodes,
@@ -82,8 +91,9 @@ def main_task(config, compute_score=None):
    mapping = {
        Role.ActorRollout: global_pool_id,
        Role.Critic: global_pool_id,
        Role.RefPolicy: global_pool_id,
    }
    if use_ref:
        mapping[Role.RefPolicy] = global_pool_id

    # we should adopt a multi-source reward function here
    # - for rule-based rm, we directly call a reward score
+1 −1
Original line number Diff line number Diff line
@@ -147,7 +147,7 @@ def compute_advantage(data: DataProto, adv_estimator, gamma=1.0, lam=1.0, num_re
            loss_mask = data.batch['loss_mask'][:, -response_length:]
            advantages, returns =core_algos.compute_gae_advantage_return_with_loss_mask(token_level_rewards=token_level_rewards,
                                                                    values=values,
                                                                    eos_mask=loss_mask,
                                                                    loss_mask=loss_mask,
                                                                    gamma=gamma,
                                                                    lam=lam)
        else: