Coverage for hopwise/model/context_aware_recommender/kd_dagfm.py: 58%
159 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 : 2023/1/20
2# @Author : Wanli Yang
3# @Email : 2013774@mail.nankai.edu.cn
5r"""KD_DAGFM
6################################################
7Reference:
8 Zhen Tian et al. "Directed Acyclic Graph Factorization Machines for CTR Prediction via Knowledge Distillation."
9 in WSDM 2023.
10Reference code:
11 https://github.com/chenyuwuxin/DAGFM
12"""
14from copy import deepcopy
16import torch
17from torch import nn
18from torch.nn.init import xavier_normal_
20from hopwise.model.abstract_recommender import ContextRecommender
21from hopwise.model.init import xavier_normal_initialization
24class KD_DAGFM(ContextRecommender):
25 r"""KD_DAGFM is a context-based recommendation model. The model is based on directed acyclic graph and knowledge
26 distillation. It can learn arbitrary feature interactions from the complex teacher networks and achieve
27 approximately lossless model performance. It can also greatly reduce the computational resource costs.
28 """
30 def __init__(self, config, dataset):
31 super().__init__(config, dataset)
33 # load parameters info
34 self.phase = config["phase"]
35 self.alpha = config["alpha"]
36 self.beta = config["beta"]
38 # add element to config for the initialization of teacher&student network
39 config["feature_num"] = self.num_feature_field
41 # initialize teacher&student network
42 self.student_network = DAGFM(config)
43 self.teacher_network = eval(f"{config['teacher']}")(self.get_teacher_config(config))
45 # initialize loss function
46 self.loss_fn = nn.BCELoss()
48 # get warm up parameters
49 if self.phase != "teacher_training":
50 if "warm_up" not in config:
51 raise ValueError("Must have warm up!")
52 else:
53 save_info = torch.load(config["warm_up"])
54 self.load_state_dict(save_info["state_dict"])
55 else:
56 self.apply(xavier_normal_initialization)
58 # get config of teacher network from config
59 def get_teacher_config(self, config):
60 teacher_cfg = deepcopy(config)
61 for key in config.final_config_dict:
62 if key.startswith("t_"):
63 teacher_cfg[key[2:]] = config[key]
64 return teacher_cfg
66 def FeatureInteraction(self, feature):
67 if self.phase == "teacher_training":
68 return self.teacher_network.FeatureInteraction(feature)
69 elif self.phase in ("distillation", "finetuning"):
70 return self.student_network.FeatureInteraction(feature)
71 else:
72 return ValueError("Phase invalid!")
74 def forward(self, interaction):
75 dagfm_all_embeddings = self.concat_embed_input_fields(interaction) # [batch_size, num_field, embed_dim]
76 if self.phase in ("teacher_training", "finetuning"):
77 return self.FeatureInteraction(dagfm_all_embeddings)
78 elif self.phase == "distillation":
79 dagfm_all_embeddings = dagfm_all_embeddings.data
80 if self.training:
81 self.t_pred = self.teacher_network(dagfm_all_embeddings)
82 return self.FeatureInteraction(dagfm_all_embeddings)
83 else:
84 raise ValueError("Phase invalid!")
86 def calculate_loss(self, interaction):
87 if self.phase in ("teacher_training", "finetuning"):
88 prediction = self.forward(interaction)
89 loss = self.loss_fn(
90 prediction.squeeze(-1),
91 interaction[self.LABEL].squeeze(-1).to(self.device),
92 )
93 elif self.phase == "distillation":
94 self.teacher_network.eval()
95 s_pred = self.forward(interaction)
96 ctr_loss = self.loss_fn(s_pred.squeeze(-1), interaction[self.LABEL].squeeze(-1).to(self.device))
97 kd_loss = torch.mean((self.teacher_network.logits.data - self.student_network.logits) ** 2)
98 loss = self.alpha * ctr_loss + self.beta * kd_loss
99 else:
100 raise ValueError("Phase invalid!")
101 return loss
103 def predict(self, interaction):
104 return self.forward(interaction)
107class DAGFM(nn.Module):
108 def __init__(self, config):
109 super().__init__()
110 if torch.cuda.is_available():
111 self.device = torch.device("cuda")
112 else:
113 self.device = torch.device("cpu")
115 # load parameters info
116 self.type = config["type"]
117 self.depth = config["depth"]
118 field_num = config["feature_num"]
119 embedding_size = config["embedding_size"]
121 # initialize parameters according to the type
122 if self.type == "inner":
123 self.p = nn.ParameterList(
124 [nn.Parameter(torch.randn(field_num, field_num, embedding_size)) for _ in range(self.depth)]
125 )
126 for _ in range(self.depth):
127 xavier_normal_(self.p[_], gain=1.414)
128 elif self.type == "outer":
129 self.p = nn.ParameterList(
130 [nn.Parameter(torch.randn(field_num, field_num, embedding_size)) for _ in range(self.depth)]
131 )
132 self.q = nn.ParameterList(
133 [nn.Parameter(torch.randn(field_num, field_num, embedding_size)) for _ in range(self.depth)]
134 )
135 for _ in range(self.depth):
136 xavier_normal_(self.p[_], gain=1.414)
137 xavier_normal_(self.q[_], gain=1.414)
138 self.adj_matrix = torch.zeros(field_num, field_num, embedding_size).to(self.device)
139 for i in range(field_num):
140 for j in range(i, field_num):
141 self.adj_matrix[i, j, :] += 1
142 self.connect_layer = nn.Parameter(torch.eye(field_num).float())
143 self.linear = nn.Linear(field_num * (self.depth + 1), 1)
145 def FeatureInteraction(self, feature):
146 init_state = self.connect_layer @ feature
147 h0, ht = init_state, init_state
148 state = [torch.sum(init_state, dim=-1)]
149 for i in range(self.depth):
150 if self.type == "inner":
151 aggr = torch.einsum("bfd,fsd->bsd", ht, self.p[i] * self.adj_matrix)
152 ht = h0 * aggr
153 elif self.type == "outer":
154 term = torch.einsum("bfd,fsd->bfs", ht, self.p[i] * self.adj_matrix)
155 aggr = torch.einsum("bfs,fsd->bsd", term, self.q[i])
156 ht = h0 * aggr
157 state.append(torch.sum(ht, dim=-1))
159 state = torch.cat(state, dim=-1)
160 self.logits = self.linear(state)
161 self.outputs = torch.sigmoid(self.logits)
162 return self.outputs
165# teacher network CrossNet
166class CrossNet(nn.Module):
167 def __init__(self, config):
168 super().__init__()
170 # load parameters info
171 self.depth = config["depth"]
172 self.embedding_size = config["embedding_size"]
173 self.feature_num = config["feature_num"]
174 self.in_feature_num = self.feature_num * self.embedding_size
175 self.cross_layer_w = nn.ParameterList(
176 nn.Parameter(torch.randn(self.in_feature_num, self.in_feature_num)) for _ in range(self.depth)
177 )
178 self.bias = nn.ParameterList(nn.Parameter(torch.zeros(self.in_feature_num, 1)) for _ in range(self.depth))
179 self.linear = nn.Linear(self.in_feature_num, 1)
180 nn.init.normal_(self.linear.weight)
182 def FeatureInteraction(self, x_0):
183 x_0 = x_0.reshape(x_0.shape[0], -1)
184 x_0 = x_0.unsqueeze(dim=2)
185 x_l = x_0 # (batch_size, in_feature_num, 1)
186 for i in range(self.depth):
187 xl_w = torch.matmul(self.cross_layer_w[i], x_l)
188 xl_w = xl_w + self.bias[i]
189 xl_dot = torch.mul(x_0, xl_w)
190 x_l = xl_dot + x_l
191 x_l = x_l.squeeze(dim=2)
192 self.logits = self.linear(x_l)
193 self.outputs = torch.sigmoid(self.logits)
194 return self.outputs
196 def forward(self, feature):
197 return self.FeatureInteraction(feature)
200class CINComp(nn.Module):
201 def __init__(self, indim, outdim, config):
202 super().__init__()
203 basedim = config["feature_num"]
204 self.conv = nn.Conv1d(indim * basedim, outdim, 1)
206 def forward(self, feature, base):
207 return self.conv(
208 (feature[:, :, None, :] * base[:, None, :, :]).reshape(
209 feature.shape[0], feature.shape[1] * base.shape[1], -1
210 )
211 )
214# teacher network CIN
215class CIN(nn.Module):
216 def __init__(self, config):
217 super().__init__()
218 self.cinlist = [config["feature_num"]] + config["cin"]
219 self.cin = nn.ModuleList(
220 [CINComp(self.cinlist[i], self.cinlist[i + 1], config) for i in range(0, len(self.cinlist) - 1)]
221 )
222 self.linear = nn.Parameter(torch.zeros(sum(self.cinlist) - self.cinlist[0], 1))
223 nn.init.normal_(self.linear, mean=0, std=0.01)
224 self.backbone = ["cin", "linear"]
225 self.loss_fn = nn.BCELoss()
226 if torch.cuda.is_available():
227 self.device = torch.device("cuda")
228 else:
229 self.device = torch.device("cpu")
231 def FeatureInteraction(self, feature):
232 base = feature
233 x = feature
234 p = []
235 for comp in self.cin:
236 x = comp(x, base)
237 p.append(torch.sum(x, dim=-1))
238 p = torch.cat(p, dim=-1)
239 self.logits = p @ self.linear
240 self.outputs = torch.sigmoid(self.logits)
241 return self.outputs
243 def forward(self, feature):
244 return self.FeatureInteraction(feature)