Coverage for hopwise/model/context_aware_recommender/afm.py: 98%
56 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/7/21
2# @Author : Zihan Lin
3# @Email : linzihan.super@foxmail.com
4# @File : afm.py
6r"""AFM
7################################################
8Reference:
9 Jun Xiao et al. "Attentional Factorization Machines: Learning the Weight of Feature Interactions via
10 Attention Networks" in IJCAI 2017.
11"""
13import torch
14from torch import nn
15from torch.nn.init import constant_, xavier_normal_
17from hopwise.model.abstract_recommender import ContextRecommender
18from hopwise.model.layers import AttLayer
21class AFM(ContextRecommender):
22 """AFM is a attention based FM model that predict the final score with the attention of input feature."""
24 def __init__(self, config, dataset):
25 super().__init__(config, dataset)
27 # load parameters info
28 self.attention_size = config["attention_size"]
29 self.dropout_prob = config["dropout_prob"]
30 self.reg_weight = config["reg_weight"]
31 self.num_pair = self.num_feature_field * (self.num_feature_field - 1) / 2
33 # define layers and loss
34 self.attlayer = AttLayer(self.embedding_size, self.attention_size)
35 self.p = nn.Parameter(torch.randn(self.embedding_size), requires_grad=True)
36 self.dropout_layer = nn.Dropout(p=self.dropout_prob)
37 self.sigmoid = nn.Sigmoid()
38 self.loss = nn.BCEWithLogitsLoss()
40 # parameters initialization
41 self.apply(self._init_weights)
43 def _init_weights(self, module):
44 if isinstance(module, nn.Embedding):
45 xavier_normal_(module.weight.data)
46 elif isinstance(module, nn.Linear):
47 xavier_normal_(module.weight.data)
48 if module.bias is not None:
49 constant_(module.bias.data, 0)
51 def build_cross(self, feat_emb):
52 """Build the cross feature columns of feature columns
54 Args:
55 feat_emb (torch.FloatTensor): input feature embedding tensor. shape of [batch_size, field_size, embed_dim].
57 Returns:
58 tuple:
59 - torch.FloatTensor: Left part of the cross feature. shape of [batch_size, num_pairs, emb_dim].
60 - torch.FloatTensor: Right part of the cross feature. shape of [batch_size, num_pairs, emb_dim].
61 """
62 # num_pairs = num_feature_field * (num_feature_field-1) / 2
63 row = []
64 col = []
65 for i in range(self.num_feature_field - 1):
66 for j in range(i + 1, self.num_feature_field):
67 row.append(i)
68 col.append(j)
69 p = feat_emb[:, row] # [batch_size, num_pairs, emb_dim]
70 q = feat_emb[:, col] # [batch_size, num_pairs, emb_dim]
71 return p, q
73 def afm_layer(self, infeature):
74 """Get the attention-based feature interaction score
76 Args:
77 infeature (torch.FloatTensor): input feature embedding tensor. shape of [batch_size, field_size, embed_dim].
79 Returns:
80 torch.FloatTensor: Result of score. shape of [batch_size, 1].
81 """ # noqa: E501
82 p, q = self.build_cross(infeature)
83 pair_wise_inter = torch.mul(p, q) # [batch_size, num_pairs, emb_dim]
85 # [batch_size, num_pairs, 1]
86 att_signal = self.attlayer(pair_wise_inter).unsqueeze(dim=2)
88 att_inter = torch.mul(att_signal, pair_wise_inter) # [batch_size, num_pairs, emb_dim]
89 att_pooling = torch.sum(att_inter, dim=1) # [batch_size, emb_dim]
90 att_pooling = self.dropout_layer(att_pooling) # [batch_size, emb_dim]
92 att_pooling = torch.mul(att_pooling, self.p) # [batch_size, emb_dim]
93 att_pooling = torch.sum(att_pooling, dim=1, keepdim=True) # [batch_size, 1]
95 return att_pooling
97 def forward(self, interaction):
98 afm_all_embeddings = self.concat_embed_input_fields(interaction) # [batch_size, num_field, embed_dim]
100 output = self.first_order_linear(interaction) + self.afm_layer(afm_all_embeddings)
101 return output.squeeze(-1)
103 def calculate_loss(self, interaction):
104 label = interaction[self.LABEL]
106 output = self.forward(interaction)
107 l2_loss = self.reg_weight * torch.norm(self.attlayer.w.weight, p=2)
108 return self.loss(output, label) + l2_loss
110 def predict(self, interaction):
111 return self.sigmoid(self.forward(interaction))