(self, hT, t)
| 63 | self.h = h |
| 64 | |
| 65 | def forward(self, hT, t): |
| 66 | # x is n x d |
| 67 | n = hT.shape[0] |
| 68 | h = hT.shape[1] |
| 69 | self.hx = torch.zeros(n, self.h).cuda() |
| 70 | self.cx = torch.zeros(n, self.h).cuda() |
| 71 | final_hs = [] |
| 72 | all_hs = [] |
| 73 | all_cs = [] |
| 74 | for i in range(t): |
| 75 | if i == 0: |
| 76 | self.hx, self.cx = self.lstm(hT, (self.hx, self.cx)) |
| 77 | else: |
| 78 | self.hx, self.cx = self.lstm(all_hs[-1], (self.hx, self.cx)) |
| 79 | all_hs.append(self.hx) |
| 80 | all_cs.append(self.cx) |
| 81 | final_hs.append(self.hx.view(1,n,h)) |
| 82 | final_hs = torch.cat(final_hs, dim=0) |
| 83 | all_recons = self.fc1(final_hs) |
| 84 | return all_recons |
| 85 | |
| 86 | |
| 87 | class MFM(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected