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
« 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
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.
12Reference code:
13 - https://github.com/CRIPAC-DIG/GraphCTR
14 - https://github.com/xue-pai/FuxiCTR
15"""
17from itertools import product
19import torch
20import torch.nn.functional as F
21from torch import nn
22from torch.nn.init import constant_, xavier_normal_, xavier_uniform_
24from hopwise.model.abstract_recommender import ContextRecommender
25from hopwise.utils import InputType
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."""
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))
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
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 """
51 input_type = InputType.POINTWISE
53 def __init__(self, config, dataset):
54 super().__init__(config, dataset)
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"]
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 )
90 self.sigmoid = nn.Sigmoid()
91 self.loss = nn.BCEWithLogitsLoss()
92 # parameters initialization
93 self.apply(self._init_weights)
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()
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
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)
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)
148 def calculate_loss(self, interaction):
149 label = interaction[self.LABEL]
150 output = self.forward(interaction)
151 return self.loss(output, label)
153 def predict(self, interaction):
154 return self.sigmoid(self.forward(interaction))