test1111111 / tests /test_batch_native.py
spitfire4794's picture
deploy c66f5aa: decode push
421b8c2
Raw History Blame Contribute Delete
6.84 kB
"""True batched decode (one weight stream, B seqs) determinism tests.
Exercises the native ``BatchSession`` (``Model.create_batch_session``) and
the :class:`cism.batch.BatchEngine` wrapper on tiny dense models: ``B=1``
must be bitwise-identical to serial ``Session`` decoding, and ``B=4`` batch
slot ``i`` must match a lone ``B=1`` run of the same prompt.
"""
import importlib
import os
import pytest
native = importlib.import_module(os.environ.get("CISM_NATIVE_TEST_MODULE", "cism._native"))
from test_native import tiny_model
class FakeTokenizer:
bos_token_id = 0
eos_token_id = None
chat_template = None
def encode(self, text, *, add_special_tokens):
return [1, 2]
def decode(self, tokens, *, skip_special_tokens, clean_up_tokenization_spaces):
assert skip_special_tokens is True
assert clean_up_tokenization_spaces is False
return "".join(f"<{t}>" for t in tokens)
def make_model(architecture="llama", precision="fp32", **overrides):
config, weights = tiny_model(architecture)
config.update(overrides)
return config, weights, native.Model(config, weights, precision)
@pytest.mark.parametrize("architecture", ["llama", "qwen3"])
@pytest.mark.parametrize("precision", ["fp32", "int8"])
def test_batch_b1_matches_serial_session(architecture, precision):
config, weights, model = make_model(architecture, precision)
prompt = [2, 7, 1, 19, 4]
expected = model.create_session(prompt, max_new_tokens=9).next_tokens(9)
batch = model.create_batch_session([prompt], max_new_tokens=9)
assert batch.batch_size == 1
actual = batch.next_tokens(9)
assert actual == [expected]
assert batch.finish_reasons == ["length"]
assert list(batch.generated_tokens) == [9]
# Chunked continuation is stateful and identical too.
serial = model.create_session(prompt, max_new_tokens=9)
batched = model.create_batch_session([prompt], max_new_tokens=9)
assert serial.next_tokens(2) == batched.next_tokens(2)[0]
assert serial.next_tokens(3) == batched.next_tokens(3)[0]
assert serial.next_tokens(100) == batched.next_tokens(100)[0]
assert serial.finish_reason == batched.finish_reasons[0] == "length"
@pytest.mark.parametrize("architecture", ["llama", "qwen3"])
def test_batch_determinism_b1_vs_b4(architecture):
config, weights, model = make_model(architecture, "fp32")
prompts = [[1, 2, 3], [2, 3, 1], [3, 3, 3], [1, 1, 2]]
serial = model.create_batch_session([prompts[0]], max_new_tokens=8).next_tokens(8)
batch = model.create_batch_session(prompts, max_new_tokens=8)
actual = batch.next_tokens(8)
assert len(actual) == 4
assert actual[0] == serial[0]
# Every slot matches its own lone run (order preserved).
for i, prompt in enumerate(prompts):
single = model.create_batch_session([prompt], max_new_tokens=8).next_tokens(8)
assert actual[i] == single[0], f"slot {i} diverged from serial"
def test_batch_sampled_seed_determinism():
config, weights, model = make_model("llama", "fp32")
prompts = [[3, 1], [1, 4], [2, 2], [4, 0]]
options = dict(max_new_tokens=12, temperature=0.9, top_p=0.8, top_k=11, seed=876)
first = model.create_batch_session([prompts[0]], **options).next_tokens(12)
batched = model.create_batch_session(prompts, **options).next_tokens(12)
assert batched[0] == first[0]
# Chunked batched calls equal one-shot.
session = model.create_batch_session(prompts, **options)
chunked = session.next_tokens(5)
rest = session.next_tokens(7)
combined = [a + b for a, b in zip(chunked, rest)]
assert combined == batched
def test_batch_eos_and_length_finish():
config, weights, model = make_model("llama", "fp32")
import numpy as np
eos = int(np.argmax(model.logits([1, 2])))
session = model.create_batch_session([[1, 2], [3, 4]], max_new_tokens=6,
eos_token_ids=[eos])
out = session.next_tokens(6)
reasons = session.finish_reasons
assert len(out) == 2 and len(reasons) == 2
for tokens, reason in zip(out, reasons):
assert reason in ("stop", "length")
if reason == "stop":
assert tokens[-1] == eos
# Serial sessions agree per slot.
for i, prompt in enumerate([[1, 2], [3, 4]]):
ref = model.create_session(prompt, max_new_tokens=6, eos_token_ids=[eos])
expected = ref.next_tokens(6)
assert out[i] == expected
assert reasons[i] == ref.finish_reason
def test_batch_wrapper_b1_vs_b4():
from cism.batch import BatchEngine
config, weights = tiny_model("llama")
engine = BatchEngine.from_weights(config, weights, FakeTokenizer(), model_id="tiny-llama")
prompts = [[1, 2, 3], [2, 3, 1], [3, 3, 3], [1, 1, 2]]
serial = engine.generate_batch([prompts[0]], max_new_tokens=8)
batch = engine.generate_batch(prompts, max_new_tokens=8)
assert len(batch) == 4
assert batch[0].token_ids == serial[0].token_ids
assert batch[0].text == serial[0].text
assert batch[0].finish_reason == serial[0].finish_reason
for i, prompt in enumerate(prompts):
single = engine.generate_batch([prompt], max_new_tokens=8)[0]
assert batch[i].token_ids == single.token_ids
def test_batch_wrapper_matches_serial_engine():
from cism.batch import BatchEngine
from cism.engine import Engine
config, weights = tiny_model("qwen3")
batched = BatchEngine.from_weights(config, weights, FakeTokenizer())
serial = Engine.from_weights(config, weights, FakeTokenizer())
prompts = [[1, 2], [2, 1], [1, 1]]
results = batched.generate_batch(prompts, max_new_tokens=6)
for prompt, result in zip(prompts, results):
expected = serial.generate(prompt, max_new_tokens=6)
assert result.token_ids == expected.token_ids
assert result.finish_reason == expected.finish_reason
def test_batch_invalid_inputs():
config, weights, model = make_model()
with pytest.raises((ValueError, TypeError)):
model.create_batch_session([], max_new_tokens=1)
with pytest.raises((ValueError, TypeError, OverflowError)):
model.create_batch_session([[]], max_new_tokens=1)
with pytest.raises((ValueError, TypeError, OverflowError)):
model.create_batch_session([[1, 999]], max_new_tokens=1)
with pytest.raises(ValueError):
model.create_batch_session([[1]], max_new_tokens=1).next_tokens(-1)
# Python wrapper rejects Surjo (dense-only in this version).
from cism.batch import BatchEngine
with pytest.raises(ValueError, match="dense-only"):
BatchEngine.from_weights({"model_type": "surjo"}, {}, FakeTokenizer())
# Empty batch returns empty.
engine = BatchEngine.from_weights(*tiny_model("llama"), FakeTokenizer())
assert engine.generate_batch([]) == []