core.inference.quantization.utils#
Module Contents#
Functions#
Resolve the canonical MXFP8 storage required by a grouped-MoE backend. |
|
Return whether a parameter or its data uses TE or MCore MXFP8 storage. |
|
Reject a selective policy that splits an MoE layer across precisions. |
|
Convert TE MXFP8 weights to mcore MXFP8Tensor format. |
|
Return True if a parameter should be converted to an MCore MXFP8 tensor. |
|
Convert a parameter value to BF16 for quantization. |
|
Record shape/dtype/device for each parameter that will be quantized. |
|
Quantize model parameters to mutable MXFP8Tensor storage. |
|
MXFP8 matmul via FlashInfer. |
|
MXFP8 matmul via torch.nn.functional.scaled_mm. |
|
Compute a matmul in MXFP8. |
API#
- core.inference.quantization.utils._verify_te_to_mcore_mxfp8_conversion(
- te_dequantized,
- fi_quantized: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Tensor,
- core.inference.quantization.utils.resolve_mxfp8_backend(
- inference_grouped_gemm_backend: str | megatron.core.inference.moe.InferenceGroupedGemmBackend,
Resolve the canonical MXFP8 storage required by a grouped-MoE backend.
- Parameters:
inference_grouped_gemm_backend – The configured backend, either as its raw string value or as the enum produced by
TransformerConfig.- Returns:
The MXFP8 quantization and storage backend to use. FlashInfer routed MoE derives its TRT-LLM Major-K weights from the canonical Triton/cuBLAS layout.
- Raises:
ValueError – If the grouped-GEMM backend does not support MXFP8.
- core.inference.quantization.utils._has_mxfp8_storage(parameter: object) bool#
Return whether a parameter or its data uses TE or MCore MXFP8 storage.
- core.inference.quantization.utils._validate_mxfp8_expert_precision_policy(
- model: torch.nn.Module,
Reject a selective policy that splits an MoE layer across precisions.
- core.inference.quantization.utils.quantize_model_to_mxfp8(
- model: torch.nn.Module,
- backend: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Backend = 'flashinfer',
- _prefix: str = '',
Convert TE MXFP8 weights to mcore MXFP8Tensor format.
Recursively converts existing TE MXFP8 parameters to MCore MXFP8Tensor. The TE per-module precision recipe selects storage during model construction; ordinary BF16 parameters are left untouched.
- Parameters:
model – The model whose TE MXFP8 parameters should be converted.
backend – ‘flashinfer’ or ‘triton’ quantization backend.
_prefix – Internal recursion prefix; callers should not set this.
- core.inference.quantization.utils._should_quantize_param(val: torch.Tensor) bool#
Return True if a parameter should be converted to an MCore MXFP8 tensor.
- core.inference.quantization.utils._to_bf16(val: torch.Tensor) torch.Tensor#
Convert a parameter value to BF16 for quantization.
- core.inference.quantization.utils.collect_mxfp8_param_metadata(
- model: torch.nn.Module,
Record shape/dtype/device for each parameter that will be quantized.
Called once before the first quantization to record the original parameter metadata (shape, dtype, device) before any format conversion.
- core.inference.quantization.utils.quantize_params_to_mxfp8(
- model: torch.nn.Module,
- persistent_buffers: Optional[Dict[str, megatron.core.inference.quantization.mxfp8_tensor.MXFP8Tensor]] = None,
- _prefix: str = '',
- backend: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Backend = 'flashinfer',
Quantize model parameters to mutable MXFP8Tensor storage.
Converts parameters already initialized with TE MXFP8 storage by the per-module precision recipe; ordinary BF16/FP16 parameters are left untouched. When persistent_buffers is provided, new quantized values are
copy_()’d into the existing MXFP8Tensor objects so that CUDA-graph device-pointer captures remain valid. Persistent buffers are deliberately created outside inference mode so later refits can update them regardless of the caller’s execution mode.- Parameters:
model – The model whose parameters should be quantized.
persistent_buffers – If not
None, a dict mapping fully-qualified parameter names to previously-createdMXFP8Tensorobjects. Updated in-place and returned._prefix – Internal recursion prefix; callers should not set this.
backend – ‘flashinfer’ or ‘triton’ quantization backend.
- Returns:
The
persistent_buffersdict (created on first call ifNone).
- core.inference.quantization.utils._mm_mxfp8_flashinfer(
- x_mxfp8: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Tensor,
- weight: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Tensor,
- out=None,
MXFP8 matmul via FlashInfer.
- core.inference.quantization.utils._mm_mxfp8_torch(
- x_mxfp8: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Tensor,
- weight: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Tensor,
- out=None,
MXFP8 matmul via torch.nn.functional.scaled_mm.
- core.inference.quantization.utils.mm_mxfp8(
- x: torch.Tensor,
- weight: megatron.core.inference.quantization.mxfp8_tensor.MXFP8Tensor,
- out: torch.Tensor = None,
Compute a matmul in MXFP8.
Quantizes the bf16 input activation tensor on the fly. Weight must be pre-quantized. Dispatches to FlashInfer or torch based on weight.backend.