Commit 3b8d8101 authored by jameskrw's avatar jameskrw
Browse files

minor update

parent 53ae16e5
Loading
Loading
Loading
Loading

vagen/env/manipulation/test.ipynb

deleted100644 → 0
+0 −133

File deleted.

Preview size limit exceeded, changes collapsed.

+6 −6
Original line number Diff line number Diff line
@@ -4,13 +4,13 @@ import copy
from typing import Dict, List, Optional, Tuple, Any
from gymnasium.utils import seeding
from vagen.env.utils.context_utils import parse_llm_raw_response, convert_numpy_to_PIL
from .env_config import ManipulationEnvConfig
from .env_config import PrimitiveSkillEnvConfig
from .maniskill.utils import build_env, handel_info, get_workspace_limits
from .prompts import system_prompt, init_observation_template, action_template
import vagen.env.manipulation.maniskill.env
import vagen.env.primitive_skill.maniskill.env

class ManipulationEnv(BaseEnv):
    def __init__(self, config: ManipulationEnvConfig):
class PrimitiveSkillEnv(BaseEnv):
    def __init__(self, config: PrimitiveSkillEnvConfig):
        self.config = config
        self.env=build_env(config.env_id,record_dir='./test')
    
@@ -193,8 +193,8 @@ if __name__ == "__main__":
    This code demonstrates how to create an instance of the environment,
    reset it, and interact with it using manual input actions.
    """
    config = ManipulationEnvConfig()
    env = ManipulationEnv(config)
    config = PrimitiveSkillEnvConfig()
    env = PrimitiveSkillEnv(config)
    
    print(env.system_prompt())
    obs, info = env.reset()
+3 −3
Original line number Diff line number Diff line
@@ -3,15 +3,15 @@ from dataclasses import dataclass, fields,field
from typing import Optional, List, Union

@dataclass
class ManipulationEnvConfig(BaseEnvConfig):
class PrimitiveSkillEnvConfig(BaseEnvConfig):
    env_id: str = "AlignTwoCube" # AlignTwoCube,PlaceTwoCube,PutAppleInDrawer,StackThreeCube
    render_mode: str = "vision" # vision, text
    
    def config_id(self) -> str:
        id_fields=["env_id","render_mode"]
        id_str = ",".join([f"{field.name}={getattr(self, field.name)}" for field in fields(self) if field.name in id_fields])
        return f"ManipulationEnvConfig({id_str})"
        return f"PrimitiveSkillEnvConfig({id_str})"

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