Implementation:NVIDIA TransformerEngine JAX Cpp Quantization
| Field | Value |
|---|---|
| Sources | TransformerEngine |
| Domains | Deep_Learning, JAX, Quantization |
| Last Updated | 2026-02-07 14:00 GMT |
Overview
Implements JAX custom primitives for tensor quantization to FP8/MXFP8/NVFP4 formats, with optional fused bias gradient computation, supporting multiple scaling modes.
Description
BaseDBiasQuantizePrimitive defines the shared abstract/lowering/partition logic for quantization. DBiasQuantizePrimitive fuses bias gradient reduction with quantization. QuantizePrimitive handles standalone quantization. GroupedQuantizePrimitive supports grouped quantization for MoE workloads. The module supports stochastic rounding for NVFP4, Randomized Hadamard Transform (RHT) fusion, and multiple quantize layouts (rowwise, colwise, both). Scale management handles amax computation and all-reduce across parallel dimensions.
This is the fundamental quantization primitive layer that all higher-level quantization operations build upon, providing the bridge between JAX tensors and the TE C++ quantization kernels with full SPMD sharding support.
Usage
Use this module indirectly through quantizer classes (CurrentScaleQuantizer, DelayedScaleQuantizer, etc.). Direct usage is needed for custom quantization workflows or when implementing operations that need fused bias gradient + quantization.
Code Reference
Source Location
- Repository
NVIDIA/TransformerEngine- File
transformer_engine/jax/cpp_extensions/quantization.py- Lines
- 1--1308
Signature
class BaseDBiasQuantizePrimitive(BasePrimitive): ...
class DBiasQuantizePrimitive(BaseDBiasQuantizePrimitive): ...
class QuantizePrimitive(BaseDBiasQuantizePrimitive): ...
class GroupedQuantizePrimitive(BasePrimitive): ...
def quantize(
x: jnp.ndarray,
quantizer: Quantizer,
flatten_axis: int = -1,
) -> ScaledTensor: ...
def quantize_dbias(
dx: jnp.ndarray,
quantizer: Quantizer,
flatten_axis: int = -1,
) -> Tuple[ScaledTensor, jnp.ndarray]: ...
def grouped_quantize(
x: jnp.ndarray,
group_sizes: jnp.ndarray,
quantizer: Quantizer,
flatten_axis: int = -1,
) -> ScaledTensor: ...
def grouped_dbias(
grad: jnp.ndarray,
group_sizes: jnp.ndarray,
) -> jnp.ndarray: ...
Import
from transformer_engine.jax.cpp_extensions.quantization import quantize, quantize_dbias, grouped_quantize
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
| x | jnp.ndarray |
Yes | Input tensor to quantize |
| quantizer | Quantizer |
Yes | Quantizer instance defining the scaling mode and parameters |
| flatten_axis | int |
No | Axis at which to flatten the tensor for 2D quantization (default -1) |
| group_sizes | jnp.ndarray |
No | Group sizes for grouped quantization (MoE) |
Outputs
| Name | Type | Description |
|---|---|---|
| quantized | ScaledTensor |
Quantized tensor with scale factors |
| dbias | jnp.ndarray |
Bias gradient (when using quantize_dbias) |
Usage Examples
from transformer_engine.jax.cpp_extensions.quantization import quantize, quantize_dbias
# Quantize a tensor to FP8
quantized_tensor = quantize(input_tensor, fp8_quantizer)
# Fused bias gradient + quantization
quantized_grad, dbias = quantize_dbias(grad_tensor, fp8_quantizer)