Commit 2a99d2f9 authored by Kangrui Wang's avatar Kangrui Wang
Browse files

dev

parent fc5b0428
Loading
Loading
Loading
Loading

vagen/env_new/base_config.py

deleted100644 → 0
+0 −15
Original line number Diff line number Diff line
from abc import ABC, abstractmethod
import re
from typing import Optional, List, Tuple, Any, Dict
from copy import deepcopy
from transformers import AutoTokenizer
import torch
from PIL import Image
import numpy as np
from dataclasses import dataclass, field


@dataclass
class BaseConifg():
    init_config:Any # for interface initialization
    reset_config:Any # for interface reset
 No newline at end of file
+3 −3
Original line number Diff line number Diff line
@@ -2,8 +2,8 @@ from abc import ABC, abstractmethod
from typing import Optional, List, Tuple, Any, Dict

class BaseEnv(ABC):
    def __init__(self, env_config):
        self.env_config = env_config    
    def __init__(self, config):
        self.config = config    
    
    
    @abstractmethod
@@ -25,7 +25,7 @@ class BaseEnv(ABC):
        pass
    
    @abstractmethod
    def reset(self, seed: Optional[int] = None) -> Tuple[Any, Dict]:
    def reset(self, seed: Optional[Any] = None) -> Tuple[Any, Dict]:
        """
        Reset the environment.
        NOTE: the environment should be same for the same seed
+12 −69
Original line number Diff line number Diff line
@@ -7,40 +7,20 @@ import torch
from PIL import Image
import numpy as np
from dataclasses import dataclass, field


IMAGE_PLACEHOLDER = "<image>"    
from .utils.io_utils import validate_reset_io,validate_step_io
           
class BaseInterface(ABC):
    image_placeholder="<image>"
    
    @classmethod
    def __init__(self, config):
    def __init__(self, config: Dict):
        self.config = config
    
    @classmethod
    def name_repr(cls) -> str:
        """Get the name of the environment."""
        return cls.__name__
        
    @abstractmethod
    def _reset(self, seed: Optional[int] = None) -> Tuple[Any, float, bool, Dict]:
        """Reset the environment."""
        pass
    
    @abstractmethod
    def _step(self, action:str) -> Tuple[Any, float, bool, Dict]:
        """Execute action string in the environment."""
        # return observation, reward, done, info
        # info must contain "llm_raw_response" key, which is a string
    def config_repr(cls, config) -> str:
        """convert config to str"""
        pass
    
    @classmethod
    @abstractmethod
    def config_repr(cls, config: Dict) -> str:
        """Get the config of the environment."""
        pass
    
    
    @abstractmethod
    def close(self):
        """Close the environment."""
@@ -51,53 +31,16 @@ class BaseInterface(ABC):
        """Get the task instruction."""
        pass
    
    
    @abstractmethod
    @validate_step_io
    def step(self, action: str) -> Tuple[Dict, float, bool, Dict]:
        """Execute action string in the environment."""
        """Please use the following assertions to validate the output, 
        then you can rewrite the step in your own class to improve the performance"""
        
        
        assert isinstance(action, str), f"action must be str, got {type(action)}"
        obs,reward,done,info = self._step(action)
        assert isinstance(reward, (int, float)), f"reward must be int or float, got {type(reward)}"
        assert isinstance(done, bool), f"done must be bool, got {type(done)}"
        assert isinstance(info, dict), f"info must be dict, got {type(info)}"
        assert isinstance(obs, dict), f"obs must be dict, got {type(obs)}"
        assert "llm_raw_response" in info, f"info must contain 'llm_raw_response' key"
        assert isinstance(info["llm_raw_response"], str), f"info['llm_raw_response'] must be str, got {type(info['llm_raw_response'])}"
        assert "text_template" in obs, f"obs must contain 'text_template' key"
        assert isinstance(obs["text_template"], str), f"obs['text_template'] must be str, got {type(obs['text_template'])}"
        
        if "multi_modal_data" in obs:
            if IMAGE_PLACEHOLDER in obs["multi_modal_data"]:
                assert isinstance(obs["multi_modal_data"][IMAGE_PLACEHOLDER], list), f"obs['multi_modal_data']['<image>'] must be list, got {type(obs['multi_modal_data'][IMAGE_PLACEHOLDER])}"
                for image in obs["multi_modal_data"][IMAGE_PLACEHOLDER]:
                    assert isinstance(image, Image.Image), f"image must be PIL.Image.Image, got {type(image)}"
                len_of_images = len(obs["multi_modal_data"][IMAGE_PLACEHOLDER])
                len_of_image_in_text_template = len(re.findall(IMAGE_PLACEHOLDER, obs["text_template"]))
                assert len_of_images == len_of_image_in_text_template, f"len_of_images must be equal to len_of_image_in_text_template, got {len_of_images} and {len_of_image_in_text_template}"
        return obs, reward, done, info
    
        pass
    
    def reset(self, seed: int):
    @abstractmethod
    @validate_reset_io    
    def reset(self, seed: int) -> Tuple[Dict, Dict]:
        """Reset the environment."""
        assert isinstance(seed, int), f"seed must be int, got {type(seed)}"
        obs, info = self._reset(seed)
        assert isinstance(info, dict), f"info must be dict, got {type(info)}"
        assert isinstance(obs, dict), f"obs must be dict, got {type(obs)}"
        assert "text_template" in obs, f"obs must contain 'text_template' key"
        assert isinstance(obs["text_template"], str), f"obs['text_template'] must be str, got {type(obs['text_template'])}"
        
        if "multi_modal_data" in obs:
            if IMAGE_PLACEHOLDER in obs["multi_modal_data"]:
                assert isinstance(obs["multi_modal_data"][IMAGE_PLACEHOLDER], list), f"obs['multi_modal_data']['<image>'] must be list, got {type(obs['multi_modal_data'][IMAGE_PLACEHOLDER])}"
                for image in obs["multi_modal_data"][IMAGE_PLACEHOLDER]:
                    assert isinstance(image, Image.Image), f"image must be PIL.Image.Image, got {type(image)}"
                len_of_images = len(obs["multi_modal_data"][IMAGE_PLACEHOLDER])
                len_of_image_in_text_template = len(re.findall(IMAGE_PLACEHOLDER, obs["text_template"]))
                assert len_of_images == len_of_image_in_text_template, f"len_of_images must be equal to len_of_image_in_text_template, got {len_of_images} and {len_of_image_in_text_template}"
        return obs, info
        pass
    
    @abstractmethod
    def get_traj_reward(self) -> float:
+3 −7
Original line number Diff line number Diff line
@@ -12,10 +12,8 @@ class DatasetCreator:
    def __init__(self, config: Dict):
        self.config = config
        self.data_dir = self.config['data_dir']

        self.env_name = self.config['name']
        self.env_config = self.config['env_config']
        self.interface_config = self.config['interface_config']
        assert "env_name" in self.interface_config
        
        

@@ -39,15 +37,13 @@ class DatasetCreator:
            
        def _create_instance(seed_idx, split: str = 'train'):
            env_settings = {
                'env_name': self.env_name,
                'env_config': self.env_config,
                'interface_config': self.interface_config,
                'config': self.interface_config,
                'seed': seed_idx
            }

            # TODO: no reward model defined here for the reward will be generated while rollout
            return {
                "data_source": self.env_name,
                "data_source": self.interface_config["env_name"],
                "prompt": [{"role": "user", "content": ''}],
                "extra_info": {"split": split, **env_settings}
            }
+93 −0
Original line number Diff line number Diff line
from vagen.env_new.base_env import BaseEnv
import gym
from gym_sokoban.envs.sokoban_env import SokobanEnv as GymSokobanEnv
from vagen.env.sokoban.room_utils import generate_room
from typing import Dict


class SokobanVisionEnv(BaseEnv, GymSokobanEnv):

    GRID_LOOKUP = {
        0: " # \t",  # wall
        1: " _ \t",  # floor
        2: " O \t",  # target
        3: " √ \t",  # box on target
        4: " X \t",  # box
        5: " P \t",  # player
        6: " S \t",  # player on target
        # Use tab separator to separate columns and \n\n to separate rows.
    }

    ACTION_LOOKUP = {
        0: "None",
        1: "Up",
        2: "Down",
        3: "Left",
        4: "Right",
    }

    def __init__(self, config: Dict):
        BaseEnv.__init__(self)
        self.config=config
        GymSokobanEnv.__init__(
            self,
            dim_room=kwargs.pop('dim_room', (6, 6)), 
            max_steps=kwargs.pop('max_steps', 100),
            num_boxes=kwargs.pop('num_boxes', 3),
            **kwargs
        )
        self.ACTION_SPACE = gym.spaces.discrete.Discrete(4, start=1)


    def reset(self, seed: int):
        with NoLoggerWarnings():
            try:
                with set_seed(seed):
                    self.room_fixed, self.room_state, self.box_mapping, action_sequence = generate_room(
                        dim=self.dim_room,
                        num_steps=self.num_gen_steps,
                        num_boxes=self.num_boxes,
                        search_depth=self.search_depth
                    )
            except (RuntimeError, RuntimeWarning) as e:
                print("[SOKOBAN] Runtime Error/Warning: {}".format(e))
                print("[SOKOBAN] Retry . . .")
                next_seed = abs(hash(str(seed))) % (2 ** 32) if seed is not None else None
                return self._reset(next_seed)
            
            # self.action_sequence = self._reverse_action_sequence(action_sequence)
            self.player_position = np.argwhere(self.room_state == 5)[0]
            self.num_env_steps = self.reward_last = self.boxes_on_target = 0
        
        return self._render(mode='text'), {}
    
    def step(self, action: int):
        """
        - Step the environment with the given action.
        - Check if the action is effective (whether player moves in the env).

        TODO modify here after definition of RolloutManager
        """
        assert not self._success()
        result = {
            'step_reward': 0,
            'done': False,
            'info': {},
        }
        
        prev_player_position = self.player_position
        obs, step_reward, done, info = GymSokobanEnv.step(self, action, observation_mode='tiny_rgb_array')
        
        info['action_is_effective'] = not np.array_equal(prev_player_position, self.player_position)
        return obs, step_reward, done, info



    def close(self):
        GymSokobanEnv.close(self)

    def _finished(self):
        return self.num_env_steps >= self.max_steps or self.success()

    def _success(self):
        return self.boxes_on_target == self.num_boxes
 No newline at end of file
Loading