| 113 | self.embed_fc = nn.Linear(300, self.hiddendim) |
| 114 | |
| 115 | def get_initial_state(self, embed, tile_times=1): |
| 116 | assert embed.shape[1] == 300 |
| 117 | state = self.embed_fc(embed) # N * sDim |
| 118 | if tile_times != 1: |
| 119 | state = state.unsqueeze(1) |
| 120 | trans_state = state.transpose(0, 1) |
| 121 | state = trans_state.tile([tile_times, 1, 1]) |
| 122 | trans_state = state.transpose(0, 1) |
| 123 | state = trans_state.reshape(-1, self.hiddendim) |
| 124 | state = state.unsqueeze(0) # 1 * N * sDim |
| 125 | return state |
| 126 | |
| 127 | def forward(self, feat, data=None): |
| 128 | # b,25,512 |