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
« 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
6# UPDATE:
7# @Time : 2022/7/16
8# @Author : Zhen Tian
9# @Email : chenyuwuxinn@gmail.com
11r"""FFM
12#####################################################
13Reference:
14 Yuchin Juan et al. "Field-aware Factorization Machines for CTR Prediction" in RecSys 2016.
16Reference code:
17 https://github.com/rixwew/pytorch-fm
18"""
20import numpy as np
21import torch
22from torch import nn
23from torch.nn.init import constant_, xavier_normal_
25from hopwise.model.abstract_recommender import ContextRecommender
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.
33 The model defines as follows:
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 """
39 def __init__(self, config, dataset):
40 super().__init__(config, dataset)
42 # load parameters info
43 self.fields = config["fields"] # a dict; key: field_id; value: feature_list
45 self.sigmoid = nn.Sigmoid()
47 self.feature2id = {}
48 self.feature2field = {}
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
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()
76 # parameters initialization
77 self.apply(self._init_weights)
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)
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
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
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
132 return (
133 token_ffm_input,
134 float_ffm_input,
135 token_seq_ffm_input,
136 float_seq_ffm_input,
137 )
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
144 return output.squeeze(-1)
146 def calculate_loss(self, interaction):
147 label = interaction[self.LABEL]
149 output = self.forward(interaction)
150 return self.loss(output, label)
152 def predict(self, interaction):
153 return self.sigmoid(self.forward(interaction))
156class FieldAwareFactorizationMachine(nn.Module):
157 r"""This is Field-Aware Factorization Machine Module for FFM."""
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__()
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]
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
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)
232 def forward(self, input_x):
233 r"""Model the different interaction strengths of different field pairs.
235 Args:
236 input_x (a tuple): (token_ffm_input, float_ffm_input, token_seq_ffm_input)
238 token_ffm_input (torch.cuda.FloatTensor): [batch_size, num_token_features] or None
240 float_ffm_input (torch.cuda.FloatTensor): [batch_size, num_float_features] or None
242 token_seq_ffm_input (list): length is num_token_seq_features or 0
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 )
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)
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 )
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]
273 return output
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]]
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)
295 for tensors in zip(*zip_args):
296 input_x_emb.append(torch.cat(tensors, dim=1))
298 return input_x_emb
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]]
309 return token_input_x_emb
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]]
322 return float_input_x_emb
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]
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]
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]
350 return token_seq_input_x_emb
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]
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]
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]
380 return float_seq_input_x_emb