| 125 | |
| 126 | |
| 127 | class Caser(nn.Module): |
| 128 | def __init__(self, hidden_size, item_num, state_size, num_filters, filter_sizes, |
| 129 | dropout_rate): |
| 130 | super(Caser, self).__init__() |
| 131 | self.hidden_size = hidden_size |
| 132 | self.item_num = int(item_num) |
| 133 | self.state_size = state_size |
| 134 | self.filter_sizes = eval(filter_sizes) |
| 135 | self.num_filters = num_filters |
| 136 | self.dropout_rate = dropout_rate |
| 137 | self.item_embeddings = nn.Embedding( |
| 138 | num_embeddings=item_num + 1, |
| 139 | embedding_dim=self.hidden_size, |
| 140 | ) |
| 141 | |
| 142 | # init embedding |
| 143 | nn.init.normal_(self.item_embeddings.weight, 0, 0.01) |
| 144 | |
| 145 | # Horizontal Convolutional Layers |
| 146 | self.horizontal_cnn = nn.ModuleList( |
| 147 | [nn.Conv2d(1, self.num_filters, (i, self.hidden_size)) for i in self.filter_sizes]) |
| 148 | # Initialize weights and biases |
| 149 | for cnn in self.horizontal_cnn: |
| 150 | nn.init.xavier_normal_(cnn.weight) |
| 151 | nn.init.constant_(cnn.bias, 0.1) |
| 152 | |
| 153 | # Vertical Convolutional Layer |
| 154 | self.vertical_cnn = nn.Conv2d(1, 1, (self.state_size, 1)) |
| 155 | nn.init.xavier_normal_(self.vertical_cnn.weight) |
| 156 | nn.init.constant_(self.vertical_cnn.bias, 0.1) |
| 157 | |
| 158 | # Fully Connected Layer |
| 159 | self.num_filters_total = self.num_filters * len(self.filter_sizes) |
| 160 | final_dim = self.hidden_size + self.num_filters_total |
| 161 | self.s_fc = nn.Linear(final_dim, item_num) |
| 162 | |
| 163 | # dropout |
| 164 | self.dropout = nn.Dropout(self.dropout_rate) |
| 165 | |
| 166 | def forward(self, states, len_states): |
| 167 | input_emb = self.item_embeddings(states) |
| 168 | mask = torch.ne(states, self.item_num).float().unsqueeze(-1) |
| 169 | input_emb *= mask |
| 170 | input_emb = input_emb.unsqueeze(1) |
| 171 | pooled_outputs = [] |
| 172 | for cnn in self.horizontal_cnn: |
| 173 | h_out = nn.functional.relu(cnn(input_emb)) |
| 174 | h_out = h_out.squeeze() |
| 175 | p_out = nn.functional.max_pool1d(h_out, h_out.shape[2]) |
| 176 | pooled_outputs.append(p_out) |
| 177 | |
| 178 | h_pool = torch.cat(pooled_outputs, 1) |
| 179 | h_pool_flat = h_pool.view(-1, self.num_filters_total) |
| 180 | |
| 181 | v_out = nn.functional.relu(self.vertical_cnn(input_emb)) |
| 182 | v_flat = v_out.view(-1, self.hidden_size) |
| 183 | |
| 184 | out = torch.cat([h_pool_flat, v_flat], 1) |