| 152 | class MemoryUnit(nn.Module): |
| 153 | # clusters_k is k keys |
| 154 | def __init__(self, input_size, output_size, emb_size, clusters_k=10): |
| 155 | super(MemoryUnit, self).__init__() |
| 156 | self.clusters_k = clusters_k |
| 157 | self.input_size = input_size |
| 158 | self.output_size = output_size |
| 159 | self.array = nn.Parameter(init.xavier_uniform_(torch.FloatTensor(self.clusters_k, input_size*output_size))) |
| 160 | self.index = nn.Parameter(init.xavier_uniform_(torch.FloatTensor(self.clusters_k, emb_size))) |
| 161 | self.softmax = nn.Softmax(dim=-1) |
| 162 | |
| 163 | def forward(self, bias_emb): |
| 164 | """ |