diff --git a/megatron/core/fp8_utils.py b/megatron/core/fp8_utils.py index 5411b676d83..e0eb53f7506 100644 --- a/megatron/core/fp8_utils.py +++ b/megatron/core/fp8_utils.py @@ -831,6 +831,29 @@ def get_fp8_context(config: TransformerConfig, layer_no: int = -1, is_init: bool return fp8_context + def get_fp8_disabled_context(config: TransformerConfig, is_init: bool = False): + """Return a context manager that disables TE quantization. + + Use this around submodule construction or execution that must stay in a higher + precision while its enclosing module uses an FP8 or FP4 context. + + Args: + config: Transformer configuration that controls quantization. + is_init: Whether to disable the parameter-initialization context instead of + the forward autocast context. + + Returns: + A disabled TE quantization context when quantization is active, otherwise a + no-op context. + """ + if is_init: + if not (config.fp8_param or config.fp4_param): + return nullcontext() + return transformer_engine.pytorch.fp8_model_init(enabled=False) + if not (config.fp8 or config.fp4): + return nullcontext() + return transformer_engine.pytorch.fp8_autocast(enabled=False) + else: def get_fp8_recipe(config: TransformerConfig): @@ -841,6 +864,10 @@ def get_fp8_context(config: TransformerConfig, layer_no: int = -1, is_init: bool """Returns dummy fp8 context manager since TE is not available.""" return nullcontext() + def get_fp8_disabled_context(config: TransformerConfig, is_init: bool = False): + """Return a no-op context manager since TE is not available.""" + return nullcontext() + if HAVE_TE: from transformer_engine.pytorch.fp8 import FP8GlobalStateManager diff --git a/megatron/core/optimizer/cpu_offloading/hybrid_optimizer.py b/megatron/core/optimizer/cpu_offloading/hybrid_optimizer.py index c87ccd5ff31..45bf910f84d 100644 --- a/megatron/core/optimizer/cpu_offloading/hybrid_optimizer.py +++ b/megatron/core/optimizer/cpu_offloading/hybrid_optimizer.py @@ -370,8 +370,11 @@ def _update_fp32_params_by_new_state(self): if not self.param_update_in_fp32: return for param, v in self.state.items(): - fp32_param = self.param_to_fp32_param[param] - fp32_param.data.copy_(v["master_param"]) + # Native FP32 params do not need a separate master parameter and are + # intentionally absent from param_to_fp32_param. + fp32_param = self.param_to_fp32_param.get(param) + if fp32_param is not None: + fp32_param.data.copy_(v["master_param"]) def update_fp32_param_by_new_param(self): """ diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index e503a16dde3..56b2fd5e843 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -1019,14 +1019,46 @@ def sharded_state_dict( state_dict = self.state_dict() + # Optimizer state ids enumerate the inner optimizer params: the fp32 main + # copies of float16 params, native fp32 params, and any frozen params, + # interleaved in the original param-group order. Map each fp32 main copy + # back to its model-side param; all other params already are model params. + main_param_id_to_model_param = { + id(main_param): model_param + for model_group, main_group in zip( + self.float16_groups, self.fp32_from_float16_groups, strict=True + ) + for model_param, main_param in zip(model_group, main_group, strict=True) + } + + def model_params_in_optimizer_order(): + for param in chain.from_iterable( + inner_group['params'] for inner_group in self.optimizer.param_groups + ): + yield main_param_id_to_model_param.get(id(param), param) + id_to_sharded_param_map = get_param_id_to_sharded_param_map( - model_sharded_state_dict, chain.from_iterable(g for g in self.float16_groups) + model_sharded_state_dict, model_params_in_optimizer_order() ) # Convert fp32_from_fp16_params assert len(state_dict['fp32_from_fp16_params']) == len( state_dict['optimizer']['param_groups'] ) + # State ids of the fp32 main copies only, skipping native fp32 and frozen params. + float16_param_ids_per_group = [] + for state_group, inner_group in zip( + state_dict['optimizer']['param_groups'], self.optimizer.param_groups, strict=True + ): + float16_param_ids_per_group.append( + [ + param_id + for param_id, param in zip( + state_group['params'], inner_group['params'], strict=True + ) + if id(param) in main_param_id_to_model_param + ] + ) state_dict['fp32_from_fp16_params'] = [ [ make_sharded_optimizer_tensor( @@ -1034,10 +1066,10 @@ def sharded_state_dict( fp32_param, prefix=f'optimizer.state.fp32_param', ) - for param_id, fp32_param in zip(state_group['params'], fp32_group) + for param_id, fp32_param in zip(param_ids, fp32_group, strict=True) ] - for fp32_group, state_group in zip( - state_dict['fp32_from_fp16_params'], state_dict['optimizer']['param_groups'] + for fp32_group, param_ids in zip( + state_dict['fp32_from_fp16_params'], float16_param_ids_per_group, strict=True ) ] diff --git a/megatron/core/transformer/module.py b/megatron/core/transformer/module.py index 558b1b07a15..bf28600a1aa 100644 --- a/megatron/core/transformer/module.py +++ b/megatron/core/transformer/module.py @@ -433,6 +433,40 @@ def float_conversion(val): return conversion_helper(val, float_conversion) +def mark_keep_in_fp32(tensor: torch.Tensor) -> torch.Tensor: + """Mark a parameter or buffer so that ``Float16Module`` keeps it in FP32. + + Args: + tensor: The parameter or buffer to mark. + + Returns: + The same tensor, for call-site convenience. + """ + tensor.keep_in_fp32 = True + return tensor + + +def convert_module_to_dtype_except_fp32_marked( + module: torch.nn.Module, dtype: torch.dtype +) -> torch.nn.Module: + """Cast floating-point parameters and buffers except those marked to stay in FP32. + + Args: + module: The module to convert in place. + dtype: The target floating-point dtype. + + Returns: + The converted module. + """ + return module._apply( + lambda tensor: ( + tensor.to(dtype) + if tensor.is_floating_point() and not getattr(tensor, 'keep_in_fp32', False) + else tensor + ) + ) + + class Float16Module(MegatronModule): """Float 16 Module. @@ -455,13 +489,17 @@ def __init__(self, config: TransformerConfig, module: torch.nn.Module): self.pg_collection = getattr(module, 'pg_collection', None) if self.fp16: - self.add_module('module', module.half()) + self.add_module( + 'module', convert_module_to_dtype_except_fp32_marked(module, torch.half) + ) def float16_convertor(val): return val.half() elif self.bf16: - self.add_module('module', module.bfloat16()) + self.add_module( + 'module', convert_module_to_dtype_except_fp32_marked(module, torch.bfloat16) + ) def float16_convertor(val): return val.bfloat16() diff --git a/tests/unit_tests/dist_checkpointing/test_optimizer.py b/tests/unit_tests/dist_checkpointing/test_optimizer.py index f93e09a43b7..7c319e0a14a 100644 --- a/tests/unit_tests/dist_checkpointing/test_optimizer.py +++ b/tests/unit_tests/dist_checkpointing/test_optimizer.py @@ -79,6 +79,27 @@ def sharded_state_dict(self): return sharded_state_dict +class NativeFp32Model(torch.nn.Module): + """Parameters for an interleaved trainable/frozen BF16 and FP32 group.""" + + def __init__(self): + super().__init__() + self.pre = torch.nn.Linear(8, 8, bias=False) + self.frozen = torch.nn.Linear(8, 8, bias=False) + self.frozen.weight.requires_grad_(False) + self.gate = torch.nn.Parameter(torch.zeros(24, dtype=torch.float32)) + self.post = torch.nn.Linear(8, 8, bias=False) + self.config = TransformerConfig( + hidden_size=8, num_attention_heads=1, num_layers=1, bf16=True + ) + + def sharded_state_dict(self): + return { + key: ShardedTensor.from_rank_offsets(key, value) + for key, value in self.state_dict(keep_vars=True).items() + } + + class SwigluFactoryModel(torch.nn.Module): def __init__(self, pp_separate_model: bool = False): super().__init__() @@ -238,6 +259,65 @@ def test_optimizer_params(self, tmp_path_dist_ckpt): ] ) + def test_float16_optimizer_with_native_fp32_and_frozen_params(self): + """Native FP32 and frozen param ids must not shift BF16 checkpoint state.""" + from megatron.core.optimizer import OptimizerConfig + from megatron.core.optimizer.optimizer import Float16OptimizerWithFloat16Params + from megatron.core.transformer.module import ( + convert_module_to_dtype_except_fp32_marked, + mark_keep_in_fp32, + ) + + Utils.initialize_model_parallel(1, 1) + model = NativeFp32Model().cuda() + model.gate = mark_keep_in_fp32(model.gate) + convert_module_to_dtype_except_fp32_marked(model, torch.bfloat16) + assert model.pre.weight.dtype == torch.bfloat16 + assert model.frozen.weight.dtype == torch.bfloat16 + assert not model.frozen.weight.requires_grad + assert model.gate.dtype == torch.float32 + assert model.post.weight.dtype == torch.bfloat16 + + # Use an explicit trainable BF16/frozen BF16/FP32/trainable BF16 order. + # Module.parameters() would yield the root gate before child parameters. + ordered_params = [model.pre.weight, model.frozen.weight, model.gate, model.post.weight] + for param in ordered_params: + if param.requires_grad: + param.grad = torch.zeros_like(param) + inner_optim = Adam(ordered_params) + inner_optim.step() + + optim = Float16OptimizerWithFloat16Params( + inner_optim, + OptimizerConfig(optimizer='adam', lr=1e-4, bf16=True), + None, + lambda opt, cfg: None, + ) + sharded_state_dict = optim.sharded_state_dict(model.sharded_state_dict()) + + # FP32 main copies pair with the BF16 params only, in optimizer order. + fp32_params = sharded_state_dict['fp32_from_fp16_params'][0] + assert [(sharded.key, tuple(sharded.data.shape)) for sharded in fp32_params] == [ + ('optimizer.state.fp32_param.pre.weight', (8, 8)), + ('optimizer.state.fp32_param.post.weight', (8, 8)), + ] + + # The frozen parameter has neither optimizer state nor an fp32 main copy. + state = sharded_state_dict['optimizer']['state'] + assert 1 not in state + + # Per-param state maps every trainable param, including native FP32, to the right key. + expected = {0: ('pre.weight', (8, 8)), 2: ('gate', (24,)), 3: ('post.weight', (8, 8))} + for param_id, (model_key, shape) in expected.items(): + for state_key in ('exp_avg', 'exp_avg_sq'): + sharded = state[param_id][state_key] + assert sharded.key == f'optimizer.state.{state_key}.{model_key}', sharded.key + assert tuple(sharded.data.shape) == shape, ( + param_id, + sharded.key, + sharded.data.shape, + ) + def initialize_pp_agnostic_model(pre_process=True, post_process=True, seed=0, **config_kwargs): torch.manual_seed(seed) diff --git a/tests/unit_tests/test_fp8_utils.py b/tests/unit_tests/test_fp8_utils.py index 5be17f03c9f..1c350d932b8 100644 --- a/tests/unit_tests/test_fp8_utils.py +++ b/tests/unit_tests/test_fp8_utils.py @@ -10,6 +10,31 @@ from tests.unit_tests.test_utilities import Utils +@pytest.mark.skipif(not fp8_utils.HAVE_TE, reason="Transformer Engine is not installed") +@pytest.mark.parametrize( + ("is_init", "config_values", "te_helper"), + [ + ( + False, + {"fp8": "hybrid", "fp4": None, "fp8_param": False, "fp4_param": False}, + "fp8_autocast", + ), + (True, {"fp8": None, "fp4": None, "fp8_param": True, "fp4_param": False}, "fp8_model_init"), + ], +) +def test_get_fp8_disabled_context_uses_disabled_te_context(is_init, config_values, te_helper): + config = Mock(**config_values) + disabled_context = Mock() + + with patch.object( + fp8_utils.transformer_engine.pytorch, te_helper, return_value=disabled_context + ) as te_context: + result = fp8_utils.get_fp8_disabled_context(config, is_init=is_init) + + assert result is disabled_context + te_context.assert_called_once_with(enabled=False) + + class MockTELinear(nn.Module): """Mock TE Linear module for testing.""" diff --git a/tests/unit_tests/test_optimizer_cpu_offloading.py b/tests/unit_tests/test_optimizer_cpu_offloading.py index 33febbb3eb0..379acc9dbda 100644 --- a/tests/unit_tests/test_optimizer_cpu_offloading.py +++ b/tests/unit_tests/test_optimizer_cpu_offloading.py @@ -17,6 +17,20 @@ from torch.optim import Adam as GPUAdam from megatron.core.optimizer.cpu_offloading import HybridDeviceOptimizer +from megatron.core.transformer.module import ( + convert_module_to_dtype_except_fp32_marked, + mark_keep_in_fp32, +) + + +class Fp32MarkedToyNet(nn.Module): + def __init__(self): + super().__init__() + self.proj = nn.Linear(4, 4, bias=False) + self.scale = mark_keep_in_fp32(nn.Parameter(torch.ones(4))) + + def forward(self, x): + return self.proj(x) * self.scale class Net(nn.Module): @@ -71,6 +85,52 @@ def setup_seed(seed): torch.backends.cudnn.benchmark = False # Disable auto-tuner for reproducibility +def test_load_state_dict_with_native_fp32_param(): + """Round-trip state for a BF16 toy net with a parameter marked to stay in FP32.""" + model = Fp32MarkedToyNet().cuda() + convert_module_to_dtype_except_fp32_marked(model, torch.bfloat16) + assert model.proj.weight.dtype == torch.bfloat16 + assert model.scale.dtype == torch.float32 + + optimizer = HybridDeviceOptimizer( + model.parameters(), + offload_fraction=1.0, + cpu_optimizer_cls=Adam, + gpu_optimizer_cls=GPUAdam, + param_update_in_fp32=True, + overlap_cpu_optimizer_d2h_h2d=False, + lr=1e-3, + ) + inputs = torch.ones(2, 4, device="cuda", dtype=torch.bfloat16) + model(inputs).sum().backward() + optimizer.step() + + restored_model = Fp32MarkedToyNet().cuda() + convert_module_to_dtype_except_fp32_marked(restored_model, torch.bfloat16) + restored_model.load_state_dict(model.state_dict()) + restored_optimizer = HybridDeviceOptimizer( + restored_model.parameters(), + offload_fraction=1.0, + cpu_optimizer_cls=Adam, + gpu_optimizer_cls=GPUAdam, + param_update_in_fp32=True, + overlap_cpu_optimizer_d2h_h2d=False, + lr=1e-3, + ) + restored_optimizer.load_state_dict(optimizer.state_dict()) + + assert set(restored_optimizer.state) == set(restored_model.parameters()) + assert restored_model.proj.weight in restored_optimizer.param_to_fp32_param + assert restored_model.scale not in restored_optimizer.param_to_fp32_param + assert torch.equal( + restored_optimizer.param_to_fp32_param[restored_model.proj.weight], + optimizer.param_to_fp32_param[model.proj.weight], + ) + + restored_model(inputs).sum().backward() + restored_optimizer.step() + + @pytest.mark.skipif( torch.__version__ < '2.3.0', reason=( diff --git a/tests/unit_tests/transformer/test_module.py b/tests/unit_tests/transformer/test_module.py index 92f15b2f46d..5faf6c81ef1 100644 --- a/tests/unit_tests/transformer/test_module.py +++ b/tests/unit_tests/transformer/test_module.py @@ -4,7 +4,7 @@ import torch from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed -from megatron.core.transformer.module import Float16Module, MegatronModule +from megatron.core.transformer.module import Float16Module, MegatronModule, mark_keep_in_fp32 from megatron.core.transformer.transformer_config import TransformerConfig from tests.unit_tests.test_utilities import Utils @@ -163,3 +163,18 @@ def test_bf16_module(self): x = torch.ones((2, 2)).cuda() # inputs are converted to bf16 then outputs are converted to fp32 assert bf16_module(x).dtype == torch.float32 + + @pytest.mark.parametrize( + ('precision', 'dtype'), [('fp16', torch.float16), ('bf16', torch.bfloat16)] + ) + def test_keep_in_fp32_params(self, precision, dtype): + transformer_config = self.transformer_config + megatron_module = self.megatron_module + megatron_module.fp32_param = mark_keep_in_fp32( + torch.nn.Parameter(torch.zeros(4, dtype=torch.float32, device='cuda')) + ) + setattr(transformer_config, precision, True) + float16_module = Float16Module(config=transformer_config, module=megatron_module) + + assert float16_module.module.linear.weight.dtype == dtype + assert float16_module.module.fp32_param.dtype == torch.float32