MMarshudinmarshud.dev·Aug 19 · 7 min readHigh Level Neural Network Modeling with FlaxIn our previous article, we broke down why JAX requires explicit state handling. We saw that wrapping our parameters, optimizer momentum, and PRNG keys into a clean carry pattern lets us use jax.jit, 71J
MMarshudinmarshud.dev·Aug 11 · 11 min readFunctional State Management and Low-Level Control Flow in JAXWell be back guys, in our previous article on managing PRNG keys, we explored how JAX abandons the conventional hidden random state in favor of stateless random number generation. We learned that to g10
MMarshudinmarshud.dev·Aug 6 · 7 min readManaging PRNG Keys in JAXIn the previous articles, we have explored the fundamental principles of JAX from functional purity to transformations to how it handles arrays to PyTrees. However, in real-world ML, when working with00
MMarshudinmarshud.dev·Aug 4 · 9 min readProgram Transformations and PyTrees in JAXWelcome back everyone, in our previous article, we covered the foundational philosophy of JAX from functional purity to array immutability to basic transformations like jax.grad, jax.vmap and jax.jit.00
MMarshudinmarshud.dev·Jul 7 · 11 min readJAX Core Fundamentals: Pure Functional ArraysWelcome everyone, in this article, I will explore the core philosophy of JAX, it's array system and the critical concept of functional purity. By the end of this article, you will understand why JAX w50