@@ -41,7 +41,7 @@ def __init__(
4141 self .policy_runner_args = _normalize_args (self .task_args .policy_runner_args )
4242 self .arena_env = self .task_args .arena_env
4343 self .arena_env_args = _normalize_args (self .task_args .arena_env_args )
44- self .env_variations = _normalize_args (self .task_args .env_variations )
44+ self .variations = _normalize_args (self .task_args .variations )
4545 self .image = image
4646
4747 @staticmethod
@@ -63,7 +63,7 @@ def add_task_arguments(parser: argparse.ArgumentParser) -> None:
6363 help = "Env-related arguments for the chosen Arena environment" ,
6464 )
6565 group .add_argument (
66- "--env_variations " ,
66+ "--variations " ,
6767 default = "" ,
6868 help = "Hydra-style variation overrides appended to the env, e.g. 'light.hdr_image.enabled=true'" ,
6969 )
@@ -90,20 +90,33 @@ def _get_policy_args(self) -> list[str]:
9090 """
9191
9292 def _get_run_script (self ) -> str :
93- policy_args_str = " " .join (self ._get_policy_args ())
94- # Skip empty segments so the command has no stray double spaces.
95- env_spec = " " .join (part for part in (self ._get_arena_env (), self .arena_env_args , self .env_variations ) if part )
96- return (
97- "set -euxo pipefail\n "
98- f"{ POLICY_RUNNER_COMMAND } "
99- f"{ policy_args_str } "
100- f"--output_base_dir { OSMO_TASK_OUTPUT_DIR } "
101- f"{ self .policy_runner_args } "
102- f"{ env_spec } \n "
103- )
104-
105- def _get_arena_env (self ) -> str :
106- """Render the env source: a ``--env_graph_spec_yaml`` flag for a YAML path, else the example-env name."""
93+ return f"set -euxo pipefail\n { self ._get_policy_runner_command ()} \n "
94+
95+ def _get_policy_runner_command (self ) -> str :
96+ """Assemble the policy_runner.py command, dropping empty segments so there are no double spaces."""
97+ parts = [
98+ POLICY_RUNNER_COMMAND ,
99+ * self ._get_policy_args (),
100+ "--output_base_dir" ,
101+ OSMO_TASK_OUTPUT_DIR ,
102+ self .policy_runner_args ,
103+ self .get_arena_env_token (),
104+ self .arena_env_args ,
105+ self .variations ,
106+ ]
107+ return " " .join (part for part in parts if part )
108+
109+ def get_arena_env_token (self ) -> str :
110+ """Render the Arena env selector token as ``policy_runner.py`` expects it.
111+
112+ The ``--arena_env`` value chooses the environment source and is resolved as follows:
113+
114+ - A registered example-environment name is passed through unchanged, e.g.
115+ ``kitchen_pick_and_place`` -> ``kitchen_pick_and_place``.
116+ - A graph-spec YAML path (ending in ``.yaml``/``.yml``) is wrapped in the
117+ ``--env_graph_spec_yaml`` flag, e.g.
118+ ``robolab/mustard_raisin_box.yaml`` -> ``--env_graph_spec_yaml robolab/mustard_raisin_box.yaml``.
119+ """
107120 if self .arena_env .endswith ((".yaml" , ".yml" )):
108121 return f"--env_graph_spec_yaml { self .arena_env } "
109122 return self .arena_env
0 commit comments