''' Dynamics Model input: 66D state + 56D action (8 step * 7D action) output: 66D state ''' import torch import torch.nn as nn import torch.nn.functional as F class Dynamics(nn.Module): def __init__(self, state_dim=66, chunk_dim=56, hidden=512): super().__init__() self.net = nn.Sequential( nn.Linear(state_dim + chunk_dim, hidden), nn.ReLU(), nn.Linear(hidden, hidden), nn.ReLU(), nn.Linear(hidden, state_dim), ) self.register_buffer("mu", torch.zeros(state_dim)) self.register_buffer("sigma", torch.ones(state_dim)) self.register_buffer("fitted", torch.zeros((), dtype=torch.bool)) def fit_norm(self, states): self.mu.copy_(states.mean(0)) self.sigma.copy_(states.std(0).clamp_min(1e-2)) # floor: near-constant dims blow up self.fitted.fill_(True) return self def _norm(self, s): return (s - self.mu) / self.sigma def forward(self, s, a): return self.net(torch.cat((self._norm(s), a), -1)) * self.sigma + self.mu def loss(self, s, a, s_next): assert self.fitted, "call fit_norm(states) first" # both sides normalised, so no dim dominates by unit alone return F.huber_loss(self.net(torch.cat((self._norm(s), a), -1)), self._norm(s_next)) if __name__ == "__main__": torch.manual_seed(0) s, a = torch.randn(2048, 66) * 0.3, torch.randn(2048, 56) s_next = s + (a @ torch.randn(56, 66)) * 0.002 m = Dynamics(hidden=64).fit_norm(s) opt = torch.optim.Adam(m.parameters(), lr=3e-3) for _ in range(300): i = torch.randint(0, len(s), (256,)) opt.zero_grad(); m.loss(s[i], a[i], s_next[i]).backward(); opt.step() print(f"loss {m.loss(s, a, s_next).item():.4f}", end="\r", flush=True) with torch.no_grad(): pred = m(s, a) err = (pred - s_next).norm(dim=-1).mean() # vs predicting the dataset mean -- the trivial baseline for ABSOLUTE prediction print(f"err/mean-baseline {err / (s_next - m.mu).norm(dim=-1).mean():.3f} " f"err/identity {err / (s - s_next).norm(dim=-1).mean():.3f}") assert err < 0.5 * (s_next - m.mu).norm(dim=-1).mean(), "did not learn the mapping"