From 52a3090bda7a48977cfab649d1eeb3919c126fcd Mon Sep 17 00:00:00 2001 From: Tim Moon Date: Fri, 7 Aug 2026 08:29:07 +0000 Subject: [PATCH] Add config to set TE quantization recipe attrs Signed-off-by: Tim Moon --- src/megatron/bridge/training/mixed_precision.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/megatron/bridge/training/mixed_precision.py b/src/megatron/bridge/training/mixed_precision.py index 1a004d5d88..92a00db573 100644 --- a/src/megatron/bridge/training/mixed_precision.py +++ b/src/megatron/bridge/training/mixed_precision.py @@ -14,7 +14,7 @@ import logging from dataclasses import dataclass, fields -from typing import TYPE_CHECKING, Callable, Optional +from typing import TYPE_CHECKING, Any, Callable, Optional import torch from megatron.core.distributed import DistributedDataParallelConfig @@ -50,6 +50,8 @@ class MixedPrecisionConfig: fp8_recipe: str = ( "tensorwise" # "tensorwise", "delayed", "mxfp8" (for Blackwell only), "blockwise" (for Hopper only) ) + fp8_recipe_attrs: Optional[dict[str, Any]] = None + fp8_quantizer_factory: Optional[str] = None first_last_layers_bf16: bool = False fp8_margin: int = 0 fp8_amax_history_len: int = 1 @@ -62,6 +64,8 @@ class MixedPrecisionConfig: # fp4 related fp4: Optional[str] = None fp4_recipe: str = "nvfp4" + fp4_recipe_attrs: Optional[dict[str, Any]] = None + fp4_quantizer_factory: Optional[str] = None fp4_param: Optional[bool] = None fp4_param_gather: bool = False # FP16 Loss scaling