FrankenJAX: Clean-Room Mathematics and the Pursuit of the Perfect Trace
How a solo developer is rebuilding JAX from first principles in Rust to achieve bit-stable, verifiable numerical sovereignty.
- FrankenJAX uses a clean-room Rust implementation to achieve bit-stable numerical parity with the original JAX Python library.
- The system employs RaptorQ erasure coding to protect a ledger of mathematical proofs against data corruption and bit-rot.
- Equality saturation via e-graphs allows the compiler to explore thousands of mathematically equivalent expressions simultaneously for optimal performance.
- A dual-runtime policy provides a hardened execution mode that generates confidence values and safety logs for adversarial environments.
In the fast-moving world of machine learning frameworks, the accepted wisdom is to build fast, wrap C++ or Python, and accept a certain level of "black box" uncertainty. FrankenJAX takes the opposite approach. It is built on the premise of mathematical paranoia. It assumes that bit-rot, silent numerical drift, and opaque execution are the enemies of reliable autonomous systems.
Rather than simply providing a Rust wrapper around Google's JAX, FrankenJAX is a clean-room reimplementation of JAX's core transform semantics. It treats the original Python library as a "Legacy Oracle"—a black box to be interrogated to achieve 1:1 "Differential Conformance." This isn't just a port; it's a compiler built like a high-security vault.
The Durability of a Proof
One of the most unusual architectural decisions in FrankenJAX is its use of RaptorQ erasure coding. Typically found in satellite communications or distributed storage, RaptorQ (RFC 6330) allows data to be reconstructed even if pieces of it are lost or corrupted.
Why does a compiler need fountain codes? Because FrankenJAX treats mathematical results not just as outputs, but as evidence. The system generates a "Trace Transform Ledger" (TTL)—an audit log that provides verifiable proofs of how an input Intermediate Representation (IR) was mutated by a stack of transforms, such as deriving grad(jit(f)).
These artifacts, including benchmark baselines and conformance fixtures, are critical for proving that the Rust implementation hasn't drifted from the Python oracle. By using RaptorQ, FrankenJAX ensures that this evidence is resilient against bit-rot, allowing the system to self-heal its test artifacts over time.
Inside the Transform Sandwich
The architecture of FrankenJAX follows a classic compiler "sandwich" pattern, but implemented entirely in safe Rust. The frontend traces Rust closures into a Canonical JAXPR-like IR. The middle-end applies higher-order transforms like grad and vmap, maintaining the Vector-Jacobian Products (VJPs) and Jacobian-Vector Products (JVPs) necessary for automatic differentiation.
Finally, the backend lowers this transformed IR into executable primitives—linear algebra, FFTs, and tensor operations—executing them via a dependency-wave parallel executor. By stripping away the Python layers, FrankenJAX exposes the raw mechanics of these transformations, providing a clearer view of how vectorization and differentiation actually work under the hood.
Optimization via Equality Saturation
Traditional compilers often rely on a fixed sequence of optimization passes, which can miss complex, multi-step simplifications. FrankenJAX takes a different route, utilizing the egg crate to perform E-graph optimization.
Instead of applying rules sequentially, equality saturation explores thousands of mathematically equivalent versions of a program simultaneously. FrankenJAX uses dozens of algebraic rules to find the most efficient form of an expression before lowering it to machine code, ensuring that the resulting mathematical operations are as optimized as possible without sacrificing correctness.
The Hardened Runtime
Because FrankenJAX is designed for autonomous agent swarms and high-stakes environments, it operates under a dual-runtime policy.
| Mode | Behavior | Use Case |
|---|---|---|
| Strict | Fails closed on unknown features or deviations from the JAX Oracle. | Ensuring 1:1 behavioral parity and preventing silent drift. |
| Hardened | Routes unknown ops through safety guards, logs confidence values, and continues with warnings. | Handling adversarial inputs or unstable edge cases in production agent swarms. |
This "Decision-Theoretic Dispatch" means the system doesn't just crash or silently produce incorrect results when faced with an anomaly. It captures a loss matrix and confidence values, making the why of a runtime decision as observable as the what.
By bringing JAX's powerful functional transformations into a secure, observable Rust environment, FrankenJAX provides the numerical backbone for the next generation of deterministic AI systems.