Jump to content

Connect SuperML | Leeroopedia MCP: Equip your AI agents with best practices, code verification, and debugging knowledge. Powered by Leeroo — building Organizational Superintelligence. Contact us at founders@leeroo.com.

Implementation:NVIDIA TransformerEngine JAX Cpp Quantization

From Leeroopedia


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)

Related Pages

Page Connections

Double-click a node to navigate. Hold to expand connections.
Principle
Implementation
Heuristic
Environment