How to Run Needle Inference on GPU (CUDA), CPU, and Apple Silicon Metal Backends

Run Needle inference on CPU, CUDA GPU, or Apple Silicon Metal by installing the appropriate JAX wheel via Pip extras and optionally setting JAX_PLATFORM_NAME.

Needle is a lightweight inference engine built on JAX that abstracts hardware accelerators through three install-time options. The cactus-needle package uses optional extras to pull in platform-specific JAX builds, enabling seamless execution across CPU, NVIDIA GPUs, and Apple Silicon without code changes. This guide shows how to configure each backend using the exact package metadata and runtime behavior defined in the Needle source code.

Installing the Correct Backend for Needle Inference

Needle declares its backend dependencies as Pip extras in pyproject.toml. Installing the base package gives you CPU support; GPU and Metal require bracketed extras.

Backend Install Command JAX Wheel Used
CPU (default) pip install cactus-needle Standard CPU-only JAX
CUDA GPU pip install "cactus-needle[gpu]" CUDA-enabled JAX (requires NVIDIA driver + CUDA toolkit)
Apple Silicon Metal pip install "cactus-needle[metal]" Metal-enabled JAX for M-series Macs

The gpu and metal extras resolve to the appropriate jax[cuda] or jax[metal] wheels. No manual JAX installation is required.

Forcing a Specific Device with JAX_PLATFORM_NAME

JAX automatically selects the first available accelerator, but you can override this with the JAX_PLATFORM_NAME environment variable. This is useful for testing, debugging, or forcing CPU execution on a GPU-equipped machine.


# Force CPU execution

export JAX_PLATFORM_NAME=cpu

# Force CUDA GPU (Linux/Windows with NVIDIA hardware)

export JAX_PLATFORM_NAME=gpu

# Force Metal GPU (macOS on Apple Silicon)

export JAX_PLATFORM_NAME=metal

Unset the variable to return to JAX's default device selection.

Running Needle Inference on CPU

CPU execution works out of the box with the base package. The needle.Needle class in needle/model/run.py creates JAX arrays on the default CPU device when no GPU is available.

import needle

agent = needle.Needle(
    weights="needle2.cact",
    tools=[...]
)
response = agent.run("Summarise the latest news")
print(response["answer"])

The forward pass executes in needle/model/architecture.py, which defines Needle's Simple Attention Network. Quantisation handling in needle/model/quantize.py reduces memory footprint regardless of backend.

Running Needle Inference on CUDA GPU

NVIDIA GPU acceleration requires the [gpu] extra and a compatible CUDA environment.


# Install CUDA-enabled JAX via Needle's extra

pip install "cactus-needle[gpu]"

# Optional: force GPU selection

export JAX_PLATFORM_NAME=gpu

# Run inference

python run_gpu.py

# run_gpu.py

import needle

agent = needle.Needle(
    weights="needle2.cact",
    tools=[...]
)
result = agent.run("Translate 'hello' to French")
print(result["answer"])

JAX automatically places arrays on the first visible CUDA device. The same needle.Needle API works unchanged—the backend switch is transparent to your application code.

Running Needle Inference on Apple Silicon Metal

M-series Macs leverage the integrated GPU through Metal, enabled by the [metal] extra.


# Install Metal-enabled JAX

pip install "cactus-needle[metal]"

# Optional: explicitly select Metal

export JAX_PLATFORM_NAME=metal

# Run inference

python run_metal.py

# run_metal.py

import needle

agent = needle.Needle(
    weights="needle2.cact",
    tools=[...]
)
result = agent.run("What's the weather in Tokyo?")
print(result["answer"])

The Metal backend in JAX routes operations through Apple's GPU frameworks. Performance characteristics differ from CUDA, but the Needle API surface remains identical.

Key Source Files for Needle Backend Implementation

Understanding the codebase helps troubleshoot device-specific issues:

These files demonstrate how Needle achieves hardware abstraction: JAX handles device placement while the model code remains purely functional.

Verifying Your Active Backend

Check which accelerator JAX has selected before running Needle inference:

import jax

print(jax.devices())           # List all available devices

print(jax.default_backend())   # Current platform: cpu, gpu, or METAL

This confirms your installation and JAX_PLATFORM_NAME settings took effect.

Summary

  • Install the correct extra: cactus-needle for CPU, [gpu] for CUDA, [metal] for Apple Silicon
  • Control device selection: Use JAX_PLATFORM_NAME environment variable or let JAX auto-detect
  • Write once, run anywhere: The needle.Needle class and run() method work identically across all backends
  • Source transparency: Backend handling lives in needle/model/run.py with architecture in needle/model/architecture.py

Frequently Asked Questions

What CUDA version does Needle require?

Needle does not specify CUDA directly—it depends on JAX's CUDA wheel requirements. Installing "cactus-needle[gpu]" pulls JAX's CUDA build, which typically supports recent CUDA versions (11.8+ or 12.x depending on JAX release). Check JAX's compatibility matrix if you encounter driver issues.

Can I switch backends without reinstalling?

No. The JAX wheel contains platform-specific compiled binaries. You must pip install the correct extra (or reinstall JAX directly) to change between CPU, CUDA, and Metal builds. However, you can force CPU on a GPU-installed system using JAX_PLATFORM_NAME=cpu.

Does Metal performance match CUDA for Needle inference?

Apple Silicon Metal gives significant speedup over CPU but generally trails NVIDIA CUDA for large-batch transformer workloads. Needle's quantisation in needle/model/quantize.py reduces memory bandwidth pressure, which helps both backends. Profile your specific model size and sequence length for accurate comparison.

Why does JAX show "METAL" instead of "metal" for the backend?

JAX reports the Metal backend as uppercase "METAL" in jax.default_backend(). This is a JAX-internal naming convention. Set JAX_PLATFORM_NAME=metal (lowercase) to activate it, but expect uppercase in runtime output.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →