Coverage for hopwise/model/sequential_recommender/rpg.py: 94%

125 statements  

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

1# Copyright (c) Meta Platforms, Inc. and affiliates. 

2# All rights reserved. 

3 

4# This source code is licensed under the license found in the 

5# LICENSE file in the root directory of this source tree. 

6 

7# UPDATE 

8# @Time : 2025/02/19 

9# @Author : Alessandro Soccol 

10# @Email : alessandro.soccol@unica.it 

11 

12r"""RPG 

13################################################ 

14 Reference: 

15 Hou Yupeng et al. "Generating Long Semantic IDs in Parallel for Recommendation". 

16""" 

17 

18import torch 

19import torch.nn.functional as F 

20from torch import nn 

21 

22from hopwise.model.abstract_recommender import SequentialRecommender 

23from hopwise.model.layers import ResidualBlock 

24 

25 

26class RPG(SequentialRecommender): 

27 r"""RPG is a recommendation model that generates each token of the next semantic ID in parallel.""" 

28 

29 def __init__(self, config, dataset): 

30 from transformers import GPT2Config, GPT2Model 

31 

32 super().__init__(config, dataset) 

33 

34 self.topk = config["topk"] 

35 self.use_gcd = config["use_gcd"] 

36 self.codebook_size = config["codebook_size"] 

37 self.embedding_size = config["embedding_size"] 

38 self.temperature = config["temperature"] 

39 self.chunk_size = config["chunk_size"] 

40 self.n_edges = config["n_edges"] 

41 self.n_beams = config["n_beams"] 

42 self.propagation_steps = config["propagation_steps"] 

43 self.loss_type = config["loss_type"] # necessary otherwise hopwise don't recognize seq recommender 

44 

45 if self.use_gcd: 

46 # the neighbors of the best beam are n_edges distinct items, so there are always enough beams to select 

47 if self.n_beams > self.n_edges: 

48 raise ValueError(f"n_beams [{self.n_beams}] should not be greater than n_edges [{self.n_edges}].") 

49 if self.n_edges > self.n_items - 1: 

50 raise ValueError(f"n_edges [{self.n_edges}] should not be greater than the number of items.") 

51 

52 # registered as a buffer so reloading never depends on FAISS reproducing the same codes 

53 self.register_buffer("item2shifted_sem_id", dataset.item2shifted_sem_id.clone()) 

54 self.n_digit = dataset.n_digit 

55 self.codebook_size = dataset.codebook_size 

56 

57 gpt2config = GPT2Config( 

58 vocab_size=dataset.vocab_size, 

59 n_positions=config["MAX_ITEM_LIST_LENGTH"], 

60 n_embd=config["embedding_size"], 

61 n_layer=config["layers"], 

62 n_head=config["heads"], 

63 n_inner=config["embedding_size_inner_mlp"], 

64 activation_function=config["activation_function"], 

65 resid_pdrop=config["resid_pdrop"], 

66 embd_pdrop=config["embd_pdrop"], 

67 attn_pdrop=config["attn_pdrop"], 

68 layer_norm_epsilon=config["layer_norm_epsilon"], 

69 initializer_range=config["initializer_range"], 

70 eos_token_id=dataset.eos_token, 

71 ) 

72 self.gpt2 = GPT2Model(gpt2config) 

73 # Number of values in a semantic id. 

74 self.n_pred_head = dataset.n_digit 

75 pred_head_list = [] 

76 

77 # Create prediction heads with a Residual Connection 

78 for _ in range(self.n_pred_head): 

79 pred_head_list.append(ResidualBlock(config["embedding_size"])) 

80 self.pred_heads = nn.Sequential(*pred_head_list) 

81 

82 self.loss = torch.nn.CrossEntropyLoss() 

83 

84 # item-item graph used for graph-constrained decoding, built once per evaluation 

85 self.adjacency = None 

86 

87 def load_state_dict(self, *args, **kwargs): 

88 # the decoding graph depends on the token embeddings, so it must be rebuilt with the loaded weights 

89 self.adjacency = None 

90 return super().load_state_dict(*args, **kwargs) 

91 

92 def train(self, mode=True): 

93 # model weights change during training, so the decoding graph must be rebuilt at the next evaluation 

94 if mode: 

95 self.adjacency = None 

96 return super().train(mode) 

97 

98 def forward(self, item_seq): 

99 input_tokens = self.item2shifted_sem_id[item_seq] 

100 attention_mask = (item_seq != 0).long() 

101 # aggregate semantic ids embeddings averaging embeddings for each item 

102 wte = self.gpt2.wte(input_tokens).mean(dim=-2) 

103 outputs = self.gpt2(inputs_embeds=wte, attention_mask=attention_mask) 

104 # outputs.last_hidden_state: shape (bs, seq_len(50), embedding_size) 

105 heads_final_states = [ 

106 self.pred_heads[i](outputs.last_hidden_state).unsqueeze(-2) for i in range(self.n_pred_head) 

107 ] # bs, 50, 1, 448 

108 heads_final_states = torch.cat(heads_final_states, dim=-2) # bs,50,32,448 

109 return heads_final_states 

110 

111 def calculate_loss(self, interaction): 

112 item_seq = interaction[self.ITEM_SEQ] 

113 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

114 pos_items = interaction[self.POS_ITEM_ID] 

115 

116 # Each augmented sequence is supervised only on its target item, predicted from the last position. 

117 # It is equivalent to the original implementation, which supervises every position of the first 

118 # max_seq_length items at once and only the last position of the following sliding windows. 

119 # shape: (bs, seq_len, n_pred_head (semantic id size), embedding_size) 

120 hidden_states = self.forward(item_seq) 

121 selected_states = hidden_states.gather( 

122 dim=1, index=(item_seq_len - 1).view(-1, 1, 1, 1).expand(-1, 1, self.n_pred_head, self.embedding_size) 

123 ).squeeze(1) # shape: (bs, n_pred_head, embedding_size) 

124 selected_states = F.normalize(selected_states, dim=-1) 

125 selected_states = torch.chunk(selected_states, self.n_pred_head, dim=1) 

126 token_emb = self.gpt2.wte.weight[1:-1] # vocab_size, emb_size -> 8192, 448 

127 token_emb = F.normalize(token_emb, dim=-1) 

128 

129 token_embs = torch.chunk(token_emb, self.n_pred_head, dim=0) 

130 # calculate the output of each head 

131 token_logits = [ 

132 torch.matmul(selected_states[i].squeeze(dim=1), token_embs[i].T) / self.temperature 

133 for i in range(self.n_pred_head) 

134 ] 

135 # convert each item to the corresponding semantic id 

136 token_labels = self.item2shifted_sem_id[pos_items] 

137 

138 # aggregate loss over the prediction heads 

139 losses = [ 

140 self.loss(token_logits[i], token_labels[:, i] - i * self.codebook_size - 1) 

141 for i in range(self.n_pred_head) 

142 ] 

143 loss = torch.mean(torch.stack(losses)) 

144 return loss 

145 

146 def predict(self, interaction): 

147 """Predict scores for the next item in the sequence. Used only in GFLOPS fn""" 

148 test_item = interaction[self.ITEM_ID] 

149 scores = self.full_sort_predict(interaction) 

150 return scores.gather(dim=1, index=test_item.unsqueeze(1)).squeeze(1) 

151 

152 def full_sort_predict(self, interaction): 

153 item_seq = interaction[self.ITEM_SEQ] 

154 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

155 hidden_states = self.forward(item_seq) 

156 hidden_states = hidden_states.gather( 

157 dim=1, index=(item_seq_len - 1).view(-1, 1, 1, 1).expand(-1, 1, self.n_pred_head, self.embedding_size) 

158 ) 

159 hidden_states = F.normalize(hidden_states, dim=-1) 

160 

161 # Do not consider PAD token and EOS token. 

162 token_emb = self.gpt2.wte.weight[1:-1] 

163 

164 token_emb = F.normalize(token_emb, dim=-1) 

165 token_embs = torch.chunk(token_emb, self.n_pred_head, dim=0) 

166 logits = [ 

167 torch.matmul(hidden_states[:, 0, i, :], token_embs[i].T) / self.temperature 

168 for i in range(self.n_pred_head) 

169 ] 

170 # create probability distribution 

171 logits = [F.log_softmax(logit, dim=-1) for logit in logits] 

172 token_logits = torch.cat(logits, dim=-1) # (batch_size, n_tokens) 

173 

174 if self.use_gcd: 

175 scores = self.graph_propagation(token_logits=token_logits) 

176 else: 

177 scores = torch.gather( 

178 # (batch_size, n_items, n_tokens) 

179 input=token_logits.unsqueeze(-2).expand(-1, self.n_items, -1), 

180 dim=-1, 

181 # (batch_size, n_items, code_dim) 

182 index=(self.item2shifted_sem_id[1:, :] - 1).unsqueeze(0).expand(token_logits.shape[0], -1, -1), 

183 ).mean(dim=-1) 

184 # account for PAD 

185 padding = torch.full((item_seq.size(0), 1), -torch.inf, device=item_seq.device) 

186 scores = torch.cat([padding, scores], dim=1) 

187 

188 return scores 

189 

190 def graph_propagation(self, token_logits): 

191 batch_size = token_logits.shape[0] 

192 

193 if self.adjacency is None: 

194 self.adjacency = self.init_graph() 

195 adjacency = self.adjacency 

196 

197 results = torch.full((batch_size, self.n_items), -torch.inf, device=self.device) 

198 

199 # Randomly sample n_beams item ids in [1, n_items) as initial beams 

200 topk_nodes_sorted = torch.randint( 

201 1, self.n_items, (batch_size, self.n_beams), dtype=torch.long, device=token_logits.device 

202 ) 

203 

204 for propagation_step in range(self.propagation_steps): 

205 # Find the neighbors of the current beams. The adjacency list is indexed by item id 

206 all_neighbors = adjacency[topk_nodes_sorted].view(batch_size, -1) 

207 

208 next_nodes = [] 

209 for batch_id in range(batch_size): 

210 neighbors_in_batch = torch.unique(all_neighbors[batch_id]) 

211 # scores for neighbors 

212 scores = torch.gather( 

213 input=token_logits[batch_id].unsqueeze(0).expand(neighbors_in_batch.shape[0], -1), 

214 dim=-1, 

215 index=(self.item2shifted_sem_id[neighbors_in_batch] - 1), 

216 ).mean(dim=-1) 

217 

218 # if it's the last propagation step, save the scores 

219 if propagation_step == self.propagation_steps - 1: 

220 topk = torch.topk(scores, min(max(self.topk), scores.shape[0])).indices 

221 results[batch_id, neighbors_in_batch[topk]] = scores[topk] 

222 else: 

223 # otherwise, select beams and propagate again 

224 topk = torch.topk(scores, self.n_beams).indices 

225 

226 next_nodes.append(neighbors_in_batch[topk]) 

227 

228 topk_nodes_sorted = torch.stack(next_nodes, dim=0) 

229 

230 return results 

231 

232 @torch.no_grad() 

233 def init_graph(self): 

234 """Builds the item-item graph used for graph-constrained decoding. 

235 

236 The similarity of two items is the average over the digits of the cosine similarity, rescaled in [0, 1], 

237 between the token embeddings of their semantic IDs. Each item is connected to its ``n_edges`` most similar 

238 items. Similarities are computed in chunks of ``chunk_size`` items, such that the full ``n_items x n_items`` 

239 similarity matrix is never materialized. 

240 

241 Returns: 

242 torch.Tensor: The adjacency list of shape ``[n_items, n_edges]``. Row 0 (PAD) is not used. 

243 """ 

244 device = self.gpt2.wte.weight.device 

245 

246 # token embeddings of each digit, ignoring PAD and EOS tokens. shape: (n_digit, codebook_size, d) 

247 wte = F.normalize(self.gpt2.wte.weight[1:-1].view(self.n_digit, self.codebook_size, -1), dim=-1) 

248 # pairwise similarities between the codewords of each digit, from [-1, 1] to [0, 1] 

249 # shape: (n_digit, codebook_size, codebook_size) 

250 token_sims = 0.5 * (torch.bmm(wte, wte.transpose(1, 2)) + 1.0) 

251 

252 # codeword index of each digit for each item, excluding PAD. shape: (n_items - 1, n_digit) 

253 digit_offsets = torch.arange(self.n_digit, device=device) * self.codebook_size + 1 

254 codes = self.item2shifted_sem_id[1:].to(device) - digit_offsets 

255 

256 adjacency = torch.zeros((self.n_items, self.n_edges), dtype=torch.long, device=device) 

257 for i_start in range(0, codes.shape[0], self.chunk_size): 

258 codes_i = codes[i_start : i_start + self.chunk_size] 

259 

260 # average similarity between the items of the chunk and all the items. shape: (chunk_size, n_items - 1) 

261 item_sims = torch.zeros((codes_i.shape[0], codes.shape[0]), device=device) 

262 for k in range(self.n_digit): 

263 item_sims += token_sims[k].index_select(0, codes_i[:, k]).index_select(1, codes[:, k]) 

264 item_sims /= self.n_digit 

265 

266 # column indices are shifted by 1 to obtain item ids, so PAD can never be a neighbor 

267 adjacency[i_start + 1 : i_start + 1 + codes_i.shape[0]] = torch.topk(item_sims, k=self.n_edges).indices + 1 

268 

269 return adjacency