"""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([]) == []