Commit 563e45e7 authored by jameskrw's avatar jameskrw
Browse files

updated device for navigation

parent 80f53af8
Loading
Loading
Loading
Loading
+0 −3
Original line number Diff line number Diff line
@@ -11,9 +11,6 @@ class BaseEnvConfig(ABC):
    def config_id(self) -> str: # config identifier, wandb and mllm rollout manager use this to identify the config
        pass
    
    def __init__(self, **kwargs):
        pass
    
    def get(self, key, default=None):
        """
        Get the value of a config key.
+0 −4
Original line number Diff line number Diff line
@@ -5,10 +5,6 @@ from typing import Optional, List, Union
class BaseServiceConfig(ABC):
    max_workers: int = 10
    
    
    def __init__(self, **kwargs):
        pass
    
    def get(self, key, default=None):
        """
        Get the value of a config key.
+10 −1
Original line number Diff line number Diff line
@@ -138,4 +138,13 @@ if __name__ == "__main__":
    test_dataset = load_dataset('parquet', data_files={"test": test_path}, split="test")
    for i in range(2):
        print(train_dataset[i])
        print(test_dataset[i])
 No newline at end of file
        env_name = train_dataset[i]["extra_info"]["env_name"]
        env_config_cls = REGISTERED_ENV[env_name]["config_cls"]
        env_config= env_config_cls(**train_dataset[i]["extra_info"]["env_config"])
        print(env_config.config_id())
    for i in range(2):
        print(train_dataset[i])
        env_name = test_dataset[i]["extra_info"]["env_name"]
        env_config_cls = REGISTERED_ENV[env_name]["config_cls"]
        env_config= env_config_cls(**test_dataset[i]["extra_info"]["env_config"])
        print(env_config.config_id())
 No newline at end of file
+0 −1
Original line number Diff line number Diff line
from .env_config import FrozenLakeEnvConfig
from .service_config import FrozenLakeServiceConfig
from .env import FrozenLakeEnv
from .service import FrozenLakeService
 No newline at end of file
+1 −1
Original line number Diff line number Diff line
@@ -13,7 +13,7 @@ class FrozenLakeEnvConfig(BaseEnvConfig):
    min_actions_to_succeed: int = 5
    
    def config_id(self) -> str:
        id_fields=["is_slippery", "size", "p", "render_mode", "max_actions_per_step", "min_actions_to_succeed"]
        id_fields=["is_slippery", "size", "p", "render_mode", "max_actions_per_step", "min_actions_to_succeed","format_reward"]
        id_str = ",".join([f"{field.name}={getattr(self, field.name)}" for field in fields(self) if field.name in id_fields])
        return f"FrozenLakeEnvConfig({id_str})"

Loading