| import numpy as np |
| import torch |
| import torch.nn as nn |
|
|
|
|
| class PointWiseFeedForward(nn.Module): |
| def __init__(self, hidden_units, dropout_rate): |
| super().__init__() |
| |
| self.conv1 = nn.Conv1d( |
| in_channels=hidden_units, out_channels=hidden_units, kernel_size=1 |
| ) |
| self.dropout1 = nn.Dropout(p=dropout_rate) |
| self.relu = nn.ReLU() |
| self.conv2 = nn.Conv1d( |
| in_channels=hidden_units, out_channels=hidden_units, kernel_size=1 |
| ) |
| self.dropout2 = nn.Dropout(p=dropout_rate) |
|
|
| def forward(self, inputs): |
| |
| |
| x = inputs.transpose(-1, -2) |
| x = self.conv1(x) |
| x = self.dropout1(x) |
| x = self.relu(x) |
| x = self.conv2(x) |
| x = self.dropout2(x) |
| x = x.transpose(-1, -2) |
| outputs = x + inputs |
| return outputs |
|
|
|
|
| class SASRec(torch.nn.Module): |
| def __init__(self, user_num, item_num, args): |
| super().__init__() |
|
|
| self.user_num = user_num |
| self.item_num = item_num |
| self.dev = args.device |
|
|
| self.norm_first = args.norm_first |
|
|
| self.item_emb = torch.nn.Embedding( |
| self.item_num + 1, args.hidden_units, padding_idx=0 |
| ) |
| self.pos_emb = torch.nn.Embedding( |
| args.maxlen + 1, args.hidden_units, padding_idx=0 |
| ) |
| self.emb_dropout = torch.nn.Dropout(p=args.dropout_rate) |
|
|
| self.attention_layernorms = torch.nn.ModuleList() |
| self.attention_layers = torch.nn.ModuleList() |
| self.forward_layernorms = torch.nn.ModuleList() |
| self.forward_layers = torch.nn.ModuleList() |
| self.last_layernorm = torch.nn.LayerNorm(args.hidden_units, eps=1e-8) |
|
|
| for _ in range(args.num_blocks): |
| self.attention_layernorms.append( |
| torch.nn.LayerNorm(args.hidden_units, eps=1e-8) |
| ) |
| self.attention_layers.append( |
| torch.nn.MultiheadAttention( |
| args.hidden_units, args.num_heads, dropout=args.dropout_rate |
| ) |
| ) |
| self.forward_layernorms.append( |
| torch.nn.LayerNorm(args.hidden_units, eps=1e-8) |
| ) |
| self.forward_layers.append( |
| PointWiseFeedForward(args.hidden_units, args.dropout_rate) |
| ) |
|
|
| |
| |
|
|
| def log2feats(self, log_seqs): |
|
|
| seqs = self.item_emb(torch.LongTensor(log_seqs).to(self.dev)) |
| seqs *= self.item_emb.embedding_dim**0.5 |
|
|
| positions = np.tile(np.arange(1, log_seqs.shape[1] + 1), [log_seqs.shape[0], 1]) |
| positions *= log_seqs != 0 |
|
|
| seqs += self.pos_emb(torch.LongTensor(positions).to(self.dev)) |
| seqs = self.emb_dropout(seqs) |
|
|
| time_len = seqs.shape[1] |
| attention_mask = ~torch.tril( |
| torch.ones((time_len, time_len), dtype=torch.bool, device=self.dev) |
| ) |
|
|
| for i in range(len(self.attention_layers)): |
| seqs = torch.transpose(seqs, 0, 1) |
| if self.norm_first: |
| x = self.attention_layernorms[i](seqs) |
| mha_outputs, _ = self.attention_layers[i]( |
| x, x, x, attn_mask=attention_mask |
| ) |
| seqs = seqs + mha_outputs |
| seqs = torch.transpose(seqs, 0, 1) |
| seqs = seqs + self.forward_layers[i](self.forward_layernorms[i](seqs)) |
| else: |
| mha_outputs, _ = self.attention_layers[i]( |
| seqs, seqs, seqs, attn_mask=attention_mask |
| ) |
| seqs = self.attention_layernorms[i](seqs + mha_outputs) |
| seqs = torch.transpose(seqs, 0, 1) |
| seqs = self.forward_layernorms[i](seqs + self.forward_layers[i](seqs)) |
|
|
| log_feats = self.last_layernorm(seqs) |
|
|
| return log_feats |
|
|
| def forward(self, user_ids, log_seqs, pos_seqs, neg_seqs): |
| log_feats = self.log2feats(log_seqs) |
|
|
| pos_embs = self.item_emb(torch.LongTensor(pos_seqs).to(self.dev)) |
| neg_embs = self.item_emb(torch.LongTensor(neg_seqs).to(self.dev)) |
|
|
| pos_logits = (log_feats * pos_embs).sum(dim=-1) |
| neg_logits = (log_feats * neg_embs).sum(dim=-1) |
|
|
| |
| |
|
|
| return pos_logits, neg_logits |
|
|
| def predict(self, user_ids, log_seqs, item_indices): |
| log_feats = self.log2feats(log_seqs) |
|
|
| final_feat = log_feats[:, -1, :] |
|
|
| item_embs = self.item_emb( |
| torch.LongTensor(item_indices).to(self.dev) |
| ) |
|
|
| logits = item_embs.matmul(final_feat.unsqueeze(-1)).squeeze(-1) |
|
|
| |
|
|
| return logits |
|
|