"""Tests for ``pp_tinygrad_worker.py`` — pure helpers plus the stdin/stdout JSON protocol exercised against a real subprocess. Names match TEST_SPEC §2 and §3 verbatim. """ from __future__ import annotations import base64 import json import os import selectors import subprocess import sys from pathlib import Path import pytest import pp_tinygrad_worker as worker WORKER = Path(__file__).parent.parent / "pp_tinygrad_worker.py" PYTHON = sys.executable STUB_HIDDEN_DIM = worker.STUB_HIDDEN_DIM STUB_VOCAB_SIZE = worker.STUB_VOCAB_SIZE STUB_HIDDEN_BYTES_PER_POS = STUB_HIDDEN_DIM * 2 # bf16 # --------------------------------------------------------------------------- # Subprocess helpers def _spawn(stage: int, num_stages: int, *, extra_env=None) -> subprocess.Popen: env = os.environ.copy() env["STAGE"] = str(stage) env["NUM_STAGES"] = str(num_stages) env["PP_WORKER_STUB"] = "1" if extra_env: env.update(extra_env) return subprocess.Popen( [PYTHON, str(WORKER), "--stub"], stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, env=env, ) def _read_reply(proc: subprocess.Popen, timeout: float = 5.0) -> dict: sel = selectors.DefaultSelector() sel.register(proc.stdout, selectors.EVENT_READ) try: if not sel.select(timeout=timeout): stderr = "" try: stderr = proc.stderr.read() or "" except Exception: pass raise TimeoutError(f"No worker reply within {timeout}s; stderr: {stderr!r}") line = proc.stdout.readline() finally: sel.close() if not line: raise EOFError("Worker closed stdout before replying") return json.loads(line.strip()) def _send(proc: subprocess.Popen, obj: dict) -> None: proc.stdin.write(json.dumps(obj) + "\n") proc.stdin.flush() def _send_raw(proc: subprocess.Popen, text: str) -> None: proc.stdin.write(text + "\n") proc.stdin.flush() def _shutdown(proc: subprocess.Popen) -> None: if proc.poll() is None: proc.terminate() try: proc.wait(timeout=5) except subprocess.TimeoutExpired: proc.kill() proc.wait(timeout=2) @pytest.fixture def stage0(): proc = _spawn(0, 2) try: yield proc finally: _shutdown(proc) @pytest.fixture def stage1(): proc = _spawn(1, 2) try: yield proc finally: _shutdown(proc) # --------------------------------------------------------------------------- # §2 — pure helper tests class TestLayerMath: def test_stage_0_layer_range_is_lower_half(self): # Explicit half-split for an even block count. assert worker.compute_layer_range(0, 2, 16) == (0, 8) # And the same shape holds for any even total. for total in (2, 4, 32, 100): start, end = worker.compute_layer_range(0, 2, total) assert start == 0 assert end == total // 2 def test_stage_1_layer_range_covers_remainder(self): # 15 blocks split across 2 stages: stage 1 must absorb the odd block. assert worker.compute_layer_range(1, 2, 15) == (7, 15) # The end of stage N-1 always equals total, regardless of remainder. for total in (2, 5, 16, 17, 101): _, end = worker.compute_layer_range(1, 2, total) assert end == total def test_layer_range_partition_is_total_coverage(self): # For 2-, 3-, 4-stage splits, the union of all stage ranges must equal # [0, total) exactly — no gaps, no overlap. Generalises early because # the only marginal cost is a few asserts. for num_stages in (2, 3, 4): # Include totals divisible by num_stages and totals that leave a # remainder, so we exercise the "last stage absorbs remainder" path. for total in ( num_stages, num_stages + 1, num_stages * 5, num_stages * 5 + (num_stages - 1), ): ranges = [ worker.compute_layer_range(s, num_stages, total) for s in range(num_stages) ] covered: list[int] = [] for start, end in ranges: assert ( start < end ), f"empty range {start}..{end} (n={num_stages}, total={total})" covered.extend(range(start, end)) assert covered == list( range(total) ), f"coverage mismatch n={num_stages} total={total}: {ranges}" def test_argmax_sampling_is_deterministic(self): logits = [0.1, 0.4, 0.2, 0.3] first = worker.argmax_sample(logits) # Determinism: many calls all produce the same id. for _ in range(10): assert worker.argmax_sample(logits) == first # And the value must actually be the argmax, not a constant — # otherwise this test would pass for ``def argmax(_): return 0``. assert first == 1 assert worker.argmax_sample([5.0, 1.0, 1.0, 1.0]) == 0 assert worker.argmax_sample([1.0, 1.0, 1.0, 9.0]) == 3 # --------------------------------------------------------------------------- # §3 — worker contract tests class TestWorkerStartup: def test_worker_emits_ready_with_pid_and_stage(self, stage0): ready = _read_reply(stage0) assert ready["status"] == "ready" assert ready["pid"] == stage0.pid assert ready["stage"] == 0 def test_worker_rejects_invalid_stage_env(self): # Every invalid configuration must exit non-zero quickly, never # reach the ready line, and never hang waiting for stdin. invalid = [ {"STAGE": "2", "NUM_STAGES": "2"}, # out of range high {"STAGE": "-1", "NUM_STAGES": "2"}, # negative {"STAGE": "abc", "NUM_STAGES": "2"}, # non-numeric {"STAGE": "", "NUM_STAGES": "2"}, # missing {"STAGE": "0", "NUM_STAGES": "0"}, # zero stages {"STAGE": "٠", "NUM_STAGES": "2"}, # arabic-indic 0 (unicode digit) ] for env_overrides in invalid: env = os.environ.copy() env["PP_WORKER_STUB"] = "1" env.update(env_overrides) proc = subprocess.Popen( [PYTHON, str(WORKER), "--stub"], stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, env=env, ) try: exit_code = proc.wait(timeout=5) finally: _shutdown(proc) assert ( exit_code != 0 ), f"expected non-zero exit for env={env_overrides}, got 0" class TestStage0Operations: def test_embed_and_forward_returns_hidden_for_prompt(self, stage0): _read_reply(stage0) # ready tokens = [1, 2, 3, 4, 5] _send( stage0, {"op": "embed_and_forward", "request_id": 1, "tokens": tokens, "position": 0}, ) reply = _read_reply(stage0) assert "error" not in reply, reply assert reply["request_id"] == 1 assert reply["seq_len"] == len(tokens) hidden = base64.b64decode(reply["hidden_b64"]) assert len(hidden) == len(tokens) * STUB_HIDDEN_BYTES_PER_POS # Payload must be non-trivial so a `return b"\0" * N` stub would not pass. assert any(b != 0 for b in hidden) def test_decode_step_returns_hidden_for_single_token(self, stage0): _read_reply(stage0) _send( stage0, {"op": "decode_step", "request_id": 2, "token_id": 42, "position": 5}, ) reply = _read_reply(stage0) assert "error" not in reply, reply assert reply["request_id"] == 2 assert reply["seq_len"] == 1 hidden = base64.b64decode(reply["hidden_b64"]) assert len(hidden) == STUB_HIDDEN_BYTES_PER_POS assert any(b != 0 for b in hidden) def test_kv_cache_grows_across_successive_decode_steps(self, stage0): _read_reply(stage0) _send( stage0, {"op": "embed_and_forward", "request_id": 1, "tokens": [10, 20, 30], "position": 0}, ) assert "error" not in _read_reply(stage0) # Successive decode_step calls at increasing positions all succeed. hidden_blobs: set[str] = set() for i, pos in enumerate((3, 4, 5, 6)): _send( stage0, {"op": "decode_step", "request_id": 100 + i, "token_id": 99, "position": pos}, ) reply = _read_reply(stage0) assert "error" not in reply, f"decode at position {pos} failed: {reply}" assert reply["seq_len"] == 1 hidden_blobs.add(reply["hidden_b64"]) # Position must actually influence output — otherwise the worker is # silently ignoring it and a real model would corrupt its KV cache. assert len(hidden_blobs) > 1 def test_stage_0_rejects_stage_1_ops(self, stage0): _read_reply(stage0) b64 = base64.b64encode(b"\x01" * STUB_HIDDEN_BYTES_PER_POS).decode() _send( stage0, { "op": "forward_and_sample", "request_id": 7, "hidden_b64": b64, "position": 0, "seq_len": 1, }, ) reply = _read_reply(stage0) assert "error" in reply, reply assert reply.get("request_id") == 7 # Worker survives — a follow-up valid request still works. _send( stage0, {"op": "decode_step", "request_id": 8, "token_id": 1, "position": 0}, ) ok = _read_reply(stage0) assert "error" not in ok, ok assert ok["request_id"] == 8 class TestStage1Operations: def test_forward_and_sample_returns_valid_token_id(self, stage1): _read_reply(stage1) hidden = base64.b64encode(b"\x42" * STUB_HIDDEN_BYTES_PER_POS).decode() _send( stage1, { "op": "forward_and_sample", "request_id": 1, "hidden_b64": hidden, "position": 0, "seq_len": 1, }, ) reply = _read_reply(stage1) assert "error" not in reply, reply token = reply["token_id"] assert isinstance(token, int) assert 0 <= token < STUB_VOCAB_SIZE def test_forward_and_sample_is_deterministic_for_same_input(self, stage1): _read_reply(stage1) hidden = base64.b64encode(bytes(range(STUB_HIDDEN_BYTES_PER_POS))).decode() observed = [] for rid in (1, 2, 3, 4): _send( stage1, { "op": "forward_and_sample", "request_id": rid, "hidden_b64": hidden, "position": 7, "seq_len": 1, }, ) reply = _read_reply(stage1) assert "error" not in reply, reply observed.append(reply["token_id"]) assert len(set(observed)) == 1, f"non-deterministic token ids: {observed}" # Diversity guard: a constant ``return 0`` implementation would # trivially satisfy determinism. Probe several distinct positions # and require at least two distinct outputs — collision across all # of these in a 32-id vocab is astronomically unlikely if the # implementation actually mixes position into the result. diverse = {observed[0]} for rid, pos in enumerate((1, 11, 101, 12345, 999_999), start=200): _send( stage1, { "op": "forward_and_sample", "request_id": rid, "hidden_b64": hidden, "position": pos, "seq_len": 1, }, ) diverse.add(_read_reply(stage1)["token_id"]) assert ( len(diverse) > 1 ), f"position appears to be ignored — all positions mapped to {observed[0]}" def test_stage_1_rejects_stage_0_ops(self, stage1): _read_reply(stage1) _send( stage1, {"op": "embed_and_forward", "request_id": 5, "tokens": [1, 2], "position": 0}, ) reply = _read_reply(stage1) assert "error" in reply, reply assert reply.get("request_id") == 5 # Survives and serves its own op. ok_hidden = base64.b64encode(b"\x00" * STUB_HIDDEN_BYTES_PER_POS).decode() _send( stage1, { "op": "forward_and_sample", "request_id": 6, "hidden_b64": ok_hidden, "position": 0, "seq_len": 1, }, ) ok = _read_reply(stage1) assert "error" not in ok, ok class TestWorkerMalformedInput: def test_malformed_json_returns_error_and_continues(self, stage0): _read_reply(stage0) _send_raw(stage0, "this is not json {{{") err = _read_reply(stage0) assert "error" in err, err # Recovery _send( stage0, {"op": "decode_step", "request_id": 99, "token_id": 1, "position": 0}, ) ok = _read_reply(stage0) assert "error" not in ok, ok assert ok["request_id"] == 99 def test_missing_op_field_returns_error(self, stage0): _read_reply(stage0) # Valid JSON object, but no "op". _send(stage0, {"request_id": 1, "tokens": [1, 2, 3]}) err = _read_reply(stage0) assert "error" in err, err # Plain JSON scalars are not objects either — they must also produce # an error, not crash the worker. _send_raw(stage0, "42") err2 = _read_reply(stage0) assert "error" in err2, err2 # Recovery _send( stage0, {"op": "decode_step", "request_id": 2, "token_id": 1, "position": 0}, ) ok = _read_reply(stage0) assert "error" not in ok, ok def test_oversized_hidden_payload_returns_error(self, stage1): _read_reply(stage1) # Declared seq_len > actual hidden length. short_hidden = base64.b64encode(b"\x00" * STUB_HIDDEN_BYTES_PER_POS).decode() _send( stage1, { "op": "forward_and_sample", "request_id": 1, "hidden_b64": short_hidden, "position": 0, "seq_len": 2, }, ) err = _read_reply(stage1) assert "error" in err, err # Inverse: declared seq_len < actual hidden length. long_hidden = base64.b64encode(b"\x00" * (STUB_HIDDEN_BYTES_PER_POS * 5)).decode() _send( stage1, { "op": "forward_and_sample", "request_id": 2, "hidden_b64": long_hidden, "position": 0, "seq_len": 1, }, ) err2 = _read_reply(stage1) assert "error" in err2, err2 # Worker still serves a well-formed follow-up. _send( stage1, { "op": "forward_and_sample", "request_id": 3, "hidden_b64": short_hidden, "position": 0, "seq_len": 1, }, ) ok = _read_reply(stage1) assert "error" not in ok, ok assert "token_id" in ok class TestWorkerEOFShutdown: def test_eof_causes_clean_exit(self, stage0): _read_reply(stage0) # consume ready stage0.stdin.close() exit_code = stage0.wait(timeout=5) assert exit_code == 0, f"worker exited with code {exit_code}, expected 0" # --------------------------------------------------------------------------- # Real (non-stub) tinygrad worker # # Loading the GGUF takes ~15s and depends on a network-fetched file. These # tests are skipped unless the developer opts in by setting # ``PP_REAL_WORKER_TESTS=1`` (or any non-empty string). The ``cargo test`` # fast tier and the default ``pytest`` invocation skip the class entirely. REAL_WORKER_GATE = "PP_REAL_WORKER_TESTS" REAL_LOAD_TIMEOUT = 180.0 # seconds; GGUF fetch + tinygrad realize REAL_OP_TIMEOUT = 120.0 # seconds; one block-range forward on CPU def _spawn_real(stage: int, num_stages: int, *, model: str = "llama3.2:1b") -> subprocess.Popen: """Spawn a real-mode (non-stub) worker. PP_WORKER_STUB is unset.""" env = os.environ.copy() env["STAGE"] = str(stage) env["NUM_STAGES"] = str(num_stages) env["MODEL"] = model env.pop("PP_WORKER_STUB", None) return subprocess.Popen( [PYTHON, str(WORKER)], stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, env=env, ) @pytest.mark.skipif( not os.environ.get(REAL_WORKER_GATE), reason=f"Set {REAL_WORKER_GATE}=1 to run real tinygrad worker tests", ) class TestRealTinygradWorker: """One prefill + one decode-step round-trip against the real GGUF. This is the slow-tier counterpart to ``TestStage0Operations`` / ``TestStage1Operations``. It does not duplicate every stub-mode test case — those exercise the protocol's error paths. Here we just establish that real-mode embeds, forwards, samples, and decodes using the actual ``llama3.2:1b`` weights without the stub. """ def test_real_worker_advertises_model_geometry_on_ready(self): stage0 = _spawn_real(0, 2) try: ready = _read_reply(stage0, timeout=REAL_LOAD_TIMEOUT) assert ready["status"] == "ready", ready assert ready["stage"] == 0 # Real-mode adds geometry fields so the orchestrator can # size hidden-state buffers without an extra handshake. assert isinstance(ready.get("hidden_dim"), int) and ready["hidden_dim"] > 0 assert isinstance(ready.get("vocab_size"), int) and ready["vocab_size"] > 0 assert isinstance(ready.get("total_blocks"), int) and ready["total_blocks"] > 0 lr = ready.get("layer_range") assert isinstance(lr, list) and len(lr) == 2 assert lr == [0, ready["total_blocks"] // 2] finally: _shutdown(stage0) def test_real_prefill_and_decode_step_round_trip(self): # Drives one full prefill + decode round-trip through both stages # on the real model. The interesting assertions are byte lengths # (catches any dtype/shape mismatch) and that stage-1 actually # samples a token in ``[0, vocab)``. stage0 = _spawn_real(0, 2) stage1 = _spawn_real(1, 2) try: ready0 = _read_reply(stage0, timeout=REAL_LOAD_TIMEOUT) ready1 = _read_reply(stage1, timeout=REAL_LOAD_TIMEOUT) hidden_dim = ready0["hidden_dim"] vocab = ready0["vocab_size"] assert ready1["hidden_dim"] == hidden_dim assert ready1["vocab_size"] == vocab # Tokenise via the real worker so the prompt actually maps to # GGUF vocab ids; saves us from duplicating the tokenizer. _send(stage0, {"op": "tokenize", "request_id": 1, "prompt": "Say hello"}) tok_reply = _read_reply(stage0, timeout=REAL_OP_TIMEOUT) assert "error" not in tok_reply, tok_reply tokens = tok_reply["tokens"] assert isinstance(tokens, list) and len(tokens) > 0 assert all(isinstance(t, int) and 0 <= t < vocab for t in tokens), tokens # Stage 0: prefill at position 0. _send( stage0, { "op": "embed_and_forward", "request_id": 2, "tokens": tokens, "position": 0, }, ) prefill = _read_reply(stage0, timeout=REAL_OP_TIMEOUT) assert "error" not in prefill, prefill assert prefill["seq_len"] == len(tokens) hidden_pref = base64.b64decode(prefill["hidden_b64"]) assert len(hidden_pref) == len(tokens) * hidden_dim * 2 # Non-trivial payload: a `b"\x00" * N` return would also pass # the length assertion, so reject that explicitly. assert any(b != 0 for b in hidden_pref) # Stage 1: forward + sample. _send( stage1, { "op": "forward_and_sample", "request_id": 3, "hidden_b64": prefill["hidden_b64"], "position": 0, "seq_len": len(tokens), }, ) sampled = _read_reply(stage1, timeout=REAL_OP_TIMEOUT) assert "error" not in sampled, sampled tok_id = sampled["token_id"] assert isinstance(tok_id, int) and 0 <= tok_id < vocab # Stage 0: decode step at position == prompt_len. Per the # plan, off-by-one position handling is the #1 risk; this # exercises the single-token branch. _send( stage0, { "op": "decode_step", "request_id": 4, "token_id": tok_id, "position": len(tokens), }, ) decode = _read_reply(stage0, timeout=REAL_OP_TIMEOUT) assert "error" not in decode, decode assert decode["seq_len"] == 1 hidden_dec = base64.b64decode(decode["hidden_b64"]) assert len(hidden_dec) == 1 * hidden_dim * 2 assert any(b != 0 for b in hidden_dec) # Stage 1: forward + sample the decode-step hidden state. _send( stage1, { "op": "forward_and_sample", "request_id": 5, "hidden_b64": decode["hidden_b64"], "position": len(tokens), "seq_len": 1, }, ) sampled2 = _read_reply(stage1, timeout=REAL_OP_TIMEOUT) assert "error" not in sampled2, sampled2 tok_id2 = sampled2["token_id"] assert isinstance(tok_id2, int) and 0 <= tok_id2 < vocab # The two sampled tokens should not both be a default # zero/special id — a constant-output implementation would # match the assertions above. Detokenise both and require # the resulting bytes to be non-empty. _send( stage1, {"op": "detokenize", "request_id": 6, "tokens": [tok_id, tok_id2]}, ) detok = _read_reply(stage1, timeout=REAL_OP_TIMEOUT) assert "error" not in detok, detok assert isinstance(detok["text"], str) assert detok["text"] != "" finally: _shutdown(stage0) _shutdown(stage1)