1212import numpy as np
1313import 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+ )
1622from train_mimic .data .dataset_lib import PRECOMPUTED_MOTION_VERSION
1723from train_mimic .scripts import train
1824from 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+
253310def 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
0 commit comments