MCPcopy Create free account
hub / github.com/AkaliKong/MiniOneRec / Caser

Class Caser

sasrec.py:127–212  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

125
126
127class 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)

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected