Skip to content

Commit 160c6fa

Browse files
authored
Support actor model averaging for ActorCriticAlgorithm (#1820)
* Support actor model averaging for ActorCriticAlgorithm model averaging is often helpful for evaluation. * Address comments
1 parent f2c844e commit 160c6fa

2 files changed

Lines changed: 53 additions & 9 deletions

File tree

alf/algorithms/actor_critic_algorithm.py

Lines changed: 47 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
"""Actor critic algorithm."""
1515

1616
import torch
17+
from typing import Callable, List
1718

1819
import alf
1920
from alf.algorithms.on_policy_algorithm import OnPolicyAlgorithm
@@ -50,6 +51,8 @@ def __init__(self,
5051
config: TrainerConfig = None,
5152
loss=None,
5253
loss_class=ActorCriticLoss,
54+
actor_avg_fns: List[Callable] = [],
55+
predict_ema_id: int = -1,
5356
optimizer=None,
5457
checkpoint=None,
5558
debug_summaries=False,
@@ -93,6 +96,30 @@ def __init__(self,
9396
None, a default loss of class loss_class will be used.
9497
loss_class (type): the class of the loss. The signature of its
9598
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
96123
optimizer (torch.optim.Optimizer): The optimizer for training
97124
checkpoint (None|str): a string in the format of "prefix@path",
98125
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,
158185

159186
self._register_load_state_dict_pre_hook(_deployment_hook)
160187

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+
161202
def convert_train_state_to_predict_state(self, state):
162203
return state._replace(value=())
163204

164205
def predict_step(self, inputs: TimeStep, state: ActorCriticState):
165206
"""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)
168213

169214
action = dist_utils.epsilon_greedy_sample(action_dist,
170215
self._epsilon_greedy)

alf/utils/model_averager.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -12,18 +12,17 @@
1212
# See the License for the specific language governing permissions and
1313
# limitations under the License.
1414

15-
from numbers import Number
1615
import torch
16+
from typing import SupportsFloat
1717
from torch.optim.swa_utils import AveragedModel as _AveragedModel
18-
from typing import Union
1918

2019
from alf.utils.schedulers import Scheduler
2120

2221

23-
def ema_avg_fn(averaged_model_parameter,
24-
model_parameter,
25-
num_averaged,
26-
ema_rate: Union[Number, Scheduler],
22+
def ema_avg_fn(averaged_model_parameter: torch.Tensor,
23+
model_parameter: torch.Tensor,
24+
num_averaged: int,
25+
ema_rate: SupportsFloat | Scheduler,
2726
starting_average_after=0,
2827
begin_with_simple_average=True):
2928
"""Exponential moving average of model parameters.
@@ -45,7 +44,7 @@ def ema_avg_fn(averaged_model_parameter,
4544
"""
4645
if num_averaged <= starting_average_after:
4746
return model_parameter
48-
if not isinstance(ema_rate, Number):
47+
if not isinstance(ema_rate, SupportsFloat):
4948
assert isinstance(
5049
ema_rate, Scheduler), ("ema_rate must be a number or a Scheduler")
5150
ema_rate = ema_rate()

0 commit comments

Comments
 (0)