Use native batched llguidance mask generation (#32412)

Co-authored-by: Alec S <10566873+alecsolder@users.noreply.github.com>
This commit is contained in:
Lianmin Zheng
2026-07-25 16:36:32 -07:00
committed by GitHub
co-authored by Alec S
parent fae84ac0f9
commit 9989077f24
5 changed files with 335 additions and 15 deletions
@@ -0,0 +1,112 @@
"""Bitwise parity tests for llguidance's regular-decode batched mask fill."""
import unittest
from unittest.mock import MagicMock
import torch
from llguidance import LLTokenizer, grammar_from
from sglang.srt.constrained.base_grammar_backend import GrammarRow
from sglang.srt.constrained.llguidance_backend import GuidanceBackend, GuidanceGrammar
from sglang.srt.runtime_context import get_resources
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
_REGEX = r"[0-9]{1,8}"
class TestLLGuidanceBatchedMask(unittest.TestCase):
@classmethod
def setUpClass(cls):
cls.template = GuidanceGrammar(
llguidance_tokenizer=LLTokenizer("byte"),
serialized_grammar=grammar_from("regex", _REGEX),
)
def _fresh(self, n):
return [self.template.copy() for _ in range(n)]
def _allocate(self, grammars):
return grammars[0].allocate_vocab_mask(
self.template.llguidance_tokenizer.vocab_size, len(grammars), "cpu"
)
def _serial(self, grammars):
mask = self._allocate(grammars)
for row, grammar in enumerate(grammars):
if not grammar.finished and not grammar.is_terminated():
grammar.fill_vocab_mask(mask, row)
return mask
def _batched(self, grammars):
mask = self._allocate(grammars)
entries = [
GrammarRow(row=row, grammar=grammar)
for row, grammar in enumerate(grammars)
if not grammar.finished and not grammar.is_terminated()
]
grammars[0].fill_vocab_mask_batched(entries, mask)
return mask
def test_batched_matches_serial(self):
for batch_size in (1, 4, 10):
with self.subTest(batch_size=batch_size):
serial = self._serial(self._fresh(batch_size))
batched = self._batched(self._fresh(batch_size))
self.assertTrue(torch.equal(serial, batched))
self.assertTrue((batched[0] != -1).any())
def test_finished_row_stays_all_allow(self):
serial_grammars = self._fresh(3)
batched_grammars = self._fresh(3)
serial_grammars[1].finished = True
batched_grammars[1].finished = True
serial = self._serial(serial_grammars)
batched = self._batched(batched_grammars)
self.assertTrue(torch.equal(serial, batched))
self.assertTrue((batched[1] == -1).all())
def test_unsupported_entry_uses_serial_fill(self):
mask = self._allocate(self._fresh(1))
fallback = MagicMock()
fallback.fill_vocab_mask.side_effect = lambda vocab_mask, row: vocab_mask[
row
].zero_()
entries = [GrammarRow(row=0, grammar=fallback)]
self.template.fill_vocab_mask_batched(entries, mask)
fallback.fill_vocab_mask.assert_called_once_with(mask, 0)
self.assertTrue((mask == 0).all())
def test_backend_initializes_fixed_mask_buffer(self):
name = "test_llguidance_vocab_mask"
get_resources().buffers.pop(name, None)
backend = object.__new__(GuidanceBackend)
backend.llguidance_tokenizer = self.template.llguidance_tokenizer
try:
mask = backend.initialize_vocab_mask_buffer(
name=name,
vocab_size=self.template.llguidance_tokenizer.vocab_size,
max_rows=4,
device="cpu",
)
same_mask = backend.initialize_vocab_mask_buffer(
name=name,
vocab_size=self.template.llguidance_tokenizer.vocab_size,
max_rows=4,
device="cpu",
)
self.assertEqual(mask.shape[0], 4)
self.assertEqual(mask.data_ptr(), same_mask.data_ptr())
finally:
get_resources().buffers.pop(name, None)
if __name__ == "__main__":
unittest.main()
@@ -43,6 +43,11 @@ def _make_info(batch_size=2, **overrides):
return SamplingBatchInfo(**defaults)
def _serial_batched_fill(entries, vocab_mask):
for entry in entries:
entry.grammar.fill_vocab_mask(vocab_mask, entry.row)
class TestMergeBiasTensor(CustomTestCase):
def test_both_none_returns_none(self):
@@ -239,12 +244,14 @@ class TestUpdateRegexVocabMask(CustomTestCase):
grammar = MagicMock()
grammar.finished = False
grammar.is_terminated.return_value = False
grammar.fill_vocab_mask_batched.side_effect = _serial_batched_fill
grammar.allocate_vocab_mask.return_value = torch.zeros(1, VOCAB_SIZE)
grammar.move_vocab_mask.return_value = torch.zeros(1, VOCAB_SIZE)
info = _make_info(batch_size=1)
info.grammars = [grammar]
info.update_regex_vocab_mask()
grammar.allocate_vocab_mask.assert_called_once()
grammar.fill_vocab_mask_batched.assert_called_once()
grammar.fill_vocab_mask.assert_called_once()
grammar.move_vocab_mask.assert_called_once()
self.assertIs(info.grammar_mask.grammar, grammar)
@@ -254,6 +261,7 @@ class TestUpdateRegexVocabMask(CustomTestCase):
active = MagicMock()
active.finished = False
active.is_terminated.return_value = False
active.fill_vocab_mask_batched.side_effect = _serial_batched_fill
active.allocate_vocab_mask.return_value = torch.zeros(3, VOCAB_SIZE)
active.move_vocab_mask.return_value = torch.zeros(3, VOCAB_SIZE)
@@ -268,10 +276,34 @@ class TestUpdateRegexVocabMask(CustomTestCase):
info.grammars = [active, finished, terminated]
info.update_regex_vocab_mask()
active.fill_vocab_mask_batched.assert_called_once()
active.fill_vocab_mask.assert_called_once()
finished.fill_vocab_mask.assert_not_called()
terminated.fill_vocab_mask.assert_not_called()
def test_batched_fill_skips_serial_loop(self):
"""A backend with a batched kernel must not also run the serial fill."""
active = MagicMock()
active.finished = False
active.is_terminated.return_value = False
active.allocate_vocab_mask.return_value = torch.zeros(2, VOCAB_SIZE)
active.move_vocab_mask.return_value = torch.zeros(2, VOCAB_SIZE)
other = MagicMock()
other.finished = False
other.is_terminated.return_value = False
info = _make_info(batch_size=2)
info.grammars = [active, other]
info.update_regex_vocab_mask()
# The batched call happened with both active rows; serial fill skipped.
active.fill_vocab_mask_batched.assert_called_once()
entries, _mask = active.fill_vocab_mask_batched.call_args.args
self.assertEqual([e.row for e in entries], [0, 1])
active.fill_vocab_mask.assert_not_called()
other.fill_vocab_mask.assert_not_called()
# filter_batch
class TestFilterBatch(CustomTestCase):