[Spec][Ngram] 6/N: Load an external corpus and construct a Suffix Automaton (#21425)
This commit is contained in:
@@ -426,5 +426,63 @@ class TestHiCacheArgs(unittest.TestCase):
|
||||
self.assertEqual(args.decode_attention_backend, "triton")
|
||||
|
||||
|
||||
class TestNgramExternalSamArgs(CustomTestCase):
|
||||
def _make_dummy_ngram_args(self, **overrides):
|
||||
args = ServerArgs(model_path="dummy")
|
||||
args.speculative_algorithm = "NGRAM"
|
||||
args.speculative_num_draft_tokens = 12
|
||||
args.device = "cuda"
|
||||
for key, value in overrides.items():
|
||||
setattr(args, key, value)
|
||||
return args
|
||||
|
||||
def test_prepare_server_args_parses_external_sam_args(self):
|
||||
server_args = prepare_server_args(
|
||||
[
|
||||
"--model-path",
|
||||
"dummy",
|
||||
"--speculative-algorithm",
|
||||
"NGRAM",
|
||||
"--speculative-ngram-external-corpus-path",
|
||||
"/tmp/ngram-corpus.jsonl",
|
||||
"--speculative-ngram-external-sam-budget",
|
||||
"4",
|
||||
"--speculative-ngram-external-corpus-max-tokens",
|
||||
"128",
|
||||
]
|
||||
)
|
||||
self.assertEqual(
|
||||
server_args.speculative_ngram_external_corpus_path,
|
||||
"/tmp/ngram-corpus.jsonl",
|
||||
)
|
||||
self.assertEqual(server_args.speculative_ngram_external_sam_budget, 4)
|
||||
self.assertEqual(server_args.speculative_ngram_external_corpus_max_tokens, 128)
|
||||
|
||||
def test_external_sam_budget_requires_path(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
self._make_dummy_ngram_args(
|
||||
speculative_ngram_external_sam_budget=2,
|
||||
)._handle_speculative_decoding()
|
||||
self.assertIn("external-sam-budget", str(context.exception))
|
||||
|
||||
def test_external_sam_budget_must_fit_draft_budget(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
self._make_dummy_ngram_args(
|
||||
speculative_num_draft_tokens=4,
|
||||
speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl",
|
||||
speculative_ngram_external_sam_budget=4,
|
||||
)._handle_speculative_decoding()
|
||||
self.assertIn("speculative_num_draft_tokens - 1", str(context.exception))
|
||||
|
||||
def test_external_corpus_max_tokens_must_be_positive(self):
|
||||
with self.assertRaises(ValueError) as context:
|
||||
self._make_dummy_ngram_args(
|
||||
speculative_ngram_external_corpus_path="/tmp/ngram-corpus.jsonl",
|
||||
speculative_ngram_external_sam_budget=2,
|
||||
speculative_ngram_external_corpus_max_tokens=0,
|
||||
)._handle_speculative_decoding()
|
||||
self.assertIn("external-corpus-max-tokens", str(context.exception))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user