Skip to main content

SGL-JAX: High-Performance LLM Inference on JAX/TPU

SGL-JAX is a high-performance, JAX-based inference engine for Large Language Models (LLMs), specifically optimized for Google TPUs. It is engineered from the ground up to deliver exceptional throughput and low latency for the most demanding LLM serving workloads.

The engine incorporates state-of-the-art techniques to maximize hardware utilization and serving efficiency, making it ideal for deploying large-scale models in production on TPUs.

Pypi License

Key Features

  • High-Throughput Continuous Batching: Implements a sophisticated scheduler that dynamically batches incoming requests, maximizing TPU utilization and overall throughput.
  • Optimized KV Cache with Radix Tree: Utilizes a Radix Tree for KV cache management (conceptually similar to PagedAttention), enabling memory-efficient prefix sharing between requests and significantly reducing computation for prompts with common prefixes.
  • FlashAttention Integration: Leverages a high-performance FlashAttention kernel for faster and more memory-efficient attention calculations, crucial for long sequences.
  • Tensor Parallelism: Natively supports tensor parallelism to distribute large models across multiple TPU devices, enabling inference for models that exceed the memory of a single accelerator.
  • OpenAI-Compatible API: Provides a drop-in replacement for the OpenAI API, allowing for seamless integration with a wide range of existing clients, SDKs, and tools (e.g., LangChain, LlamaIndex).
  • Native Qwen Support: Includes first-class, optimized support for the Qwen model family, including recent Mixture-of-Experts (MoE) variants.

Architecture Overview

SGL-JAX operates on a distributed architecture designed for scalability and performance:

  1. HTTP Server: The entry point for all requests, compatible with the OpenAI API standard.
  2. Scheduler: The core of the engine. It receives requests, manages prompts, and schedules token generation in batches. It intelligently groups requests to form optimal batches for the model executor.
  3. TP Worker (Tensor Parallel Worker): A set of distributed workers that host the model weights, distributed via tensor parallelism. They execute the forward pass for the model.
  4. Model Runner: Manages the actual JAX-based model execution, including the forward pass, attention computation, and KV cache operations.
  5. Radix Cache: A global, memory-efficient KV cache that is shared across all requests, enabling prefix reuse and reducing the memory footprint.

Getting Started

Documentation

For more features and usage details, please read the documents in the docs directory.

Supported Models

SGL-JAX is designed for easy extension to new model architectures. It currently provides first-class, optimized support for:

  • Qwen
  • Qwen 3
  • Qwen 3 MoE

Performance and Benchmarking

For detailed performance evaluation and to run the benchmarks yourself, please see the scripts located in the benchmark/ and python/sgl_jax/ directories (e.g., bench_serving.py).

Testing

The project includes a comprehensive test suite to ensure correctness and stability. To run the full suite of tests:

cd test/srt
python run_suite.py

Contributing

Contributions are welcome! If you would like to contribute, please feel free to open an issue to discuss your ideas or submit a pull request.

Metadata

Release files for sglang-jax 0.0.2

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for sglang-jax 0.0.2
File Size Uploaded
sglang_jax-0.0.2.tar.gz 329.3 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for sglang-jax 0.0.2
File Interpreter ABI Platform
sglang_jax-0.0.2-py3-none-any.whl Python 3 none any Details

Total release size: 723.8 kB

Release files / sglang_jax-0.0.2.tar.gz

Download URL sglang_jax-0.0.2.tar.gz
Size 329.3 kB
Tags Source
SHA-256 checksum
How to use checksums
70e38d3513797ea208c4d33e957b6950b47ff069492b8099b65152a70e40416e
BLAKE2b-256 checksum
How to use checksums
c7941b525f5936c296b877ff7fe4c8c8b5f9088a73dd0197b624105a6b9f68ce
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.12.11

Release files / sglang_jax-0.0.2-py3-none-any.whl

Download URL sglang_jax-0.0.2-py3-none-any.whl
Size 394.5 kB
Tags Python 3
SHA-256 checksum
How to use checksums
1ba0d3f7801a47c3961e50995a6ee832d8d094c5595d9ad865f4f9ebbe3fdfcb
BLAKE2b-256 checksum
How to use checksums
4619053f86f545869badfb5164df6e314f5cfb87dd22ec89b71efe093a85fbf0
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/6.2.0 CPython/3.12.11
Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page