MCPcopy Create free account
hub / github.com/MotrixLab/insactor / train

Function train

diffmimic/brax_lib/agent_diffmimic.py:69–377  ·  view source on GitHub ↗

Direct trajectory optimization training.

(environment: envs.Env,
          episode_length: int,
          action_repeat: int = 1,
          num_envs: int = 1,
          max_devices_per_host: Optional[int] = None,
          num_eval_envs: int = 128,
          learning_rate: float = 1e-4,
          seed: int = 0,
          truncation_length: Optional[int] = None,
          max_gradient_norm: float = 1e9,
          num_evals: int = 1,
          normalize_observations: bool = False,
          deterministic_eval: bool = False,
          network_factory: types.NetworkFactory[
              apg_networks.APGNetworks] = apg_networks.make_apg_networks,
          progress_fn: Callable[[int, Metrics], None] = lambda *args: None,
          eval_environment: Optional[envs.Env] = None,
          eval_episode_length: Optional[int] = None,
          save_dir: Optional[str] = None,
          use_linear_scheduler: Optional[bool] = False,
          train_loader: Optional[data.DataLoader] = None,
          test_loader: Optional[data.DataLoader] = None,
          latent_dim: int=64,
          beta: float=0.01,
          pretrain=None,
          deterministic=False,
          large=False,
          conditional=False,
          weight_decay=0.,
          skip_encoder=False
          )

Source from the content-addressed store, hash-verified

67
68
69def train(environment: envs.Env,
70 episode_length: int,
71 action_repeat: int = 1,
72 num_envs: int = 1,
73 max_devices_per_host: Optional[int] = None,
74 num_eval_envs: int = 128,
75 learning_rate: float = 1e-4,
76 seed: int = 0,
77 truncation_length: Optional[int] = None,
78 max_gradient_norm: float = 1e9,
79 num_evals: int = 1,
80 normalize_observations: bool = False,
81 deterministic_eval: bool = False,
82 network_factory: types.NetworkFactory[
83 apg_networks.APGNetworks] = apg_networks.make_apg_networks,
84 progress_fn: Callable[[int, Metrics], None] = lambda *args: None,
85 eval_environment: Optional[envs.Env] = None,
86 eval_episode_length: Optional[int] = None,
87 save_dir: Optional[str] = None,
88 use_linear_scheduler: Optional[bool] = False,
89 train_loader: Optional[data.DataLoader] = None,
90 test_loader: Optional[data.DataLoader] = None,
91 latent_dim: int=64,
92 beta: float=0.01,
93 pretrain=None,
94 deterministic=False,
95 large=False,
96 conditional=False,
97 weight_decay=0.,
98 skip_encoder=False
99 ):
100 """Direct trajectory optimization training."""
101 # best_pose_error = 1e8
102 best_reward = -1e8
103
104 xt = time.time()
105
106 process_count = jax.process_count()
107 process_id = jax.process_index()
108 local_device_count = jax.local_device_count()
109 local_devices_to_use = local_device_count
110 if max_devices_per_host:
111 local_devices_to_use = min(local_devices_to_use, max_devices_per_host)
112 logging.info(
113 'Device count: %d, process count: %d (id %d), local device count: %d, '
114 'devices to be used count: %d', jax.device_count(), process_count,
115 process_id, local_device_count, local_devices_to_use)
116 device_count = local_devices_to_use * process_count
117
118 if truncation_length is not None:
119 assert truncation_length > 0
120
121 num_evals_after_init = max(num_evals - 1, 1)
122
123 assert num_envs % device_count == 0
124 env = environment
125 env = wrappers.EpisodeWrapper(env, episode_length, action_repeat)
126 env = wrappers.VmapWrapper(env)

Callers

nothing calls this directly

Calls 6

run_evaluationMethod · 0.95
serialize_qpFunction · 0.90
TrainingStateClass · 0.85
_unpmapFunction · 0.85
_data_pmapFunction · 0.85

Tested by

no test coverage detected