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

1# @Time : 2023/1/20 

2# @Author : Wanli Yang 

3# @Email : 2013774@mail.nankai.edu.cn 

4 

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

13 

14from copy import deepcopy 

15 

16import torch 

17from torch import nn 

18from torch.nn.init import xavier_normal_ 

19 

20from hopwise.model.abstract_recommender import ContextRecommender 

21from hopwise.model.init import xavier_normal_initialization 

22 

23 

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

29 

30 def __init__(self, config, dataset): 

31 super().__init__(config, dataset) 

32 

33 # load parameters info 

34 self.phase = config["phase"] 

35 self.alpha = config["alpha"] 

36 self.beta = config["beta"] 

37 

38 # add element to config for the initialization of teacher&student network 

39 config["feature_num"] = self.num_feature_field 

40 

41 # initialize teacher&student network 

42 self.student_network = DAGFM(config) 

43 self.teacher_network = eval(f"{config['teacher']}")(self.get_teacher_config(config)) 

44 

45 # initialize loss function 

46 self.loss_fn = nn.BCELoss() 

47 

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) 

57 

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 

65 

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

73 

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

85 

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 

102 

103 def predict(self, interaction): 

104 return self.forward(interaction) 

105 

106 

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

114 

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

120 

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) 

144 

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

158 

159 state = torch.cat(state, dim=-1) 

160 self.logits = self.linear(state) 

161 self.outputs = torch.sigmoid(self.logits) 

162 return self.outputs 

163 

164 

165# teacher network CrossNet 

166class CrossNet(nn.Module): 

167 def __init__(self, config): 

168 super().__init__() 

169 

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) 

181 

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 

195 

196 def forward(self, feature): 

197 return self.FeatureInteraction(feature) 

198 

199 

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) 

205 

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 ) 

212 

213 

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

230 

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 

242 

243 def forward(self, feature): 

244 return self.FeatureInteraction(feature)