Coverage for hopwise/model/context_aware_recommender/autoint.py: 95%
59 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2020/09/01
2# @Author : Shuqing Bian
3# @Email : shuqingbian@gmail.com
4# @File : autoint.py
6r"""AutoInt
7################################################
8Reference:
9 Weiping Song et al. "AutoInt: Automatic Feature Interaction Learning via Self-Attentive Neural Networks"
10 in CIKM 2018.
11"""
13import torch
14import torch.nn.functional as F
15from torch import nn
16from torch.nn.init import constant_, xavier_normal_
18from hopwise.model.abstract_recommender import ContextRecommender
19from hopwise.model.layers import MLPLayers
22class AutoInt(ContextRecommender):
23 """AutoInt is a novel CTR prediction model based on self-attention mechanism,
24 which can automatically learn high-order feature interactions in an explicit fashion.
26 """
28 def __init__(self, config, dataset):
29 super().__init__(config, dataset)
31 # load parameters info
32 self.attention_size = config["attention_size"]
33 self.dropout_probs = config["dropout_probs"]
34 self.n_layers = config["n_layers"]
35 self.num_heads = config["num_heads"]
36 self.mlp_hidden_size = config["mlp_hidden_size"]
37 self.has_residual = config["has_residual"]
39 # define layers and loss
40 self.att_embedding = nn.Linear(self.embedding_size, self.attention_size)
41 self.embed_output_dim = self.num_feature_field * self.embedding_size
42 self.atten_output_dim = self.num_feature_field * self.attention_size
43 size_list = [self.embed_output_dim] + self.mlp_hidden_size
44 self.mlp_layers = MLPLayers(size_list, dropout=self.dropout_probs[1])
45 # multi-head self-attention network
46 self.self_attns = nn.ModuleList(
47 [
48 nn.MultiheadAttention(self.attention_size, self.num_heads, dropout=self.dropout_probs[0])
49 for _ in range(self.n_layers)
50 ]
51 )
52 self.attn_fc = torch.nn.Linear(self.atten_output_dim, 1)
53 self.deep_predict_layer = nn.Linear(self.mlp_hidden_size[-1], 1)
54 if self.has_residual:
55 self.v_res_embedding = torch.nn.Linear(self.embedding_size, self.attention_size)
57 self.dropout_layer = nn.Dropout(p=self.dropout_probs[2])
58 self.sigmoid = nn.Sigmoid()
59 self.loss = nn.BCEWithLogitsLoss()
61 # parameters initialization
62 self.apply(self._init_weights)
64 def _init_weights(self, module):
65 if isinstance(module, nn.Embedding):
66 xavier_normal_(module.weight.data)
67 elif isinstance(module, nn.Linear):
68 xavier_normal_(module.weight.data)
69 if module.bias is not None:
70 constant_(module.bias.data, 0)
72 def autoint_layer(self, infeature):
73 """Get the attention-based feature interaction score
75 Args:
76 infeature (torch.FloatTensor): input feature embedding tensor. shape of[batch_size,field_size,embed_dim].
78 Returns:
79 torch.FloatTensor: Result of score. shape of [batch_size,1] .
80 """
81 att_infeature = self.att_embedding(infeature)
82 cross_term = att_infeature.transpose(0, 1)
83 for self_attn in self.self_attns:
84 cross_term, _ = self_attn(cross_term, cross_term, cross_term)
85 cross_term = cross_term.transpose(0, 1)
86 # Residual connection
87 if self.has_residual:
88 v_res = self.v_res_embedding(infeature)
89 cross_term += v_res
90 # Interacting layer
91 cross_term = F.relu(cross_term).contiguous().view(-1, self.atten_output_dim)
92 batch_size = infeature.shape[0]
93 att_output = self.attn_fc(cross_term) + self.deep_predict_layer(
94 self.mlp_layers(infeature.view(batch_size, -1))
95 )
96 return att_output
98 def forward(self, interaction):
99 autoint_all_embeddings = self.concat_embed_input_fields(interaction) # [batch_size, num_field, embed_dim]
100 output = self.first_order_linear(interaction) + self.autoint_layer(autoint_all_embeddings)
101 return output.squeeze(1)
103 def calculate_loss(self, interaction):
104 label = interaction[self.LABEL]
105 output = self.forward(interaction)
106 return self.loss(output, label)
108 def predict(self, interaction):
109 return self.sigmoid(self.forward(interaction))