Skip to content

Commit 4470ff7

Browse files
committed
Add history encoder ablation option
1 parent 6d26580 commit 4470ff7

5 files changed

Lines changed: 156 additions & 1 deletion

File tree

tests/test_train_script.py

Lines changed: 58 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,13 @@
1212
import numpy as np
1313
import pytest
1414

15-
from train_mimic.app import DEFAULT_TASK, validate_checkpoint_path, validate_motion_file
15+
from train_mimic.app import (
16+
DEFAULT_TASK,
17+
HISTORY_ENCODER_TEMPORAL_CNN,
18+
apply_history_encoder_config,
19+
validate_checkpoint_path,
20+
validate_motion_file,
21+
)
1622
from train_mimic.data.dataset_lib import PRECOMPUTED_MOTION_VERSION
1723
from train_mimic.scripts import train
1824
from train_mimic.tasks.tracking.config.rl import make_general_tracking_ppo_runner_cfg
@@ -42,6 +48,7 @@ def _args(**overrides: object) -> argparse.Namespace:
4248
"rewind_prob": None,
4349
"rewind_min_steps": None,
4450
"rewind_max_steps": None,
51+
"history_encoder": HISTORY_ENCODER_TEMPORAL_CNN,
4552
"device": None,
4653
"gpu_ids": None,
4754
"master_port": 29500,
@@ -57,6 +64,7 @@ class TestTrainLauncherHelpers:
5764
def test_parse_args_defaults_to_tensorboard_logger(self) -> None:
5865
args = train.parse_args([])
5966
assert args.logger == "tensorboard"
67+
assert args.history_encoder == HISTORY_ENCODER_TEMPORAL_CNN
6068

6169
def test_parse_args_accepts_logger_choice(self) -> None:
6270
args = train.parse_args(["--logger", "swanlab"])
@@ -75,6 +83,10 @@ def test_parse_args_accepts_robot_xml(self) -> None:
7583
args = train.parse_args(["--robot_xml", "assets/robots/unitree_g1/g1_29dof_dex3.xml"])
7684
assert args.robot_xml == "assets/robots/unitree_g1/g1_29dof_dex3.xml"
7785

86+
def test_parse_args_accepts_history_encoder_none(self) -> None:
87+
args = train.parse_args(["--history_encoder", "none"])
88+
assert args.history_encoder == "none"
89+
7890
def test_should_launch_multi_gpu(self) -> None:
7991
args = _args(gpu_ids=[0, 1, 2, 3])
8092
assert train._should_launch_multi_gpu(args, env={"WORLD_SIZE": "1"}) is True
@@ -219,6 +231,7 @@ def test_configure_swanlab_logger_syncs_tensorboard(self, monkeypatch: pytest.Mo
219231
"rewind_prob": 0.8,
220232
"rewind_min_steps": 25,
221233
"rewind_max_steps": 75,
234+
"history_encoder": HISTORY_ENCODER_TEMPORAL_CNN,
222235
},
223236
},
224237
),
@@ -250,6 +263,50 @@ def test_tracking_runner_configs_disable_model_upload() -> None:
250263
assert make_general_tracking_ppo_runner_cfg().upload_model is False
251264

252265

266+
def test_apply_history_encoder_none_switches_runner_to_current_frame_mlp() -> None:
267+
actor_marker = object()
268+
critic_marker = object()
269+
env_cfg = types.SimpleNamespace(
270+
observations={
271+
"actor": actor_marker,
272+
"critic": critic_marker,
273+
"actor_history": object(),
274+
"critic_history": object(),
275+
}
276+
)
277+
agent_cfg = make_general_tracking_ppo_runner_cfg()
278+
279+
apply_history_encoder_config(env_cfg, agent_cfg, "none")
280+
281+
assert env_cfg.observations == {
282+
"actor": actor_marker,
283+
"critic": critic_marker,
284+
}
285+
assert agent_cfg.history_encoder == "none"
286+
assert agent_cfg.obs_groups == {
287+
"actor": ("actor",),
288+
"critic": ("critic",),
289+
}
290+
assert agent_cfg.actor.class_name == "rsl_rl.models.mlp_model:MLPModel"
291+
assert agent_cfg.actor.cnn_cfg is None
292+
assert agent_cfg.critic.class_name == "rsl_rl.models.mlp_model:MLPModel"
293+
assert agent_cfg.critic.cnn_cfg is None
294+
295+
296+
def test_apply_history_encoder_temporal_cnn_keeps_default_runner_cfg() -> None:
297+
env_cfg = types.SimpleNamespace(observations={"actor_history": object()})
298+
agent_cfg = make_general_tracking_ppo_runner_cfg()
299+
actor_class_name = agent_cfg.actor.class_name
300+
obs_groups = agent_cfg.obs_groups
301+
302+
apply_history_encoder_config(env_cfg, agent_cfg, HISTORY_ENCODER_TEMPORAL_CNN)
303+
304+
assert env_cfg.observations.keys() == {"actor_history"}
305+
assert agent_cfg.history_encoder == HISTORY_ENCODER_TEMPORAL_CNN
306+
assert agent_cfg.actor.class_name == actor_class_name
307+
assert agent_cfg.obs_groups == obs_groups
308+
309+
253310
def test_make_g1_training_robot_cfg_uses_requested_xml() -> None:
254311
from train_mimic.tasks.tracking.config.env import make_g1_training_robot_cfg
255312

train_mimic/app.py

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,10 @@
1414
from train_mimic.data.dataset_lib import find_precomputed_motion_shards, validate_precomputed_motion_dataset
1515

1616
DEFAULT_TASK = GENERAL_TRACKING_TASK
17+
HISTORY_ENCODER_TEMPORAL_CNN = "temporal_cnn"
18+
HISTORY_ENCODER_NONE = "none"
19+
HISTORY_ENCODER_CHOICES = (HISTORY_ENCODER_TEMPORAL_CNN, HISTORY_ENCODER_NONE)
20+
_MLP_MODEL_CLASS = "rsl_rl.models.mlp_model:MLPModel"
1721

1822

1923
def validate_motion_file(motion_file: str) -> None:
@@ -93,6 +97,37 @@ def build_runner_cfg_dict(agent_cfg: Any, *, force_tensorboard: bool = False) ->
9397
return agent_dict
9498

9599

100+
def apply_history_encoder_config(
101+
env_cfg: Any,
102+
agent_cfg: Any,
103+
history_encoder: str,
104+
) -> None:
105+
"""Apply the training/eval observation-model variant for history ablations."""
106+
if history_encoder not in HISTORY_ENCODER_CHOICES:
107+
raise ValueError(
108+
f"Unsupported history_encoder={history_encoder!r}. "
109+
f"Supported values are: {', '.join(HISTORY_ENCODER_CHOICES)}."
110+
)
111+
112+
agent_cfg.history_encoder = history_encoder
113+
if history_encoder == HISTORY_ENCODER_TEMPORAL_CNN:
114+
return
115+
116+
observations = getattr(env_cfg, "observations", None)
117+
if observations is not None:
118+
observations.pop("actor_history", None)
119+
observations.pop("critic_history", None)
120+
121+
agent_cfg.obs_groups = {
122+
"actor": ("actor",),
123+
"critic": ("critic",),
124+
}
125+
agent_cfg.actor.class_name = _MLP_MODEL_CLASS
126+
agent_cfg.actor.cnn_cfg = None
127+
agent_cfg.critic.class_name = _MLP_MODEL_CLASS
128+
agent_cfg.critic.cnn_cfg = None
129+
130+
96131
def resolve_device(requested_device: str | None, torch_module: Any) -> str:
97132
if requested_device is not None:
98133
return requested_device

train_mimic/scripts/benchmark.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,9 @@
3131

3232
from train_mimic.app import (
3333
DEFAULT_TASK,
34+
HISTORY_ENCODER_CHOICES,
35+
HISTORY_ENCODER_TEMPORAL_CNN,
36+
apply_history_encoder_config,
3437
build_runner_cfg_dict,
3538
import_training_stack,
3639
load_task_components,
@@ -71,6 +74,16 @@ def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
7174
default=DEFAULT_TASK,
7275
help="Task id to benchmark (default: %(default)s)",
7376
)
77+
parser.add_argument(
78+
"--history_encoder",
79+
type=str,
80+
default=HISTORY_ENCODER_TEMPORAL_CNN,
81+
choices=HISTORY_ENCODER_CHOICES,
82+
help=(
83+
"Policy history encoder variant used by the checkpoint. Use none for "
84+
"MLP checkpoints trained without actor_history/critic_history."
85+
),
86+
)
7487
return parser.parse_args(argv)
7588

7689

@@ -359,6 +372,7 @@ def main(argv: Sequence[str] | None = None) -> int:
359372
load_rl_cfg=_load_rl_cfg,
360373
load_runner_cls=_load_runner_cls,
361374
)
375+
apply_history_encoder_config(base_env_cfg, agent_cfg, args.history_encoder)
362376
base_env_cfg.commands["motion"].motion_file = args.motion_file
363377
benchmark_env_cfg = _configure_benchmark_env_cfg(
364378
base_env_cfg,
@@ -429,6 +443,7 @@ def main(argv: Sequence[str] | None = None) -> int:
429443
"motion_file": args.motion_file,
430444
"seed": args.seed,
431445
"num_envs": args.num_envs,
446+
"history_encoder": args.history_encoder,
432447
},
433448
plan=plan,
434449
results=results,

train_mimic/scripts/play.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,10 +28,19 @@
2828

2929
import argparse
3030
import os
31+
import sys
32+
from pathlib import Path
33+
34+
REPO_ROOT = Path(__file__).resolve().parents[2]
35+
if str(REPO_ROOT) not in sys.path:
36+
sys.path.insert(0, str(REPO_ROOT))
3137

3238
from mjlab.viewer import NativeMujocoViewer, ViserPlayViewer
3339
from train_mimic.app import (
3440
DEFAULT_TASK,
41+
HISTORY_ENCODER_CHOICES,
42+
HISTORY_ENCODER_TEMPORAL_CNN,
43+
apply_history_encoder_config,
3544
build_runner_cfg_dict,
3645
import_training_stack,
3746
load_task_components,
@@ -54,6 +63,16 @@ def parse_args() -> argparse.Namespace:
5463
parser.add_argument("--device", type=str, default=None)
5564
parser.add_argument("--task", type=str, default=DEFAULT_TASK,
5665
help="Task id to play (default: %(default)s)")
66+
parser.add_argument(
67+
"--history_encoder",
68+
type=str,
69+
default=HISTORY_ENCODER_TEMPORAL_CNN,
70+
choices=HISTORY_ENCODER_CHOICES,
71+
help=(
72+
"Policy history encoder variant used by the checkpoint. Use none for "
73+
"MLP checkpoints trained without actor_history/critic_history."
74+
),
75+
)
5776
return parser.parse_args()
5877

5978

@@ -90,6 +109,7 @@ def main() -> None:
90109
)
91110

92111
# Override for playback
112+
apply_history_encoder_config(env_cfg, agent_cfg, args.history_encoder)
93113
env_cfg.scene.num_envs = args.num_envs
94114
env_cfg.commands["motion"].motion_file = args.motion_file
95115

train_mimic/scripts/train.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,11 @@
2323
--motion_file data/datasets_precomputed \
2424
--logger swanlab
2525
26+
# Ablate temporal history encoder
27+
python train_mimic/scripts/train.py \
28+
--history_encoder none \
29+
--motion_file data/datasets_precomputed
30+
2631
# Resume for additional iterations
2732
python train_mimic/scripts/train.py \
2833
--resume logs/rsl_rl/g1_general_tracking/<run>/model_12000.pt \
@@ -40,10 +45,18 @@
4045
import sys
4146
import time
4247
from datetime import datetime
48+
from pathlib import Path
4349
from typing import Any, Sequence
4450

51+
REPO_ROOT = Path(__file__).resolve().parents[2]
52+
if str(REPO_ROOT) not in sys.path:
53+
sys.path.insert(0, str(REPO_ROOT))
54+
4555
from train_mimic.app import (
4656
DEFAULT_TASK,
57+
HISTORY_ENCODER_CHOICES,
58+
HISTORY_ENCODER_TEMPORAL_CNN,
59+
apply_history_encoder_config,
4760
build_runner_cfg_dict,
4861
import_training_stack,
4962
load_task_components,
@@ -106,6 +119,17 @@ def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
106119
help="Minimum policy steps to rewind for rewind sampling")
107120
parser.add_argument("--rewind_max_steps", type=int, default=None,
108121
help="Maximum policy steps to rewind for rewind sampling")
122+
parser.add_argument(
123+
"--history_encoder",
124+
type=str,
125+
default=HISTORY_ENCODER_TEMPORAL_CNN,
126+
choices=HISTORY_ENCODER_CHOICES,
127+
help=(
128+
"History observation encoder variant. temporal_cnn uses actor_history/"
129+
"critic_history with the TemporalCNN model; none uses current-frame "
130+
"actor/critic observations with a single-input MLP."
131+
),
132+
)
109133
parser.add_argument("--device", type=str, default=None)
110134
parser.add_argument(
111135
"--gpu_ids",
@@ -307,6 +331,9 @@ def _configure_experiment_logger(
307331
"rewind_prob": env_cfg.commands["motion"].rewind_prob,
308332
"rewind_min_steps": env_cfg.commands["motion"].rewind_min_steps,
309333
"rewind_max_steps": env_cfg.commands["motion"].rewind_max_steps,
334+
"history_encoder": getattr(
335+
agent_cfg, "history_encoder", HISTORY_ENCODER_TEMPORAL_CNN
336+
),
310337
},
311338
)
312339
swanlab.sync_tensorboard_torch(types=["scalar", "scalars", "image", "text"])
@@ -404,6 +431,7 @@ def _handle_shutdown(signum: int, _frame: Any) -> None:
404431
agent_cfg.max_iterations = args.max_iterations
405432
if args.experiment_name is not None:
406433
agent_cfg.experiment_name = args.experiment_name
434+
apply_history_encoder_config(env_cfg, agent_cfg, args.history_encoder)
407435

408436
device = _resolve_device(args, torch)
409437

0 commit comments

Comments
 (0)