Commit f4f9f190 authored by jameskrw's avatar jameskrw
Browse files

updated service structrue

parent 04aa2f8d
Loading
Loading
Loading
Loading
+15 −15
Original line number Diff line number Diff line
from .sokoban import SokobanEnv,SokobanConfig
from .frozenlake import FrozenLakeEnv,FrozenLakeConfig, FrozenLakeService
# from .navigation import NavigationEnv, NavigationConfig
# from .svg import SVGEnv, SVGConfig, SVGService
from .sokoban import SokobanEnv,SokobanEnvConfig
from .frozenlake import FrozenLakeEnv,FrozenLakeEnvConfig, FrozenLakeService
from .navigation import NavigationEnv, NavigationEnvConfig
from .svg import SVGEnv, SvgEnvConfig, SVGService

REGISTERED_ENV = {
    "sokoban": {
        "env_cls": SokobanEnv,
        "config_cls": SokobanConfig,
        "config_cls": SokobanEnvConfig,
    },
    "frozenlake": {
        "env_cls": FrozenLakeEnv,
        "config_cls": FrozenLakeConfig,
        "config_cls": FrozenLakeEnvConfig,
        "service_cls": FrozenLakeService
    },
    # "navigation": {
    #     "env_cls": NavigationEnv,
    #     "config_cls": NavigationConfig
    # },
    # "svg": {
    #     "env_cls": SVGEnv,
    #     "config_cls": SVGConfig,
    #     "service_cls": SVGService
    # },
    "navigation": {
        "env_cls": NavigationEnv,
        "config_cls": NavigationEnvConfig
    },
    "svg": {
        "env_cls": SVGEnv,
        "config_cls": SvgEnvConfig,
        "service_cls": SVGService
    },
}
 No newline at end of file
+1 −1
Original line number Diff line number Diff line
@@ -32,7 +32,7 @@ def create_dataset_from_yaml(yaml_file_path: str, force_gen=False,seed=42,train_
        test_size:100
    ```
    
    If the environment config class (e.g., SokobanConfig, FrozenLakeConfig) has a 
    If the environment config class (e.g., SokobanEnvConfig, FrozenLakeEnvConfig) has a 
    generate_seeds(size) method, it will be used to generate seeds for that environment.
    """
    
+2 −1
Original line number Diff line number Diff line
from .env_config import FrozenLakeConfig
from .env_config import FrozenLakeEnvConfig
from .service_config import FrozenLakeServiceConfig
from .env import FrozenLakeEnv
from .service import FrozenLakeService
 No newline at end of file
+3 −3
Original line number Diff line number Diff line
@@ -7,7 +7,7 @@ from gymnasium.envs.toy_text.frozen_lake import FrozenLakeEnv as GymFrozenLakeEn
from vagen.env.utils.env_utils import NoLoggerWarnings, set_seed
from vagen.env.utils.context_utils import parse_llm_raw_response, convert_numpy_to_PIL
from .prompt import system_prompt_text, system_prompt_vision, init_observation_template, action_template
from .env_config import FrozenLakeConfig
from .env_config import FrozenLakeEnvConfig
from .utils import generate_random_map, is_valid

class FrozenLakeEnv(BaseEnv):
@@ -36,7 +36,7 @@ class FrozenLakeEnv(BaseEnv):
        "Up": 3,
    }

    def __init__(self, config: FrozenLakeConfig):
    def __init__(self, config: FrozenLakeEnvConfig):
        BaseEnv.__init__(self)
        self.config = config
       
@@ -208,7 +208,7 @@ class FrozenLakeEnv(BaseEnv):


if __name__ == "__main__":
    config = FrozenLakeConfig()
    config = FrozenLakeEnvConfig()
    env = FrozenLakeEnv(config)
    print(env.system_prompt())
    obs, info = env.reset()
+4 −4
Original line number Diff line number Diff line
from vagen.env.base_env_config import BaseConfig
from vagen.env.base_env_config import BaseEnvConfig
from dataclasses import dataclass, fields,field
from typing import Optional, List, Union

@dataclass
class FrozenLakeConfig(BaseConfig):
class FrozenLakeEnvConfig(BaseEnvConfig):
    desc: Optional[List[str]] = None  # environment map
    is_slippery: bool = False
    size: int = 4
@@ -15,8 +15,8 @@ class FrozenLakeConfig(BaseConfig):
    def config_id(self) -> str:
        id_fields=["is_slippery", "size", "p", "render_mode", "max_actions_per_step", "min_actions_to_succeed"]
        id_str = ",".join([f"{field.name}={getattr(self, field.name)}" for field in fields(self) if field.name in id_fields])
        return f"FrozenLakeConfig({id_str})"
        return f"FrozenLakeEnvConfig({id_str})"

if __name__ == "__main__":
    config = FrozenLakeConfig()
    config = FrozenLakeEnvConfig()
    print(config.config_id())
 No newline at end of file
Loading