Coverage for hopwise/model/general_recommender/line.py: 91%

101 statements  

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

1# @Time : 2020/12/8 

2# @Author : Yihong Guo 

3# @Email : gyihong@hotmail.com 

4 

5r"""LINE 

6################################################ 

7Reference: 

8 Jian Tang et al. "LINE: Large-scale Information Network Embedding." in WWW 2015. 

9 

10Reference code: 

11 https://github.com/shenweichen/GraphEmbedding 

12""" 

13 

14import random 

15 

16import numpy as np 

17import torch 

18from torch import nn 

19 

20from hopwise.model.abstract_recommender import GeneralRecommender 

21from hopwise.model.init import xavier_normal_initialization 

22from hopwise.utils import InputType 

23 

24 

25class NegSamplingLoss(nn.Module): 

26 def __init__(self): 

27 super().__init__() 

28 

29 def forward(self, sign, score): 

30 return -torch.mean(torch.log(torch.sigmoid(sign * score))) 

31 

32 

33class LINE(GeneralRecommender): 

34 r"""LINE is a graph embedding model. 

35 

36 We implement the model to train users and items embedding for recommendation. 

37 """ 

38 

39 input_type = InputType.PAIRWISE 

40 

41 def __init__(self, config, dataset): 

42 super().__init__(config, dataset) 

43 

44 self.embedding_size = config["embedding_size"] 

45 self.order = config["order"] 

46 self.second_order_loss_weight = config["second_order_loss_weight"] 

47 

48 self.interaction_feat = dataset.inter_feat 

49 

50 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size) 

51 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size) 

52 

53 if self.order == 2: # noqa: PLR2004 

54 self.user_context_embedding = nn.Embedding(self.n_users, self.embedding_size) 

55 self.item_context_embedding = nn.Embedding(self.n_items, self.embedding_size) 

56 

57 self.loss_fct = NegSamplingLoss() 

58 

59 self.used_ids = dataset.get_item_used_ids() 

60 self.random_list = self.get_user_id_list() 

61 np.random.shuffle(self.random_list) 

62 self.random_pr = 0 

63 self.random_list_length = len(self.random_list) 

64 

65 self.apply(xavier_normal_initialization) 

66 

67 def sampler(self, key_ids): 

68 key_ids = np.array(key_ids.cpu()) 

69 key_num = len(key_ids) 

70 total_num = key_num 

71 value_ids = np.zeros(total_num, dtype=np.int64) 

72 check_list = np.arange(total_num) 

73 key_ids = np.tile(key_ids, 1) 

74 while len(check_list) > 0: 

75 value_ids[check_list] = self.random_num(len(check_list)) 

76 check_list = np.array( 

77 [ 

78 i 

79 for i, used, v in zip( 

80 check_list, 

81 self.used_ids[key_ids[check_list]], 

82 value_ids[check_list], 

83 ) 

84 if v in used 

85 ] 

86 ) 

87 

88 return torch.tensor(value_ids, device=self.device) 

89 

90 def random_num(self, num): 

91 value_id = [] 

92 self.random_pr %= self.random_list_length 

93 while True: 

94 if self.random_pr + num <= self.random_list_length: 

95 value_id.append(self.random_list[self.random_pr : self.random_pr + num]) 

96 self.random_pr += num 

97 break 

98 else: 

99 value_id.append(self.random_list[self.random_pr :]) 

100 num -= self.random_list_length - self.random_pr 

101 self.random_pr = 0 

102 np.random.shuffle(self.random_list) 

103 return np.concatenate(value_id) 

104 

105 def get_user_id_list(self): 

106 return np.arange(1, self.n_users) 

107 

108 def forward(self, h, t): 

109 h_embedding = self.user_embedding(h) 

110 t_embedding = self.item_embedding(t) 

111 

112 return torch.sum(h_embedding.mul(t_embedding), dim=1) 

113 

114 def context_forward(self, h, t, field): 

115 if field == "uu": 

116 h_embedding = self.user_embedding(h) 

117 t_embedding = self.item_context_embedding(t) 

118 else: 

119 h_embedding = self.item_embedding(h) 

120 t_embedding = self.user_context_embedding(t) 

121 

122 return torch.sum(h_embedding.mul(t_embedding), dim=1) 

123 

124 def calculate_loss(self, interaction): 

125 user = interaction[self.USER_ID] 

126 pos_item = interaction[self.ITEM_ID] 

127 neg_item = interaction[self.NEG_ITEM_ID] 

128 

129 score_pos = self.forward(user, pos_item) 

130 

131 ones = torch.ones(len(score_pos), device=self.device) 

132 

133 if self.order == 1: 

134 if random.random() < 0.5: # noqa: PLR2004 

135 score_neg = self.forward(user, neg_item) 

136 else: 

137 neg_user = self.sampler(pos_item) 

138 score_neg = self.forward(neg_user, pos_item) 

139 return self.loss_fct(ones, score_pos) + self.loss_fct(-1 * ones, score_neg) 

140 

141 else: 

142 # randomly train i-i relation and u-u relation with u-i relation 

143 if random.random() < 0.5: # noqa: PLR2004 

144 score_neg = self.forward(user, neg_item) 

145 score_pos_con = self.context_forward(user, pos_item, "uu") 

146 score_neg_con = self.context_forward(user, neg_item, "uu") 

147 else: 

148 # sample negative user for item 

149 neg_user = self.sampler(pos_item) 

150 score_neg = self.forward(neg_user, pos_item) 

151 score_pos_con = self.context_forward(pos_item, user, "ii") 

152 score_neg_con = self.context_forward(pos_item, neg_user, "ii") 

153 

154 return ( 

155 self.loss_fct(ones, score_pos) 

156 + self.loss_fct(-1 * ones, score_neg) 

157 + self.loss_fct(ones, score_pos_con) * self.second_order_loss_weight 

158 + self.loss_fct(-1 * ones, score_neg_con) * self.second_order_loss_weight 

159 ) 

160 

161 def predict(self, interaction): 

162 user = interaction[self.USER_ID] 

163 item = interaction[self.ITEM_ID] 

164 

165 scores = self.forward(user, item) 

166 

167 return scores 

168 

169 def full_sort_predict(self, interaction): 

170 user = interaction[self.USER_ID] 

171 

172 # get user embedding from storage variable 

173 u_embeddings = self.user_embedding(user) 

174 i_embedding = self.item_embedding.weight 

175 # dot with all item embedding to accelerate 

176 scores = torch.matmul(u_embeddings, i_embedding.transpose(0, 1)) 

177 

178 return scores.view(-1)