diff --git a/primus/backends/megatron/core/models/diffusion/common/config.py b/primus/backends/megatron/core/models/diffusion/common/config.py index de4074744..17f3cc930 100644 --- a/primus/backends/megatron/core/models/diffusion/common/config.py +++ b/primus/backends/megatron/core/models/diffusion/common/config.py @@ -11,6 +11,7 @@ from dataclasses import dataclass from typing import Optional +import torch from megatron.core.enums import Fp8Recipe from megatron.core.transformer.transformer_config import TransformerConfig @@ -38,6 +39,12 @@ class BaseDiffusionConfig(TransformerConfig): sensitive_layers_start: Number of sensitive layers at start (default: 0) sensitive_layers_end: Number of sensitive layers at end (default: 0) sensitive_layer_precision: Precision for sensitive layers (default: 'bf16') + outer_sensitive_layers_start: Number of outermost starting layers that + override the sensitive-layer precision (default: 0) + outer_sensitive_layers_end: Number of outermost ending layers that + override the sensitive-layer precision (default: 0) + outer_sensitive_layer_precision: Precision for the outer override + (default: 'bf16') Inherited from TransformerConfig: hidden_size: Hidden dimension size @@ -82,6 +89,9 @@ class BaseDiffusionConfig(TransformerConfig): sensitive_layers_start: int = 0 sensitive_layers_end: int = 0 sensitive_layer_precision: str = "bf16" # "bf16", "tw_fp8", or "mxfp8" (future) + outer_sensitive_layers_start: int = 0 + outer_sensitive_layers_end: int = 0 + outer_sensitive_layer_precision: str = "bf16" def __post_init__(self): """Post-initialization processing.""" @@ -97,6 +107,25 @@ def __post_init__(self): "Set pipeline_model_parallel_size=1." ) + layer_count_fields = ( + "sensitive_layers_start", + "sensitive_layers_end", + "outer_sensitive_layers_start", + "outer_sensitive_layers_end", + ) + for field_name in layer_count_fields: + value = getattr(self, field_name) + if type(value) is not int: + raise ValueError(f"{field_name} must be an integer, got {value!r}") + if value < 0: + raise ValueError(f"{field_name} must be non-negative, got {value}") + + sensitive_count = self.sensitive_layers_start + self.sensitive_layers_end + outer_sensitive_count = self.outer_sensitive_layers_start + self.outer_sensitive_layers_end + if (sensitive_count or outer_sensitive_count) and not self.sensitive_layers_enabled: + raise ValueError("sensitive layer counts require sensitive_layers_enabled=True") + + active_precisions = set() if self.sensitive_layers_enabled: if self.num_layers <= 1: raise ValueError( @@ -112,17 +141,106 @@ def __post_init__(self): f"sensitive_layers_end ({self.sensitive_layers_end}) exceeds " f"num_layers ({self.num_layers})" ) + if self.outer_sensitive_layers_start > self.sensitive_layers_start: + raise ValueError( + "outer_sensitive_layers_start " + f"({self.outer_sensitive_layers_start}) exceeds " + f"sensitive_layers_start ({self.sensitive_layers_start})" + ) + if self.outer_sensitive_layers_end > self.sensitive_layers_end: + raise ValueError( + "outer_sensitive_layers_end " + f"({self.outer_sensitive_layers_end}) exceeds " + f"sensitive_layers_end ({self.sensitive_layers_end})" + ) + + if self.transformer_impl != "local": + raise ValueError( + "sensitive layer routing requires transformer_impl='local'; " + f"got {self.transformer_impl!r}" + ) + if self.fp4 != "mxfp4" or self.fp4_recipe != "mxfp4": + raise ValueError( + "sensitive layer routing requires fp4='mxfp4' and " + f"fp4_recipe='mxfp4'; got fp4={self.fp4!r}, " + f"fp4_recipe={self.fp4_recipe!r}" + ) + + inner_sensitive_count = ( + self.sensitive_layers_start + self.sensitive_layers_end - outer_sensitive_count + ) + if inner_sensitive_count > 0: + active_precisions.add(self.sensitive_layer_precision) + if outer_sensitive_count > 0: + active_precisions.add(self.outer_sensitive_layer_precision) + unsupported_precisions = active_precisions - {"bf16", "tw_fp8"} + if unsupported_precisions: + raise ValueError( + "sensitive layer precision must be 'bf16' or 'tw_fp8'; " + f"got {sorted(unsupported_precisions)!r}" + ) + + if outer_sensitive_count > 0: + collapsed_start = ( + self.outer_sensitive_layers_start > 0 + and self.outer_sensitive_layers_start == self.sensitive_layers_start + ) + collapsed_end = ( + self.outer_sensitive_layers_end > 0 + and self.outer_sensitive_layers_end == self.sensitive_layers_end + ) + if collapsed_start or collapsed_end: + raise ValueError( + "each graduated outer boundary requires at least one " + "inner sensitive layer on the same side" + ) + if ( + self.sensitive_layer_precision != "tw_fp8" + or self.outer_sensitive_layer_precision != "bf16" + ): + raise ValueError( + "graduated sensitive routing requires inner precision " + "'tw_fp8' and outer precision 'bf16'" + ) + + if "tw_fp8" in active_precisions: + if self.fp8 not in (None, "e4m3"): + raise ValueError( + "tw_fp8 sensitive layers require fp8=None or fp8='e4m3'; " f"got {self.fp8!r}" + ) + if self.fp8_recipe not in ( + None, + Fp8Recipe.delayed, + Fp8Recipe.tensorwise, + ): + raise ValueError( + "tw_fp8 sensitive layers require fp8_recipe='tensorwise' " + f"or the deferred 'delayed' default; got {self.fp8_recipe!r}" + ) + if "bf16" in active_precisions and ( + not self.bf16 or self.fp16 or self.params_dtype != torch.bfloat16 + ): + raise ValueError( + "bf16 sensitive layers require bf16=True, fp16=False, " + f"and params_dtype=torch.bfloat16; got bf16={self.bf16!r}, " + f"fp16={self.fp16!r}, params_dtype={self.params_dtype!r}" + ) + + # The FP4 context uses these legacy fields to exclude the complete + # heterogeneous boundary from MXFP4. The Flux layer spec chooses + # BF16 versus FP8 within that excluded boundary. self.first_last_layers_bf16 = True self.num_layers_at_start_in_bf16 = self.sensitive_layers_start self.num_layers_at_end_in_bf16 = self.sensitive_layers_end - if self.sensitive_layers_enabled and self.sensitive_layer_precision == "tw_fp8": - _deferred_fp8 = "e4m3" if self.fp8 is None else None - _deferred_fp8_recipe = ( - Fp8Recipe.tensorwise - if self.fp8_recipe is None or self.fp8_recipe == Fp8Recipe.delayed - else None - ) + uses_tw_fp8 = self.sensitive_layers_enabled and "tw_fp8" in active_precisions + if uses_tw_fp8: + # Megatron rejects global FP4+FP8 before it knows these formats are + # assigned to disjoint layers. Hide validated FP8 settings during + # parent validation, then restore the canonical local FP8 state. + _deferred_fp8 = "e4m3" + _deferred_fp8_recipe = Fp8Recipe.tensorwise + self.fp8 = None else: _deferred_fp8 = None _deferred_fp8_recipe = None diff --git a/primus/backends/megatron/core/models/diffusion/flux/layer_spec.py b/primus/backends/megatron/core/models/diffusion/flux/layer_spec.py index 9a1932b09..c637395df 100644 --- a/primus/backends/megatron/core/models/diffusion/flux/layer_spec.py +++ b/primus/backends/megatron/core/models/diffusion/flux/layer_spec.py @@ -527,20 +527,120 @@ def get_flux_layer_spec( ) from megatron.core.transformer.transformer_layer import get_transformer_layer_offset - # Default backend selection based on config - sensitive_backend = None + # Resolve heterogeneous routing before selecting a homogeneous fallback. + sensitive_enabled = getattr(config, "sensitive_layers_enabled", False) + total = config.num_joint_layers + config.num_single_layers + layer_counts = { + "sensitive_layers_start": getattr(config, "sensitive_layers_start", 0), + "sensitive_layers_end": getattr(config, "sensitive_layers_end", 0), + "outer_sensitive_layers_start": getattr(config, "outer_sensitive_layers_start", 0), + "outer_sensitive_layers_end": getattr(config, "outer_sensitive_layers_end", 0), + } + for field_name, value in layer_counts.items(): + if type(value) is not int: + raise ValueError(f"{field_name} must be an integer, got {value!r}") + if value < 0: + raise ValueError(f"{field_name} must be non-negative, got {value}") + if not sensitive_enabled and any(layer_counts.values()): + raise ValueError("sensitive layer counts require sensitive_layers_enabled=True") + + num_start = layer_counts["sensitive_layers_start"] if sensitive_enabled else 0 + num_end = layer_counts["sensitive_layers_end"] if sensitive_enabled else 0 + outer_num_start = layer_counts["outer_sensitive_layers_start"] if sensitive_enabled else 0 + outer_num_end = layer_counts["outer_sensitive_layers_end"] if sensitive_enabled else 0 + + if sensitive_enabled: + if num_start + num_end <= 0: + raise ValueError("sensitive_layers_enabled=True requires a non-empty boundary") + if num_start + num_end > total: + raise ValueError( + "sensitive layer counts exceed the number of Flux layers: " + f"{num_start} + {num_end} > {total}" + ) + if outer_num_start > num_start: + raise ValueError("outer_sensitive_layers_start cannot exceed sensitive_layers_start") + if outer_num_end > num_end: + raise ValueError("outer_sensitive_layers_end cannot exceed sensitive_layers_end") + + inner_sensitive_count = num_start + num_end - outer_num_start - outer_num_end + outer_sensitive_count = outer_num_start + outer_num_end - if backend is None: + sensitive_backend = None + outer_sensitive_backend = None + + if sensitive_enabled: + active_precisions = set() + if inner_sensitive_count > 0: + active_precisions.add(getattr(config, "sensitive_layer_precision", "bf16")) + if outer_sensitive_count > 0: + active_precisions.add(getattr(config, "outer_sensitive_layer_precision", "bf16")) + collapsed_start = outer_num_start > 0 and outer_num_start == num_start + collapsed_end = outer_num_end > 0 and outer_num_end == num_end + if collapsed_start or collapsed_end: + raise ValueError( + "each graduated outer boundary requires at least one " + "inner sensitive layer on the same side" + ) + if ( + getattr(config, "sensitive_layer_precision", None) != "tw_fp8" + or getattr(config, "outer_sensitive_layer_precision", None) != "bf16" + ): + raise ValueError( + "graduated sensitive routing requires inner precision " + "'tw_fp8' and outer precision 'bf16'" + ) + if "tw_fp8" in active_precisions and (config.fp8 != "e4m3" or config.fp8_recipe != "tensorwise"): + raise ValueError( + "tw_fp8 sensitive layers require normalized fp8='e4m3' " "and fp8_recipe='tensorwise'" + ) + if "bf16" in active_precisions and ( + not config.bf16 or config.fp16 or config.params_dtype != torch.bfloat16 + ): + raise ValueError( + "bf16 sensitive layers require bf16=True, fp16=False, " "and params_dtype=torch.bfloat16" + ) + if backend is not None: + raise ValueError( + "sensitive layer routing does not support an explicit backend; " + "pass backend=None so each precision region is selected explicitly" + ) + if config.transformer_impl != "local": + raise ValueError( + "sensitive layer routing requires transformer_impl='local'; " + f"got {config.transformer_impl!r}" + ) + if config.fp4 != "mxfp4" or config.fp4_recipe != "mxfp4": + raise ValueError( + "sensitive layer routing requires the local MXFP4 backend; " + f"got fp4={config.fp4!r}, fp4_recipe={config.fp4_recipe!r}" + ) + if PrimusTurboMXFP4LocalSpecProvider is None: + raise RuntimeError("sensitive layer routing requires the MXFP4 local provider") + + def resolve_sensitive_backend(precision): + if precision == "tw_fp8": + if PrimusTurboFloat8LocalSpecProvider is None: + raise RuntimeError("tw_fp8 sensitive layers require the FP8 local provider") + return PrimusTurboFloat8LocalSpecProvider() + if precision == "bf16": + if PrimusTurboLocalSpecProvider is None: + raise RuntimeError("bf16 sensitive layers require the native local provider") + return PrimusTurboLocalSpecProvider() + raise ValueError(f"unsupported sensitive layer precision: {precision!r}") + + backend = PrimusTurboMXFP4LocalSpecProvider() + if inner_sensitive_count > 0: + sensitive_backend = resolve_sensitive_backend( + getattr(config, "sensitive_layer_precision", "bf16") + ) + if outer_sensitive_count > 0: + outer_sensitive_backend = resolve_sensitive_backend( + getattr(config, "outer_sensitive_layer_precision", "bf16") + ) + elif backend is None: if config.transformer_impl == "local": if config.fp4 is not None and PrimusTurboMXFP4LocalSpecProvider is not None: backend = PrimusTurboMXFP4LocalSpecProvider() - - # Resolve sensitive layer backend - sensitive_precision = getattr(config, "sensitive_layer_precision", "bf16") - if sensitive_precision == "tw_fp8": - sensitive_backend = PrimusTurboFloat8LocalSpecProvider() - elif sensitive_precision == "bf16": - sensitive_backend = PrimusTurboLocalSpecProvider() elif ( config.fp8 is not None and HAVE_PRIMUS_TURBO_LOCAL @@ -562,16 +662,19 @@ def get_flux_layer_spec( backend = LocalSpecProvider() - # Build per-layer specs with optional sensitive-layer heterogeneity - sensitive_enabled = getattr(config, "sensitive_layers_enabled", False) - num_start = getattr(config, "sensitive_layers_start", 0) if sensitive_enabled else 0 - num_end = getattr(config, "sensitive_layers_end", 0) if sensitive_enabled else 0 - total = config.num_joint_layers + config.num_single_layers - + # Build per-layer specs with optional sensitive-layer heterogeneity. layer_specs = [] for i in range(total): + is_outer_sensitive = outer_sensitive_backend is not None and ( + (i < outer_num_start) or (i >= total - outer_num_end) + ) is_sensitive = sensitive_backend is not None and ((i < num_start) or (i >= total - num_end)) - layer_backend = sensitive_backend if is_sensitive else backend + if is_outer_sensitive: + layer_backend = outer_sensitive_backend + elif is_sensitive: + layer_backend = sensitive_backend + else: + layer_backend = backend if i < config.num_joint_layers: layer_specs.append(get_flux_double_transformer_spec_for_backend(layer_backend)) diff --git a/primus/backends/megatron/flux_pretrain_trainer.py b/primus/backends/megatron/flux_pretrain_trainer.py index 602a21d85..33961d50d 100644 --- a/primus/backends/megatron/flux_pretrain_trainer.py +++ b/primus/backends/megatron/flux_pretrain_trainer.py @@ -540,6 +540,9 @@ def _build_flux_config_from_yaml(self): "sensitive_layers_start": getattr(params, "sensitive_layers_start", 0), "sensitive_layers_end": getattr(params, "sensitive_layers_end", 0), "sensitive_layer_precision": getattr(params, "sensitive_layer_precision", "bf16"), + "outer_sensitive_layers_start": getattr(params, "outer_sensitive_layers_start", 0), + "outer_sensitive_layers_end": getattr(params, "outer_sensitive_layers_end", 0), + "outer_sensitive_layer_precision": getattr(params, "outer_sensitive_layer_precision", "bf16"), "mxfp4_gradient_stochastic_rounding": getattr( params, "mxfp4_gradient_stochastic_rounding", False ), @@ -732,6 +735,9 @@ def _log_flux_config(self, config, args): "sensitive_layers_start", "sensitive_layers_end", "sensitive_layer_precision", + "outer_sensitive_layers_start", + "outer_sensitive_layers_end", + "outer_sensitive_layer_precision", ], "Recomputation": [ "recompute_granularity", diff --git a/tests/unit_tests/backends/megatron/diffusion/test_flux_config.py b/tests/unit_tests/backends/megatron/diffusion/test_flux_config.py index 2a4348832..037ad231c 100644 --- a/tests/unit_tests/backends/megatron/diffusion/test_flux_config.py +++ b/tests/unit_tests/backends/megatron/diffusion/test_flux_config.py @@ -8,6 +8,7 @@ """ import pytest +import torch from tests.utils import skip_if_no_cuda @@ -17,6 +18,27 @@ from primus.backends.megatron.core.models.diffusion.flux.config import FluxConfig from tests.utils import PrimusUT + +def graduated_config_kwargs(**overrides): + kwargs = { + "transformer_impl": "local", + "fp4": "mxfp4", + "fp4_recipe": "mxfp4", + "bf16": True, + "fp16": False, + "params_dtype": torch.bfloat16, + "sensitive_layers_enabled": True, + "sensitive_layers_start": 4, + "sensitive_layers_end": 4, + "sensitive_layer_precision": "tw_fp8", + "outer_sensitive_layers_start": 1, + "outer_sensitive_layers_end": 1, + "outer_sensitive_layer_precision": "bf16", + } + kwargs.update(overrides) + return kwargs + + # ======================================================================== # Base Diffusion Configuration Tests # ======================================================================== @@ -89,6 +111,155 @@ def test_validation_axes_dim_positive_values(self): config.validate() self.assertIn("All axes_dim values must be positive", str(cm.exception)) + def test_sensitive_layer_counts_require_sensitive_routing(self): + """Boundary counts are invalid when heterogeneous routing is disabled.""" + for field_name in ("sensitive_layers_start", "outer_sensitive_layers_start"): + with self.subTest(field_name=field_name): + with self.assertRaisesRegex( + ValueError, + "sensitive layer counts require sensitive_layers_enabled=True", + ): + FluxConfig.flux_12b(**{field_name: 1}) + + def test_outer_sensitive_layers_must_fit_inside_boundary(self): + """The outer override cannot extend into the MXFP4 middle.""" + with self.assertRaisesRegex( + ValueError, "outer_sensitive_layers_start .* exceeds sensitive_layers_start" + ): + FluxConfig.flux_12b( + sensitive_layers_enabled=True, + sensitive_layers_start=4, + sensitive_layers_end=4, + outer_sensitive_layers_start=5, + ) + + def test_outer_sensitive_layer_counts_must_be_non_negative(self): + with self.assertRaisesRegex(ValueError, "outer_sensitive_layers_end must be non-negative"): + FluxConfig.flux_12b(outer_sensitive_layers_end=-1) + + def test_sensitive_layer_counts_must_be_integers(self): + cases = [ + ("sensitive_layers_start", 0.5), + ("sensitive_layers_end", True), + ("outer_sensitive_layers_start", "1"), + ("outer_sensitive_layers_end", 1.0), + ] + for field_name, value in cases: + with self.subTest(field_name=field_name, value=value): + with self.assertRaisesRegex(ValueError, f"{field_name} must be an integer"): + FluxConfig.flux_12b(**{field_name: value}) + + def test_outer_sensitive_end_must_fit_inside_boundary(self): + with self.assertRaisesRegex(ValueError, "outer_sensitive_layers_end .* exceeds sensitive_layers_end"): + FluxConfig.flux_12b( + sensitive_layers_enabled=True, + sensitive_layers_start=4, + sensitive_layers_end=4, + outer_sensitive_layers_end=5, + ) + + def test_sensitive_routing_defaults_are_backward_compatible(self): + config = FluxConfig.flux_12b() + assert config.sensitive_layers_enabled is False + assert config.outer_sensitive_layers_start == 0 + assert config.outer_sensitive_layers_end == 0 + + def test_graduated_routing_normalizes_tensorwise_fp8(self): + config = FluxConfig.flux_12b(**graduated_config_kwargs()) + assert config.fp8 == "e4m3" + assert config.fp8_recipe == "tensorwise" + assert config.bf16 is True + assert config.params_dtype == torch.bfloat16 + + def test_graduated_routing_accepts_explicit_normalized_fp8(self): + config = FluxConfig.flux_12b( + **graduated_config_kwargs( + fp8="e4m3", + fp8_recipe="tensorwise", + ) + ) + assert config.fp8 == "e4m3" + assert config.fp8_recipe == "tensorwise" + + def test_sensitive_routing_requires_local_transformer(self): + with self.assertRaisesRegex(ValueError, "sensitive layer routing requires transformer_impl='local'"): + FluxConfig.flux_12b(**graduated_config_kwargs(transformer_impl="transformer_engine")) + + def test_sensitive_routing_requires_local_mxfp4(self): + cases = [ + {"fp4": None}, + {"fp4_recipe": None}, + {"fp4": "nvfp4", "fp4_recipe": "nvfp4"}, + ] + for overrides in cases: + with self.subTest(overrides=overrides): + with self.assertRaisesRegex( + ValueError, + "sensitive layer routing requires fp4='mxfp4' and fp4_recipe='mxfp4'", + ): + FluxConfig.flux_12b(**graduated_config_kwargs(**overrides)) + + def test_tw_fp8_sensitive_layers_require_tensorwise_recipe(self): + for fp8_recipe in ("blockwise", "mxfp8", "custom"): + with self.subTest(fp8_recipe=fp8_recipe): + with self.assertRaisesRegex( + ValueError, + "tw_fp8 sensitive layers require fp8_recipe='tensorwise'", + ): + FluxConfig.flux_12b(**graduated_config_kwargs(fp8_recipe=fp8_recipe)) + + def test_tw_fp8_sensitive_layers_require_e4m3_format(self): + with self.assertRaisesRegex(ValueError, "tw_fp8 sensitive layers require fp8=None or fp8='e4m3'"): + FluxConfig.flux_12b(**graduated_config_kwargs(fp8="hybrid")) + + def test_bf16_sensitive_layers_require_bf16_model_dtype(self): + cases = [ + {"bf16": False}, + {"bf16": False, "fp16": True, "params_dtype": torch.float16}, + {"params_dtype": torch.float32}, + ] + for overrides in cases: + with self.subTest(overrides=overrides): + with self.assertRaisesRegex(ValueError, "bf16 sensitive layers require bf16=True"): + FluxConfig.flux_12b(**graduated_config_kwargs(**overrides)) + + def test_sensitive_routing_rejects_unknown_precision(self): + for field_name in ( + "sensitive_layer_precision", + "outer_sensitive_layer_precision", + ): + with self.subTest(field_name=field_name): + with self.assertRaisesRegex( + ValueError, + "sensitive layer precision must be 'bf16' or 'tw_fp8'", + ): + FluxConfig.flux_12b(**graduated_config_kwargs(**{field_name: "mxfp8"})) + + def test_graduated_routing_rejects_precision_role_inversion(self): + with self.assertRaisesRegex( + ValueError, + "requires inner precision 'tw_fp8' and outer precision 'bf16'", + ): + FluxConfig.flux_12b( + **graduated_config_kwargs( + sensitive_layer_precision="bf16", + outer_sensitive_layer_precision="tw_fp8", + ) + ) + + def test_graduated_routing_requires_inner_layer_on_each_outer_side(self): + cases = [ + {"sensitive_layers_start": 1}, + {"sensitive_layers_end": 1}, + ] + for overrides in cases: + with self.subTest(overrides=overrides): + with self.assertRaisesRegex( + ValueError, + "requires at least one inner sensitive layer on the same side", + ): + FluxConfig.flux_12b(**graduated_config_kwargs(**overrides)) + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/unit_tests/backends/megatron/diffusion/test_flux_layer_spec_backend_selection.py b/tests/unit_tests/backends/megatron/diffusion/test_flux_layer_spec_backend_selection.py index d1af8b474..4028540b0 100644 --- a/tests/unit_tests/backends/megatron/diffusion/test_flux_layer_spec_backend_selection.py +++ b/tests/unit_tests/backends/megatron/diffusion/test_flux_layer_spec_backend_selection.py @@ -8,7 +8,11 @@ This ensures alignment between backend selection and FSDP2 wrapping decisions. """ +from collections import Counter +from unittest.mock import patch + import pytest +import torch from tests.utils import skip_if_no_cuda @@ -21,6 +25,25 @@ from tests.utils import PrimusUT +def graduated_config(): + return FluxConfig.flux_12b( + transformer_impl="local", + fp4="mxfp4", + fp4_recipe="mxfp4", + mxfp4_backward_precision="mxfp4", + bf16=True, + fp16=False, + params_dtype=torch.bfloat16, + sensitive_layers_enabled=True, + sensitive_layers_start=4, + sensitive_layers_end=4, + sensitive_layer_precision="tw_fp8", + outer_sensitive_layers_start=1, + outer_sensitive_layers_end=1, + outer_sensitive_layer_precision="bf16", + ) + + class TestFluxLayerSpecBackendSelection(PrimusUT): """Tests for backend selection in get_flux_layer_spec().""" @@ -77,3 +100,162 @@ def test_backend_selection_local_no_fp8_uses_native_linear(self): attn_spec.submodules.linear_qkv == ColumnParallelLinear ), f"Expected native ColumnParallelLinear when fp8=None, got {attn_spec.submodules.linear_qkv}" assert found_any, "No attention linear_qkv specs found to validate backend selection" + + def test_graduated_boundary_routes_outer_bf16_inner_fp8_middle_mxfp4(self): + """Every declared Flux linear follows the exact three-precision layout.""" + from megatron.core.tensor_parallel import ( + ColumnParallelLinear, + RowParallelLinear, + ) + + from primus.backends.megatron.core.extensions.primus_turbo_float8_local import ( + Float8ColumnParallelLinear, + Float8RowParallelLinear, + ) + from primus.backends.megatron.core.extensions.primus_turbo_mxfp4_local import ( + MXFP4ColumnParallelLinear, + MXFP4RowParallelLinear, + ) + + expected_classes = { + "bf16": (ColumnParallelLinear, RowParallelLinear), + "fp8": (Float8ColumnParallelLinear, Float8RowParallelLinear), + "mxfp4": (MXFP4ColumnParallelLinear, MXFP4RowParallelLinear), + } + + block_submodules = get_flux_layer_spec(graduated_config(), backend=None) + precision_counts = Counter() + observed_slots = 0 + for index, layer_spec in enumerate(block_submodules.layer_specs): + expected_precision = ( + "bf16" if index in (0, 56) else "fp8" if index in (1, 2, 3, 53, 54, 55) else "mxfp4" + ) + column_class, row_class = expected_classes[expected_precision] + attention = layer_spec.submodules.self_attention.submodules + mlp = layer_spec.submodules.mlp.submodules + slots = { + "linear_qkv": (attention.linear_qkv, column_class), + "linear_proj": (attention.linear_proj, row_class), + "linear_fc1": (mlp.linear_fc1, column_class), + "linear_fc2": (mlp.linear_fc2, row_class), + } + if index < 19: + slots["added_linear_qkv"] = ( + attention.added_linear_qkv, + column_class, + ) + + for slot, (actual_class, expected_class) in slots.items(): + assert actual_class is expected_class, ( + f"block {index} {slot}: expected {expected_class}, " f"got {actual_class}" + ) + precision_counts[expected_precision] += 1 + observed_slots += 1 + + assert observed_slots == 247 + assert precision_counts == {"bf16": 9, "fp8": 27, "mxfp4": 211} + + def test_graduated_routing_rejects_explicit_backend(self): + from primus.backends.megatron.core.extensions.primus_turbo_local_spec import ( + PrimusTurboMXFP4LocalSpecProvider, + ) + + with self.assertRaisesRegex(ValueError, "does not support an explicit backend"): + get_flux_layer_spec( + graduated_config(), + backend=PrimusTurboMXFP4LocalSpecProvider(), + ) + + def test_graduated_routing_rejects_missing_mxfp4_provider(self): + with patch( + "primus.backends.megatron.core.models.diffusion.flux.layer_spec." + "PrimusTurboMXFP4LocalSpecProvider", + None, + ): + with self.assertRaisesRegex(RuntimeError, "requires the MXFP4 local provider"): + get_flux_layer_spec(graduated_config(), backend=None) + + def test_graduated_routing_rejects_missing_fp8_provider(self): + with patch( + "primus.backends.megatron.core.models.diffusion.flux.layer_spec." + "PrimusTurboFloat8LocalSpecProvider", + None, + ): + with self.assertRaisesRegex(RuntimeError, "requires the FP8 local provider"): + get_flux_layer_spec(graduated_config(), backend=None) + + def test_graduated_routing_rejects_missing_native_provider(self): + with patch( + "primus.backends.megatron.core.models.diffusion.flux.layer_spec." "PrimusTurboLocalSpecProvider", + None, + ): + with self.assertRaisesRegex(RuntimeError, "requires the native local provider"): + get_flux_layer_spec(graduated_config(), backend=None) + + def test_graduated_routing_revalidates_mutated_config(self): + cases = [ + ( + "sensitive_layers_enabled", + False, + "sensitive layer counts require sensitive_layers_enabled=True", + ), + ( + "fp8_recipe", + "blockwise", + "require normalized fp8='e4m3'", + ), + ( + "outer_sensitive_layer_precision", + "tw_fp8", + "requires inner precision 'tw_fp8' and outer precision 'bf16'", + ), + ( + "outer_sensitive_layers_start", + 4, + "requires at least one inner sensitive layer on the same side", + ), + ( + "outer_sensitive_layers_start", + 5, + "outer_sensitive_layers_start cannot exceed sensitive_layers_start", + ), + ( + "outer_sensitive_layers_end", + 5, + "outer_sensitive_layers_end cannot exceed sensitive_layers_end", + ), + ( + "sensitive_layers_start", + "4", + "sensitive_layers_start must be an integer", + ), + ( + "sensitive_layers_end", + -1, + "sensitive_layers_end must be non-negative", + ), + ( + "sensitive_layers_start", + 54, + "sensitive layer counts exceed the number of Flux layers", + ), + ( + "params_dtype", + torch.float32, + "bf16 sensitive layers require bf16=True", + ), + ] + for field_name, value, message in cases: + with self.subTest(field_name=field_name, value=value): + config = graduated_config() + setattr(config, field_name, value) + with self.assertRaisesRegex(ValueError, message): + get_flux_layer_spec(config, backend=None) + + config = graduated_config() + config.sensitive_layers_start = 0 + config.sensitive_layers_end = 0 + config.outer_sensitive_layers_start = 0 + config.outer_sensitive_layers_end = 0 + with self.assertRaisesRegex(ValueError, "requires a non-empty boundary"): + get_flux_layer_spec(config, backend=None) diff --git a/tests/unit_tests/backends/megatron/diffusion/test_flux_model.py b/tests/unit_tests/backends/megatron/diffusion/test_flux_model.py index b8d16c6d3..47e1ab6d1 100644 --- a/tests/unit_tests/backends/megatron/diffusion/test_flux_model.py +++ b/tests/unit_tests/backends/megatron/diffusion/test_flux_model.py @@ -21,6 +21,7 @@ pack_latents, unpack_latents, ) +from tests.unit_tests.backends.megatron.conftest import requires_mxfp4 from tests.unit_tests.backends.megatron.diffusion.constants import ( CLIP_L_EMBEDDING_DIM, T5_XXL_EMBEDDING_DIM, @@ -37,6 +38,27 @@ class TestFluxModel(PrimusUT): def setup_parallel(self, init_parallel_state): """Initialize parallel state for model tests.""" + @pytest.fixture(autouse=True) + def pin_fp4_aiter(self, monkeypatch): + """Pin the backend required by MXFP4 module construction.""" + import collections + import os + + from primus_turbo.pytorch.core.backend import ( + BackendType, + GlobalBackendManager, + PrecisionType, + ) + + if os.environ.get("PRIMUS_TURBO_GEMM_BACKEND", None) == "": + monkeypatch.delenv("PRIMUS_TURBO_GEMM_BACKEND", raising=False) + pinned = collections.defaultdict(lambda: None) + if GlobalBackendManager._gemm_backend: + pinned.update(GlobalBackendManager._gemm_backend) + pinned[PrecisionType.FP4] = BackendType.AITER + monkeypatch.setattr(GlobalBackendManager, "_gemm_backend", pinned) + monkeypatch.setattr(GlobalBackendManager, "_auto_tune", False) + def test_forward_pass_small(self): """Test forward pass with small inputs.""" if not torch.cuda.is_available(): @@ -79,6 +101,53 @@ def test_forward_pass_small(self): assert not torch.isnan(output).any() assert not torch.isinf(output).any() + @requires_mxfp4 + def test_build_graduated_precision_model(self): + """Build a small model that instantiates BF16, FP8, and MXFP4 blocks.""" + from megatron.core.tensor_parallel import ColumnParallelLinear + + from primus.backends.megatron.core.extensions.primus_turbo_float8_local import ( + Float8ColumnParallelLinear, + ) + from primus.backends.megatron.core.extensions.primus_turbo_mxfp4_local import ( + MXFP4ColumnParallelLinear, + ) + + config = FluxConfig( + num_joint_layers=2, + num_single_layers=3, + hidden_size=128, + num_attention_heads=4, + ffn_hidden_size=256, + context_dim=128, + vec_in_dim=64, + model_channels=64, + axes_dim=(4, 14, 14), + transformer_impl="local", + fp4="mxfp4", + fp4_recipe="mxfp4", + bf16=True, + fp16=False, + params_dtype=torch.bfloat16, + sensitive_layers_enabled=True, + sensitive_layers_start=2, + sensitive_layers_end=2, + sensitive_layer_precision="tw_fp8", + outer_sensitive_layers_start=1, + outer_sensitive_layers_end=1, + outer_sensitive_layer_precision="bf16", + ) + + model = Flux(config).cuda() + layers = model.transformer.layers + + assert len(layers) == 5 + assert type(layers[0].self_attention.linear_qkv) is ColumnParallelLinear + assert type(layers[1].self_attention.linear_qkv) is Float8ColumnParallelLinear + assert type(layers[2].self_attention.linear_qkv) is MXFP4ColumnParallelLinear + assert type(layers[3].self_attention.linear_qkv) is Float8ColumnParallelLinear + assert type(layers[4].self_attention.linear_qkv) is ColumnParallelLinear + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/unit_tests/backends/megatron/diffusion/training/test_flux_model_creation.py b/tests/unit_tests/backends/megatron/diffusion/training/test_flux_model_creation.py index 0e876b128..c825e4129 100644 --- a/tests/unit_tests/backends/megatron/diffusion/training/test_flux_model_creation.py +++ b/tests/unit_tests/backends/megatron/diffusion/training/test_flux_model_creation.py @@ -161,6 +161,69 @@ def test_build_flux_config_from_yaml_extracts_parameters(self, monkeypatch: pyte assert config.params_dtype == torch.bfloat16 assert config.transformer_impl == "local" + def test_build_flux_config_from_yaml_forwards_graduated_routing(self, monkeypatch: pytest.MonkeyPatch): + backend_args = SimpleNamespace( + mock_data=True, + transformer_impl="local", + fp4="mxfp4", + fp4_recipe="mxfp4", + bf16=True, + fp16=False, + params_dtype=torch.bfloat16, + sensitive_layers_enabled=True, + sensitive_layers_start=4, + sensitive_layers_end=4, + sensitive_layer_precision="tw_fp8", + outer_sensitive_layers_start=1, + outer_sensitive_layers_end=1, + outer_sensitive_layer_precision="bf16", + ) + + trainer = _build_flux_trainer(monkeypatch, backend_args) + config = trainer._build_flux_config_from_yaml() + + assert config.sensitive_layers_start == 4 + assert config.sensitive_layers_end == 4 + assert config.sensitive_layer_precision == "tw_fp8" + assert config.outer_sensitive_layers_start == 1 + assert config.outer_sensitive_layers_end == 1 + assert config.outer_sensitive_layer_precision == "bf16" + assert config.fp8 == "e4m3" + assert config.fp8_recipe == "tensorwise" + + def test_log_flux_config_includes_graduated_routing(self, monkeypatch: pytest.MonkeyPatch): + backend_args = SimpleNamespace( + mock_data=True, + transformer_impl="local", + fp4="mxfp4", + fp4_recipe="mxfp4", + bf16=True, + fp16=False, + params_dtype=torch.bfloat16, + sensitive_layers_enabled=True, + sensitive_layers_start=4, + sensitive_layers_end=4, + sensitive_layer_precision="tw_fp8", + outer_sensitive_layers_start=1, + outer_sensitive_layers_end=1, + outer_sensitive_layer_precision="bf16", + ) + trainer = _build_flux_trainer(monkeypatch, backend_args) + config = trainer._build_flux_config_from_yaml() + messages = [] + monkeypatch.setattr( + "primus.backends.megatron.flux_pretrain_trainer.log_rank_0", + messages.append, + ) + + trainer._log_flux_config(config, SimpleNamespace(rank=0)) + + rendered = "\n".join(messages) + assert "sensitive_layer_precision" in rendered + assert "outer_sensitive_layers_start" in rendered + assert "outer_sensitive_layers_end" in rendered + assert "outer_sensitive_layer_precision" in rendered + def test_build_flux_config_from_yaml_torch_compile_settings(self, monkeypatch: pytest.MonkeyPatch): """Test that torch_compile settings are extracted from backend_args.""" backend_args = SimpleNamespace(