File size: 5,048 Bytes
ff33497 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 | import numpy as np
import torch
import torch.nn as nn
class PointWiseFeedForward(nn.Module):
def __init__(self, hidden_units, dropout_rate):
super().__init__()
# 1x1 Convolution
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):
# inputs shape: (batch size, length, channels)
# conv1d expected shape: (batch size, channels, length)
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)
)
# self.pos_sigmoid = torch.nn.Sigmoid()
# self.neg_sigmoid = torch.nn.Sigmoid()
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) # (U, T, C) -> (U, -1, C)
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)
# pos_pred = self.pos_sigmoid(pos_logits)
# neg_pred = self.neg_sigmoid(neg_logits)
return pos_logits, neg_logits # pos_pred, neg_pred
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)
) # (U, I, C)
logits = item_embs.matmul(final_feat.unsqueeze(-1)).squeeze(-1)
# preds = self.pos_sigmoid(logits) # rank same item list for different users
return logits # (U, I)
|