Migrate ngram corpus from torch cpp_extension to TVM FFI jit_kernel (#21920)

Co-authored-by: DarkSharpness <2040703891@qq.com>
This commit is contained in:
Liangsheng Yin
2026-04-02 02:18:11 -07:00
committed by GitHub
co-authored by DarkSharpness
parent b684b0b72f
commit 9d9537fbd3
14 changed files with 270 additions and 115 deletions
+14 -16
View File
@@ -480,41 +480,39 @@ class TestDraftBudgetSaturation(CustomTestCase):
class TestTruncate(CustomTestCase):
"""Verify the Result.truncate method via the Python binding."""
"""Verify truncation logic on batch_get output."""
def test_truncate_reduces_output(self):
corpus = _make_corpus("BFS", draft_token_num=8)
corpus.batch_put(SEED_SEQUENCES)
corpus.synchronize()
result = corpus._ngram.batchMatch([[1, 2, 3]])
original_len = len(result.token)
self.assertEqual(original_len, 8)
ids, masks = corpus.batch_get([[1, 2, 3]])
ids = ids.reshape(8)
self.assertEqual(len(ids), 8)
result.truncate(4)
self.assertEqual(len(result.token), 4)
self.assertEqual(len(result.mask), 4 * 4)
# Simulate truncate to 4
trunc_n = 4
trunc_ids = ids[:trunc_n]
self.assertEqual(len(trunc_ids), trunc_n)
def test_truncate_preserves_mask_structure(self):
corpus = _make_corpus("BFS", draft_token_num=8)
corpus.batch_put(SEED_SEQUENCES)
corpus.synchronize()
result = corpus._ngram.batchMatch([[1, 2, 3]])
full_ids = list(result.token)
full_mask = list(result.mask)
n = len(full_ids)
ids, masks = corpus.batch_get([[1, 2, 3]])
n = 8
full_mask = masks.reshape(n, n)
result_copy = corpus._ngram.batchMatch([[1, 2, 3]])
trunc_n = 4
result_copy.truncate(trunc_n)
trunc_mask = list(result_copy.mask)
trunc_mask = full_mask[:trunc_n, :trunc_n]
for i in range(trunc_n):
for j in range(trunc_n):
self.assertEqual(
trunc_mask[i * trunc_n + j],
full_mask[i * n + j],
trunc_mask[i, j],
full_mask[i, j],
f"Mask mismatch at ({i},{j})",
)