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
« 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
5r"""LINE
6################################################
7Reference:
8 Jian Tang et al. "LINE: Large-scale Information Network Embedding." in WWW 2015.
10Reference code:
11 https://github.com/shenweichen/GraphEmbedding
12"""
14import random
16import numpy as np
17import torch
18from torch import nn
20from hopwise.model.abstract_recommender import GeneralRecommender
21from hopwise.model.init import xavier_normal_initialization
22from hopwise.utils import InputType
25class NegSamplingLoss(nn.Module):
26 def __init__(self):
27 super().__init__()
29 def forward(self, sign, score):
30 return -torch.mean(torch.log(torch.sigmoid(sign * score)))
33class LINE(GeneralRecommender):
34 r"""LINE is a graph embedding model.
36 We implement the model to train users and items embedding for recommendation.
37 """
39 input_type = InputType.PAIRWISE
41 def __init__(self, config, dataset):
42 super().__init__(config, dataset)
44 self.embedding_size = config["embedding_size"]
45 self.order = config["order"]
46 self.second_order_loss_weight = config["second_order_loss_weight"]
48 self.interaction_feat = dataset.inter_feat
50 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
51 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size)
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)
57 self.loss_fct = NegSamplingLoss()
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)
65 self.apply(xavier_normal_initialization)
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 )
88 return torch.tensor(value_ids, device=self.device)
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)
105 def get_user_id_list(self):
106 return np.arange(1, self.n_users)
108 def forward(self, h, t):
109 h_embedding = self.user_embedding(h)
110 t_embedding = self.item_embedding(t)
112 return torch.sum(h_embedding.mul(t_embedding), dim=1)
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)
122 return torch.sum(h_embedding.mul(t_embedding), dim=1)
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]
129 score_pos = self.forward(user, pos_item)
131 ones = torch.ones(len(score_pos), device=self.device)
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)
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")
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 )
161 def predict(self, interaction):
162 user = interaction[self.USER_ID]
163 item = interaction[self.ITEM_ID]
165 scores = self.forward(user, item)
167 return scores
169 def full_sort_predict(self, interaction):
170 user = interaction[self.USER_ID]
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))
178 return scores.view(-1)