$npx skillfedfor your agent

tokamax

A Pallas Custom Kernel Library.

With conditionsPyPI Artificial IntelligenceReleased Mar 2026242.0K downloads / mopermissive licensePure Python

Decision gist · record as of 2026-08-14

pure-Python wheel — tokamax-0.0.12-py3-none-any.whl
v0.0.12 · released 2026-03-19 · Python >=3.11 · 12 runtime deps: absl-py, einshape, jax, jaxlib, jaxtyping, pydantic, qwix, tqdm

Yes, if you are training or serving large models on NVIDIA GPUs or TPUs and want to use pre-optimized kernels for attention, normalization, and mixture-of-experts without rewriting core operations. The low install friction and permissive license make adoption straightforward. However, expect API changes and incomplete features—suitable for research and experimentation, less so for production systems requiring stability guarantees. No known vulnerabilities.AI-flagged interpretation of the facts on this page — verify before relying

Before you install

  • Requires Python >=3.11 and JAX with jaxlib installed; GPU or TPU hardware needed for actual acceleration.
  • Low install friction with a pure Python wheel.
  • Active maintenance status as of 148 days since latest release.

License · maintenance · safety

permissive license (permissive) — Apache License 2.0 is permissive and allows commercial use, modification, and distribution with minimal restrictions—primarily requiring license and copyright notice preservation.

last release 2026-03-19 (148 days)

0 known vulnerabilities (OSV.dev, 2026-08-14) · 241,972 downloads/mo, #8,865 on PyPI

Verify before relying

pip install tokamax

import jax
import jax.numpy as jnp
import tokamax

def loss(x, scale):
  x = tokamax.layer_norm(x, scale=scale, offset=None)
  x = tokamax.dot_product_attention(x, x, x)
  return jnp.sum(x)

f_grad = jax.jit(jax.grad(loss))
  • Stability and API compatibility guarantees beyond the stated 'heavily under development' warning
  • Performance benchmarks comparing implementations against standard JAX operations
  • Supported GPU models and compute capabilities beyond the H100 example
Same gist for agents: .md · .json

What it is and what it does

Tokamax is a JAX library that wraps hand-tuned accelerator kernels for common deep-learning operations, built on top of JAX's Pallas framework. It provides optimized implementations of dot-product attention (FlashAttention), gated linear units, layer and RMS normalization, mixture-of-experts routing, and linear softmax cross-entropy loss, with support for both NVIDIA GPUs and Google TPUs. The library also exposes an autotuning system to discover optimal kernel configurations for your hardware and input shapes, and utilities for benchmarking kernel execution time separately from Python overhead.

Tokamax is designed for researchers and practitioners building large-scale models who need fine-grained control over kernel selection and performance. It lets you choose between multiple implementations (e.g., Triton, Mosaic, XLA) for each operation, or allow automatic selection. The package is young—still in active development with API changes expected—but offers a path to both use pre-optimized kernels and build custom ones by inheriting from the Op class.

Use it for

  • Accelerate transformer attention layers on H100 or TPU hardware by swapping standard JAX attention with tokamax.dot_product_attention
  • Optimize layer normalization and RMS norm in deep networks using tokamax.layer_norm with hardware-specific implementations
  • Build mixture-of-experts models with efficient ragged tensor routing via tokamax.ragged_dot on GPU or TPU
  • Autotune kernel configurations for your specific hardware and input shapes to maximize throughput
  • Benchmark actual kernel execution time (excluding Python overhead) using tokamax.benchmark with CUPTI or default profiling
  • Export JAX functions containing custom kernels to StableHLO with device-specific guarantees via tokamax.DISABLE_JAX_EXPORT_CHECKS

Worth the install?

AI-flagged interpretation of the facts on this page. Verify before relying on it.

With conditions

Yes, if you are training or serving large models on NVIDIA GPUs or TPUs and want to use pre-optimized kernels for attention, normalization, and mixture-of-experts without rewriting core operations.

The low install friction and permissive license make adoption straightforward. However, expect API changes and incomplete features—suitable for research and experimentation, less so for production systems requiring stability guarantees. No known vulnerabilities.

Install

tokamax on PyPI

Before you install

Low install friction with a pure Python wheel. Active maintenance status as of 148 days since latest release. However, the package declares itself still heavily under development with incomplete features and expected API changes.

Requires Python >=3.11 and JAX with jaxlib installed; GPU or TPU hardware needed for actual acceleration.

License in practice

Apache License 2.0 is permissive and allows commercial use, modification, and distribution with minimal restrictions—primarily requiring license and copyright notice preservation.

Quickstart

pip install tokamax

import jax
import jax.numpy as jnp
import tokamax

def loss(x, scale):
  x = tokamax.layer_norm(x, scale=scale, offset=None)
  x = tokamax.dot_product_attention(x, x, x)
  return jnp.sum(x)

f_grad = jax.jit(jax.grad(loss))

Verify before relying

  • Stability and API compatibility guarantees beyond the stated 'heavily under development' warning
  • Performance benchmarks comparing implementations against standard JAX operations
  • Supported GPU models and compute capabilities beyond the H100 example

Package facts

Licensepermissive license permissive
Python supportSupports the current Python release >=3.11
Install frictionLow. Pure-Python wheel
Runtime dependencies
12 packages
absl-pyeinshapejaxjaxlibjaxtypingpydanticqwixtqdmtyping_extensionstypeguardimmutabledicttensorboardx
MaintenanceActively maintained 148 days since the last release
First released
Downloads241,972 / month, #8,865 on PyPI 30-day window, as of 2026-08-14
Known vulnerabilitiesNone known OSV.dev, checked 2026-08-14

Evidence: tokamax-0.0.12-py3-none-any.whl

Tags

Capabilities
JAX custom kernels GPU TPUflash attention implementationaccelerator kernel optimizationJAX Pallas kernelslayer norm gated linear unitmixture of experts kernelkernel autotuning JAX
Topics
jax-ecosystemgpu-optimizationkernel-tuning

Let your AI agent find packages like this

Example. Real query, live index.

You found this page by searching. An agent finds it by wishing: SkillFed indexes 14,416 PyPI packages by what they can do, searchable in plain language.

wish › “JAX custom kernels GPU TPU”

  • tokamaxTokamax provides custom accelerator kernels for JAX, including…
  • blackjaxBlackJAX provides a collection of MCMC and other sampling algorithms…
  • tpu-inferencetpu-inference is a hardware plugin for vLLM that enables…

Give your agent the search over MCP, or paste the wish link into any chat.

More Artificial Intelligence packages

litellm With conditions
PyPI · Artificial Intelligence · released Aug 2026

LiteLLM provides a unified Python interface to call 100+ LLM providers (OpenAI, Anthropic, Gemini, Bedrock, Azure, and others) using OpenAI-compatible API format, available as both a Python SDK and a self-hosted AI Gateway proxy server.

Install it if you need to work with multiple LLM providers or want to centralize LLM routing in your organization.

MITcompiled wheel
682.8Mdownloads / mo
huggingface-hub Worth it
PyPI · Artificial Intelligence · released Aug 2026

Client library and CLI tool for downloading, uploading, and managing models, datasets, and repositories on the Hugging Face Hub platform.

Install it if you work with Hugging Face Hub models or datasets.

Apache-2.0pure Python · 3.10.0+
442.4Mdownloads / mo
langchain Worth it
PyPI · Python Modules · released Aug 2026

LangChain provides a framework for building agents and LLM-powered applications by composing language models, tools, and memory through a unified API that abstracts over multiple model providers.

MITpure Python
315.4Mdownloads / mo
hf-xet With conditions
PyPI · Artificial Intelligence · released Aug 2026

hf-xet provides chunk-based deduplication and efficient file transfer for the Hugging Face Hub, enabling faster uploads and downloads of large files with local disk caching.

Apache-2.0compiled wheel · 3.8+
258.4Mdownloads / mo
tokenizers Worth it
PyPI · Artificial Intelligence · released Apr 2026

Tokenizers converts raw text into token sequences for NLP models, with support for training custom vocabularies and using pre-built tokenizers (BPE, WordPiece) optimized for speed via Rust.

Apache-2.0compiled wheel · 3.10+
222.9Mdownloads / mo
transformers Worth it
PyPI · Artificial Intelligence · released Aug 2026

Transformers provides a unified framework for loading, fine-tuning, and running state-of-the-art pretrained models across text, vision, audio, video, and multimodal tasks using PyTorch, JAX, or TensorFlow.

Install it if you need to run or train any transformer-based model for NLP, vision, audio, or multimodal tasks.

permissive licensepure Python · 3.10.0+
186.6Mdownloads / mo

See also jax-cuda12-pjrt · jax-cuda13-pjrt · jax-cuda13-plugin · jax-cuda12-plugin · jaxlib · torchax · pathwaysutils · google-tunix · augmax · libtpu

Further reading