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
« 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
5"""hopwise.model.logits_processor
6#############################
7Common logits processor in recommender system
8"""
10import inspect
12import numpy as np
13import torch
14from cachetools import LFUCache
16from hopwise.utils import KnowledgeEvaluationType
19class LogitsProcessor:
20 """
21 Abstract base class for all logit processors that can be applied during generation.
22 Copy of HuggingFace's LogitsProcessor.
23 """
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 )
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 """
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.
50 Return:
51 `torch.FloatTensor` of shape `(batch_size, config.vocab_size)`:
52 The processed prediction scores.
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)
67 return scores
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 """
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
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
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()
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)
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]
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)
134 banned_mask = self.get_banned_mask(key, candidate_tokens)
136 if banned_mask.all():
137 banned_mask[self.tokenizer.pad_token_id] = False
139 full_mask[idx] = banned_mask
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
146 return scores
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)
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)
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)
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))
173 return key, candidate_tokens
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)
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)
185 return key, candidate_tokens
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)
191 # bos_token determines if the current length is even or odd
192 return current_len % 2 == has_bos_token
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(),)
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")
220 def get_candidates_lp(self, key):
221 return list(self.tokenized_used_ids[key]) + self.special_tokens_ids
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
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
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)
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
266 return scores
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 """
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
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
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 )
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)
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]
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
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
332 return scores
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()
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)
342 # bos_token determines if the current length is even or odd
343 return current_len % 2 == has_bos_token
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)
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))
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
366 return candidate_tokens