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
« 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.
4# This source code is licensed under the license found in the
5# LICENSE file in the root directory of this source tree.
7# UPDATE
8# @Time : 2025/02/19
9# @Author : Alessandro Soccol
10# @Email : alessandro.soccol@unica.it
12r"""RPG
13################################################
14 Reference:
15 Hou Yupeng et al. "Generating Long Semantic IDs in Parallel for Recommendation".
16"""
18import torch
19import torch.nn.functional as F
20from torch import nn
22from hopwise.model.abstract_recommender import SequentialRecommender
23from hopwise.model.layers import ResidualBlock
26class RPG(SequentialRecommender):
27 r"""RPG is a recommendation model that generates each token of the next semantic ID in parallel."""
29 def __init__(self, config, dataset):
30 from transformers import GPT2Config, GPT2Model
32 super().__init__(config, dataset)
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
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.")
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
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 = []
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)
82 self.loss = torch.nn.CrossEntropyLoss()
84 # item-item graph used for graph-constrained decoding, built once per evaluation
85 self.adjacency = None
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)
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)
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
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]
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)
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]
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
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)
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)
161 # Do not consider PAD token and EOS token.
162 token_emb = self.gpt2.wte.weight[1:-1]
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)
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)
188 return scores
190 def graph_propagation(self, token_logits):
191 batch_size = token_logits.shape[0]
193 if self.adjacency is None:
194 self.adjacency = self.init_graph()
195 adjacency = self.adjacency
197 results = torch.full((batch_size, self.n_items), -torch.inf, device=self.device)
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 )
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)
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)
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
226 next_nodes.append(neighbors_in_batch[topk])
228 topk_nodes_sorted = torch.stack(next_nodes, dim=0)
230 return results
232 @torch.no_grad()
233 def init_graph(self):
234 """Builds the item-item graph used for graph-constrained decoding.
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.
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
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)
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
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]
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
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
269 return adjacency