Coverage for hopwise/model/sequence_postprocessor.py: 71%
116 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 : 2025/06
2# @Author : Giacomo Medda, Alessandro Soccol
3# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it
5"""hopwise.model.postprocessor
6#######################
7Common post-processors for path sequences in path language modeling recommender systems.
8"""
10from collections import defaultdict
12import torch
14from hopwise.utils import PathLanguageModelingTokenType
17class BaseSequencePostProcessor:
18 """
19 Base class for sequence score post-processors.
20 """
22 def __init__(self, tokenizer, used_ids, item_num, topk=10):
23 self.tokenizer = tokenizer
24 self.used_ids = used_ids
25 self.item_num = item_num
26 self.topk = topk
28 def get_sequences(self, generation_outputs, max_new_tokens=24):
29 """
30 This method should be implemented by subclasses to extract sequences and their scores.
32 Args:
33 generation_outputs: A mapping containing the generated sequences and their scores.
34 max_new_tokens: The maximum number of new tokens to consider for scoring.
36 """
37 raise NotImplementedError("Subclasses must implement this method.")
39 def parse_sequences(self, user_index, sequences, sequences_scores):
40 """
41 Parses the sequences and scores to extract user IDs, recommended items, and their scores.
43 Args:
44 user_index (torch.Tensor): A tensor containing user indices.
45 sequences (torch.Tensor): A tensor containing the generated sequences.
46 sequences_scores (torch.Tensor): A tensor containing the scores for each sequence.
48 Returns:
49 tuple: A tuple containing:
51 - scores (torch.Tensor): A tensor of shape (user_num, item_num) containing the scores.
52 - user_topk_sequences (list): A list of lists with
53 [user_id, recommended_item, score, decoded_sequence].
54 """
55 user_num = user_index.unique().numel()
56 scores = torch.full((user_num, self.item_num), -torch.inf)
57 user_topk_sequences = list()
59 for batch_uidx, sequence, sequence_score in zip(user_index, sequences, sequences_scores):
60 parsed_seq = self._parse_single_sequence(scores, batch_uidx, sequence)
61 if parsed_seq is None:
62 continue
63 uid, recommended_item, decoded_seq = parsed_seq
65 scores[batch_uidx, recommended_item] = sequence_score
66 user_topk_sequences.append([uid, recommended_item, sequence_score.item(), decoded_seq])
68 return scores, user_topk_sequences
70 def _parse_single_sequence(self, scores, batch_uidx, sequence):
71 """Parses a single sequence to extract user ID, recommended item, and the decoded sequence."""
72 seq = self.tokenizer.decode(sequence).split(" ")
74 uid_token = seq[1]
75 recommended_token = seq[-1]
77 if (
78 not (
79 uid_token.startswith(PathLanguageModelingTokenType.USER.token)
80 and recommended_token.startswith(PathLanguageModelingTokenType.ITEM.token)
81 )
82 or recommended_token == self.tokenizer.pad_token
83 ):
84 return
86 uid = int(uid_token[1:])
87 recommended_item = int(recommended_token[1:])
89 if torch.isfinite(scores[batch_uidx, recommended_item]) or recommended_item in self.used_ids[uid]:
90 return
92 return uid, recommended_item, seq
95class SequencePostProcessorLP:
96 def __init__(self, tokenizer, kg_positives, K=10, max_new_tokens=24):
97 self.tokenizer = tokenizer
98 self.kg_positives = kg_positives
99 self.topk = defaultdict(list)
100 self.topk_sequences = defaultdict(list)
101 self.max_new_tokens = max_new_tokens
102 self.K = K
104 def update_topk(self, generate_outputs):
105 sorted_scores = generate_outputs.sequences_scores.argsort(descending=True)
106 generate_outputs.sequences = generate_outputs.sequences[sorted_scores]
107 for sequence in generate_outputs.sequences:
108 seq = self.tokenizer.decode(sequence).split(" ")
109 head_eid = int(seq[1][1:])
110 rel_rid = int(seq[2][1:])
111 if len(self.topk[head_eid, rel_rid]) >= self.K:
112 continue
113 recommended_token = seq[-1]
114 recommended_item = int(recommended_token[1:])
115 if (
116 recommended_item in self.kg_positives[(head_eid, rel_rid)]
117 or recommended_item in self.topk[head_eid, rel_rid]
118 ):
119 continue
120 self.topk[head_eid, rel_rid].append(recommended_item)
121 self.topk_sequences[head_eid, rel_rid].append(seq)
123 def reset_topks(self):
124 del self.topk
125 del self.topk_sequences
126 self.topk = defaultdict(list)
127 self.topk_sequences = defaultdict(list)
130class CumulativeSequenceScorePostProcessor(BaseSequencePostProcessor):
131 """
132 Post-processor that uses the cumulative sequence score of the final
133 `max_new_tokens` predicted tokens to rank sequences.
134 """
136 def calculate_sequence_scores(self, normalized_tuple, sequences, max_new_tokens=24):
137 new_sequence_tokens = sequences[:, -max_new_tokens - 1 : -1]
138 sequence_scores = []
139 # Iterate over each tensor in the normalized tuple
140 for i in range(max_new_tokens):
141 # Get the probabilities corresponding to the ith token in new_sequence_tokens
142 probs = normalized_tuple[i].gather(1, new_sequence_tokens[:, [i]])
143 sequence_scores.append(probs)
144 # Convert the list of tensors into a single tensor
145 sequence_scores = torch.cat(sequence_scores, dim=-1)
146 # Calculate the average score over the last 5 positions for each sequence
147 sequence_scores = sequence_scores.mean(dim=-1)
148 return sequence_scores
150 def normalize_tuple(self, logits_tuple):
151 # Normalize each tensor in the tuple
152 normalized_tuple = tuple(torch.softmax(logits, dim=-1) for logits in logits_tuple)
153 return normalized_tuple
155 def get_sequences(self, generation_outputs, max_new_tokens=24):
156 user_num = generation_outputs["sequences"][:, 1].unique().numel()
158 normalized_scores = self.normalize_tuple(generation_outputs["scores"])
159 normalized_sequences_scores = self.calculate_sequence_scores(
160 normalized_scores, generation_outputs["sequences"], max_new_tokens=max_new_tokens
161 )
163 sequences = generation_outputs["sequences"]
164 num_return_sequences = sequences.shape[0] // user_num
165 batch_user_index = torch.arange(user_num, device=sequences.device).repeat_interleave(num_return_sequences)
167 valid_sequences_mask = torch.logical_not(torch.isfinite(normalized_sequences_scores)) # false if finite
168 normalized_sequences_scores = torch.where(valid_sequences_mask, -torch.inf, normalized_sequences_scores)
170 sorted_indices = normalized_sequences_scores.argsort(descending=True)
171 sorted_sequences = sequences[sorted_indices]
172 sorted_sequences_scores = normalized_sequences_scores[sorted_indices]
173 sorted_batch_user_index = batch_user_index[sorted_indices]
175 return self.parse_sequences(sorted_batch_user_index, sorted_sequences, sorted_sequences_scores)
178class SampleSearchSequenceScorePostProcessor(BaseSequencePostProcessor):
179 """
180 Post-processor that uses the sequence score of the beam search to rank sequences.
182 To use only if do_sample = True and if topk and topp are set.
183 """
185 def get_scores(self, sequences, scores):
186 sequences_scores = None
188 for i, tstep in enumerate(scores):
189 # tstep is a tensor for logits at time t
190 score = torch.softmax(tstep, dim=-1)
191 if sequences_scores is None:
192 sequences_scores = score[:, sequences[:, i]].sum(-1)
193 else:
194 sequences_scores += score[:, sequences[:, i]].sum(-1)
196 return sequences_scores
198 def get_sequences(self, generation_outputs, max_new_tokens=24):
199 user_num = generation_outputs["sequences"][:, 1].unique().numel()
201 sequences = generation_outputs["sequences"]
202 num_return_sequences = sequences.shape[0] // user_num
203 batch_user_index = torch.arange(user_num, device=sequences.device).repeat_interleave(num_return_sequences)
205 sequences_score = self.get_scores(sequences[:, -max_new_tokens:], generation_outputs["scores"])
206 return self.parse_sequences(batch_user_index, sequences, sequences_score)
209class BeamSearchSequenceScorePostProcessor(BaseSequencePostProcessor):
210 """
211 Post-processor that uses the sequence score of the beam search to rank sequences.
212 """
214 def get_sequences(self, generation_outputs, max_new_tokens=24):
215 user_num = generation_outputs["sequences"][:, 1].unique().numel()
217 sequences = generation_outputs["sequences"]
218 num_return_sequences = sequences.shape[0] // user_num
219 batch_user_index = torch.arange(user_num, device=sequences.device).repeat_interleave(num_return_sequences)
221 sorted_indices = generation_outputs["sequences_scores"].argsort(descending=True)
222 sorted_sequences = sequences[sorted_indices]
223 sorted_batch_user_index = batch_user_index[sorted_indices]
224 sorted_sequences_scores = generation_outputs["sequences_scores"][sorted_indices]
226 return self.parse_sequences(sorted_batch_user_index, sorted_sequences, sorted_sequences_scores)