|
14 | 14 | """Actor critic algorithm.""" |
15 | 15 |
|
16 | 16 | import torch |
| 17 | +from typing import Callable, List |
17 | 18 |
|
18 | 19 | import alf |
19 | 20 | from alf.algorithms.on_policy_algorithm import OnPolicyAlgorithm |
@@ -50,6 +51,8 @@ def __init__(self, |
50 | 51 | config: TrainerConfig = None, |
51 | 52 | loss=None, |
52 | 53 | loss_class=ActorCriticLoss, |
| 54 | + actor_avg_fns: List[Callable] = [], |
| 55 | + predict_ema_id: int = -1, |
53 | 56 | optimizer=None, |
54 | 57 | checkpoint=None, |
55 | 58 | debug_summaries=False, |
@@ -93,6 +96,30 @@ def __init__(self, |
93 | 96 | None, a default loss of class loss_class will be used. |
94 | 97 | loss_class (type): the class of the loss. The signature of its |
95 | 98 | constructor: ``loss_class(debug_summaries)`` |
| 99 | + actor_avg_fns: a list of functions for doing model averaging for |
| 100 | + the actor network. Each function will be responsible for |
| 101 | + performing one model averaging. The function must take in the |
| 102 | + current value of the `AveragedModel` parameter, the |
| 103 | + current value of `model` parameter, and the number of models |
| 104 | + already averaged. See alf.utils/model_averager.ema_avg_fn for an example. |
| 105 | +
|
| 106 | + Example: |
| 107 | +
|
| 108 | + .. code-block:: python |
| 109 | +
|
| 110 | + actor_avg_fns=[ |
| 111 | + partial(ema_avg_fn, |
| 112 | + ema_rate=1e-3, |
| 113 | + starting_average_after=300_000), |
| 114 | + partial(ema_avg_fn, |
| 115 | + ema_rate=1e-2, |
| 116 | + starting_average_after=300_000), |
| 117 | + ]) |
| 118 | +
|
| 119 | + predict_ema_id: the index of the actor average model to be used |
| 120 | + for prediction. -1 means the original actor network is used. |
| 121 | + 0 means the first averaged model in actor_avg_fns is used, |
| 122 | + and so on |
96 | 123 | optimizer (torch.optim.Optimizer): The optimizer for training |
97 | 124 | checkpoint (None|str): a string in the format of "prefix@path", |
98 | 125 | where the "prefix" is the multi-step path to the contents in the |
@@ -158,13 +185,31 @@ def _deployment_hook(state_dict, prefix: str, unused_loacl_metadata, |
158 | 185 |
|
159 | 186 | self._register_load_state_dict_pre_hook(_deployment_hook) |
160 | 187 |
|
| 188 | + self._actor_emas = torch.nn.ModuleList() |
| 189 | + self._predict_ema_id = predict_ema_id |
| 190 | + for actor_avg_fn in actor_avg_fns: |
| 191 | + self._actor_emas.append( |
| 192 | + alf.utils.model_averager.AveragedModel(self._actor_network, |
| 193 | + avg_fn=actor_avg_fn)) |
| 194 | + |
| 195 | + def after_update(self, root_inputs: TimeStep, info: ActorCriticInfo): |
| 196 | + for actor_ema in self._actor_emas: |
| 197 | + actor_ema.update_parameters(self._actor_network) |
| 198 | + |
| 199 | + def _trainable_attributes_to_ignore(self): |
| 200 | + return ['_actor_emas'] |
| 201 | + |
161 | 202 | def convert_train_state_to_predict_state(self, state): |
162 | 203 | return state._replace(value=()) |
163 | 204 |
|
164 | 205 | def predict_step(self, inputs: TimeStep, state: ActorCriticState): |
165 | 206 | """Predict for one step.""" |
166 | | - action_dist, actor_state = self._actor_network(inputs.observation, |
167 | | - state=state.actor) |
| 207 | + if self._predict_ema_id == -1: |
| 208 | + predict_model = self._actor_network |
| 209 | + else: |
| 210 | + predict_model = self._actor_emas[self._predict_ema_id] |
| 211 | + action_dist, actor_state = predict_model(inputs.observation, |
| 212 | + state=state.actor) |
168 | 213 |
|
169 | 214 | action = dist_utils.epsilon_greedy_sample(action_dist, |
170 | 215 | self._epsilon_greedy) |
|
0 commit comments