Skip to content

Commit 244dbe5

Browse files
committed
Address review: rename flags, assemble command via a helper
- Rename --env_variations to --variations (and self.variations). - Rename the env resolver to get_arena_env_token and expand its docstring with registered-name and YAML examples. - Assemble the policy_runner command in _get_policy_runner_command, filtering empty segments so the command never has double spaces. Signed-off-by: alex <amillane@nvidia.com>
1 parent 3a14df3 commit 244dbe5

1 file changed

Lines changed: 29 additions & 16 deletions

File tree

osmo/tasks/policy_runner_task.py

Lines changed: 29 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)