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

1# @Time : 2025/06 

2# @Author : Giacomo Medda, Alessandro Soccol 

3# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it 

4 

5"""hopwise.model.postprocessor 

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

7Common post-processors for path sequences in path language modeling recommender systems. 

8""" 

9 

10from collections import defaultdict 

11 

12import torch 

13 

14from hopwise.utils import PathLanguageModelingTokenType 

15 

16 

17class BaseSequencePostProcessor: 

18 """ 

19 Base class for sequence score post-processors. 

20 """ 

21 

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 

27 

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. 

31 

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. 

35 

36 """ 

37 raise NotImplementedError("Subclasses must implement this method.") 

38 

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. 

42 

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. 

47 

48 Returns: 

49 tuple: A tuple containing: 

50 

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() 

58 

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 

64 

65 scores[batch_uidx, recommended_item] = sequence_score 

66 user_topk_sequences.append([uid, recommended_item, sequence_score.item(), decoded_seq]) 

67 

68 return scores, user_topk_sequences 

69 

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(" ") 

73 

74 uid_token = seq[1] 

75 recommended_token = seq[-1] 

76 

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 

85 

86 uid = int(uid_token[1:]) 

87 recommended_item = int(recommended_token[1:]) 

88 

89 if torch.isfinite(scores[batch_uidx, recommended_item]) or recommended_item in self.used_ids[uid]: 

90 return 

91 

92 return uid, recommended_item, seq 

93 

94 

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 

103 

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) 

122 

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) 

128 

129 

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 """ 

135 

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 

149 

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 

154 

155 def get_sequences(self, generation_outputs, max_new_tokens=24): 

156 user_num = generation_outputs["sequences"][:, 1].unique().numel() 

157 

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 ) 

162 

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) 

166 

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) 

169 

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] 

174 

175 return self.parse_sequences(sorted_batch_user_index, sorted_sequences, sorted_sequences_scores) 

176 

177 

178class SampleSearchSequenceScorePostProcessor(BaseSequencePostProcessor): 

179 """ 

180 Post-processor that uses the sequence score of the beam search to rank sequences. 

181 

182 To use only if do_sample = True and if topk and topp are set. 

183 """ 

184 

185 def get_scores(self, sequences, scores): 

186 sequences_scores = None 

187 

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) 

195 

196 return sequences_scores 

197 

198 def get_sequences(self, generation_outputs, max_new_tokens=24): 

199 user_num = generation_outputs["sequences"][:, 1].unique().numel() 

200 

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) 

204 

205 sequences_score = self.get_scores(sequences[:, -max_new_tokens:], generation_outputs["scores"]) 

206 return self.parse_sequences(batch_user_index, sequences, sequences_score) 

207 

208 

209class BeamSearchSequenceScorePostProcessor(BaseSequencePostProcessor): 

210 """ 

211 Post-processor that uses the sequence score of the beam search to rank sequences. 

212 """ 

213 

214 def get_sequences(self, generation_outputs, max_new_tokens=24): 

215 user_num = generation_outputs["sequences"][:, 1].unique().numel() 

216 

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) 

220 

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] 

225 

226 return self.parse_sequences(sorted_batch_user_index, sorted_sequences, sorted_sequences_scores)