Port TRPO, PPO, GAIL
This commit is contained in:
162
scratch/etienne/trpo/core/reparam_module.py
Normal file
162
scratch/etienne/trpo/core/reparam_module.py
Normal file
@@ -0,0 +1,162 @@
|
||||
# Source: https://github.com/SsnL/PyTorch-Reparam-Module
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import warnings
|
||||
import types
|
||||
from collections import namedtuple
|
||||
from contextlib import contextmanager
|
||||
|
||||
class ReparamModule(nn.Module):
|
||||
def __init__(self, module):
|
||||
super(ReparamModule, self).__init__()
|
||||
self.module = module
|
||||
|
||||
param_infos = []
|
||||
shared_param_memo = {}
|
||||
shared_param_infos = []
|
||||
params = []
|
||||
param_numels = []
|
||||
param_shapes = []
|
||||
for m in self.modules():
|
||||
for n, p in m.named_parameters(recurse=False):
|
||||
if p is not None:
|
||||
if p in shared_param_memo:
|
||||
shared_m, shared_n = shared_param_memo[p]
|
||||
shared_param_infos.append((m, n, shared_m, shared_n))
|
||||
else:
|
||||
shared_param_memo[p] = (m, n)
|
||||
param_infos.append((m, n))
|
||||
params.append(p.detach())
|
||||
param_numels.append(p.numel())
|
||||
param_shapes.append(p.size())
|
||||
|
||||
assert len(set(p.dtype for p in params)) <= 1, \
|
||||
"expects all parameters in module to have same dtype"
|
||||
|
||||
# store the info for unflatten
|
||||
self._param_infos = tuple(param_infos)
|
||||
self._shared_param_infos = tuple(shared_param_infos)
|
||||
self._param_numels = tuple(param_numels)
|
||||
self._param_shapes = tuple(param_shapes)
|
||||
|
||||
# flatten
|
||||
flat_param = nn.Parameter(torch.cat([p.reshape(-1) for p in params], 0))
|
||||
self.register_parameter('flat_param', flat_param)
|
||||
self.param_numel = flat_param.numel()
|
||||
del params
|
||||
del shared_param_memo
|
||||
|
||||
# deregister the names as parameters
|
||||
for m, n in self._param_infos:
|
||||
delattr(m, n)
|
||||
for m, n, _, _ in self._shared_param_infos:
|
||||
delattr(m, n)
|
||||
|
||||
# register the views as plain attributes
|
||||
self._unflatten_param(self.flat_param)
|
||||
|
||||
# now buffers
|
||||
# they are not reparametrized. just store info as (module, name, buffer)
|
||||
buffer_infos = []
|
||||
for m in self.modules():
|
||||
for n, b in m.named_buffers(recurse=False):
|
||||
if b is not None:
|
||||
buffer_infos.append((m, n, b))
|
||||
|
||||
self._buffer_infos = tuple(buffer_infos)
|
||||
self._traced_self = None
|
||||
|
||||
def trace(self, example_input, **trace_kwargs):
|
||||
assert self._traced_self is None, 'This ReparamModule is already traced'
|
||||
|
||||
if isinstance(example_input, torch.Tensor):
|
||||
example_input = (example_input,)
|
||||
example_input = tuple(example_input)
|
||||
example_param = (self.flat_param.detach().clone(),)
|
||||
example_buffers = (tuple(b.detach().clone() for _, _, b in self._buffer_infos),)
|
||||
|
||||
self._traced_self = torch.jit.trace_module(
|
||||
self,
|
||||
inputs=dict(
|
||||
_forward_with_param=example_param + example_input,
|
||||
_forward_with_param_and_buffers=example_param + example_buffers + example_input,
|
||||
),
|
||||
**trace_kwargs,
|
||||
)
|
||||
|
||||
# replace forwards with traced versions
|
||||
self._forward_with_param = self._traced_self._forward_with_param
|
||||
self._forward_with_param_and_buffers = self._traced_self._forward_with_param_and_buffers
|
||||
return self
|
||||
|
||||
def clear_views(self):
|
||||
for m, n in self._param_infos:
|
||||
setattr(m, n, None) # This will set as plain attr
|
||||
|
||||
def _apply(self, *args, **kwargs):
|
||||
if self._traced_self is not None:
|
||||
self._traced_self._apply(*args, **kwargs)
|
||||
return self
|
||||
return super(ReparamModule, self)._apply(*args, **kwargs)
|
||||
|
||||
def _unflatten_param(self, flat_param):
|
||||
ps = (t.view(s) for (t, s) in zip(flat_param.split(self._param_numels), self._param_shapes))
|
||||
for (m, n), p in zip(self._param_infos, ps):
|
||||
setattr(m, n, p) # This will set as plain attr
|
||||
for (m, n, shared_m, shared_n) in self._shared_param_infos:
|
||||
setattr(m, n, getattr(shared_m, shared_n))
|
||||
|
||||
@contextmanager
|
||||
def unflattened_param(self, flat_param):
|
||||
saved_views = [getattr(m, n) for m, n in self._param_infos]
|
||||
self._unflatten_param(flat_param)
|
||||
yield
|
||||
# Why not just `self._unflatten_param(self.flat_param)`?
|
||||
# 1. because of https://github.com/pytorch/pytorch/issues/17583
|
||||
# 2. slightly faster since it does not require reconstruct the split+view
|
||||
# graph
|
||||
for (m, n), p in zip(self._param_infos, saved_views):
|
||||
setattr(m, n, p)
|
||||
for (m, n, shared_m, shared_n) in self._shared_param_infos:
|
||||
setattr(m, n, getattr(shared_m, shared_n))
|
||||
|
||||
@contextmanager
|
||||
def replaced_buffers(self, buffers):
|
||||
for (m, n, _), new_b in zip(self._buffer_infos, buffers):
|
||||
setattr(m, n, new_b)
|
||||
yield
|
||||
for m, n, old_b in self._buffer_infos:
|
||||
setattr(m, n, old_b)
|
||||
|
||||
def _forward_with_param_and_buffers(self, flat_param, buffers, *inputs, **kwinputs):
|
||||
with self.unflattened_param(flat_param):
|
||||
with self.replaced_buffers(buffers):
|
||||
return self.module(*inputs, **kwinputs)
|
||||
|
||||
def _forward_with_param(self, flat_param, *inputs, **kwinputs):
|
||||
with self.unflattened_param(flat_param):
|
||||
return self.module(*inputs, **kwinputs)
|
||||
|
||||
def forward(self, *inputs, flat_param=None, buffers=None, **kwinputs):
|
||||
if flat_param is None:
|
||||
flat_param = self.flat_param
|
||||
if buffers is None:
|
||||
return self._forward_with_param(flat_param, *inputs, **kwinputs)
|
||||
else:
|
||||
return self._forward_with_param_and_buffers(flat_param, tuple(buffers), *inputs, **kwinputs)
|
||||
|
||||
|
||||
class ReparamPolicy(ReparamModule):
|
||||
|
||||
def sample(self, *args, **kwargs):
|
||||
return self.module.sample(*args, **kwargs)
|
||||
|
||||
def log_prob(self, *args, **kwargs):
|
||||
return self.module.log_prob(*args, **kwargs)
|
||||
|
||||
def kl_divergence(self, *args, **kwargs):
|
||||
return self.module.kl_divergence(*args, **kwargs)
|
||||
|
||||
def predict(self, *args, **kwargs):
|
||||
return self.module.predict(*args, **kwargs)
|
||||
Reference in New Issue
Block a user