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
)
| 67 | |
| 68 | |
| 69 | def 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) |
nothing calls this directly
no test coverage detected