{ "cells": [ { "cell_type": "markdown", "id": "00954fac", "metadata": {}, "source": [ "# Swapping a data structure at runtime with coheriq\n", "\n", "**coheriq** lets a library mark certain functions *and data structures* as\n", "candidates for acceleration, and lets a separate provider swap in alternative\n", "implementations at runtime — without the calling code changing at all.\n", "\n", "There are three parties in this example, as is typical:\n", "\n", "1. **The domain** — a library (`coheriq_ragged_domain`) that marks candidates\n", " and ships working pure-Python defaults.\n", "2. **The engine** — a separate package (`coheriq_ragged_engine`) that provides\n", " a faster, compiled implementation (C++ via nanobind, in this case), discovered through a\n", " Python *entry point*.\n", "3. **User code** — activates an engine once, then calls the library exactly as\n", " before.\n", "\n", "This notebook demonstrates the **mechanism**: a container class *and* the functions\n", "that operate on it are swapped, and the calling code is unchanged.\n", "\n", "It is written to run **either way** — against the pure-Python default *or* the\n", "compiled engine — without editing any calling code. You choose the backend by\n", "activating the engine (see below) or by setting the `RAGGED_ENGINE=fused`\n", "environment variable before launching. Run it both ways and compare the timing\n", "printed near the end: same code, different backend.\n", "\n", "The source code to the domain and engine are available in the `example/` directory\n", "of the coheriq repository." ] }, { "cell_type": "markdown", "id": "d3a6153b", "metadata": {}, "source": [ "## The shape: ragged data\n", "\n", "Our computation is a **per-row softmax** over a batch of rows that have\n", "*different lengths* — a *ragged* batch. Ragged data is everywhere:\n", "variable-length sequences, a different number of candidates per item, the\n", "neighbors of each node in a graph.\n", "\n", "This is exactly the shape that a rectangular array library like NumPy handles\n", "awkwardly. You have two options, both unpleasant:\n", "\n", "- **A Python list of arrays** — back to per-row Python overhead, which defeats\n", " the vectorization that made NumPy worth reaching for.\n", "- **Pad every row to the longest and carry a mask** — memory and compute become\n", " `O(rows × max_len)` instead of `O(total elements)`, which is wasteful when\n", " row lengths vary a lot.\n", "\n", "This is why libraries grew dedicated ragged/segmented machinery —\n", "`tf.RaggedTensor`, `torch_scatter`, JAX's `segment_*` operations. The shape is\n", "real and general. coheriq gives *your own* library a clean way to swap in such a\n", "backend for *its own* data type, without your users rewriting anything." ] }, { "cell_type": "markdown", "id": "31168234", "metadata": {}, "source": [ "## Using the library (pure Python, no engine yet)\n", "\n", "A user installs `coheriq_ragged_domain` and imports its container and functions." ] }, { "cell_type": "code", "execution_count": 1, "id": "feac613f", "metadata": {}, "outputs": [], "source": [ "import inspect\n", "\n", "from coheriq_ragged_domain import RaggedBatch, segmented_softmax, segmented_topk_softmax" ] }, { "cell_type": "markdown", "id": "41127d52", "metadata": {}, "source": [ "## What the default looks like\n", "\n", "coheriq resolves each candidate's implementation on its **first call** and then\n", "freezes it for the rest of the process. So we must choose the backend *before*\n", "calling anything — which is why this notebook decides up front and then runs one\n", "way through. (If you call a candidate first and *then* try to enable an engine,\n", "coheriq deliberately raises to tell you it is too late.)\n", "\n", "Here is the **default** implementation, shown by *reading* it (this always shows\n", "the domain's own source, whichever backend is active). The container stores all\n", "rows in one flat buffer plus an `offsets` array (CSR-style), and\n", "`segmented_softmax` does the obvious per-row work — find the max, exponentiate,\n", "sum, divide — materializing a Python list for each row:" ] }, { "cell_type": "code", "execution_count": 2, "id": "3a5c20b3", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "@_domain.acceleration_candidate\n", "def segmented_softmax(batch: RaggedBatch) -> RaggedBatch:\n", " \"\"\"Return a new batch holding the per-row softmax of ``batch``.\n", "\n", " The default implementation is the obvious, readable version: for each row it\n", " makes several passes -- find the max (for numerical stability), subtract and\n", " exponentiate, sum, then divide. Each row is materialized as a Python list\n", " along the way.\n", " \"\"\"\n", " result_rows: list[list[float]] = []\n", " for row in batch:\n", " if not row:\n", " result_rows.append([])\n", " continue\n", " m = max(row)\n", " exps = [math.exp(x - m) for x in row]\n", " total = sum(exps)\n", " result_rows.append([e / total for e in exps])\n", " return RaggedBatch(result_rows)\n", "\n" ] } ], "source": [ "print(inspect.getsource(segmented_softmax))" ] }, { "cell_type": "markdown", "id": "f62e7e35", "metadata": {}, "source": "## Activating the engine — two ways, without importing it\n\nA *separate* package, `coheriq_ragged_engine`, is installed. It declares itself\nunder the entry-point group `coheriq.engines.ragged` with the name `\"fused\"`.\nIts implementation is compiled C++ (via nanobind) — but you never `import` it.\n\nThere are two equivalent ways to activate it, and **neither** imports the engine\npackage:\n\n1. **Explicitly**, by name — uncomment the line in the next cell:\n `coheriq.enable_engine(\"ragged\", \"fused\")`.\n2. **Via the environment** — set `RAGGED_ENGINE=fused` before launching; coheriq\n activates it automatically on first use.\n\nEither way, coheriq looks up the entry point, imports the compiled engine for\nus, and enables it. Leave the line commented and the variable unset to run\nagainst the pure-Python default instead. **Notice there is no\n`import coheriq_ragged_engine` anywhere in this notebook.**" }, { "cell_type": "code", "execution_count": null, "id": "84a0c126", "metadata": {}, "outputs": [], "source": "# Uncomment to activate the compiled engine explicitly (or set RAGGED_ENGINE=fused):\n# import coheriq\n# coheriq.enable_engine(\"ragged\", \"fused\")" }, { "cell_type": "markdown", "id": "5ecf3d78", "metadata": {}, "source": [ "## The swap: same calling code, whichever backend is active\n", "\n", "We build a batch with wildly varying row lengths (including an empty row). The\n", "calling code is identical no matter which backend is active — the object we get\n", "back reports its backend via `.backend`, so we can see which implementation is\n", "live. If you activated the engine, it says `fused`; otherwise `default`." ] }, { "cell_type": "code", "execution_count": 4, "id": "9963a919", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "RaggedBatch(4 rows, default backend)\n", "backend: default\n", "len: 4\n", "row_lengths: [3, 1, 4, 0]\n", "batch[0] (a copy): [2.0, 1.0, 0.1]\n" ] } ], "source": [ "batch = RaggedBatch([[2.0, 1.0, 0.1], [1.0], [3.0, 3.0, 0.0, 0.0], []])\n", "\n", "print(repr(batch))\n", "print(\"backend: \", batch.backend)\n", "print(\"len: \", len(batch))\n", "print(\"row_lengths:\", batch.row_lengths())\n", "print(\"batch[0] (a copy):\", batch[0])" ] }, { "cell_type": "markdown", "id": "ff6e74b4", "metadata": {}, "source": [ "The operation is swapped too. `segmented_softmax` runs whichever backend is\n", "active, and the result it hands back reports the same backend — so a result from\n", "the compiled engine is itself a `fused` batch." ] }, { "cell_type": "code", "execution_count": 5, "id": "48935119", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "result backend: default\n", "row 0: sum=1.000000 [0.659, 0.2424, 0.0986]\n", "row 1: sum=1.000000 [1.0]\n", "row 2: sum=1.000000 [0.4763, 0.4763, 0.0237, 0.0237]\n", "row 3: sum=0.000000 []\n" ] } ], "source": [ "probs = segmented_softmax(batch)\n", "\n", "print(\"result backend:\", probs.backend)\n", "for i in range(len(probs)):\n", " row = probs[i]\n", " print(f\"row {i}: sum={sum(row):.6f} {[round(x, 4) for x in row]}\")" ] }, { "cell_type": "markdown", "id": "76d67e22", "metadata": {}, "source": [ "## Correctness by invariant\n", "\n", "Different engines need not produce *bit-identical* output (a compiled backend\n", "might sum in a different order). So we check **invariants** rather than equality:\n", "the batch shape is preserved and every non-empty row is a valid probability\n", "distribution." ] }, { "cell_type": "code", "execution_count": 6, "id": "62eb97d6", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "invariants hold; active backend is default\n" ] } ], "source": [ "def check_softmax_invariants(inp, out, tol=1e-9):\n", " assert len(out) == len(inp)\n", " assert out.row_lengths() == inp.row_lengths()\n", " for i in range(len(out)):\n", " row = out[i]\n", " if row:\n", " assert abs(sum(row) - 1.0) < tol, (i, sum(row))\n", " assert all(-tol <= p <= 1.0 + tol for p in row)\n", " return True\n", "\n", "\n", "assert check_softmax_invariants(batch, segmented_softmax(batch))\n", "print(\"invariants hold; active backend is\", probs.backend)" ] }, { "cell_type": "markdown", "id": "59ef6f2c", "metadata": {}, "source": [ "## Where the compiled engine wins\n", "\n", "When the `fused` engine is active, `segmented_softmax` runs a compiled C++\n", "implementation (via nanobind) instead of the Python default. It exploits exactly\n", "the structure this data shape affords:\n", "\n", "- **A fused pass over contiguous storage** — the flat `values` buffer is walked\n", " directly; no per-row temporary Python list is allocated.\n", "- **No Python object overhead** — no boxed floats, tuples, or dynamic dispatch\n", " in the inner loop.\n", "- **Room for SIMD** on the exponential and the arithmetic.\n", "- **Parallelism across rows** — the rows are independent; the computation is\n", " embarrassingly parallel.\n", "\n", "From the caller's point of view nothing changes: the same\n", "`segmented_softmax(batch)` call dispatches to it. The benchmark near the end\n", "lets you measure the difference by running this notebook once with the default\n", "and once with the engine active." ] }, { "cell_type": "markdown", "id": "63ff5fda", "metadata": {}, "source": [ "## Elaboration: top-k softmax\n", "\n", "The same swap pattern applies to other operations on the container. A common\n", "variant keeps only the `k` largest entries per row and softmaxes over just\n", "those. The default sorts each row (`O(n log n)`); the compiled engine *selects*\n", "the top `k` with `std::nth_element` (`O(n)`) — an *algorithmic* win, not just a\n", "constant factor. Whichever backend is active, the call site is the same:" ] }, { "cell_type": "code", "execution_count": 7, "id": "dc1b6345", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "row 0: [(0, 0.7311), (1, 0.2689)]\n", "row 1: [(0, 1.0)]\n", "row 2: [(0, 0.5), (1, 0.5)]\n", "row 3: []\n" ] } ], "source": [ "topk = segmented_topk_softmax(batch, k=2)\n", "for i, row in enumerate(topk):\n", " print(f\"row {i}: {[(idx, round(p, 4)) for idx, p in row]}\")" ] }, { "cell_type": "markdown", "id": "aafd764c", "metadata": {}, "source": [ "## Benchmark: same code, two backends\n", "\n", "The honest way to compare is to run *this exact notebook* twice — once with the\n", "default backend and once with the engine active (uncomment the activation cell,\n", "or set `RAGGED_ENGINE=fused`) — and compare the number printed below. The cell\n", "only prints its timing; it makes no assertion about it (hardware varies, and a\n", "compiled backend only pulls ahead once the rows are large enough to outweigh the\n", "Python call overhead), so it never fails a test run either way.\n", "\n", "We build a large ragged batch with a fixed seed so both runs use identical\n", "input, then time `segmented_softmax`." ] }, { "cell_type": "code", "execution_count": 8, "id": "c1675c36", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "backend=default rows=20000 elapsed=182.8 ms\n" ] } ], "source": [ "import random\n", "import time\n", "\n", "rng = random.Random(20260710)\n", "big_rows = [[rng.gauss(0.0, 1.0) for _ in range(rng.randint(1, 128))] for _ in range(20_000)]\n", "big = RaggedBatch(big_rows)\n", "\n", "start = time.perf_counter()\n", "result = segmented_softmax(big)\n", "elapsed = time.perf_counter() - start\n", "\n", "# Correctness is checked regardless of backend; timing is only reported.\n", "assert check_softmax_invariants(big, result)\n", "print(f\"backend={big.backend} rows={len(big)} elapsed={elapsed * 1e3:.1f} ms\")" ] }, { "cell_type": "markdown", "id": "ab36a703", "metadata": {}, "source": [ "## Recap\n", "\n", "- A library (**domain**) marked a container and some functions as acceleration\n", " candidates and shipped working pure-Python defaults.\n", "- A separate **engine** package provided a compiled C++ implementation and was\n", " discovered via a Python entry point — we activated it *by name* (or via an\n", " environment variable) and never imported it.\n", "- The **user code** was unchanged across the swap; both the `RaggedBatch`\n", " container and the functions operating on it were replaced underneath it.\n", "- We proved the swap with an observable marker (`.backend`), checked correctness\n", " by invariant, and measured the difference by running the same notebook against\n", " each backend." ] } ], "metadata": { "kernelspec": { "display_name": "Python 3 (ipykernel)", "language": "python", "name": "python3" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.14.6" } }, "nbformat": 4, "nbformat_minor": 5 }