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.

• View on GitHub • More from Dicklesworthstone

A massive, ancient stone vault where the door is a glowing, complex mathematical lattice. Inside, a single, tiny, perfect gear is being measured by laser-precise calipers.
FrankenJAX treats mathematical transformations not as a library feature, but as high-stakes evidence that must be protected and verified.

Key Takeaways

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)).

A document being shattered into a hundred floating shards, caught by a net of thin threads that pull them back together.
RaptorQ erasure coding ensures that the 'evidence' of mathematical correctness survives bit-rot.

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.

The transform pipeline: stripping away Python magic to build a deterministic IR in safe Rust.

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.

A single tree trunk splitting into a shimmering cloud of hundreds of possible branch configurations, with one highlighted path.
Equality saturation explores thousands of mathematically equivalent expressions simultaneously to find the optimal path.

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.

ModeBehaviorUse Case
StrictFails closed on unknown features or deviations from the JAX Oracle.Ensuring 1:1 behavioral parity and preventing silent drift.
HardenedRoutes 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.

Strict mode enforces absolute parity, while Hardened mode provides bounded recovery for adversarial inputs.

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.