Commit 178907d5 authored by jameskrw's avatar jameskrw
Browse files

refined primitive skills

parent 6e707706
Loading
Loading
Loading
Loading
+6 −1
Original line number Diff line number Diff line
@@ -2,7 +2,7 @@ from .sokoban import SokobanEnv,SokobanEnvConfig
from .frozenlake import FrozenLakeEnv,FrozenLakeEnvConfig, FrozenLakeService
from .navigation import NavigationEnv, NavigationEnvConfig, NavigationServiceConfig, NavigationService
from .svg import SVGEnv, SvgEnvConfig, SVGService

from .primitive_skill import PrimitiveSkillEnv, PrimitiveSkillEnvConfig, PrimitiveSkillService
REGISTERED_ENV = {
    "sokoban": {
        "env_cls": SokobanEnv,
@@ -24,4 +24,9 @@ REGISTERED_ENV = {
        "config_cls": SvgEnvConfig,
        "service_cls": SVGService
    },
    "primitive_skill": {
        "env_cls": PrimitiveSkillEnv,
        "config_cls": PrimitiveSkillEnvConfig,
        "service_cls": PrimitiveSkillService
    }
}
 No newline at end of file
+1 −1
Original line number Diff line number Diff line
@@ -347,7 +347,7 @@ class NavigationEnv(BaseEnv):
        Returns:
            System prompt string
        """
        if self.config.visual_env:
        if self.config.render_mode == "vision":
            return system_prompt_vision.format(
                max_actions_per_step=self.config.max_actions_per_step,
                action_sep=self.config.action_sep
+2 −2
Original line number Diff line number Diff line
@@ -10,7 +10,7 @@ class NavigationEnvConfig(BaseEnvConfig):
    down_sample_ratio: float = 1.0
    fov: int = 100
    multiview: bool = False
    visual_env: bool = True
    render_mode: str= 'vision'
    max_actions_per_step: int = 10
    max_action_penalty: float = -0.1
    format_reward: float = 0.5
@@ -19,7 +19,7 @@ class NavigationEnvConfig(BaseEnvConfig):
    def config_id(self) -> str:
        """Generate a unique identifier for this configuration."""
        id_fields = ["resolution", "eval_set", "exp_name", "down_sample_ratio", 
                    "fov", "multiview", "visual_env", "max_actions_per_step"]
                    "fov", "multiview", "render_mode", "max_actions_per_step"]
        id_str = ",".join([f"{field.name}={getattr(self, field.name)}" for field in fields(self) if field.name in id_fields])
        return f"NavigationEnvConfig({id_str})"

+3 −0
Original line number Diff line number Diff line
from .env import PrimitiveSkillEnv
from .env_config import PrimitiveSkillEnvConfig
from .service import PrimitiveSkillService
 No newline at end of file
+12 −6
Original line number Diff line number Diff line
@@ -12,14 +12,18 @@ import vagen.env.primitive_skill.maniskill.env
class PrimitiveSkillEnv(BaseEnv):
    def __init__(self, config: PrimitiveSkillEnvConfig):
        self.config = config
        self.env=build_env(config.env_id,record_dir='./test')
        if self.config.record_video:
            record_dir = self.config.video_record_dir
        else:
            record_dir = None
        self.env=build_env(config.env_id,record_dir=record_dir)
    
    def reset(self, seed: Optional[int] = None) -> Tuple[Dict[str, Any], Dict[str, Any]]:
        self.total_reward = 0
        _, info=self.env.reset(seed=seed)
        obs=self._render(info,init_obs=True)
        self.last_info=info
        return obs, info
        return obs, {}
    
    def step(self,action_str):
        reward=0
@@ -34,12 +38,14 @@ class PrimitiveSkillEnv(BaseEnv):
        valid_actions = []
        metrics = {
            "turn_metrics": {
                "action_is_valid": True,  # True if at least one valid action was parsed
                "action_is_valid": False,  # True if at least one valid action was parsed
            },
            "traj_metrics": {
                "success": False,  # Will be set to True if agent reaches goal
            },
        }
        info=self.last_info
        terminated, truncated = False, False
        for action in rst['actions']:
            parsed_action = self._parse_action(action)
            if parsed_action is not None:
@@ -49,10 +55,10 @@ class PrimitiveSkillEnv(BaseEnv):
            else:
                info=self.last_info
                terminated, truncated = False, False
                metrics["turn_metrics"]['action_is_valid'] = False
                break
            if truncated or terminated:
                break
        metrics["turn_metrics"]['action_is_valid'] = len(valid_actions) > 0 and len(valid_actions)==len(rst['actions'])
        if metrics["turn_metrics"]['action_is_valid']:
            reward += self.config.format_reward
        if info['is_success']:
@@ -162,7 +168,7 @@ class PrimitiveSkillEnv(BaseEnv):
                action_array[2] = 1
            else:
                # Invalid action name
                return np.zeros(9)
                return None
            
            # Extract parameters
            params_str = action_str.split('(')[1].split(')')[0]
@@ -204,7 +210,7 @@ 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 = PrimitiveSkillEnvConfig()
    config = PrimitiveSkillEnvConfig(record_video=True, video_record_dir="./test_manipulation_video")
    env = PrimitiveSkillEnv(config)
    
    print(env.system_prompt())
Loading