Commit c3fe3ffc authored by williamzhangNU's avatar williamzhangNU
Browse files

update rollout

parent 81f21f4c
Loading
Loading
Loading
Loading
+6 −2
Changes for README.md: 6 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -21,5 +21,9 @@ git submodule update --init --recursive
1. Implement the validation in ray_trainer
2. Transfer the old ray_trainer file from RAGEN to VAGEN
    - Implement the metric for wandb
3. Transfer the train.py from RAGEN to VAGEN
4. Change recorder to a logging class
 No newline at end of file
3. Add environment specific metrics for wandb logging
4. Remove <box_start> and <box_end> in the rollout
5. Loss mask

## NOTE
Does not support use_dynamic_bsz for now
 No newline at end of file
+21 −25
Changes for vagen/env/base.py: 21 added lines, 25 removed lines.
Original line number Diff line number Diff line
@@ -181,20 +181,16 @@ class EnvFeedbackSingleStep:
        step_action_str: String representation of the action taken.
        step_reward: Reward received for taking the action.
        step_done: Flag indicating if the episode is done after this step.
        step_env_finished_before: Flag indicating if the environment was already 
                                 finished before this step.
        step_info: Additional information about the step.
    """
    step_env_observation: EnvObservation = field(default_factory=EnvObservation)
    step_action_str: str = ""
    step_observation: EnvObservation = field(default_factory=EnvObservation)
    step_reward: float = 0.0
    step_done: bool = False
    step_env_finished_before: bool = False
    step_info: Dict[str, Any] = field(default_factory=dict)
    
    def is_terminal(self) -> bool:
        """Check if this step resulted in a terminal state."""
        return self.step_done or self.step_env_finished_before
        return self.step_done


@dataclass
@@ -207,17 +203,12 @@ class EnvFeedback:
        info: Additional information about the overall feedback.
    """
    env_feedbacks: List[EnvFeedbackSingleStep] = field(default_factory=list)
    info: Dict[str, Any] = field(default_factory=dict)
    llm_raw_response: str = ""

    @property
    def env_observation(self) -> EnvObservation:
    def observation(self) -> EnvObservation:
        """Get the merged observation from all steps."""
        return EnvObservation.merge_observation([feedback.step_env_observation for feedback in self.env_feedbacks])
    
    @property
    def action_str(self) -> List[str]:
        """Get the action string from all steps."""
        return [feedback.step_action_str for feedback in self.env_feedbacks]
        return EnvObservation.merge_observation([feedback.step_observation for feedback in self.env_feedbacks])
    
    @property
    def reward(self) -> float:
@@ -234,6 +225,18 @@ class EnvFeedback:
        """Check if any step resulted in a terminal state."""
        return any(feedback.is_terminal() for feedback in self.env_feedbacks)
    
    @property
    def info(self) -> Dict[str, Any]:
        """Get the info from all steps."""
        llm_raw_response = self.llm_raw_response
        merged_info = []
        for feedback in self.env_feedbacks:
            merged_info.append(feedback.step_info)
        return {
            'llm_raw_response': llm_raw_response,
            'info_each_step': merged_info,
        }
    
    def add_step(self, step: EnvFeedbackSingleStep) -> None:
        """Add a new step feedback to the collection."""
        self.env_feedbacks.append(step)
@@ -339,19 +342,12 @@ class BaseGame(ABC):

    @staticmethod
    def convert_numpy_to_PIL(numpy_array: np.ndarray) -> Image.Image:
        """Convert a numpy array to a PIL RGBA image."""
        """Convert a numpy array to a PIL RGB image."""
        if numpy_array.shape[-1] == 3:
            # Convert RGB to RGBA by adding an alpha channel
            height, width, _ = numpy_array.shape
            rgba_array = np.zeros((height, width, 4), dtype=numpy_array.dtype)
            rgba_array[:, :, 0:3] = numpy_array
            rgba_array[:, :, 3] = 255  # Set alpha channel to fully opaque
            return Image.fromarray(rgba_array, mode='RGBA')
        elif numpy_array.shape[-1] == 4:
            # Already has 4 channels, assume it's RGBA
            return Image.fromarray(numpy_array, mode='RGBA')
            # Convert numpy array to RGB PIL Image
            return Image.fromarray(numpy_array, mode='RGB')
        else:
            raise ValueError(f"Unsupported number of channels: {numpy_array.shape[-1]}. Expected 3 (RGB) or 4 (RGBA).")
            raise ValueError(f"Unsupported number of channels: {numpy_array.shape[-1]}. Expected 3 (RGB).")

    @abstractmethod
    def _preprocess(self, text: str) -> Dict:
+1 −0
Changes for vagen/env/config/sokoban.yaml: 1 added line, 0 removed lines.
Original line number Diff line number Diff line
@@ -7,3 +7,4 @@ env:
    num_boxes: 1
    max_steps: 100
    search_depth: 30 # this will change the starting position of the player
    visual_env: true
 No newline at end of file
+28 −23
Changes for vagen/env/sokoban/env.py: 28 added lines, 23 removed lines.
Original line number Diff line number Diff line
@@ -238,17 +238,17 @@ class SokobanGame(BaseGame):
        ):
        super().__init__(**env_config)

        dim_room = self.env_config.get('dim_room', (6, 6))
        num_boxes = self.env_config.get('num_boxes', 1)
        max_steps = self.env_config.get('max_steps', 100)
        search_depth = self.env_config.get('search_depth', 30)
        dim_room = self.env_config['dim_room']
        num_boxes = self.env_config['num_boxes']
        max_steps = self.env_config['max_steps']
        search_depth = self.env_config['search_depth']
        self.env = SokobanEnv(
            dim_room=dim_room,
            num_boxes=num_boxes,
            max_steps=max_steps,
            search_depth=search_depth
        )
        self.use_visual = self.env_config.get('use_visual', False)
        self.visual_env = self.env_config.get('visual_env', True)
        
    @classmethod
    def _extract_action(cls, text):
@@ -281,9 +281,9 @@ class SokobanGame(BaseGame):
    def _get_observation(self):
        """
        Get the observation of the environment.
        If use_visual is True, return the visual observation (PIL RGBA image).
        If visual_env is True, return the visual observation (PIL RGBA image).
        """
        if self.use_visual:
        if self.visual_env:
            visual_observation = self.env.render('rgb_array')
            if isinstance(visual_observation, np.ndarray):
                visual_observation = self.convert_numpy_to_PIL(visual_observation)
@@ -356,10 +356,8 @@ class SokobanGame(BaseGame):
    @classmethod
    def _postprocess(
        cls, 
        env_finished_before: bool = False,
        env_init: bool = False,
        action_valid: bool = True,
        last_action_str: str = "",
        observation: Dict = {},
        reward: float = 0,
        done: bool = False,
@@ -370,27 +368,32 @@ class SokobanGame(BaseGame):
        The returned observation_template is a string with placeholder,
            and multi_modal_observation defines mapping from placeholder to multi-modal observation.
        """
        if env_finished_before:
            return EnvFeedbackSingleStep(step_env_finished_before=True)
        env_observation = EnvObservation()

        if env_init:
            observation_template = cls.PROMPT_TEMPLATE.init_observation_template
            env_observation.create_observation(
                template=observation_template,
                contents=[observation['visual']],
                replace_keys=['{observation}']
            )
        else:
            if not action_valid:
                observation_template = cls.PROMPT_TEMPLATE.invalid_action_template
            else:
                observation_template = cls.PROMPT_TEMPLATE.valid_action_template
            
        env_observation = EnvObservation()
            env_observation.create_observation(
                template=observation_template,
                contents=[observation['visual'], reward, done],
            replace_keys=['observation', 'reward', 'done']
                replace_keys=['{observation}', '{reward}', '{done}']
            )

        
        
        
        return EnvFeedbackSingleStep(
            step_env_observation = env_observation,
            step_action_str = last_action_str,
            step_observation = env_observation,
            step_reward = reward,
            step_done = done,
            step_info = info,
@@ -399,14 +402,13 @@ class SokobanGame(BaseGame):

    def step(self, raw_text: str) -> EnvFeedback:

        if self.finished():
            return EnvFeedback(env_feedbacks=[self._postprocess(env_finished_before=True)])
        assert not self.finished(), "Environment finished before step"

        preprocess_result = self._preprocess(raw_text)
        env_feedback = EnvFeedback()
        actions = preprocess_result.action
        action_valid = preprocess_result.action_valid
        env_feedback.info['llm_raw_response'] = raw_text
        env_feedback.llm_raw_response = raw_text


        for action, valid in zip(actions, action_valid):
@@ -415,22 +417,24 @@ class SokobanGame(BaseGame):
                reward = step_result['step_reward']
                done = step_result['done']
                info = step_result['info']
                last_action = action
            else:
                reward = self.PENALTY_FOR_INVALID
                done = False
                info = {} # TODO
                last_action = self.INVALID_ACTION
                action = self.INVALID_ACTION
            self.traj_reward += reward

            observation = self._get_observation()
            env_feedback_single_step = self._postprocess(
                action_valid=valid,
                last_action_str=self.env.ACTION_LOOKUP[last_action],
                observation=observation, 
                reward=reward, 
                done=done, 
                info=info
                info={
                    'action_valid': valid,
                    'action_str': self.env.ACTION_LOOKUP[action],
                    **info,
                }
            )
            env_feedback.add_step(env_feedback_single_step)
            if done or self.finished():
@@ -445,10 +449,11 @@ class SokobanGame(BaseGame):
        self.env.reset(seed=seed)
        self.traj_reward = 0
        observation = self._get_observation()
        return self._postprocess(
        step_feedback = self._postprocess(
            env_init=True,
            observation=observation,
        )
        return EnvFeedback(env_feedbacks=[step_feedback])

    def finished(self) -> bool:
        return self.env.finished()
+3 −2
Changes for vagen/examples/sokoban/debug_qwen2_5_vl.sh: 3 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -2,7 +2,7 @@ set -x

export VLLM_ATTENTION_BACKEND=XFORMERS

python3 -m main_ppo \
python3 -m vagen.trainer.main_ppo \
    algorithm.adv_estimator=grpo \
    data.train_files=data/sokoban/train.parquet \
    data.val_files=data/sokoban/test.parquet \
@@ -41,4 +41,5 @@ python3 -m main_ppo \
    trainer.save_freq=-1 \
    trainer.test_freq=5 \
    trainer.total_epochs=15 \
    +max_turns=1 2>&1 | tee debug.log
    +max_turns=2 \
    2>&1 | tee debug.log
Loading