Coverage for hopwise/model/abstract_recommender.py: 80%
352 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/6/25
2# @Author : Shanlei Mu
3# @Email : slmu@ruc.edu.cn
5# UPDATE:
6# @Time : 2022/7/16, 2020/8/6, 2020/8/25, 2023/4/24
7# @Author : Zhen Tian, Shanlei Mu, Yupeng Hou, Chenglong Ma
8# @Email : chenyuwuxinn@gmail.com, slmu@ruc.edu.cn, houyupeng@ruc.edu.cn, chenglong.m@outlook.com
10"""hopwise.model.abstract_recommender
11##################################
12"""
14from logging import getLogger
16import numpy as np
17import torch
18from torch import nn
20from hopwise.model.layers import FLEmbedding, FMEmbedding, FMFirstOrderLinear
21from hopwise.model.logits_processor import LogitsProcessorList
22from hopwise.utils import (
23 FeatureSource,
24 FeatureType,
25 GenerationOutputs,
26 InputType,
27 KnowledgeEvaluationType,
28 ModelType,
29 PathLanguageModelingTokenType,
30 get_logits_processor,
31 get_sequence_postprocessor,
32 set_color,
33)
36class AbstractRecommender(nn.Module):
37 r"""Base class for all models"""
39 def __init__(self, _skip_nn_module_init=False):
40 self.logger = getLogger()
42 if not _skip_nn_module_init:
43 super().__init__()
45 def calculate_loss(self, interaction):
46 r"""Calculate the training loss for a batch data.
48 Args:
49 interaction (Interaction): Interaction class of the batch.
51 Returns:
52 torch.Tensor: Training loss, shape: []
53 """
54 raise NotImplementedError
56 def predict(self, interaction):
57 r"""Predict the scores between users and items.
59 Args:
60 interaction (Interaction): Interaction class of the batch.
62 Returns:
63 torch.Tensor: Predicted scores for given users and items, shape: [batch_size]
64 """
65 raise NotImplementedError
67 def full_sort_predict(self, interaction):
68 r"""Full sort prediction function.
69 Given users, calculate the scores between users and all candidate items.
71 Args:
72 interaction (Interaction): Interaction class of the batch.
74 Returns:
75 torch.Tensor: Predicted scores for given users and all candidate items,
76 shape: [n_batch_users * n_candidate_items]
77 """
78 raise NotImplementedError
80 def full_sort_predict_kg(self, interaction):
81 r"""Full sort prediction KG function.
82 Given heads, calculate the scores between heads and all candidate tails.
84 Args:
85 interaction (Interaction): Interaction class of the batch.
87 Returns:
88 torch.Tensor: Predicted scores for given heads and all candidate tails,
89 shape: [n_batch_heads * n_candidate_tails]
90 """
91 raise NotImplementedError
93 def other_parameter(self):
94 if hasattr(self, "other_parameter_name"):
95 return {key: getattr(self, key) for key in self.other_parameter_name}
96 return dict()
98 def load_other_parameter(self, para):
99 if para is None:
100 return
101 for key, value in para.items():
102 setattr(self, key, value)
104 def __str__(self):
105 """Model prints with number of trainable parameters"""
106 model_parameters = filter(lambda p: p.requires_grad, self.parameters())
107 params = sum([np.prod(p.size()) for p in model_parameters])
108 return super().__str__() + set_color("\nTrainable parameters", "blue") + f": {params}"
111class GeneralRecommender(AbstractRecommender):
112 """This is a abstract general recommender. All the general model should implement this class.
113 The base general recommender class provide the basic dataset and parameters information.
114 """
116 type = ModelType.GENERAL
118 def __init__(self, config, dataset):
119 super().__init__()
121 # load dataset info
122 self.USER_ID = config["USER_ID_FIELD"]
123 self.ITEM_ID = config["ITEM_ID_FIELD"]
124 self.NEG_ITEM_ID = config["NEG_PREFIX"] + self.ITEM_ID
125 self.n_users = dataset.num(self.USER_ID)
126 self.n_items = dataset.num(self.ITEM_ID)
128 # load parameters info
129 self.device = config["device"]
132class AutoEncoderMixin:
133 """This is a common part of auto-encoders. All the auto-encoder models should inherit this class,
134 including CDAE, MacridVAE, MultiDAE, MultiVAE, RaCT and RecVAE.
135 The base AutoEncoderMixin class provides basic dataset information and rating matrix function.
136 """
138 def build_histroy_items(self, dataset):
139 self.history_item_id, self.history_item_value, _ = dataset.history_item_matrix()
140 self.history_item_id = self.history_item_id.to(self.device)
141 self.history_item_value = self.history_item_value.to(self.device)
143 def get_rating_matrix(self, user):
144 r"""Get a batch of user's feature with the user's id and history interaction matrix.
146 Args:
147 user (torch.LongTensor): The input tensor that contains user's id, shape: [batch_size, ]
149 Returns:
150 torch.FloatTensor: The user's feature of a batch of user, shape: [batch_size, n_items]
151 """
152 # Following lines construct tensor of shape [B,n_items] using the tensor of shape [B,H]
153 col_indices = self.history_item_id[user].flatten()
154 row_indices = torch.arange(user.shape[0]).repeat_interleave(self.history_item_id.shape[1], dim=0)
155 rating_matrix = torch.zeros(1, device=self.device).repeat(user.shape[0], self.n_items)
156 rating_matrix.index_put_((row_indices, col_indices), self.history_item_value[user].flatten())
157 return rating_matrix
160class SequentialRecommender(AbstractRecommender):
161 """This is a abstract sequential recommender. All the sequential model should implement This class."""
163 type = ModelType.SEQUENTIAL
165 def __init__(self, config, dataset):
166 super().__init__()
168 # load dataset info
169 self.USER_ID = config["USER_ID_FIELD"]
170 self.ITEM_ID = config["ITEM_ID_FIELD"]
171 self.ITEM_SEQ = self.ITEM_ID + config["LIST_SUFFIX"]
172 self.ITEM_SEQ_LEN = config["ITEM_LIST_LENGTH_FIELD"]
173 self.POS_ITEM_ID = self.ITEM_ID
174 self.NEG_ITEM_ID = config["NEG_PREFIX"] + self.ITEM_ID
175 self.max_seq_length = config["MAX_ITEM_LIST_LENGTH"]
176 self.n_items = dataset.num(self.ITEM_ID)
178 # load parameters info
179 self.device = config["device"]
181 def gather_indexes(self, output, gather_index):
182 """Gathers the vectors at the specific positions over a minibatch"""
183 gather_index = gather_index.view(-1, 1, 1).expand(-1, -1, output.shape[-1])
184 output_tensor = output.gather(dim=1, index=gather_index)
185 return output_tensor.squeeze(1)
187 def get_attention_mask(self, item_seq, bidirectional=False):
188 """Generate left-to-right uni-directional or bidirectional attention mask for multi-head attention."""
189 attention_mask = item_seq != 0
190 extended_attention_mask = attention_mask.unsqueeze(1).unsqueeze(2) # torch.bool
191 if not bidirectional:
192 extended_attention_mask = torch.tril(extended_attention_mask.expand((-1, -1, item_seq.size(-1), -1)))
193 extended_attention_mask = torch.where(extended_attention_mask, 0.0, -10000.0)
194 return extended_attention_mask
197class KnowledgeRecommender(AbstractRecommender):
198 """This is a abstract knowledge-based recommender. All the knowledge-based model should implement this class.
199 The base knowledge-based recommender class provide the basic dataset and parameters information.
200 """
202 type = ModelType.KNOWLEDGE
204 def __init__(self, config, dataset, _skip_nn_module_init=False):
205 super().__init__(_skip_nn_module_init=_skip_nn_module_init)
207 # load dataset info
208 self.USER_ID = config["USER_ID_FIELD"]
209 self.ITEM_ID = config["ITEM_ID_FIELD"]
210 self.NEG_ITEM_ID = config["NEG_PREFIX"] + self.ITEM_ID
211 self.ENTITY_ID = config["ENTITY_ID_FIELD"]
212 self.RELATION_ID = config["RELATION_ID_FIELD"]
213 self.HEAD_ENTITY_ID = config["HEAD_ENTITY_ID_FIELD"]
214 self.TAIL_ENTITY_ID = config["TAIL_ENTITY_ID_FIELD"]
215 self.NEG_TAIL_ENTITY_ID = config["NEG_PREFIX"] + self.TAIL_ENTITY_ID
216 self.n_users = dataset.num(self.USER_ID)
217 self.n_items = dataset.num(self.ITEM_ID)
218 self.n_entities = dataset.num(self.ENTITY_ID)
219 self.n_relations = dataset.num(self.RELATION_ID)
221 # load parameters info
222 if not _skip_nn_module_init:
223 self.device = config["device"]
226class ExplainableRecommender:
227 """This is a abstract explainable-based recommender. All the explainable-based model should implement this class.
228 This class use templates to make the explanation more interpretable.
230 """
232 def explain(self, interaction):
233 r"""
234 Explain the prediction function.
236 Given users, calculate the scores and paths between users and all candidate items,
237 then return the templates filled with path data.
239 Args:
240 interaction (Interaction): The interaction batch.
242 Returns:
243 torch.Tensor: Predicted scores for given users and all candidate items,
244 with shape [n_batch_users * n_candidate_items].
245 pandas.DataFrame: Explanation of the prediction, containing paths and corresponding templates,
246 with shape [n_paths * [uid, pid, score, template1, template2, ..., #templates]].
247 """
248 raise NotImplementedError("explain is not implemented")
250 def decode_path(self, path):
251 r"""
252 Decode the path into a string. Path decoding is specific to each model.
254 Args:
255 path (list): The path data.
257 Returns:
258 str: The decoded path string.
259 """
260 raise NotImplementedError("decode_path is not implemented")
263class PathLanguageModelingRecommender(KnowledgeRecommender):
264 """This is an abstract path-language-modeling recommender.
265 All the path-language-modeling model should implement this class.
266 The base path-language-modeling recommender class inherits the knowledge-aware recommender class to
267 learn from knowledge graph paths defined by a chain of entity-relation triplets.
268 """
270 type = ModelType.PATH_LANGUAGE_MODELING
271 input_type = InputType.PATHWISE
273 def __init__(self, config, dataset, _skip_nn_module_init=True):
274 super().__init__(config, dataset, _skip_nn_module_init=_skip_nn_module_init)
276 self.n_tokens = len(dataset.tokenizer)
277 self.token_sequence_length = dataset.token_sequence_length - 1 # EOS token is not included
279 logits_processor = get_logits_processor(config["model"])(
280 tokenized_ckg=dataset.get_tokenized_ckg(),
281 tokenized_used_ids=dataset.get_tokenized_used_ids(),
282 max_sequence_length=self.token_sequence_length,
283 tokenizer=dataset.tokenizer,
284 task=KnowledgeEvaluationType.REC,
285 )
286 self.logits_processor_list = LogitsProcessorList([logits_processor])
288 self.sequence_postprocessor = get_sequence_postprocessor(config["sequence_postprocessor"])(
289 dataset.tokenizer,
290 dataset.get_user_used_ids(),
291 dataset.item_num,
292 topk=config["topk"],
293 )
295 @torch.no_grad()
296 def generate(self, inputs, top_k=None, paths_per_user=1, **kwargs):
297 """
298 Take a conditioning sequence of indices idx (LongTensor of shape (b,t)) and complete
299 the sequence max_new_tokens times, feeding the predictions back into the model each time.
300 Most likely you'll want to make sure to be in model.eval() mode of operation for this.
302 Args:
303 inputs (dict): A dictionary containing the input_ids tensor with shape (b, t).
304 top_k (int, optional): If specified, only the top k logits will be considered
305 for sampling at each step. Defaults to None.
306 paths_per_user (int, optional): How many paths to return for each user.
307 **kwargs: Additional keyword arguments for the model. In future, it can be used to pass
308 other generation parameters such as temperature, repetition penalty, etc.
309 """
310 max_new_tokens = self.token_sequence_length - inputs["input_ids"].size(1)
312 # How many paths to return?
313 inputs["input_ids"] = inputs["input_ids"].repeat_interleave(paths_per_user, dim=0)
314 scores = torch.full((inputs["input_ids"].size(0), max_new_tokens, self.n_tokens), -torch.inf).to(self.device)
315 for i in range(max_new_tokens):
316 # forward the model to get the logits for the index in the sequence
317 logits = self.predict(inputs)
318 # pluck the logits at the final step and scale by desired temperature
319 logits = logits[:, -1, :] / self.temperature
321 # KGCD
322 logits = self.logits_processor_list(inputs["input_ids"], logits)
324 # optionally crop the logits to only the top k options
325 if top_k is not None:
326 v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
327 logits[logits < v[:, [-1]]] = -torch.inf
328 # apply softmax to convert logits to (normalized) probabilities
329 probs = torch.nn.functional.softmax(logits, dim=-1)
330 scores[:, i] = probs
331 # sample from the distribution
332 path_next = torch.multinomial(probs, num_samples=1)
333 # append sampled index to the running sequence and continue
334 inputs["input_ids"] = torch.cat((inputs["input_ids"], path_next), dim=1)
336 return GenerationOutputs(sequences=inputs["input_ids"], scores=torch.unbind(scores, dim=1))
339class ExplainablePathLanguageModelingRecommender(PathLanguageModelingRecommender, ExplainableRecommender):
340 """This is an abstract explainable path-language-modeling recommender.
341 All the explainable path-language-modeling model should implement this class.
342 The base explainable path-language-modeling recommender class inherits the path-language-modeling recommender class
343 to learn from knowledge graph paths defined by a chain of entity-relation triplets.
344 """
346 def __init__(self, config, dataset, _skip_nn_module_init=True):
347 super().__init__(config, dataset, _skip_nn_module_init=_skip_nn_module_init)
349 def explain(self, inputs, **kwargs):
350 kwargs["max_length"] = self.token_sequence_length
351 kwargs["min_length"] = self.token_sequence_length
352 outputs = self.generate(inputs, **kwargs)
354 max_new_tokens = self.token_sequence_length - inputs["input_ids"].size(1)
356 scores, sequences = self.sequence_postprocessor.get_sequences(outputs, max_new_tokens=max_new_tokens)
358 for seq in sequences:
359 seq[-1] = self.decode_path(seq[-1])
361 return scores, sequences
363 def decode_path(self, path):
364 """Standardize path format"""
365 new_path = []
366 # Process the path
367 # [BOS] U R I R E/I R I
368 for node_idx in range(1, len(path) + 1, 2):
369 if path[node_idx].startswith(PathLanguageModelingTokenType.USER.token):
370 user_id = int(path[node_idx][1:])
371 if node_idx - 1 == 0:
372 relation = "self_loop"
373 else:
374 relation = int(path[node_idx - 1][1:])
376 new_node = (relation, "user", user_id)
377 elif path[node_idx].startswith(PathLanguageModelingTokenType.ITEM.token):
378 relation = int(path[node_idx - 1][1:])
379 item_id = int(path[node_idx][1:])
380 new_node = (relation, "item", item_id)
381 else:
382 # Is an entity
383 relation = int(path[node_idx - 1][1:])
384 entity_id = int(path[node_idx][1:])
385 new_node = (relation, "entity", entity_id)
386 new_path.append(new_node)
387 return new_path
390class ContextRecommender(AbstractRecommender):
391 """This is a abstract context-aware recommender. All the context-aware model should implement this class.
392 The base context-aware recommender class provide the basic embedding function of feature fields which also
393 contains a first-order part of feature fields.
394 """
396 type = ModelType.CONTEXT
397 input_type = InputType.POINTWISE
399 def __init__(self, config, dataset):
400 super().__init__()
402 self.field_names = dataset.fields(
403 source=[
404 FeatureSource.INTERACTION,
405 FeatureSource.USER,
406 FeatureSource.USER_ID,
407 FeatureSource.ITEM,
408 FeatureSource.ITEM_ID,
409 ]
410 )
411 self.LABEL = config["LABEL_FIELD"]
412 self.embedding_size = config["embedding_size"]
413 self.device = config["device"]
414 self.double_tower = config["double_tower"]
415 self.numerical_features = config["numerical_features"]
416 if self.double_tower is None:
417 self.double_tower = False
418 self.token_field_names = []
419 self.token_field_dims = []
420 self.float_field_names = []
421 self.float_field_dims = []
422 self.token_seq_field_names = []
423 self.token_seq_field_dims = []
424 self.float_seq_field_names = []
425 self.float_seq_field_dims = []
426 self.num_feature_field = 0
428 if self.double_tower:
429 self.user_field_names = dataset.fields(source=[FeatureSource.USER, FeatureSource.USER_ID])
430 self.item_field_names = dataset.fields(source=[FeatureSource.ITEM, FeatureSource.ITEM_ID])
431 self.field_names = self.user_field_names + self.item_field_names
432 self.user_token_field_num = 0
433 self.user_float_field_num = 0
434 self.user_token_seq_field_num = 0
435 for field_name in self.user_field_names:
436 if dataset.field2type[field_name] == FeatureType.TOKEN:
437 self.user_token_field_num += 1
438 elif dataset.field2type[field_name] == FeatureType.TOKEN_SEQ:
439 self.user_token_seq_field_num += 1
440 else:
441 self.user_float_field_num += 1
442 self.item_token_field_num = 0
443 self.item_float_field_num = 0
444 self.item_token_seq_field_num = 0
445 for field_name in self.item_field_names:
446 if dataset.field2type[field_name] == FeatureType.TOKEN:
447 self.item_token_field_num += 1
448 elif dataset.field2type[field_name] == FeatureType.TOKEN_SEQ:
449 self.item_token_seq_field_num += 1
450 else:
451 self.item_float_field_num += 1
453 for field_name in self.field_names:
454 if field_name == self.LABEL:
455 continue
456 if dataset.field2type[field_name] == FeatureType.TOKEN:
457 self.token_field_names.append(field_name)
458 self.token_field_dims.append(dataset.num(field_name))
459 elif dataset.field2type[field_name] == FeatureType.TOKEN_SEQ:
460 self.token_seq_field_names.append(field_name)
461 self.token_seq_field_dims.append(dataset.num(field_name))
462 elif dataset.field2type[field_name] == FeatureType.FLOAT and field_name in self.numerical_features:
463 self.float_field_names.append(field_name)
464 self.float_field_dims.append(dataset.num(field_name))
465 elif dataset.field2type[field_name] == FeatureType.FLOAT_SEQ and field_name in self.numerical_features:
466 self.float_seq_field_names.append(field_name)
467 self.float_seq_field_dims.append(dataset.num(field_name))
468 else:
469 continue
471 self.num_feature_field += 1
472 if len(self.token_field_dims) > 0:
473 self.token_field_offsets = np.array((0, *np.cumsum(self.token_field_dims)[:-1]), dtype=np.long)
474 self.token_embedding_table = FMEmbedding(
475 self.token_field_dims, self.token_field_offsets, self.embedding_size
476 )
477 if len(self.float_field_dims) > 0:
478 self.float_field_offsets = np.array((0, *np.cumsum(self.float_field_dims)[:-1]), dtype=np.long)
479 self.float_embedding_table = FLEmbedding(
480 self.float_field_dims, self.float_field_offsets, self.embedding_size
481 )
482 if len(self.token_seq_field_dims) > 0:
483 self.token_seq_embedding_table = nn.ModuleList()
484 for token_seq_field_dim in self.token_seq_field_dims:
485 self.token_seq_embedding_table.append(nn.Embedding(token_seq_field_dim, self.embedding_size))
486 if len(self.float_seq_field_dims) > 0:
487 self.float_seq_embedding_table = nn.ModuleList()
488 for float_seq_field_dim in self.float_seq_field_dims:
489 self.float_seq_embedding_table.append(nn.Embedding(float_seq_field_dim, self.embedding_size))
491 self.first_order_linear = FMFirstOrderLinear(config, dataset)
493 def embed_float_fields(self, float_fields):
494 """Embed the float feature columns
496 Args:
497 float_fields (torch.FloatTensor): The input dense tensor. shape of [batch_size, num_float_field]
499 Returns:
500 torch.FloatTensor: The result embedding tensor of float columns.
501 """
502 # input Tensor shape : [batch_size, num_float_field]
503 if float_fields is None:
504 return None
505 # [batch_size, num_float_field, embed_dim]
506 float_embedding = self.float_embedding_table(float_fields)
508 return float_embedding
510 def embed_float_seq_fields(self, float_seq_fields, mode="mean"):
511 """Embed the float feature columns
513 Args:
514 float_seq_fields (torch.LongTensor): The input tensor. shape of [batch_size, seq_len]
515 mode (str): How to aggregate the embedding of feature in this field. default=mean
517 Returns:
518 torch.FloatTensor: The result embedding tensor of token sequence columns.
519 """
520 # input is a list of Tensor shape of [batch_size, seq_len, 2]
521 fields_result = []
522 for i, float_seq_field in enumerate(float_seq_fields):
523 embedding_table = self.float_seq_embedding_table[i]
524 base, index = torch.split(float_seq_field, [1, 1], dim=-1)
525 index = index.squeeze(-1)
526 mask = index != 0 # [batch_size, seq_len]
527 mask = mask.float()
528 value_cnt = torch.sum(mask, dim=1, keepdim=True) # [batch_size, 1]
530 float_seq_embedding = base * embedding_table(index.long()) # [batch_size, seq_len, embed_dim]
532 mask = mask.unsqueeze(2).expand_as(float_seq_embedding) # [batch_size, seq_len, embed_dim]
533 if mode == "max":
534 masked_float_seq_embedding = float_seq_embedding - (1 - mask) * 1e9 # [batch_size, seq_len, embed_dim]
535 result = torch.max(masked_float_seq_embedding, dim=1, keepdim=True) # [batch_size, 1, embed_dim]
536 elif mode == "sum":
537 masked_float_seq_embedding = float_seq_embedding * mask.float()
538 result = torch.sum(masked_float_seq_embedding, dim=1, keepdim=True) # [batch_size, 1, embed_dim]
539 else:
540 masked_float_seq_embedding = float_seq_embedding * mask.float()
541 result = torch.sum(masked_float_seq_embedding, dim=1) # [batch_size, embed_dim]
542 eps = torch.FloatTensor([1e-8]).to(self.device)
543 result = torch.div(result, value_cnt + eps) # [batch_size, embed_dim]
544 result = result.unsqueeze(1) # [batch_size, 1, embed_dim]
545 fields_result.append(result)
546 if len(fields_result) == 0:
547 return None
548 else:
549 return torch.cat(fields_result, dim=1) # [batch_size, num_token_seq_field, embed_dim]
551 def embed_token_fields(self, token_fields):
552 """Embed the token feature columns
554 Args:
555 token_fields (torch.LongTensor): The input tensor. shape of [batch_size, num_token_field]
557 Returns:
558 torch.FloatTensor: The result embedding tensor of token columns.
559 """
560 # input Tensor shape : [batch_size, num_token_field]
561 if token_fields is None:
562 return None
563 # [batch_size, num_token_field, embed_dim]
564 token_embedding = self.token_embedding_table(token_fields)
566 return token_embedding
568 def embed_token_seq_fields(self, token_seq_fields, mode="mean"):
569 """Embed the token feature columns
571 Args:
572 token_seq_fields (torch.LongTensor): The input tensor. shape of [batch_size, seq_len]
573 mode (str): How to aggregate the embedding of feature in this field. default=mean
575 Returns:
576 torch.FloatTensor: The result embedding tensor of token sequence columns.
577 """
578 # input is a list of Tensor shape of [batch_size, seq_len]
579 fields_result = []
580 for i, token_seq_field in enumerate(token_seq_fields):
581 embedding_table = self.token_seq_embedding_table[i]
582 mask = token_seq_field != 0 # [batch_size, seq_len]
583 mask = mask.float()
584 value_cnt = torch.sum(mask, dim=1, keepdim=True) # [batch_size, 1]
586 token_seq_embedding = embedding_table(token_seq_field) # [batch_size, seq_len, embed_dim]
588 mask = mask.unsqueeze(2).expand_as(token_seq_embedding) # [batch_size, seq_len, embed_dim]
589 if mode == "max":
590 masked_token_seq_embedding = token_seq_embedding - (1 - mask) * 1e9 # [batch_size, seq_len, embed_dim]
591 result = torch.max(masked_token_seq_embedding, dim=1, keepdim=True) # [batch_size, 1, embed_dim]
592 elif mode == "sum":
593 masked_token_seq_embedding = token_seq_embedding * mask.float()
594 result = torch.sum(masked_token_seq_embedding, dim=1, keepdim=True) # [batch_size, 1, embed_dim]
595 else:
596 masked_token_seq_embedding = token_seq_embedding * mask.float()
597 result = torch.sum(masked_token_seq_embedding, dim=1) # [batch_size, embed_dim]
598 eps = torch.FloatTensor([1e-8]).to(self.device)
599 result = torch.div(result, value_cnt + eps) # [batch_size, embed_dim]
600 result = result.unsqueeze(1) # [batch_size, 1, embed_dim]
601 fields_result.append(result)
602 if len(fields_result) == 0:
603 return None
604 else:
605 return torch.cat(fields_result, dim=1) # [batch_size, num_token_seq_field, embed_dim]
607 def double_tower_embed_input_fields(self, interaction):
608 """Embed the whole feature columns in a double tower way.
610 Args:
611 interaction (Interaction): The input data collection.
613 Returns:
614 torch.FloatTensor: The embedding tensor of token sequence columns in the first part.
615 torch.FloatTensor: The embedding tensor of float sequence columns in the first part.
616 torch.FloatTensor: The embedding tensor of token sequence columns in the second part.
617 torch.FloatTensor: The embedding tensor of float sequence columns in the second part.
619 """
620 if not self.double_tower:
621 raise RuntimeError("Please check your model hyper parameters and set 'double tower' as True")
622 sparse_embedding, dense_embedding = self.embed_input_fields(interaction)
623 if dense_embedding is not None:
624 first_dense_embedding, second_dense_embedding = torch.split(
625 dense_embedding,
626 [self.user_float_field_num, self.item_float_field_num],
627 dim=1,
628 )
629 else:
630 first_dense_embedding, second_dense_embedding = None, None
632 if sparse_embedding is not None:
633 sizes = [
634 self.user_token_seq_field_num,
635 self.item_token_seq_field_num,
636 self.user_token_field_num,
637 self.item_token_field_num,
638 ]
639 (
640 first_token_seq_embedding,
641 second_token_seq_embedding,
642 first_token_embedding,
643 second_token_embedding,
644 ) = torch.split(sparse_embedding, sizes, dim=1)
645 first_sparse_embedding = torch.cat([first_token_seq_embedding, first_token_embedding], dim=1)
646 second_sparse_embedding = torch.cat([second_token_seq_embedding, second_token_embedding], dim=1)
647 else:
648 first_sparse_embedding, second_sparse_embedding = None, None
650 return (
651 first_sparse_embedding,
652 first_dense_embedding,
653 second_sparse_embedding,
654 second_dense_embedding,
655 )
657 def concat_embed_input_fields(self, interaction):
658 sparse_embedding, dense_embedding = self.embed_input_fields(interaction)
659 all_embeddings = []
660 if sparse_embedding is not None:
661 all_embeddings.append(sparse_embedding)
662 if dense_embedding is not None and len(dense_embedding.shape) == 3: # noqa: PLR2004
663 all_embeddings.append(dense_embedding)
664 return torch.cat(all_embeddings, dim=1) # [batch_size, num_field, embed_dim]
666 def embed_input_fields(self, interaction):
667 """Embed the whole feature columns.
669 Args:
670 interaction (Interaction): The input data collection.
672 Returns:
673 torch.FloatTensor: The embedding tensor of token sequence columns.
674 torch.FloatTensor: The embedding tensor of float sequence columns.
675 """
676 float_fields = []
677 for field_name in self.float_field_names:
678 if len(interaction[field_name].shape) == 3: # noqa: PLR2004
679 float_fields.append(interaction[field_name])
680 else:
681 float_fields.append(interaction[field_name].unsqueeze(1))
682 if len(float_fields) > 0:
683 float_fields = torch.cat(float_fields, dim=1) # [batch_size, num_float_field, 2]
684 else:
685 float_fields = None
686 # [batch_size, num_float_field] or [batch_size, num_float_field, embed_dim] or None
687 float_fields_embedding = self.embed_float_fields(float_fields)
689 float_seq_fields = []
690 for field_name in self.float_seq_field_names:
691 float_seq_fields.append(interaction[field_name])
693 float_seq_fields_embedding = self.embed_float_seq_fields(float_seq_fields)
695 if float_fields_embedding is None:
696 dense_embedding = float_seq_fields_embedding
697 elif float_seq_fields_embedding is None:
698 dense_embedding = float_fields_embedding
699 else:
700 dense_embedding = torch.cat([float_seq_fields_embedding, float_fields_embedding], dim=1)
702 token_fields = []
703 for field_name in self.token_field_names:
704 token_fields.append(interaction[field_name].unsqueeze(1))
705 if len(token_fields) > 0:
706 token_fields = torch.cat(token_fields, dim=1) # [batch_size, num_token_field, 2]
707 else:
708 token_fields = None
709 # [batch_size, num_token_field, embed_dim] or None
710 token_fields_embedding = self.embed_token_fields(token_fields)
712 token_seq_fields = []
713 for field_name in self.token_seq_field_names:
714 token_seq_fields.append(interaction[field_name])
715 # [batch_size, num_token_seq_field, embed_dim] or None
716 token_seq_fields_embedding = self.embed_token_seq_fields(token_seq_fields)
718 if token_fields_embedding is None:
719 sparse_embedding = token_seq_fields_embedding
720 elif token_seq_fields_embedding is None:
721 sparse_embedding = token_fields_embedding
722 else:
723 sparse_embedding = torch.cat([token_seq_fields_embedding, token_fields_embedding], dim=1)
725 # sparse_embedding shape: [batch_size, num_token_seq_field+num_token_field, embed_dim] or None
726 # dense_embedding shape: [batch_size, num_float_field, 2] or [batch_size, num_float_field, embed_dim] or None
727 return sparse_embedding, dense_embedding