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

1# @Time : 2020/6/25 

2# @Author : Shanlei Mu 

3# @Email : slmu@ruc.edu.cn 

4 

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 

9 

10"""hopwise.model.abstract_recommender 

11################################## 

12""" 

13 

14from logging import getLogger 

15 

16import numpy as np 

17import torch 

18from torch import nn 

19 

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) 

34 

35 

36class AbstractRecommender(nn.Module): 

37 r"""Base class for all models""" 

38 

39 def __init__(self, _skip_nn_module_init=False): 

40 self.logger = getLogger() 

41 

42 if not _skip_nn_module_init: 

43 super().__init__() 

44 

45 def calculate_loss(self, interaction): 

46 r"""Calculate the training loss for a batch data. 

47 

48 Args: 

49 interaction (Interaction): Interaction class of the batch. 

50 

51 Returns: 

52 torch.Tensor: Training loss, shape: [] 

53 """ 

54 raise NotImplementedError 

55 

56 def predict(self, interaction): 

57 r"""Predict the scores between users and items. 

58 

59 Args: 

60 interaction (Interaction): Interaction class of the batch. 

61 

62 Returns: 

63 torch.Tensor: Predicted scores for given users and items, shape: [batch_size] 

64 """ 

65 raise NotImplementedError 

66 

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. 

70 

71 Args: 

72 interaction (Interaction): Interaction class of the batch. 

73 

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 

79 

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. 

83 

84 Args: 

85 interaction (Interaction): Interaction class of the batch. 

86 

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 

92 

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

97 

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) 

103 

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

109 

110 

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

115 

116 type = ModelType.GENERAL 

117 

118 def __init__(self, config, dataset): 

119 super().__init__() 

120 

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) 

127 

128 # load parameters info 

129 self.device = config["device"] 

130 

131 

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

137 

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) 

142 

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. 

145 

146 Args: 

147 user (torch.LongTensor): The input tensor that contains user's id, shape: [batch_size, ] 

148 

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 

158 

159 

160class SequentialRecommender(AbstractRecommender): 

161 """This is a abstract sequential recommender. All the sequential model should implement This class.""" 

162 

163 type = ModelType.SEQUENTIAL 

164 

165 def __init__(self, config, dataset): 

166 super().__init__() 

167 

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) 

177 

178 # load parameters info 

179 self.device = config["device"] 

180 

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) 

186 

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 

195 

196 

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

201 

202 type = ModelType.KNOWLEDGE 

203 

204 def __init__(self, config, dataset, _skip_nn_module_init=False): 

205 super().__init__(_skip_nn_module_init=_skip_nn_module_init) 

206 

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) 

220 

221 # load parameters info 

222 if not _skip_nn_module_init: 

223 self.device = config["device"] 

224 

225 

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. 

229 

230 """ 

231 

232 def explain(self, interaction): 

233 r""" 

234 Explain the prediction function. 

235 

236 Given users, calculate the scores and paths between users and all candidate items, 

237 then return the templates filled with path data. 

238 

239 Args: 

240 interaction (Interaction): The interaction batch. 

241 

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

249 

250 def decode_path(self, path): 

251 r""" 

252 Decode the path into a string. Path decoding is specific to each model. 

253 

254 Args: 

255 path (list): The path data. 

256 

257 Returns: 

258 str: The decoded path string. 

259 """ 

260 raise NotImplementedError("decode_path is not implemented") 

261 

262 

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

269 

270 type = ModelType.PATH_LANGUAGE_MODELING 

271 input_type = InputType.PATHWISE 

272 

273 def __init__(self, config, dataset, _skip_nn_module_init=True): 

274 super().__init__(config, dataset, _skip_nn_module_init=_skip_nn_module_init) 

275 

276 self.n_tokens = len(dataset.tokenizer) 

277 self.token_sequence_length = dataset.token_sequence_length - 1 # EOS token is not included 

278 

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

287 

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 ) 

294 

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. 

301 

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) 

311 

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 

320 

321 # KGCD 

322 logits = self.logits_processor_list(inputs["input_ids"], logits) 

323 

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) 

335 

336 return GenerationOutputs(sequences=inputs["input_ids"], scores=torch.unbind(scores, dim=1)) 

337 

338 

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

345 

346 def __init__(self, config, dataset, _skip_nn_module_init=True): 

347 super().__init__(config, dataset, _skip_nn_module_init=_skip_nn_module_init) 

348 

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) 

353 

354 max_new_tokens = self.token_sequence_length - inputs["input_ids"].size(1) 

355 

356 scores, sequences = self.sequence_postprocessor.get_sequences(outputs, max_new_tokens=max_new_tokens) 

357 

358 for seq in sequences: 

359 seq[-1] = self.decode_path(seq[-1]) 

360 

361 return scores, sequences 

362 

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

375 

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 

388 

389 

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

395 

396 type = ModelType.CONTEXT 

397 input_type = InputType.POINTWISE 

398 

399 def __init__(self, config, dataset): 

400 super().__init__() 

401 

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 

427 

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 

452 

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 

470 

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

490 

491 self.first_order_linear = FMFirstOrderLinear(config, dataset) 

492 

493 def embed_float_fields(self, float_fields): 

494 """Embed the float feature columns 

495 

496 Args: 

497 float_fields (torch.FloatTensor): The input dense tensor. shape of [batch_size, num_float_field] 

498 

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) 

507 

508 return float_embedding 

509 

510 def embed_float_seq_fields(self, float_seq_fields, mode="mean"): 

511 """Embed the float feature columns 

512 

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 

516 

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] 

529 

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

531 

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] 

550 

551 def embed_token_fields(self, token_fields): 

552 """Embed the token feature columns 

553 

554 Args: 

555 token_fields (torch.LongTensor): The input tensor. shape of [batch_size, num_token_field] 

556 

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) 

565 

566 return token_embedding 

567 

568 def embed_token_seq_fields(self, token_seq_fields, mode="mean"): 

569 """Embed the token feature columns 

570 

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 

574 

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] 

585 

586 token_seq_embedding = embedding_table(token_seq_field) # [batch_size, seq_len, embed_dim] 

587 

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] 

606 

607 def double_tower_embed_input_fields(self, interaction): 

608 """Embed the whole feature columns in a double tower way. 

609 

610 Args: 

611 interaction (Interaction): The input data collection. 

612 

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. 

618 

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 

631 

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 

649 

650 return ( 

651 first_sparse_embedding, 

652 first_dense_embedding, 

653 second_sparse_embedding, 

654 second_dense_embedding, 

655 ) 

656 

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] 

665 

666 def embed_input_fields(self, interaction): 

667 """Embed the whole feature columns. 

668 

669 Args: 

670 interaction (Interaction): The input data collection. 

671 

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) 

688 

689 float_seq_fields = [] 

690 for field_name in self.float_seq_field_names: 

691 float_seq_fields.append(interaction[field_name]) 

692 

693 float_seq_fields_embedding = self.embed_float_seq_fields(float_seq_fields) 

694 

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) 

701 

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) 

711 

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) 

717 

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) 

724 

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