Coverage for hopwise/model/context_aware_recommender/dssm.py: 96%

53 statements  

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

1# @Time : 2020/9/2 

2# @Author : Yingqian Min 

3# @Email : gmqszyq@qq.com 

4# @File : dssm.py 

5 

6"""DSSM 

7################################################ 

8Reference: 

9 PS Huang et al. "Learning Deep Structured Semantic Models for Web Search using Clickthrough Data" in CIKM 2013. 

10""" 

11 

12import torch 

13from torch import nn 

14from torch.nn.init import constant_, xavier_normal_ 

15 

16from hopwise.model.abstract_recommender import ContextRecommender 

17from hopwise.model.layers import MLPLayers 

18 

19 

20class DSSM(ContextRecommender): 

21 """DSSM respectively expresses user and item as low dimensional vectors with mlp layers, 

22 and uses cosine distance to calculate the distance between the two semantic vectors. 

23 

24 """ 

25 

26 def __init__(self, config, dataset): 

27 super().__init__(config, dataset) 

28 

29 # load parameters info 

30 self.mlp_hidden_size = config["mlp_hidden_size"] 

31 self.dropout_prob = config["dropout_prob"] 

32 

33 self.user_feature_num = self.user_token_field_num + self.user_float_field_num + self.user_token_seq_field_num 

34 self.item_feature_num = self.item_token_field_num + self.item_float_field_num + self.item_token_seq_field_num 

35 user_size_list = [self.embedding_size * self.user_feature_num] + self.mlp_hidden_size 

36 item_size_list = [self.embedding_size * self.item_feature_num] + self.mlp_hidden_size 

37 

38 # define layers and loss 

39 self.user_mlp_layers = MLPLayers(user_size_list, self.dropout_prob, activation="tanh", bn=True) 

40 self.item_mlp_layers = MLPLayers(item_size_list, self.dropout_prob, activation="tanh", bn=True) 

41 

42 self.loss = nn.BCEWithLogitsLoss() 

43 self.sigmoid = nn.Sigmoid() 

44 

45 # parameters initialization 

46 self.apply(self._init_weights) 

47 

48 def _init_weights(self, module): 

49 if isinstance(module, nn.Embedding): 

50 xavier_normal_(module.weight.data) 

51 elif isinstance(module, nn.Linear): 

52 xavier_normal_(module.weight.data) 

53 if module.bias is not None: 

54 constant_(module.bias.data, 0) 

55 

56 def forward(self, interaction): 

57 # user_sparse_embedding shape: [batch_size, user_token_seq_field_num + user_token_field_num , embed_dim] or None # noqa: E501 

58 # user_dense_embedding shape: [batch_size, user_float_field_num] or [batch_size, user_float_field_num, embed_dim] or None # noqa: E501 

59 # item_sparse_embedding shape: [batch_size, item_token_seq_field_num + item_token_field_num , embed_dim] or None # noqa: E501 

60 # item_dense_embedding shape: [batch_size, item_float_field_num] or [batch_size, item_float_field_num, embed_dim] or None # noqa: E501 

61 embed_result = self.double_tower_embed_input_fields(interaction) 

62 user_sparse_embedding, user_dense_embedding = embed_result[:2] 

63 item_sparse_embedding, item_dense_embedding = embed_result[2:] 

64 

65 user = [] 

66 if user_sparse_embedding is not None: 

67 user.append(user_sparse_embedding) 

68 if user_dense_embedding is not None and len(user_dense_embedding.shape) == 3: # noqa: PLR2004 

69 user.append(user_dense_embedding) 

70 

71 embed_user = torch.cat(user, dim=1) 

72 

73 item = [] 

74 if item_sparse_embedding is not None: 

75 item.append(item_sparse_embedding) 

76 if item_dense_embedding is not None and len(item_dense_embedding.shape) == 3: # noqa: PLR2004 

77 item.append(item_dense_embedding) 

78 

79 embed_item = torch.cat(item, dim=1) 

80 

81 batch_size = embed_item.shape[0] 

82 user_dnn_out = self.user_mlp_layers(embed_user.view(batch_size, -1)) 

83 item_dnn_out = self.item_mlp_layers(embed_item.view(batch_size, -1)) 

84 score = torch.cosine_similarity(user_dnn_out, item_dnn_out, dim=1) 

85 return score.squeeze(-1) 

86 

87 def calculate_loss(self, interaction): 

88 label = interaction[self.LABEL] 

89 output = self.forward(interaction) 

90 return self.loss(output, label) 

91 

92 def predict(self, interaction): 

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