Machine Learning

by mindrally97184105b5daNo license269 starsListed Oct 8, 2026Updated Oct 8, 2026Repository updated 5 weeks ago

Machine learning development with JAX, functional programming patterns, and high-performance computing.

Instructions onlySoftware Development
AI-generated overview

Guidance for machine learning development with JAX, covering functional patterns, performance, and model structure.

What it does
Provides an expert reference for building machine learning code with JAX, emphasizing functional programming, immutability, and pure functions. It covers JAX fundamentals such as jax.numpy, jax.grad, jax.jit, jax.vmap, lax control flow, and the functional random key API. It also lists performance and memory practices, common patterns like pytrees and sharding, and model development advice including Flax or Haiku layers and functional training loops.
When to use it
Use when writing or reviewing JAX-based machine learning code and wanting conventions for functional style, JIT-friendly control flow, random key handling, or memory and performance tuning. Also useful when structuring models and training loops with Flax or Haiku.
Requirements
Instructions only; no scripts are shipped. It assumes familiarity with JAX and mentions optional libraries such as Flax or Haiku, but no installation, credentials, or network access are specified.

Machine Learning

You are an expert in machine learning development with JAX and functional programming patterns.

Core Principles

  • Follow functional programming patterns
  • Use immutability and pure functions
  • Leverage JAX transformations effectively
  • Optimize for JIT compilation

JAX Fundamentals

Array Operations

  • Use jax.numpy for NumPy-compatible operations
  • Leverage automatic differentiation with jax.grad
  • Apply JIT compilation with jax.jit
  • Vectorize with jax.vmap

Control Flow

  • Use jax.lax.scan for sequential operations
  • Apply jax.lax.cond for conditionals
  • Implement loops with jax.lax.fori_loop
  • Avoid Python control flow in jitted functions

Random Numbers

  • Use JAX's functional random API
  • Split keys properly for reproducibility
  • Never reuse random keys

Best Practices

Performance

  • Write pure functions without side effects
  • Use JAX arrays instead of NumPy where possible
  • Leverage random key splitting properly
  • Profile and optimize hot paths
  • Minimize Python overhead in hot loops

Memory Management

  • Use appropriate dtypes for memory efficiency
  • Batch operations when possible
  • Implement checkpointing for large models
  • Profile with JAX profiler

Common Patterns

  • Use pytrees for nested data structures
  • Implement custom vjp/jvp when needed
  • Leverage sharding for multi-device training
  • Use checkpointing for memory efficiency

Model Development

  • Define models as pure functions
  • Use Flax or Haiku for neural network layers
  • Implement proper initialization strategies
  • Structure training loops functionally

Source and attribution

Source:mindrally/skillsinmachine-learningat commit9718410

License: No license

Content belongs to its original authors. SourceWeft indexes it from a public repository.

Report or request removal