Coverage for hopwise/model/context_aware_recommender/ffm.py: 76%

209 statements  

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

1# @Time : 2020/10/04 

2# @Author : Xinyan Fan 

3# @Email : xinyan.fan@ruc.edu.cn 

4# @File : ffm.py 

5 

6# UPDATE: 

7# @Time : 2022/7/16 

8# @Author : Zhen Tian 

9# @Email : chenyuwuxinn@gmail.com 

10 

11r"""FFM 

12##################################################### 

13Reference: 

14 Yuchin Juan et al. "Field-aware Factorization Machines for CTR Prediction" in RecSys 2016. 

15 

16Reference code: 

17 https://github.com/rixwew/pytorch-fm 

18""" 

19 

20import numpy as np 

21import torch 

22from torch import nn 

23from torch.nn.init import constant_, xavier_normal_ 

24 

25from hopwise.model.abstract_recommender import ContextRecommender 

26 

27 

28class FFM(ContextRecommender): 

29 r"""FFM is a context-based recommendation model. It aims to model the different feature interactions 

30 between different fields. Each feature has several latent vectors :math:`v_{i,F(j)}`, 

31 which depend on the field of other features, and one of them is used to do the inner product. 

32 

33 The model defines as follows: 

34 

35 .. math:: 

36 y = w_0 + \sum_{i=1}^{m}x_{i}w_{i} + \sum_{i=1}^{m}\sum_{j=i+1}^{m}x_{i}x_{j}<v_{i,F(j)}, v_{j,F(i)}> 

37 """ 

38 

39 def __init__(self, config, dataset): 

40 super().__init__(config, dataset) 

41 

42 # load parameters info 

43 self.fields = config["fields"] # a dict; key: field_id; value: feature_list 

44 

45 self.sigmoid = nn.Sigmoid() 

46 

47 self.feature2id = {} 

48 self.feature2field = {} 

49 

50 self.feature_names = ( 

51 self.token_field_names, 

52 self.float_field_names, 

53 self.token_seq_field_names, 

54 self.float_seq_field_names, 

55 ) 

56 self.feature_dims = ( 

57 self.token_field_dims, 

58 self.float_field_dims, 

59 self.token_seq_field_dims, 

60 self.float_seq_field_dims, 

61 ) 

62 self._get_feature2field() 

63 self.num_fields = len(set(self.feature2field.values())) # the number of fields 

64 

65 self.ffm = FieldAwareFactorizationMachine( 

66 self.feature_names, 

67 self.feature_dims, 

68 self.feature2id, 

69 self.feature2field, 

70 self.num_fields, 

71 self.embedding_size, 

72 self.device, 

73 ) 

74 self.loss = nn.BCEWithLogitsLoss() 

75 

76 # parameters initialization 

77 self.apply(self._init_weights) 

78 

79 def _init_weights(self, module): 

80 if isinstance(module, nn.Embedding): 

81 xavier_normal_(module.weight.data) 

82 elif isinstance(module, nn.Linear): 

83 xavier_normal_(module.weight.data) 

84 if module.bias is not None: 

85 constant_(module.bias.data, 0) 

86 

87 def _get_feature2field(self): 

88 r"""Create a mapping between features and fields.""" 

89 fea_id = 0 

90 for names in self.feature_names: 

91 if names is not None: 

92 for name in names: 

93 self.feature2id[name] = fea_id 

94 fea_id += 1 

95 

96 if self.fields is None: 

97 field_id = 0 

98 for key, value in self.feature2id.items(): 

99 self.feature2field[self.feature2id[key]] = field_id 

100 field_id += 1 

101 else: 

102 for key, value in self.fields.items(): 

103 for v in value: 

104 try: 

105 self.feature2field[self.feature2id[v]] = key 

106 except IndexError: 

107 pass 

108 

109 def get_ffm_input(self, interaction): 

110 r"""Get different types of ffm layer's input.""" 

111 token_ffm_input = [] 

112 if self.token_field_names is not None: 

113 for tn in self.token_field_names: 

114 token_ffm_input.append(torch.unsqueeze(interaction[tn], 1)) 

115 if len(token_ffm_input) > 0: 

116 token_ffm_input = torch.cat(token_ffm_input, dim=1) # [batch_size, num_token_features] 

117 float_ffm_input = [] 

118 if self.float_field_names is not None: 

119 for fn in self.float_field_names: 

120 float_ffm_input.append(torch.unsqueeze(interaction[fn], 1)) 

121 if len(float_ffm_input) > 0: 

122 float_ffm_input = torch.cat(float_ffm_input, dim=1) # [batch_size, num_float_features] 

123 token_seq_ffm_input = [] 

124 if self.token_seq_field_names is not None: 

125 for tsn in self.token_seq_field_names: 

126 token_seq_ffm_input.append(interaction[tsn]) # a list 

127 float_seq_ffm_input = [] 

128 if self.float_seq_field_names is not None: 

129 for tsn in self.float_seq_field_names: 

130 float_seq_ffm_input.append(interaction[tsn]) # a list 

131 

132 return ( 

133 token_ffm_input, 

134 float_ffm_input, 

135 token_seq_ffm_input, 

136 float_seq_ffm_input, 

137 ) 

138 

139 def forward(self, interaction): 

140 ffm_input = self.get_ffm_input(interaction) 

141 ffm_output = torch.sum(torch.sum(self.ffm(ffm_input), dim=1), dim=1, keepdim=True) 

142 output = self.first_order_linear(interaction) + ffm_output 

143 

144 return output.squeeze(-1) 

145 

146 def calculate_loss(self, interaction): 

147 label = interaction[self.LABEL] 

148 

149 output = self.forward(interaction) 

150 return self.loss(output, label) 

151 

152 def predict(self, interaction): 

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

154 

155 

156class FieldAwareFactorizationMachine(nn.Module): 

157 r"""This is Field-Aware Factorization Machine Module for FFM.""" 

158 

159 def __init__( 

160 self, 

161 feature_names, 

162 feature_dims, 

163 feature2id, 

164 feature2field, 

165 num_fields, 

166 embed_dim, 

167 device, 

168 ): 

169 super().__init__() 

170 

171 self.token_feature_names = feature_names[0] 

172 self.float_feature_names = feature_names[1] 

173 self.token_seq_feature_names = feature_names[2] 

174 self.float_seq_feature_names = feature_names[3] 

175 self.token_feature_dims = feature_dims[0] 

176 self.float_feature_dims = feature_dims[1] 

177 self.token_seq_feature_dims = feature_dims[2] 

178 self.float_seq_feature_dims = feature_dims[3] 

179 

180 self.feature2id = feature2id 

181 self.feature2field = feature2field 

182 self.num_features = ( 

183 len(self.token_feature_names) 

184 + len(self.float_feature_names) 

185 + len(self.token_seq_feature_names) 

186 + len(self.float_seq_feature_names) 

187 ) 

188 self.num_fields = num_fields 

189 self.embed_dim = embed_dim 

190 self.device = device 

191 

192 # init token field-aware embeddings if there is token type of features. 

193 if len(self.token_feature_names) > 0: 

194 self.num_token_features = len(self.token_feature_names) 

195 self.token_embeddings = torch.nn.ModuleList( 

196 [nn.Embedding(sum(self.token_feature_dims), self.embed_dim) for _ in range(self.num_fields)] 

197 ) 

198 self.token_offsets = np.array((0, *np.cumsum(self.token_feature_dims)[:-1]), dtype=np.long) 

199 for embedding in self.token_embeddings: 

200 nn.init.xavier_uniform_(embedding.weight.data) 

201 # init float field-aware embeddings if there is float type of features. 

202 if len(self.float_feature_names) > 0: 

203 self.num_float_features = len(self.float_feature_names) 

204 self.float_offsets = np.array((0, *np.cumsum(self.float_feature_dims)[:-1]), dtype=np.long) 

205 self.float_embeddings = torch.nn.ModuleList( 

206 [nn.Embedding(sum(self.float_feature_dims), self.embed_dim) for _ in range(self.num_fields)] 

207 ) 

208 for embedding in self.float_embeddings: 

209 nn.init.xavier_uniform_(embedding.weight.data) 

210 # init token_seq field-aware embeddings if there is token_seq type of features. 

211 if len(self.token_seq_feature_names) > 0: 

212 self.num_token_seq_features = len(self.token_seq_feature_names) 

213 self.token_seq_embeddings = torch.nn.ModuleList() 

214 self.token_seq_embedding = torch.nn.ModuleList() 

215 for i in range(self.num_fields): 

216 for token_seq_feature_dim in self.token_seq_feature_dims: 

217 self.token_seq_embedding.append(nn.Embedding(token_seq_feature_dim, self.embed_dim)) 

218 for embedding in self.token_seq_embedding: 

219 nn.init.xavier_uniform_(embedding.weight.data) 

220 self.token_seq_embeddings.append(self.token_seq_embedding) 

221 if len(self.float_seq_feature_names) > 0: 

222 self.num_float_seq_features = len(self.float_seq_feature_names) 

223 self.float_seq_embeddings = torch.nn.ModuleList() 

224 self.float_seq_embedding = torch.nn.ModuleList() 

225 for i in range(self.num_fields): 

226 for float_seq_feature_dim in self.float_seq_feature_dims: 

227 self.float_seq_embedding.append(nn.Embedding(float_seq_feature_dim, self.embed_dim)) 

228 for embedding in self.float_seq_embedding: 

229 nn.init.xavier_uniform_(embedding.weight.data) 

230 self.float_seq_embeddings.append(self.float_seq_embedding) 

231 

232 def forward(self, input_x): 

233 r"""Model the different interaction strengths of different field pairs. 

234 

235 Args: 

236 input_x (a tuple): (token_ffm_input, float_ffm_input, token_seq_ffm_input) 

237 

238 token_ffm_input (torch.cuda.FloatTensor): [batch_size, num_token_features] or None 

239 

240 float_ffm_input (torch.cuda.FloatTensor): [batch_size, num_float_features] or None 

241 

242 token_seq_ffm_input (list): length is num_token_seq_features or 0 

243 

244 Returns: 

245 torch.cuda.FloatTensor: The results of all features' field-aware interactions. 

246 shape: [batch_size, num_fields, emb_dim] 

247 """ 

248 token_ffm_input, float_ffm_input, token_seq_ffm_input, float_seq_ffm_input = ( 

249 input_x[0], 

250 input_x[1], 

251 input_x[2], 

252 input_x[3], 

253 ) 

254 

255 token_input_x_emb = self._emb_token_ffm_input(token_ffm_input) 

256 float_input_x_emb = self._emb_float_ffm_input(float_ffm_input) 

257 token_seq_input_x_emb = self._emb_token_seq_ffm_input(token_seq_ffm_input) 

258 float_seq_input_x_emb = self._emb_float_seq_ffm_input(float_seq_ffm_input) 

259 

260 input_x_emb = self._get_input_x_emb( 

261 token_input_x_emb, 

262 float_input_x_emb, 

263 token_seq_input_x_emb, 

264 float_seq_input_x_emb, 

265 ) 

266 

267 output = list() 

268 for i in range(self.num_features - 1): 

269 for j in range(i + 1, self.num_features): 

270 output.append(input_x_emb[self.feature2field[j]][:, i] * input_x_emb[self.feature2field[i]][:, j]) 

271 output = torch.stack(output, dim=1) # [batch_size, num_fields, emb_dim] 

272 

273 return output 

274 

275 def _get_input_x_emb( 

276 self, 

277 token_input_x_emb, 

278 float_input_x_emb, 

279 token_seq_input_x_emb, 

280 float_seq_input_x_emb, 

281 ): 

282 # merge different types of field-aware embeddings 

283 input_x_emb = [] # [num_fields: [batch_size, num_fields, emb_dim]] 

284 

285 zip_args = [] 

286 if len(self.token_feature_names) > 0: 

287 zip_args.append(token_input_x_emb) 

288 if len(self.float_feature_names) > 0: 

289 zip_args.append(float_input_x_emb) 

290 if len(self.token_seq_feature_names) > 0: 

291 zip_args.append(token_seq_input_x_emb) 

292 if len(self.float_seq_feature_names) > 0: 

293 zip_args.append(float_seq_input_x_emb) 

294 

295 for tensors in zip(*zip_args): 

296 input_x_emb.append(torch.cat(tensors, dim=1)) 

297 

298 return input_x_emb 

299 

300 def _emb_token_ffm_input(self, token_ffm_input): 

301 # get token field-aware embeddings 

302 token_input_x_emb = [] 

303 if len(self.token_feature_names) > 0: 

304 token_input_x = token_ffm_input + token_ffm_input.new_tensor(self.token_offsets).unsqueeze(0) 

305 token_input_x_emb = [ 

306 self.token_embeddings[i](token_input_x) for i in range(self.num_fields) 

307 ] # [num_fields: [batch_size, num_token_features, emb_dim]] 

308 

309 return token_input_x_emb 

310 

311 def _emb_float_ffm_input(self, float_ffm_input): 

312 # get float field-aware embeddings 

313 float_input_x_emb = [] 

314 if len(self.float_feature_names) > 0: 

315 base, index = torch.split(float_ffm_input, [1, 1], dim=-1) 

316 index = index.squeeze(-1).long() 

317 index = index + index.new_tensor(self.float_offsets).unsqueeze(0) 

318 float_input_x_emb = [ 

319 self.float_embeddings[i](index) * base for i in range(self.num_fields) 

320 ] # [num_fields: [batch_size, num_float_features, emb_dim]] 

321 

322 return float_input_x_emb 

323 

324 def _emb_token_seq_ffm_input(self, token_seq_ffm_input): 

325 # get token_seq field-aware embeddings 

326 token_seq_input_x_emb = [] 

327 if len(self.token_seq_feature_names) > 0: 

328 for i in range(self.num_fields): 

329 token_seq_result = [] 

330 for j, token_seq in enumerate(token_seq_ffm_input): 

331 embedding_table = self.token_seq_embeddings[i][j] 

332 mask = token_seq != 0 # [batch_size, seq_len] 

333 mask = mask.float() 

334 value_cnt = torch.sum(mask, dim=1, keepdim=True) # [batch_size, 1] 

335 

336 token_seq_embedding = embedding_table(token_seq) # [batch_size, seq_len, embed_dim] 

337 mask = mask.unsqueeze(2).expand_as(token_seq_embedding) # [batch_size, seq_len, embed_dim] 

338 # mean 

339 masked_token_seq_embedding = token_seq_embedding * mask.float() 

340 result = torch.sum(masked_token_seq_embedding, dim=1) # [batch_size, embed_dim] 

341 eps = torch.FloatTensor([1e-8]).to(self.device) 

342 result = torch.div(result, value_cnt + eps) # [batch_size, embed_dim] 

343 result = result.unsqueeze(1) # [batch_size, 1, embed_dim] 

344 

345 token_seq_result.append(result) 

346 token_seq_input_x_emb.append( 

347 torch.cat(token_seq_result, dim=1) 

348 ) # [num_fields: batch_size, num_token_seq_features, embed_dim] 

349 

350 return token_seq_input_x_emb 

351 

352 def _emb_float_seq_ffm_input(self, float_seq_ffm_input): 

353 # get float_seq field-aware embeddings 

354 float_seq_input_x_emb = [] 

355 if len(self.float_seq_feature_names) > 0: 

356 for i in range(self.num_fields): 

357 float_seq_result = [] 

358 for j, float_seq in enumerate(float_seq_ffm_input): 

359 embedding_table = self.float_seq_embeddings[i][j] 

360 base, index = torch.split(float_seq, [1, 1], dim=-1) 

361 index = index.squeeze(-1) 

362 mask = index != 0 # [batch_size, seq_len] 

363 mask = mask.float() 

364 value_cnt = torch.sum(mask, dim=1, keepdim=True) # [batch_size, 1] 

365 

366 float_seq_embedding = base * embedding_table(index.long()) # [batch_size, seq_len, embed_dim] 

367 mask = mask.unsqueeze(2).expand_as(float_seq_embedding) # [batch_size, seq_len, embed_dim] 

368 # mean 

369 masked_float_seq_embedding = float_seq_embedding * mask.float() 

370 result = torch.sum(masked_float_seq_embedding, dim=1) # [batch_size, embed_dim] 

371 eps = torch.FloatTensor([1e-8]).to(self.device) 

372 result = torch.div(result, value_cnt + eps) # [batch_size, embed_dim] 

373 result = result.unsqueeze(1) # [batch_size, 1, embed_dim] 

374 

375 float_seq_result.append(result) 

376 float_seq_input_x_emb.append( 

377 torch.cat(float_seq_result, dim=1) 

378 ) # [num_fields: batch_size, num_token_seq_features, embed_dim] 

379 

380 return float_seq_input_x_emb