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

1# @Time : 2020/7/21 

2# @Author : Zihan Lin 

3# @Email : linzihan.super@foxmail.com 

4# @File : afm.py 

5 

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""" 

12 

13import torch 

14from torch import nn 

15from torch.nn.init import constant_, xavier_normal_ 

16 

17from hopwise.model.abstract_recommender import ContextRecommender 

18from hopwise.model.layers import AttLayer 

19 

20 

21class AFM(ContextRecommender): 

22 """AFM is a attention based FM model that predict the final score with the attention of input feature.""" 

23 

24 def __init__(self, config, dataset): 

25 super().__init__(config, dataset) 

26 

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 

32 

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() 

39 

40 # parameters initialization 

41 self.apply(self._init_weights) 

42 

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) 

50 

51 def build_cross(self, feat_emb): 

52 """Build the cross feature columns of feature columns 

53 

54 Args: 

55 feat_emb (torch.FloatTensor): input feature embedding tensor. shape of [batch_size, field_size, embed_dim]. 

56 

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 

72 

73 def afm_layer(self, infeature): 

74 """Get the attention-based feature interaction score 

75 

76 Args: 

77 infeature (torch.FloatTensor): input feature embedding tensor. shape of [batch_size, field_size, embed_dim]. 

78 

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] 

84 

85 # [batch_size, num_pairs, 1] 

86 att_signal = self.attlayer(pair_wise_inter).unsqueeze(dim=2) 

87 

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] 

91 

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] 

94 

95 return att_pooling 

96 

97 def forward(self, interaction): 

98 afm_all_embeddings = self.concat_embed_input_fields(interaction) # [batch_size, num_field, embed_dim] 

99 

100 output = self.first_order_linear(interaction) + self.afm_layer(afm_all_embeddings) 

101 return output.squeeze(-1) 

102 

103 def calculate_loss(self, interaction): 

104 label = interaction[self.LABEL] 

105 

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 

109 

110 def predict(self, interaction): 

111 return self.sigmoid(self.forward(interaction))