From 00f605522dab2efcdf1690d991f0dd054ca9c30c Mon Sep 17 00:00:00 2001 From: Benjamin-Walker Date: Wed, 15 Oct 2025 11:54:45 +0100 Subject: [PATCH] Fixed dplr to stop cross terms --- models/slcde.py | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/models/slcde.py b/models/slcde.py index 92297b4..bd0ce7d 100644 --- a/models/slcde.py +++ b/models/slcde.py @@ -178,14 +178,16 @@ def forward(self, X: torch.Tensor) -> torch.Tensor: if self.diagonal: if self.rank > 0: diag = (inp[:, i] @ self.vf_A) * y - U = (inp[:, i] @ self.vf_A_u).reshape( - -1, self.hidden_dim, self.rank - ) - V = (inp[:, i] @ self.vf_A_v).reshape( - -1, self.hidden_dim, self.rank - ) - z = torch.bmm(V.transpose(1, 2), y.unsqueeze(-1)).squeeze(-1) - state_transition = diag + torch.bmm(U, z.unsqueeze(-1)).squeeze(-1) + U = self.vf_A_u.view( + self.input_dim + 1, self.hidden_dim, self.rank + ) # (C, H, R) + V = self.vf_A_v.view( + self.input_dim + 1, self.hidden_dim, self.rank + ) # (C, H, R) + vTy = (V.unsqueeze(0) * y.unsqueeze(1).unsqueeze(-1)).sum(dim=2) + vTy = inp[:, i].unsqueeze(-1) * vTy + lowrank = (U.unsqueeze(0) * vTy.unsqueeze(2)).sum(dim=(1, 3)) + state_transition = diag + lowrank elif self.fwht: state_transition = (inp[:, i] @ torch.tanh(self.vf_A)) * y state_transition = hadamard_transform(