Coverage for hopwise/model/logits_processor.py: 83%

175 statements  

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

1# @Time : 2025 

2# @Author : Giacomo Medda, Alessandro Soccol 

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

4 

5"""hopwise.model.logits_processor 

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

7Common logits processor in recommender system 

8""" 

9 

10import inspect 

11 

12import numpy as np 

13import torch 

14from cachetools import LFUCache 

15 

16from hopwise.utils import KnowledgeEvaluationType 

17 

18 

19class LogitsProcessor: 

20 """ 

21 Abstract base class for all logit processors that can be applied during generation. 

22 Copy of HuggingFace's LogitsProcessor. 

23 """ 

24 

25 def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor: 

26 raise NotImplementedError( 

27 f"{self.__class__} is an abstract class. Only classes inheriting this class can be called." 

28 ) 

29 

30 

31class LogitsProcessorList(list): 

32 """ 

33 This class can be used to create a list of [`LogitsProcessor`] to subsequently process a `scores` input tensor. 

34 This class inherits from list and adds a specific *__call__* method to apply each [`LogitsProcessor`] to the 

35 inputs. 

36 Copy of HuggingFace's LogitsProcessorList. 

37 """ 

38 

39 def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> torch.FloatTensor: 

40 r""" 

41 Args: 

42 input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`): 

43 Indices of input sequence tokens in the vocabulary. [What are input IDs?](../glossary#input-ids) 

44 scores (`torch.FloatTensor` of shape `(batch_size, config.vocab_size)`): 

45 Prediction scores of a language modeling head. These can be logits for each vocabulary when not using 

46 beam search or log softmax for each vocabulary token when using beam search 

47 kwargs (`Dict[str, Any]`, *optional*): 

48 Additional kwargs that are specific to a logits processor. 

49 

50 Return: 

51 `torch.FloatTensor` of shape `(batch_size, config.vocab_size)`: 

52 The processed prediction scores. 

53 

54 """ 

55 for processor in self: 

56 function_args = inspect.signature(processor.__call__).parameters 

57 if len(function_args) > 2: # noqa: PLR2004 

58 if not all(arg in kwargs for arg in list(function_args.keys())[2:]): 

59 raise ValueError( 

60 f"Make sure that all the required parameters: {list(function_args.keys())} for " 

61 f"{processor.__class__} are passed to the logits processor." 

62 ) 

63 scores = processor(input_ids, scores, **kwargs) 

64 else: 

65 scores = processor(input_ids, scores) 

66 

67 return scores 

68 

69 

70class ConstrainedLogitsProcessorWordLevel(LogitsProcessor): 

71 """ 

72 Force the last token to be one of the force_tokens if the total length is reached, in the path generation stage 

73 this means to limit the hop size. This is a word-level constraint, does not work with piece tokenizers. 

74 If task is link prediction (LP) logit processor forces last token to reachable ones 

75 """ 

76 

77 def __init__( 

78 self, 

79 tokenized_ckg, 

80 tokenized_used_ids, 

81 max_sequence_length, 

82 tokenizer, 

83 mask_cache_size=3 * 10**4, 

84 pos_candidates_cache_size=1 * 10**5, 

85 task=KnowledgeEvaluationType.REC, 

86 **kwargs, 

87 ): 

88 super().__init__(**kwargs) 

89 self.tokenized_ckg = tokenized_ckg 

90 self.tokenized_used_ids = tokenized_used_ids 

91 self.max_sequence_length = max_sequence_length 

92 self.tokenizer = tokenizer 

93 self.bos_token_id = self.tokenizer.convert_tokens_to_ids(self.tokenizer.bos_token) 

94 self.pos_candidates_cache = LFUCache(pos_candidates_cache_size) 

95 self.mask_cache = LFUCache(mask_cache_size) 

96 self.task = task 

97 

98 if self.task == KnowledgeEvaluationType.LP: 

99 self.special_tokens_ids = [ 

100 self.tokenizer.encode(x, add_special_tokens=False)[0] 

101 for x in self.tokenizer.all_special_tokens_extended 

102 ] 

103 else: 

104 self.special_tokens_ids = None 

105 

106 def is_bos_token_in_input(self, input_ids): 

107 """Check if the input contains a BOS token. Checking the first sequence is enough.""" 

108 return (input_ids[0, 0] == self.bos_token_id).item() 

109 

110 def __call__(self, input_ids, scores): 

111 current_len = input_ids.shape[-1] 

112 has_bos_token = self.is_bos_token_in_input(input_ids) 

113 

114 unique_input_ids = input_ids 

115 if self.task == KnowledgeEvaluationType.REC and current_len < self.max_sequence_length - 2 + has_bos_token: 

116 # Determine whether the next token to generate is a relation or an entity: 

117 # - relation: only the last entity is needed (1 token) → for [user123] last_n_tokens = 1 

118 # - entity: the last 2 tokens are needed (entity, relation) → for [user123, watched] last_n_tokens = 2 

119 # Apply deduplication: select unique sequences based only on the relevant last tokens (1 or 2) 

120 # This avoids recomputing the same mask for sequences that share the same context 

121 last_n_tokens = 2 if self.is_next_token_entity(input_ids) else 1 

122 _, input_ids_indices, input_ids_inv = np.unique( 

123 input_ids.cpu().numpy()[:, -last_n_tokens:], axis=0, return_index=True, return_inverse=True 

124 ) 

125 unique_input_ids = input_ids[input_ids_indices] 

126 

127 full_mask = np.zeros((unique_input_ids.shape[0], len(self.tokenizer)), dtype=bool) 

128 for idx in range(unique_input_ids.shape[0]): 

129 if self.task == KnowledgeEvaluationType.REC: 

130 key, candidate_tokens = self.process_scores_rec(unique_input_ids, idx) 

131 elif self.task == KnowledgeEvaluationType.LP: 

132 key, candidate_tokens = self.process_scores_lp(unique_input_ids, idx) 

133 

134 banned_mask = self.get_banned_mask(key, candidate_tokens) 

135 

136 if banned_mask.all(): 

137 banned_mask[self.tokenizer.pad_token_id] = False 

138 

139 full_mask[idx] = banned_mask 

140 

141 if self.task == KnowledgeEvaluationType.REC and current_len < self.max_sequence_length - 2 + has_bos_token: 

142 scores[full_mask[input_ids_inv]] = -torch.inf 

143 else: 

144 scores[full_mask] = -torch.inf 

145 

146 return scores 

147 

148 def process_scores_rec(self, input_ids, idx): 

149 """Process each score based on input length and update mask list.""" 

150 current_len = input_ids.shape[-1] 

151 has_bos_token = self.is_bos_token_in_input(input_ids) 

152 

153 key = self.get_current_key(input_ids, idx) 

154 # Last content token (the recommended item) is generated when the prefix holds all but one 

155 # content token, i.e. (max_sequence_length - 1) content tokens minus 1, plus the BOS token if present. 

156 if current_len == self.max_sequence_length - 2 + has_bos_token: 

157 current_uid = input_ids[idx, int(has_bos_token)].item() 

158 uid_cond_key = (current_uid, *key) 

159 

160 candidate_tokens = self.pos_candidates_cache.get(uid_cond_key) 

161 if candidate_tokens is None: 

162 candidate_tokens = self.get_candidates_rec(*key) 

163 

164 # Get user positives 

165 user_used_ids = self.tokenized_used_ids[current_uid] 

166 # Select negatives 

167 candidate_tokens = list(candidate_tokens - user_used_ids) 

168 # Useless if during evaluation a user is seen once 

169 self.pos_candidates_cache[uid_cond_key] = candidate_tokens 

170 else: 

171 candidate_tokens = list(self.get_candidates_rec(*key)) 

172 

173 return key, candidate_tokens 

174 

175 def process_scores_lp(self, input_ids, idx): 

176 """Process each score based on input length or skip.""" 

177 current_len = input_ids.shape[-1] 

178 has_bos_token = self.is_bos_token_in_input(input_ids) 

179 

180 key, candidate_tokens = None, None 

181 if current_len % 2 == has_bos_token: 

182 key = self.get_current_key(input_ids, idx) 

183 candidate_tokens = self.get_candidates_lp(key) 

184 

185 return key, candidate_tokens 

186 

187 def is_next_token_entity(self, input_ids): 

188 current_len = input_ids.shape[-1] 

189 has_bos_token = self.is_bos_token_in_input(input_ids) 

190 

191 # bos_token determines if the current length is even or odd 

192 return current_len % 2 == has_bos_token 

193 

194 def get_current_key(self, input_ids, idx): 

195 if self.is_next_token_entity(input_ids): 

196 return input_ids[idx, -2].item(), input_ids[idx, -1].item() 

197 else: 

198 # The next token is a relation 

199 return (input_ids[idx, -1].item(),) 

200 

201 def get_candidates_rec(self, key1, key2=None): 

202 """ 

203 :param key1: 

204 :param key2: if key2 is not None, it returns entity candidates, otherwise relation candidates 

205 """ 

206 if key1 in self.tokenized_ckg: 

207 if key2 is not None and key2 in self.tokenized_ckg[key1]: 

208 # return tail given head + relation 

209 return self.tokenized_ckg[key1][key2] 

210 else: 

211 # return relations given head 

212 return set(self.tokenized_ckg[key1].keys()) 

213 elif key1 in self.tokenizer.all_special_ids: 

214 # A special token (e.g. the pad emitted by the empty-candidate fallback) is a dead-end path, 

215 # not a graph node: it has no continuation, which lets the fallback terminate the path. 

216 return set() 

217 else: 

218 raise ValueError(f"Key {key1} ('{self.tokenizer.convert_ids_to_tokens(key1)}') not found in tokenized_ckg") 

219 

220 def get_candidates_lp(self, key): 

221 return list(self.tokenized_used_ids[key]) + self.special_tokens_ids 

222 

223 def get_banned_mask(self, key, candidate_tokens): 

224 """Retrieve or cache the banned token mask for a specific key.""" 

225 banned_mask = self.mask_cache.get(key) 

226 if banned_mask is None: 

227 banned_mask = np.ones(len(self.tokenizer), dtype=bool) 

228 banned_mask[candidate_tokens] = False 

229 self.mask_cache[key] = banned_mask 

230 return banned_mask 

231 

232 

233class PrefixConstrainedLogitsProcessorWordLevel(ConstrainedLogitsProcessorWordLevel): 

234 def __init__( 

235 self, 

236 tokenized_ckg, 

237 tokenized_used_ids, 

238 max_sequence_length, 

239 tokenizer, 

240 **kwargs, 

241 ): 

242 super().__init__( 

243 tokenized_ckg, 

244 tokenized_used_ids, 

245 max_sequence_length, 

246 tokenizer, 

247 **kwargs, 

248 ) 

249 self.mask_cache = None 

250 

251 def __call__(self, input_ids, scores): 

252 current_len = input_ids.shape[-1] 

253 if current_len == self.max_sequence_length - 1: 

254 self.mask_non_eos_tokens(scores) 

255 else: 

256 indices = [] 

257 masked_scores = torch.full_like(scores, -torch.inf) 

258 for idx in range(scores.shape[0]): 

259 _, candidate_tokens = self.process_scores(input_ids, idx, current_len) 

260 

261 candidate_tokens = torch.LongTensor(candidate_tokens, device=scores.device) 

262 indices.append(candidate_tokens) 

263 masked_scores[idx].scatter_(dim=-1, index=candidate_tokens, src=scores[idx]) 

264 scores = masked_scores 

265 

266 return scores 

267 

268 

269class PLMLogitsProcessorWordLevel(LogitsProcessor): 

270 """ 

271 https://dl.acm.org/doi/pdf/10.1145/3485447.3511937 

272 Constraint decoding strategy for PLM, it forces the model to generate alternatively entities and relations 

273 """ 

274 

275 def __init__( 

276 self, 

277 tokenized_ckg, 

278 tokenized_used_ids, 

279 max_sequence_length, 

280 tokenizer, 

281 pos_candidates_cache_size=1 * 10**5, 

282 task=KnowledgeEvaluationType.REC, 

283 **kwargs, 

284 ): 

285 super().__init__(**kwargs) 

286 self.tokenized_ckg = tokenized_ckg 

287 self.tokenized_used_ids = tokenized_used_ids 

288 self.max_sequence_length = max_sequence_length 

289 self.tokenizer = tokenizer 

290 self.bos_token_id = self.tokenizer.convert_tokens_to_ids(self.tokenizer.bos_token) 

291 self.pos_candidates_cache = LFUCache(pos_candidates_cache_size) 

292 self.task = task 

293 

294 if self.task == KnowledgeEvaluationType.LP: 

295 self.special_tokens_ids = [ 

296 self.tokenizer.encode(x, add_special_tokens=False)[0] 

297 for x in self.tokenizer.all_special_tokens_extended 

298 ] 

299 else: 

300 self.special_tokens_ids = None 

301 

302 self.entity_token_ids = torch.LongTensor(list(set(self.tokenized_ckg.keys()))) 

303 self.relation_token_ids = torch.LongTensor( 

304 list(set([rel for rel_dict in self.tokenized_ckg.values() for rel in rel_dict.keys()])) 

305 ) 

306 

307 def __call__(self, input_ids, scores): 

308 current_len = input_ids.shape[-1] 

309 has_bos_token = self.is_bos_token_in_input(input_ids) 

310 

311 unique_input_ids = input_ids 

312 if self.task == KnowledgeEvaluationType.REC and current_len == (self.max_sequence_length - 2 + has_bos_token): 

313 user_idx = int(has_bos_token) 

314 _, input_ids_indices, input_ids_inv = np.unique( 

315 input_ids.cpu().numpy()[:, [user_idx]], axis=0, return_index=True, return_inverse=True 

316 ) 

317 unique_input_ids = input_ids[input_ids_indices] 

318 

319 full_mask = np.ones((unique_input_ids.shape[0], len(self.tokenizer)), dtype=bool) 

320 for idx in range(unique_input_ids.shape[0]): 

321 candidate_tokens = self.process_scores(unique_input_ids, idx) 

322 full_mask[idx, candidate_tokens] = False 

323 

324 scores[full_mask[input_ids_inv]] = -torch.inf 

325 else: 

326 # Paths are expected to be the same type and length, so we can use the same mask for all 

327 full_mask = np.ones((unique_input_ids.shape[0], len(self.tokenizer)), dtype=bool) 

328 candidate_tokens = self.process_scores(input_ids, 0) 

329 full_mask[:, candidate_tokens] = False 

330 scores[full_mask] = -torch.inf 

331 

332 return scores 

333 

334 def is_bos_token_in_input(self, input_ids): 

335 """Check if the input contains a BOS token. Checking the first sequence is enough.""" 

336 return (input_ids[0, 0] == self.bos_token_id).item() 

337 

338 def is_next_token_entity(self, input_ids): 

339 current_len = input_ids.shape[-1] 

340 has_bos_token = self.is_bos_token_in_input(input_ids) 

341 

342 # bos_token determines if the current length is even or odd 

343 return current_len % 2 == has_bos_token 

344 

345 def process_scores(self, input_ids, idx): 

346 """Process each score based on input length and update mask to allow only entities or only relations.""" 

347 current_len = input_ids.shape[-1] 

348 has_bos_token = self.is_bos_token_in_input(input_ids) 

349 

350 # Last content token (the recommended item) is generated when the prefix holds all but one 

351 # content token, i.e. (max_sequence_length - 1) content tokens minus 1, plus the BOS token if present. 

352 if current_len == self.max_sequence_length - 2 + has_bos_token: 

353 current_uid = input_ids[idx, int(has_bos_token)].item() 

354 candidate_tokens = self.pos_candidates_cache.get(current_uid) 

355 if candidate_tokens is None: 

356 candidate_tokens = np.arange(len(self.tokenizer)) 

357 

358 user_used_ids = self.tokenized_used_ids[current_uid] 

359 candidate_tokens = np.setdiff1d(candidate_tokens, list(user_used_ids), assume_unique=True) 

360 self.pos_candidates_cache[current_uid] = candidate_tokens 

361 elif self.is_next_token_entity(input_ids): 

362 candidate_tokens = self.entity_token_ids 

363 else: 

364 candidate_tokens = self.relation_token_ids 

365 

366 return candidate_tokens