Commit 50e6682c authored by jameskrw's avatar jameskrw
Browse files

improve primitive skill training speed

parent e28432f2
Loading
Loading
Loading
Loading
+3 −2
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
from .primitive_skill import PrimitiveSkillEnv, PrimitiveSkillEnvConfig, PrimitiveSkillService, PrimitiveSkillConfig
REGISTERED_ENV = {
    "sokoban": {
        "env_cls": SokobanEnv,
@@ -27,6 +27,7 @@ REGISTERED_ENV = {
    "primitive_skill": {
        "env_cls": PrimitiveSkillEnv,
        "config_cls": PrimitiveSkillEnvConfig,
        "service_cls": PrimitiveSkillService
        "service_cls": PrimitiveSkillService,
        "service_config_cls": PrimitiveSkillConfig
    }
}
 No newline at end of file
+1 −0
Original line number Diff line number Diff line
from .env import PrimitiveSkillEnv
from .env_config import PrimitiveSkillEnvConfig
from .service import PrimitiveSkillService
from .service_config import PrimitiveSkillConfig
 No newline at end of file
+1 −1
Original line number Diff line number Diff line
@@ -11,7 +11,7 @@ import os


def build_env(env_id, control_mode="pd_ee_delta_pose", stage=0, record_dir='./test'):
    env_kwargs = dict(obs_mode="state", control_mode=control_mode, render_mode="rgb_array", sim_backend="cpu")
    env_kwargs = dict(obs_mode="state", control_mode=control_mode, render_mode="rgb_array", sim_backend="cpu",render_backend="gpu")
    env = gym.make(env_id, num_envs=1, enable_shadow=True, stage=stage, **env_kwargs)
    env = CPUGymWrapper(env)
    env = SkillGymWrapper(env,
+559 −151

File changed.

Preview size limit exceeded, changes collapsed.

+10 −0
Original line number Diff line number Diff line
from vagen.env.base.base_service_config import BaseServiceConfig
from dataclasses import dataclass, fields,field

@dataclass
class PrimitiveSkillConfig(BaseServiceConfig):
    max_process_workers: int = field(default=8)
    max_thread_workers: int = field(default=4)
    devices: list = field(default_factory=lambda: [0,1])
    spawn_method: str = field(default="fork")
    timeout: int = field(default=120)
 No newline at end of file
Loading