diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-21 08:26:40 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-21 08:26:40 -0500 |
| commit | d6a00aab979a50394cd4d5423f7304bb7819ce25 (patch) | |
| tree | d2feb5affd443147b57f3b31c30fd28149f4f722 /sdil/local_baselines.py | |
| parent | c8a372cae2612839bdbdafcf9c114f52cb4a6de3 (diff) | |
fix: match published baseline protocols
Diffstat (limited to 'sdil/local_baselines.py')
| -rw-r--r-- | sdil/local_baselines.py | 46 |
1 files changed, 33 insertions, 13 deletions
diff --git a/sdil/local_baselines.py b/sdil/local_baselines.py index 23b1b5c..dc33fea 100644 --- a/sdil/local_baselines.py +++ b/sdil/local_baselines.py @@ -41,26 +41,41 @@ class FANet(SDILNet): self.B.append((torch.randn(*shape, generator=g) * (b_scale / math.sqrt(shape[1])) ).to(self.W[l].device, self.dtype)) + def fa_hidden_gradients(self, fwd, yoh): + """Return FA estimates of dL/dh for each hidden block output. + + The returned tensors include the residual-branch scale used by the + corresponding local weight update. Identity skips transport ``dh`` + exactly; only the nonlinear branch is transported through fixed B. + Keeping this calculation separate makes the signal directly + measurable without using autograd for learning. + """ + h, u = fwd["h"], fwd["u"] + e = torch.softmax(h[-1], 1) - yoh + dh = e @ self.B[self.L - 1] + estimates = [None] * (self.L - 1) + for l in range(self.L - 2, -1, -1): + branch_scale = self.res_alpha if self.residual and l >= 1 else 1.0 + estimates[l] = branch_scale * dh + du = estimates[l] * self.act_prime(u[l]) + if l > 0: + feedback = du @ self.B[l] + dh = dh + feedback if self.residual and l >= 1 else feedback + return estimates + def fa_step(self, x, y, yoh, eta, momentum=0.0): with torch.no_grad(): fwd = self.forward(x) h, u = fwd["h"], fwd["u"] e = torch.softmax(h[-1], 1) - yoh # grad wrt logits loss = F.cross_entropy(h[-1], y).item() + hidden_grad = self.fa_hidden_gradients(fwd, yoh) # output layer (exact local) self._upd(self.L - 1, e, h[self.L - 1], eta, momentum) - # dh is the error wrt the hidden *state*. On a residual block, - # h'=h+alpha*phi(Wh), the identity branch transports dh exactly and - # only the nonlinear branch uses a fixed random feedback matrix. - dh = e @ self.B[self.L - 1] - for l in range(self.L - 2, -1, -1): - branch_scale = self.res_alpha if self.residual and l >= 1 else 1.0 - du = branch_scale * dh * self.act_prime(u[l]) + for l in range(self.L - 1): + du = hidden_grad[l] * self.act_prime(u[l]) self._upd(l, du, h[l], eta, momentum) - if l > 0: - feedback = du @ self.B[l] - dh = dh + feedback if self.residual and l >= 1 else feedback return loss def _upd(self, l, delta, pre, eta, momentum): @@ -321,9 +336,13 @@ class EPNet: s = new return s - def train_step(self, x, y, yoh, eta): + def train_step(self, x, y, yoh, eta, free_state=None, return_free_state=False): with torch.no_grad(): - s0 = self._settle(x, beta=0.0, T=self.T_free) # free phase + s0 = self._settle(x, beta=0.0, s=free_state, T=self.T_free) # free phase + # The authors keep one free-phase particle per training example + # and reuse it on the next epoch. Save the free state, not the + # subsequently nudged state. + persistent_state = [t.detach().clone() for t in s0] beta = self.beta if self.random_beta_sign and torch.randint(0, 2, ()).item() == 0: beta = -beta @@ -338,7 +357,8 @@ class EPNet: self.W[k] += layer_eta * dW self.b[k] += layer_eta * db # free-phase output as prediction proxy for loss logging - return F.mse_loss(s0[-1], yoh).item() + loss = F.mse_loss(s0[-1], yoh).item() + return (loss, persistent_state) if return_free_state else loss @torch.no_grad() def predict(self, x): |
