Spaces:
Sleeping
Sleeping
Download tests/test_batch_native.py from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 6.84 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/tests/test_batch_native.py
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/tests/test_batch_native.py
-
curl -L -o test_batch_native.py https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/tests/test_batch_native.py
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) | |
| 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" | |
| 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([]) == [] | |