sasrec-pytorch / model.py
MongrelIntruder's picture
Upload model.py with huggingface_hub
ff33497 verified
Raw
History Blame Contribute Delete
5.05 kB
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)