Download code/0810_dyanmics_model/model.py from chomeed/threading-d0-dynamics-and-value: direct link, hf CLI and curl.
- Browser
- Download file 2.24 kB
-
https://huggingface.co/chomeed/threading-d0-dynamics-and-value/resolve/main/code/0810_dyanmics_model/model.py
- Command line
-
hf download hf://chomeed/threading-d0-dynamics-and-value/code/0810_dyanmics_model/model.py
-
curl -L -o model.py https://huggingface.co/chomeed/threading-d0-dynamics-and-value/resolve/main/code/0810_dyanmics_model/model.py
2.24 kB
| ''' | |
| 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" | |