chomeed's picture
add svg_term.py + dynamics model modules
7738b02 verified
Raw History Blame Contribute Delete
2.73 kB
'''
Dynamics Model -- ABSOLUTE control for model_residual.
input: 66D state + 56D action (8 step * 7D action)
output: 66D state (absolute: predicts s' directly, NOT s + delta)
Built 2026-08-30 to test one claim. Every state dynamics model in this directory predicts a DELTA,
and model_flow.py:31 gives the reason: "the state barely moves over one 8-step chunk, so absolute
prediction is dominated by copying s". That was asserted, never measured against an absolute arm.
It is worth measuring because the DINO-WM side found the OPPOSITE: on patch features at the same
8-step horizon, absolute beats residual at 67 of 69 matched steps, and a per-window motion sweep
put the crossover at copy_mse ~0.83 -- residual wins only the quietest ~25% of windows. If motion
is what decides the parameterization, then the state, which barely moves, should sit on the far
side of that line and residual should win here by a clear margin.
Everything except the parameterisation is copied from model_residual: same MLP, same hidden width,
same normalisation buffers, same Huber loss on NORMALISED targets. The only difference is what the
network's output means.
'''
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))
self.fitted.fill_(True)
return self
def _norm(self, s):
return (s - self.mu) / self.sigma
def forward(self, s, a):
# de-normalise the predicted next state; NO `s +`, that is the whole point
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"
# target is the normalised ABSOLUTE next state, where model_residual uses (s_next - s)/sigma
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)
m = Dynamics().fit_norm(s)
s_next = s + torch.randn_like(s) * 0.01
print("loss", float(m.loss(s, a, s_next)), "out", tuple(m(s, a).shape))