fix: prevent stale bitmask leakage in LLGuidance grammar backend (#23653)
Signed-off-by: jellysnack <oleg.jellysnack@gmail.com> Co-authored-by: YorkSu <45256779+YorkSu@users.noreply.github.com> Co-authored-by: jellysnack <oleg.jellysnack@gmail.com>
This commit is contained in:
co-authored by
YorkSu
parent
47fabb52ed
commit
8a9f2aa116
@@ -51,7 +51,6 @@ class GuidanceGrammar(BaseGrammarObject):
|
|||||||
)
|
)
|
||||||
self._check_err()
|
self._check_err()
|
||||||
|
|
||||||
self.bitmask = None
|
|
||||||
self.eos_token = self.llguidance_tokenizer.eos_token
|
self.eos_token = self.llguidance_tokenizer.eos_token
|
||||||
|
|
||||||
def accept_token(self, token: int):
|
def accept_token(self, token: int):
|
||||||
@@ -83,16 +82,7 @@ class GuidanceGrammar(BaseGrammarObject):
|
|||||||
def allocate_vocab_mask(
|
def allocate_vocab_mask(
|
||||||
self, vocab_size: int, batch_size: int, device
|
self, vocab_size: int, batch_size: int, device
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if self.bitmask is None or self.bitmask.shape[0] < batch_size:
|
return allocate_token_bitmask(batch_size, self.llguidance_tokenizer.vocab_size)
|
||||||
# only create bitmask when batch gets larger
|
|
||||||
self.bitmask = allocate_token_bitmask(
|
|
||||||
batch_size, self.llguidance_tokenizer.vocab_size
|
|
||||||
)
|
|
||||||
bitmask = self.bitmask
|
|
||||||
else:
|
|
||||||
bitmask = self.bitmask[:batch_size]
|
|
||||||
|
|
||||||
return bitmask
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def move_vocab_mask(vocab_mask: torch.Tensor, device) -> torch.Tensor:
|
def move_vocab_mask(vocab_mask: torch.Tensor, device) -> torch.Tensor:
|
||||||
|
|||||||
Reference in New Issue
Block a user