MCPcopy Create free account
hub / github.com/WarmCongee/SDUMC / forward

Method forward

toolkit/models/mfm.py:65–84  ·  view source on GitHub ↗
(self, hT, t)

Source from the content-addressed store, hash-verified

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
87class MFM(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected