惯性聚合 高效追踪和阅读你感兴趣的博客、新闻、科技资讯
阅读原文 在惯性聚合中打开

推荐订阅源

Know Your Adversary
Know Your Adversary
C
CERT Recently Published Vulnerability Notes
V
Vulnerabilities – Threatpost
N
News | PayPal Newsroom
O
OpenAI News
A
About on SuperTechFans
月光博客
月光博客
Martin Fowler
Martin Fowler
L
LangChain Blog
OSCHINA 社区最新新闻
OSCHINA 社区最新新闻
aimingoo的专栏
aimingoo的专栏
freeCodeCamp Programming Tutorials: Python, JavaScript, Git & More
V2EX - 技术
V2EX - 技术
博客园 - 司徒正美
大猫的无限游戏
大猫的无限游戏
K
KPMG report finds enterprise disconnect between AI and its ROI | CIO
人人都是产品经理
人人都是产品经理
T
Tor Project blog
D
DataBreaches.Net
Cloudbric
Cloudbric
让小产品的独立变现更简单 - ezindie.com
让小产品的独立变现更简单 - ezindie.com
云风的 BLOG
云风的 BLOG
F
Full Disclosure
S
SegmentFault 最新的问题
Vercel News
Vercel News
T
Tailwind CSS Blog
Schneier on Security
Schneier on Security
宝玉的分享
宝玉的分享
M
MIT News - Artificial intelligence
博客园 - 【当耐特】
P
Privacy International News Feed
美团技术团队
S
Secure Thoughts
P
Privacy & Cybersecurity Law Blog
Google DeepMind News
Google DeepMind News
F
Fortinet All Blogs
Scott Helme
Scott Helme
Forbes - Security
Forbes - Security
D
Darknet – Hacking Tools, Hacker News & Cyber Security
Recent Announcements
Recent Announcements
AWS News Blog
AWS News Blog
Stack Overflow Blog
Stack Overflow Blog
S
Security Affairs
CTFtime.org: upcoming CTF events
CTFtime.org: upcoming CTF events
The GitHub Blog
The GitHub Blog
T
Tenable Blog
Cyberwarzone
Cyberwarzone
奇客Solidot–传递最新科技情报
奇客Solidot–传递最新科技情报
J
Java Code Geeks
T
Threatpost

MarkTechPost

A Coding Implementation of End-to-End Brain Decoding from MEG Signals Using NeuralSet and Deep Learning for Predicting Linguistic Features Meta Introduces Autodata: An Agentic Framework That Turns AI Models into Autonomous Data Scientists for High-Quality Training Data Creation A Coding Guide on LLM Post Training with TRL from Supervised Fine Tuning to DPO and GRPO Reasoning Qwen AI Releases Qwen-Scope: An Open-Source Sparse AutoEncoders (SAE) Suite That Turns LLM Internal Features into Practical Development Tools A Coding Deep Dive into Agentic UI, Generative UI, State Synchronization, and Interrupt-Driven Approval Flows Moonshot AI Open-Sources FlashKDA: CUTLASS Kernels for Kimi Delta Attention with Variable-Length Batching and H20 Benchmarks Microsoft Research’s World-R1 Uses Flow-GRPO and 3D-Aware Rewards to Inject Geometric Consistency Into Wan 2.1 Without Architectural Changes A Coding Implementation on Pyright Type Checking Covering Generics, Protocols, Strict Mode, Type Narrowing, and Modern Python Typing IBM Releases Two Granite Speech 4.1 2B Models: Autoregressive ASR with Translation and Non-Autoregressive Editing for Fast Inference Top 10 KV Cache Compression Techniques for LLM Inference: Reducing Memory Overhead Across Eviction, Quantization, and Low-Rank Methods Qwen Team Releases FlashQLA: a High-Performance Linear Attention Kernel Library That Achieves Up to 3× Speedup on NVIDIA Hopper GPUs Step by Step Guide to Build a Complete PII Detection and Redaction Pipeline with OpenAI Privacy Filter Meta FAIR Releases NeuralSet: A Python Package for Neuro-AI That Supports fMRI, M/EEG, Spikes, and HuggingFace Embeddings smol-audio: A Colab-Friendly Notebook Collection for Fine-Tuning Whisper, Parakeet, Voxtral, Granite Speech, and Audio Flamingo 3 A Coding Implementation on Document Parsing Benchmarking with LlamaIndex ParseBench Using Python, Hugging Face, and Evaluation Metrics Poolside AI Introduces Laguna XS.2 and M.1: Agentic Coding Models Reaching 68.2% and 72.5% on SWE-bench Verified How to Build Traceable and Evaluated LLM Workflows Using Promptflow, Prompty, and OpenAI OpenAI Releases Privacy Filter: A 1.5B-Parameter Open-Source PII Redaction Model with 50M Active Parameters Top 10 Physical AI Models Powering Real-World Robots in 2026 How to Build a Lightweight Vision-Language-Action-Inspired Embodied Agent with Latent World Modeling and Model Predictive Control Meet Talkie-1930: A 13B Open-Weight LLM Trained on Pre-1931 English Text for Historical Reasoning and Generalization Research Build a Reinforcement Learning Powered Agent that Learns to Retrieve Relevant Long-Term Memories for Accurate LLM Question Answering OpenMOSS Releases MOSS-Audio: An Open-Source Foundation Model for Speech, Sound, Music, and Time-Aware Audio Reasoning Meta AI Releases Sapiens2: A High-Resolution Human-Centric Vision Model for Pose, Segmentation, Normals, Pointmap, and Albedo The LoRA Assumption That Breaks in Production How to Build a Fully Searchable AI Knowledge Base with OpenKB, OpenRouter, and Llama How to Build Smarter Multilingual Text Wrapping with BudouX Through Parsing, HTML Rendering, Model Introspection, and Toy Training Top 7 Benchmarks That Actually Matter for Agentic Reasoning in Large Language Models RAG Without Vectors: How PageIndex Retrieves by Reasoning A Coding Tutorial on Datashader on Rendering Massive Datasets with High-Performance Python Visual Analytics xAI Launches grok-voice-think-fast-1.0: Topping τ-voice Bench at 67.3%, Outperforming Gemini, GPT Realtime, and More A Coding Implementation on kvcached for Elastic KV Cache Memory, Bursty LLM Serving, and Multi-Model GPU Sharing Google DeepMind Introduces Vision Banana: An Instruction-Tuned Image Generator That Beats SAM 3 on Segmentation and Depth Anything V3 on Metric Depth Estimation Meet GitNexus: An Open-Source MCP-Native Knowledge Graph Engine That Gives Claude Code and Cursor Full Codebase Structural Awareness A Coding Implementation on Deepgram Python SDK for Transcription, Text-to-Speech, Async Audio Processing, and Text Intelligence A Coding Implementation on Microsoft’s OpenMementos with Trace Structure Analysis, Context Compression, and Fine-Tuning Data Preparation DeepSeek AI Releases DeepSeek-V4: Compressed Sparse Attention and Heavily Compressed Attention Enable One-Million-Token Contexts Google DeepMind Introduces Decoupled DiLoCo: An Asynchronous Training Architecture Achieving 88% Goodput Under High Hardware Failure Rates Mend Releases AI Security Governance Framework: Covering Asset Inventory, Risk Tiering, AI Supply Chain Security, and Maturity Model Mend.io Releases AI Security Governance Framework Covering Asset Inventory, Risk Tiering, AI Supply Chain Security, and Maturity Model OpenAI Releases GPT-5.5, a Fully Retrained Agentic Model That Scores 82.7% on Terminal-Bench 2.0 and 84.9% on GDPval A Coding Tutorial on OpenMythos on Recurrent-Depth Transformers with Depth Extrapolation, Adaptive Computation, and Mixture-of-Experts Routing Google Cloud AI Research Introduces ReasoningBank: A Memory Framework that Distills Reasoning Strategies from Agent Successes and Failures Xiaomi Releases MiMo-V2.5-Pro and MiMo-V2.5: Matching Frontier Model Benchmarks at Significantly Lower Token Cost How to Design a Production-Grade CAMEL Multi-Agent System with Planning, Tool Use, Self-Consistency, and Critique-Driven Refinement Alibaba Qwen Team Releases Qwen3.6-27B: A Dense Open-Weight Model Outperforming 397B MoE on Agentic Coding Benchmarks Next Leap to Harness Engineering: JiuwenClaw Pioneers ‘Coordination Engineering’ Photon Releases Spectrum: An Open-Source TypeScript Framework that Deploys AI Agents Directly to iMessage, WhatsApp, and Telegram OpenAI Open-Sources Euphony: A Browser-Based Visualization Tool for Harmony Chat Data and Codex Session Logs Hugging Face Releases ml-intern: An Open-Source AI Agent that Automates the LLM Post-Training Workflow A Coding Implementation to Build a Conditional Bayesian Hyperparameter Optimization Pipeline with Hyperopt, TPE, and Early Stopping Google Introduces Simula: A Reasoning-First Framework for Generating Controllable, Scalable Synthetic Datasets Across Specialized AI Domains A Coding Implementation on Qwen 3.6-35B-A3B Covering Multimodal Inference, Thinking Control, Tool Calling, MoE Routing, RAG, and Session Persistence Moonshot AI Releases Kimi K2.6 with Long-Horizon Coding, Agent Swarm Scaling to 300 Sub-Agents and 4,000 Coordinated Steps A Coding Implementation on Microsoft’s Phi-4-Mini for Quantized Inference Reasoning Tool Use RAG and LoRA Fine-Tuning OpenAI Scales Trusted Access for Cyber Defense With GPT-5.4-Cyber: a Fine-Tuned Model Built for Verified Security Defenders Moonshot AI and Tsinghua Researchers Propose PrfaaS: A Cross-Datacenter KVCache Architecture that Rethinks How LLMs are Served at Scale Meet OpenMythos: An Open-Source PyTorch Reconstruction of Claude Mythos Where 770M Parameters Match a 1.3B Transformer How TabPFN Leverages In-Context Learning to Achieve Superior Accuracy on Tabular Datasets Compared to Random Forest and CatBoost A Coding Implementation to Build an AI-Powered File Type Detection and Security Analysis Pipeline with Magika and OpenAI NVIDIA Releases Ising: the First Open Quantum AI Model Family for Hybrid Quantum-Classical Systems xAI Launches Standalone Grok Speech-to-Text and Text-to-Speech APIs, Targeting Enterprise Voice Developers A Coding Tutorial for Running PrismML Bonsai 1-Bit LLM on CUDA with GGUF, Benchmarking, Chat, JSON, and RAG A Coding Guide for Property-Based Testing Using Hypothesis with Stateful, Differential, and Metamorphic Test Design Anthropic Releases Claude Opus 4.7: A Major Upgrade for Agentic Coding, High-Resolution Vision, and Long-Horizon Autonomous Tasks Google AI Releases Auto-Diagnose: An Large Language Model LLM-Based System to Diagnose Integration Test Failures at Scale A End-to-End Coding Guide to Running OpenAI GPT-OSS Open-Weight Models with Advanced Inference Workflows Top 19 AI Red Teaming Tools (2026): Secure Your ML Models A Coding Guide to Build a Production-Grade Background Task Processing System Using Huey with SQLite, Scheduling, Retries, Pipelines, and Concurrency Control Qwen Team Open-Sources Qwen3.6-35B-A3B: A Sparse MoE Vision-Language Model with 3B Active Parameters and Agentic Coding Capabilities OpenAI Launches GPT-Rosalind: Its First Life Sciences AI Model Built to Accelerate Drug Discovery and Genomics Research Building Transformer-Based NQS for Frustrated Spin Systems with NetKet UCSD and Together AI Research Introduces Parcae: A Stable Architecture for Looped Language Models That Achieves the Quality of a Transformer Twice the Size How to Build a Universal Long-Term Memory Layer for AI Agents Using Mem0 and OpenAI A Coding Implementation to Build Multi-Agent AI Systems with SmolAgents Using Code Execution, Tool Calling, and Dynamic Orchestration A Technical Deep Dive into the Essential Stages of Modern Large Language Model Training, Alignment, and Deployment Google AI Launches Gemini 3.1 Flash TTS: A New Benchmark in Expressive and Controllable AI Voice Google DeepMind Releases Gemini Robotics-ER 1.6: Bringing Enhanced Embodied Reasoning and Instrument Reading to Physical AI Google Launches ‘Skills’ in Chrome: Turning Reusable AI Prompts into One-Click Browser Workflows A Coding Implementation of Crawl4AI for Web Crawling, Markdown Generation, JavaScript Execution, and LLM-Based Structured Extraction TinyFish AI Releases Full Web Infrastructure Platform for AI Agents: Search, Fetch, Browser, and Agent Under One API Key NVIDIA and the University of Maryland Researchers Released Audio Flamingo Next (AF-Next): A Super Powerful and Open Large Audio-Language Model A Hands-On Coding Tutorial for Microsoft VibeVoice Covering Speaker-Aware ASR, Real-Time TTS, and Speech-to-Speech Pipelines Meta AI and KAUST Researchers Propose Neural Computers That Fold Computation, Memory, and I/O Into One Learned Model A Coding Implementation of MolmoAct for Depth-Aware Spatial Reasoning, Visual Trajectory Tracing, and Robotic Action Prediction MiniMax Just Open Sourced MiniMax M2.7: A Self-Evolving Agent Model that Scores 56.22% on SWE-Pro and 57.0% on Terminal Bench 2 Liquid AI Releases LFM2.5-VL-450M: a 450M-Parameter Vision-Language Model with Bounding Box Prediction, Multilingual Support, and Sub-250ms Edge Inference Researchers from MIT, NVIDIA, and Zhejiang University Propose TriAttention: A KV Cache Compression Method That Matches Full Attention at 2.5× Higher Throughput How to Build a Secure Local-First Agent Runtime with OpenClaw Gateway, Skills, and Controlled Tool Execution How Knowledge Distillation Compresses Ensemble Intelligence into a Single Deployable AI Model Alibaba’s Tongyi Lab Releases VimRAG: a Multimodal RAG Framework that Uses a Memory Graph to Navigate Massive Visual Contexts A Coding Guide to Markerless 3D Human Kinematics with Pose2Sim, RTMPose, and OpenSim NVIDIA Releases AITune: An Open-Source Inference Toolkit That Automatically Finds the Fastest Inference Backend for Any PyTorch Model Five AI Compute Architectures Every Engineer Should Know: CPUs, GPUs, TPUs, NPUs, and LPUs Compared An End-to-End Coding Guide to NVIDIA KVPress for Long-Context LLM Inference, KV Cache Compression, and Memory-Efficient Generation Meta Superintelligence Lab Releases Muse Spark: A Multimodal Reasoning Model With Thought Compression and Parallel Agents Sigmoid vs ReLU Activation Functions: The Inference Cost of Losing Geometric Context A Coding Guide to Build Advanced Document Intelligence Pipelines with Google LangExtract, OpenAI Models, Structured Extraction, and Interactive Visualization Google AI Research Introduces PaperOrchestra: A Multi-Agent Framework for Automated AI Research Paper Writing A Comprehensive Implementation Guide to ModelScope for Model Search, Inference, Fine-Tuning, Evaluation, and Export
A Detailed Implementation on Equinox with JAX Native Modules, Filtered Transforms, Stateful Layers, and End-to-End Training Workflows
Sana Hassan · 2026-04-23 · via MarkTechPost

In this tutorial, we explore Equinox, a lightweight and elegant neural network library built on JAX, and show how to use it. We begin by understanding how eqx.Module treats models as PyTrees, which makes parameter handling, transformation, and serialization feel simple and explicit. As we move forward, we work through static fields, filtered transformations such as filter_jit and filter_grad, PyTree manipulation utilities, stateful layers such as BatchNorm, and a complete end-to-end training workflow for a toy regression problem. Throughout the tutorial, we focus on writing clear, executable code that demonstrates not only how Equinox works but also why it fits so well into the JAX ecosystem for research and practical experimentation.

!pip install equinox optax jaxtyping matplotlib -q


import jax
import jax.numpy as jnp
import equinox as eqx
import optax
from jaxtyping import Array, Float, Int, PRNGKeyArray
from typing import Optional
import matplotlib.pyplot as plt
import time


print(f"JAX version   : {jax.__version__}")
print(f"Equinox version: {eqx.__version__}")
print(f"Devices       : {jax.devices()}")


print("\n" + "="*60)
print("SECTION 1: eqx.Module basics")
print("="*60)


class Linear(eqx.Module):
   weight: Float[Array, "out in"]
   bias:   Float[Array, "out"]


   def __init__(self, in_size: int, out_size: int, *, key: PRNGKeyArray):
       wkey, bkey = jax.random.split(key)
       self.weight = jax.random.normal(wkey, (out_size, in_size)) * 0.1
       self.bias = jax.random.normal(bkey, (out_size,)) * 0.01


   def __call__(self, x: Float[Array, "in"]) -> Float[Array, "out"]:
       return self.weight @ x + self.bias




key = jax.random.PRNGKey(0)
lin = Linear(4, 2, key=key)


leaves, treedef = jax.tree_util.tree_flatten(lin)
print("Leaves shapes:", [l.shape for l in leaves])
print("Treedef:", treedef)


print("\n" + "="*60)
print("SECTION 2: Static fields")
print("="*60)


class Conv1dBlock(eqx.Module):
   conv:        eqx.nn.Conv1d
   norm:        eqx.nn.LayerNorm
   activation:  str = eqx.field(static=True)


   def __init__(self, channels: int, kernel: int, activation: str, *, key: PRNGKeyArray):
       self.conv       = eqx.nn.Conv1d(channels, channels, kernel, padding="same", key=key)
       self.norm       = eqx.nn.LayerNorm((channels,))
       self.activation = activation


   def __call__(self, x: Float[Array, "C L"]) -> Float[Array, "C L"]:
       x = self.conv(x)
       x = jax.vmap(self.norm)(x.T).T
       if self.activation == "relu":
           return jax.nn.relu(x)
       elif self.activation == "gelu":
           return jax.nn.gelu(x)
       return x




key, subkey = jax.random.split(key)
block = Conv1dBlock(8, 3, "gelu", key=subkey)
x_seq = jnp.ones((8, 16))
out = block(x_seq)
print(f"Conv1dBlock output shape: {out.shape}")

We set up the full Equinox environment by installing the required libraries and importing JAX, Equinox, Optax, Jaxtyping, Matplotlib, and other essentials. We immediately verify the runtime by printing the JAX and Equinox versions and the available devices, which helps us confirm that our Colab environment is ready for execution. We then begin with the foundations of Equinox by defining a simple Linear module, creating an instance of it, and inspecting its PyTree leaves and structure before introducing a Conv1dBlock that demonstrates how static fields and learnable layers work together in practice.

print("\n" + "="*60)
print("SECTION 3: Filtered transforms")
print("="*60)


class MLP(eqx.Module):
   layers: list
   dropout: eqx.nn.Dropout


   def __init__(self, in_size, hidden, out_size, *, key: PRNGKeyArray):
       k1, k2, k3 = jax.random.split(key, 3)
       self.layers  = [
           eqx.nn.Linear(in_size, hidden, key=k1),
           eqx.nn.Linear(hidden,  hidden, key=k2),
           eqx.nn.Linear(hidden,  out_size, key=k3),
       ]
       self.dropout = eqx.nn.Dropout(p=0.1)


   def __call__(self, x: Float[Array, "in"], *, key: Optional[PRNGKeyArray] = None) -> Float[Array, "out"]:
       for layer in self.layers[:-1]:
           x = jax.nn.relu(layer(x))
           if key is not None:
               key, subkey = jax.random.split(key)
               x = self.dropout(x, key=subkey)
       return self.layers[-1](x)




key, mk = jax.random.split(key)
mlp = MLP(8, 32, 4, key=mk)


@eqx.filter_jit
def forward(model, x, *, key):
   return model(x, key=key)


x_in  = jnp.ones((8,))
key, fk = jax.random.split(key)
y_out = forward(mlp, x_in, key=fk)
print(f"MLP output: {y_out}")


@eqx.filter_jit
def loss_fn(model: MLP,
           x: Float[Array, "B in"],
           y: Float[Array, "B out"],
           key: PRNGKeyArray) -> Float[Array, ""]:
   keys  = jax.random.split(key, x.shape[0])
   preds = jax.vmap(model)(x)
   return jnp.mean((preds - y) ** 2)


grad_fn = eqx.filter_grad(loss_fn)


key, dk = jax.random.split(key)
X = jax.random.normal(dk, (16, 8))
Y = jax.random.normal(dk, (16, 4))
grads = grad_fn(mlp, X, Y, dk)
print(f"Grad of first layer weight: shape={grads.layers[0].weight.shape}, norm={jnp.linalg.norm(grads.layers[0].weight):.4f}")

We focus on Equinox’s filtered transformations by building an MLP that includes both linear layers and dropout. We use filter_jit to compile the forward pass while allowing the model to contain non-array fields, and we use filter_grad to compute gradients only for array leaves that should actually participate in learning. By running a forward pass and then evaluating gradients on synthetic data, we see how Equinox cleanly bridges model definition and differentiable computation in a JAX-friendly way.

print("\n" + "="*60)
print("SECTION 4: PyTree manipulation")
print("="*60)


arrays, non_arrays = eqx.partition(mlp, eqx.is_array)
print("Non-array leaves (structure only):", jax.tree_util.tree_leaves(non_arrays))


trainable_filter = jax.tree_util.tree_map(
   lambda _: True, mlp
)
trainable_filter = eqx.tree_at(
   lambda m: (m.layers[0].weight, m.layers[0].bias),
   trainable_filter,
   replace=(False, False),
)
trainable, frozen = eqx.partition(mlp, trainable_filter)
print("Frozen params (first layer weight shape):", frozen.layers[0].weight.shape)
print("Trainable first-layer weight is sentinel:", trainable.layers[0].weight)


key, nk = jax.random.split(key)
new_weight = jax.random.normal(nk, mlp.layers[0].weight.shape)
mlp_updated = eqx.tree_at(lambda m: m.layers[0].weight, mlp, new_weight)
print("Updated first-layer weight norm:", jnp.linalg.norm(mlp_updated.layers[0].weight).item())


print("\n" + "="*60)
print("SECTION 5: Stateful layers — BatchNorm with inference mode")
print("="*60)


class BNModel(eqx.Module):
   linear1: eqx.nn.Linear
   bn:      eqx.nn.BatchNorm
   linear2: eqx.nn.Linear


   def __init__(self, in_f, hidden, out_f, *, key: PRNGKeyArray):
       k1, k2 = jax.random.split(key)
       self.linear1 = eqx.nn.Linear(in_f, hidden, key=k1)
       self.bn      = eqx.nn.BatchNorm(hidden, axis_name="batch")
       self.linear2 = eqx.nn.Linear(hidden, out_f, key=k2)


   def __call__(self, x, state, *, inference: bool = False):
       x, state = self.bn(jax.nn.relu(self.linear1(x)), state, inference=inference)
       return self.linear2(x), state




key, bk = jax.random.split(key)
bn_model, bn_state = eqx.nn.make_with_state(BNModel)(4, 16, 2, key=bk)


@eqx.filter_jit
def train_step_bn(model, state, x):
   def single(x):
       return model(x, state)
   outs, states = jax.vmap(single, axis_name="batch", out_axes=(0, None))(x)
   return outs, states


x_batch = jax.random.normal(key, (8, 4))
preds, bn_state = train_step_bn(bn_model, bn_state, x_batch)
print(f"BNModel output shape: {preds.shape}")

We explore PyTree manipulation utilities that make Equinox especially flexible for research workflows. We partition the model into array and non-array parts, create a trainable filter to freeze the first layer, and use tree_at to perform an immutable update on a specific parameter without rewriting the whole model. We then extend the tutorial to stateful layers by defining a BatchNorm-based model, creating both the model and its state, and running a batched training-style pass that returns updated state information.

print("\n" + "="*60)
print("SECTION 6: Full training loop (ResNet MLP on noisy sine)")
print("="*60)


class ResBlock(eqx.Module):
   fc1: eqx.nn.Linear
   fc2: eqx.nn.Linear
   proj: Optional[eqx.nn.Linear]


   def __init__(self, size: int, *, key: PRNGKeyArray):
       k1, k2 = jax.random.split(key)
       self.fc1  = eqx.nn.Linear(size, size, key=k1)
       self.fc2  = eqx.nn.Linear(size, size, key=k2)
       self.proj = None


   def __call__(self, x):
       residual = x
       x = jax.nn.gelu(self.fc1(x))
       x = self.fc2(x)
       return jax.nn.gelu(x + residual)




class ResNetMLP(eqx.Module):
   embed:   eqx.nn.Linear
   blocks:  list
   head:    eqx.nn.Linear


   def __init__(self, in_size, hidden, out_size, n_blocks, *, key: PRNGKeyArray):
       keys = jax.random.split(key, n_blocks + 2)
       self.embed  = eqx.nn.Linear(in_size, hidden, key=keys[0])
       self.blocks = [ResBlock(hidden, key=keys[i+1]) for i in range(n_blocks)]
       self.head   = eqx.nn.Linear(hidden, out_size, key=keys[-1])


   def __call__(self, x):
       x = jax.nn.gelu(self.embed(x))
       for block in self.blocks:
           x = block(x)
       return self.head(x)




def make_dataset(n: int, key: PRNGKeyArray):
   xk, nk = jax.random.split(key)
   x = jax.random.uniform(xk, (n, 1), minval=-1.0, maxval=1.0)
   y = jnp.sin(2 * jnp.pi * x) + 0.1 * jax.random.normal(nk, (n, 1))
   return x, y


key, dk = jax.random.split(key)
X_train, Y_train = make_dataset(2048, dk)
key, dk = jax.random.split(key)
X_val,   Y_val   = make_dataset(512,  dk)


key, mk = jax.random.split(key)
model = ResNetMLP(1, 64, 1, n_blocks=4, key=mk)


schedule = optax.warmup_cosine_decay_schedule(
   init_value=0.0, peak_value=3e-3,
   warmup_steps=200, decay_steps=2000
)
optimiser = optax.chain(optax.clip_by_global_norm(1.0), optax.adam(schedule))
opt_state = optimiser.init(eqx.filter(model, eqx.is_array))


@eqx.filter_jit
def train_step(model, opt_state, x, y):
   def compute_loss(model, x, y):
       preds = jax.vmap(model)(x)
       return jnp.mean((preds - y) ** 2)


   loss, grads = eqx.filter_value_and_grad(compute_loss)(model, x, y)
   updates, opt_state_new = optimiser.update(
       grads, opt_state, eqx.filter(model, eqx.is_array)
   )
   model_new = eqx.apply_updates(model, updates)
   return model_new, opt_state_new, loss




@eqx.filter_jit
def evaluate(model, x, y):
   preds = jax.vmap(model)(x)
   return jnp.mean((preds - y) ** 2)

We build a deeper end-to-end learning example by defining a residual block and a ResNetMLP model for a noisy sine regression task. We generate synthetic training and validation datasets, initialize the model, configure a warmup cosine learning rate schedule, and prepare the optimizer state using only the model’s array leaves. We also define the jitted train_step and evaluate functions, which provide the core training and validation mechanics for the full workflow.

BATCH  = 128
EPOCHS = 30
steps_per_epoch = len(X_train) // BATCH
train_losses, val_losses = [], []


t0 = time.time()
for epoch in range(EPOCHS):
   key, sk = jax.random.split(key)
   perm = jax.random.permutation(sk, len(X_train))
   X_s, Y_s = X_train[perm], Y_train[perm]


   epoch_loss = 0.0
   for step in range(steps_per_epoch):
       xb = X_s[step*BATCH:(step+1)*BATCH]
       yb = Y_s[step*BATCH:(step+1)*BATCH]
       model, opt_state, loss = train_step(model, opt_state, xb, yb)
       epoch_loss += loss.item()


   val_loss = evaluate(model, X_val, Y_val).item()
   train_losses.append(epoch_loss / steps_per_epoch)
   val_losses.append(val_loss)


   if (epoch + 1) % 5 == 0:
       print(f"Epoch {epoch+1:3d}/{EPOCHS}  "
             f"train_loss={train_losses[-1]:.5f}  "
             f"val_loss={val_losses[-1]:.5f}")


print(f"\nTotal training time: {time.time()-t0:.1f}s")


print("\n" + "="*60)
print("SECTION 7: Save & load model weights")
print("="*60)


eqx.tree_serialise_leaves("model_weights.eqx", model)


key, mk2 = jax.random.split(key)
model_skeleton = ResNetMLP(1, 64, 1, n_blocks=4, key=mk2)
model_loaded   = eqx.tree_deserialise_leaves("model_weights.eqx", model_skeleton)


diff = jnp.max(jnp.abs(
   jax.tree_util.tree_leaves(eqx.filter(model, eqx.is_array))[0]
 - jax.tree_util.tree_leaves(eqx.filter(model_loaded, eqx.is_array))[0]
))
print(f"Max weight difference after reload: {diff:.2e}  (should be 0.0)")


fig, axes = plt.subplots(1, 2, figsize=(12, 4))


axes[0].plot(train_losses, label="Train MSE", color="#4C72B0")
axes[0].plot(val_losses,   label="Val MSE",   color="#DD8452", linestyle="--")
axes[0].set_xlabel("Epoch")
axes[0].set_ylabel("MSE")
axes[0].set_title("Training curves")
axes[0].legend()
axes[0].grid(True, alpha=0.3)


x_plot  = jnp.linspace(-1, 1, 300).reshape(-1, 1)
y_true  = jnp.sin(2 * jnp.pi * x_plot)
y_pred  = jax.vmap(model)(x_plot)


axes[1].scatter(X_val[:100], Y_val[:100], s=10, alpha=0.4, color="gray", label="Data")
axes[1].plot(x_plot, y_true, color="#4C72B0",  linewidth=2, label="True f(x)")
axes[1].plot(x_plot, y_pred, color="#DD8452", linewidth=2, linestyle="--", label="Predicted")
axes[1].set_xlabel("x")
axes[1].set_ylabel("y")
axes[1].set_title("Sine regression fit")
axes[1].legend()
axes[1].grid(True, alpha=0.3)


plt.tight_layout()
plt.savefig("equinox_tutorial.png", dpi=150)
plt.show()
print("\nDone! Plot saved to equinox_tutorial.png")


print("\n" + "="*60)
print("BONUS: eqx.filter_jit + shape inference debug tip")
print("="*60)


jaxpr = jax.make_jaxpr(jax.vmap(model))(x_plot)
n_eqns = len(jaxpr.jaxpr.eqns)
print(f"Compiled ResNetMLP jaxpr has {n_eqns} equations (ops) for batch input {x_plot.shape}")
BATCH  = 128
EPOCHS = 30
steps_per_epoch = len(X_train) // BATCH
train_losses, val_losses = [], []


t0 = time.time()
for epoch in range(EPOCHS):
   key, sk = jax.random.split(key)
   perm = jax.random.permutation(sk, len(X_train))
   X_s, Y_s = X_train[perm], Y_train[perm]


   epoch_loss = 0.0
   for step in range(steps_per_epoch):
       xb = X_s[step*BATCH:(step+1)*BATCH]
       yb = Y_s[step*BATCH:(step+1)*BATCH]
       model, opt_state, loss = train_step(model, opt_state, xb, yb)
       epoch_loss += loss.item()


   val_loss = evaluate(model, X_val, Y_val).item()
   train_losses.append(epoch_loss / steps_per_epoch)
   val_losses.append(val_loss)


   if (epoch + 1) % 5 == 0:
       print(f"Epoch {epoch+1:3d}/{EPOCHS}  "
             f"train_loss={train_losses[-1]:.5f}  "
             f"val_loss={val_losses[-1]:.5f}")


print(f"\nTotal training time: {time.time()-t0:.1f}s")


print("\n" + "="*60)
print("SECTION 7: Save & load model weights")
print("="*60)


eqx.tree_serialise_leaves("model_weights.eqx", model)


key, mk2 = jax.random.split(key)
model_skeleton = ResNetMLP(1, 64, 1, n_blocks=4, key=mk2)
model_loaded   = eqx.tree_deserialise_leaves("model_weights.eqx", model_skeleton)


diff = jnp.max(jnp.abs(
   jax.tree_util.tree_leaves(eqx.filter(model, eqx.is_array))[0]
 - jax.tree_util.tree_leaves(eqx.filter(model_loaded, eqx.is_array))[0]
))
print(f"Max weight difference after reload: {diff:.2e}  (should be 0.0)")


fig, axes = plt.subplots(1, 2, figsize=(12, 4))


axes[0].plot(train_losses, label="Train MSE", color="#4C72B0")
axes[0].plot(val_losses,   label="Val MSE",   color="#DD8452", linestyle="--")
axes[0].set_xlabel("Epoch")
axes[0].set_ylabel("MSE")
axes[0].set_title("Training curves")
axes[0].legend()
axes[0].grid(True, alpha=0.3)


x_plot  = jnp.linspace(-1, 1, 300).reshape(-1, 1)
y_true  = jnp.sin(2 * jnp.pi * x_plot)
y_pred  = jax.vmap(model)(x_plot)


axes[1].scatter(X_val[:100], Y_val[:100], s=10, alpha=0.4, color="gray", label="Data")
axes[1].plot(x_plot, y_true, color="#4C72B0",  linewidth=2, label="True f(x)")
axes[1].plot(x_plot, y_pred, color="#DD8452", linewidth=2, linestyle="--", label="Predicted")
axes[1].set_xlabel("x")
axes[1].set_ylabel("y")
axes[1].set_title("Sine regression fit")
axes[1].legend()
axes[1].grid(True, alpha=0.3)


plt.tight_layout()
plt.savefig("equinox_tutorial.png", dpi=150)
plt.show()
print("\nDone! Plot saved to equinox_tutorial.png")


print("\n" + "="*60)
print("BONUS: eqx.filter_jit + shape inference debug tip")
print("="*60)


jaxpr = jax.make_jaxpr(jax.vmap(model))(x_plot)
n_eqns = len(jaxpr.jaxpr.eqns)
print(f"Compiled ResNetMLP jaxpr has {n_eqns} equations (ops) for batch input {x_plot.shape}")

We run the complete training loop across multiple epochs, shuffle the data, process mini-batches, and track both training and validation losses over time. We then serialize the trained model with Equinox utilities, reconstruct a matching skeleton model, verify that deserialization restores the weights correctly, and visualize the learned fit and loss curves. Also, we inspect the compiled computation graph using jax.make_jaxpr, which provides a useful debugging and introspection view of how the trained Equinox model is executed under JAX.

In conclusion, we built a strong practical understanding of how Equinox helps us write clean, modular, and JAX-native deep learning code without adding unnecessary abstraction. We saw how to define custom modules, manage static and trainable components, apply filtered transformations safely, work with stateful layers, train a residual MLP, save and reload model weights, and inspect compiled computations. In doing so, we experienced how Equinox gives us the flexibility of raw JAX while still providing the structure needed for modern model development. As a result, we came away with a complete hands-on foundation that prepares us to use Equinox confidently for more advanced machine learning experiments and research workflows.


Check out the Full Codes with Notebook here. Also, feel free to follow us on Twitter and don’t forget to join our 130k+ ML SubReddit and Subscribe to our Newsletter. Wait! are you on telegram? now you can join us on telegram as well.

Need to partner with us for promoting your GitHub Repo OR Hugging Face Page OR Product Release OR Webinar etc.? Connect with us

Sana Hassan, a consulting intern at Marktechpost and dual-degree student at IIT Madras, is passionate about applying technology and AI to address real-world challenges. With a keen interest in solving practical problems, he brings a fresh perspective to the intersection of AI and real-life solutions.