Coverage for hopwise/model/context_aware_recommender/fignn.py: 97%

92 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2022/10/27 

2# @Author : Yuyan Zhang 

3# @Email : 2019308160102@cau.edu.cn 

4# @File : fignn.py 

5 

6r"""FiGNN 

7################################################ 

8Reference: 

9 Li, Zekun, et al. "Fi-GNN: Modeling Feature Interactions via Graph Neural Networks for CTR Prediction" 

10 in CIKM 2019. 

11 

12Reference code: 

13 - https://github.com/CRIPAC-DIG/GraphCTR 

14 - https://github.com/xue-pai/FuxiCTR 

15""" 

16 

17from itertools import product 

18 

19import torch 

20import torch.nn.functional as F 

21from torch import nn 

22from torch.nn.init import constant_, xavier_normal_, xavier_uniform_ 

23 

24from hopwise.model.abstract_recommender import ContextRecommender 

25from hopwise.utils import InputType 

26 

27 

28class GraphLayer(nn.Module): 

29 """The implementations of the GraphLayer part and the Attentional Edge Weights part are adapted from https://github.com/xue-pai/FuxiCTR.""" 

30 

31 def __init__(self, num_fields, embedding_size): 

32 super().__init__() 

33 self.W_in = nn.Parameter(torch.Tensor(num_fields, embedding_size, embedding_size)) 

34 self.W_out = nn.Parameter(torch.Tensor(num_fields, embedding_size, embedding_size)) 

35 xavier_normal_(self.W_in) 

36 xavier_normal_(self.W_out) 

37 self.bias_p = nn.Parameter(torch.zeros(embedding_size)) 

38 

39 def forward(self, g, h): 

40 h_out = torch.matmul(self.W_out, h.unsqueeze(-1)).squeeze(-1) 

41 aggr = torch.bmm(g, h_out) 

42 a = torch.matmul(self.W_in, aggr.unsqueeze(-1)).squeeze(-1) + self.bias_p 

43 return a 

44 

45 

46class FiGNN(ContextRecommender): 

47 """FiGNN is a CTR prediction model based on GGNN, 

48 which can model sophisticated interactions among feature fields on the graph-structured features. 

49 """ 

50 

51 input_type = InputType.POINTWISE 

52 

53 def __init__(self, config, dataset): 

54 super().__init__(config, dataset) 

55 

56 # load parameters info 

57 self.attention_size = config["attention_size"] 

58 self.n_layers = config["n_layers"] 

59 self.num_heads = config["num_heads"] 

60 self.hidden_dropout_prob = config["hidden_dropout_prob"] 

61 self.attn_dropout_prob = config["attn_dropout_prob"] 

62 

63 # define layers and loss 

64 self.dropout_layer = nn.Dropout(p=self.hidden_dropout_prob) 

65 self.att_embedding = nn.Linear(self.embedding_size, self.attention_size) 

66 # multi-head self-attention network 

67 self.self_attn = nn.MultiheadAttention( 

68 self.attention_size, 

69 self.num_heads, 

70 dropout=self.attn_dropout_prob, 

71 batch_first=True, 

72 ) 

73 self.v_res_embedding = torch.nn.Linear(self.embedding_size, self.attention_size) 

74 # FiGNN 

75 self.src_nodes, self.dst_nodes = zip(*list(product(range(self.num_feature_field), repeat=2))) 

76 self.gnn = nn.ModuleList( 

77 [GraphLayer(self.num_feature_field, self.attention_size) for _ in range(self.n_layers - 1)] 

78 ) 

79 self.leaky_relu = nn.LeakyReLU(negative_slope=0.01) 

80 self.W_attn = nn.Linear(self.attention_size * 2, 1, bias=False) 

81 self.gru_cell = nn.GRUCell(self.attention_size, self.attention_size) 

82 # Attentional Scoring Layer 

83 self.mlp1 = nn.Linear(self.attention_size, 1, bias=False) 

84 self.mlp2 = nn.Linear( 

85 self.num_feature_field * self.attention_size, 

86 self.num_feature_field, 

87 bias=False, 

88 ) 

89 

90 self.sigmoid = nn.Sigmoid() 

91 self.loss = nn.BCEWithLogitsLoss() 

92 # parameters initialization 

93 self.apply(self._init_weights) 

94 

95 def fignn_layer(self, in_feature): 

96 emb_feature = self.att_embedding(in_feature) 

97 emb_feature = self.dropout_layer(emb_feature) 

98 # multi-head self-attention network 

99 att_feature, _ = self.self_attn(emb_feature, emb_feature, emb_feature) # [batch_size, num_field, att_dim] 

100 # Residual connection 

101 v_res = self.v_res_embedding(in_feature) 

102 att_feature += v_res 

103 att_feature = F.relu(att_feature).contiguous() 

104 

105 # init graph 

106 src_emb = att_feature[:, self.src_nodes, :] 

107 dst_emb = att_feature[:, self.dst_nodes, :] 

108 concat_emb = torch.cat([src_emb, dst_emb], dim=-1) 

109 alpha = self.leaky_relu(self.W_attn(concat_emb)) 

110 alpha = alpha.view(-1, self.num_feature_field, self.num_feature_field) 

111 mask = torch.eye(self.num_feature_field).to(self.device) 

112 alpha = alpha.masked_fill(mask.bool(), float("-inf")) 

113 self.graph = F.softmax(alpha, dim=-1) 

114 # message passing 

115 if self.n_layers > 1: 

116 h = att_feature 

117 for i in range(self.n_layers - 1): 

118 a = self.gnn[i](self.graph, h) 

119 a = a.view(-1, self.attention_size) 

120 h = h.view(-1, self.attention_size) 

121 h = self.gru_cell(a, h) 

122 h = h.view(-1, self.num_feature_field, self.attention_size) 

123 h += att_feature 

124 else: 

125 h = att_feature 

126 # Attentional Scoring Layer 

127 score = self.mlp1(h).squeeze(-1) 

128 weight = self.mlp2(h.flatten(start_dim=1)) 

129 logit = (weight * score).sum(dim=1).unsqueeze(-1) 

130 return logit 

131 

132 def _init_weights(self, module): 

133 if isinstance(module, nn.Embedding): 

134 xavier_normal_(module.weight.data) 

135 elif isinstance(module, nn.Linear): 

136 xavier_normal_(module.weight.data) 

137 if module.bias is not None: 

138 constant_(module.bias.data, 0) 

139 elif isinstance(module, nn.GRU): 

140 xavier_uniform_(module.weight_hh_l0) 

141 xavier_uniform_(module.weight_ih_l0) 

142 

143 def forward(self, interaction): 

144 fignn_all_embeddings = self.concat_embed_input_fields(interaction) # [batch_size, num_field, embed_dim] 

145 output = self.fignn_layer(fignn_all_embeddings) 

146 return output.squeeze(1) 

147 

148 def calculate_loss(self, interaction): 

149 label = interaction[self.LABEL] 

150 output = self.forward(interaction) 

151 return self.loss(output, label) 

152 

153 def predict(self, interaction): 

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