Skip to content
Draft
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,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
Expand Down Expand Up @@ -82,6 +88,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."""
Expand All @@ -97,6 +106,12 @@ def __post_init__(self):
"Set pipeline_model_parallel_size=1."
)

if self.outer_sensitive_layers_start < 0 or self.outer_sensitive_layers_end < 0:
raise ValueError("outer sensitive layer counts must be non-negative")
outer_sensitive_count = self.outer_sensitive_layers_start + self.outer_sensitive_layers_end
if outer_sensitive_count and not self.sensitive_layers_enabled:
raise ValueError("outer sensitive layers require sensitive_layers_enabled=True")

if self.sensitive_layers_enabled:
if self.num_layers <= 1:
raise ValueError(
Expand All @@ -112,11 +127,31 @@ 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})"
)

# 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":
uses_tw_fp8 = self.sensitive_layers_enabled and (
self.sensitive_layer_precision == "tw_fp8"
or (outer_sensitive_count > 0 and self.outer_sensitive_layer_precision == "tw_fp8")
)
if uses_tw_fp8:
_deferred_fp8 = "e4m3" if self.fp8 is None else None
_deferred_fp8_recipe = (
Fp8Recipe.tensorwise
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -529,18 +529,30 @@ def get_flux_layer_spec(

# Default backend selection based on config
sensitive_backend = None
outer_sensitive_backend = None

if 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()
def resolve_sensitive_backend(precision):
if precision == "tw_fp8":
return PrimusTurboFloat8LocalSpecProvider()
if precision == "bf16":
return PrimusTurboLocalSpecProvider()
raise ValueError(f"unsupported sensitive layer precision: {precision!r}")

sensitive_backend = resolve_sensitive_backend(
getattr(config, "sensitive_layer_precision", "bf16")
)
if (
getattr(config, "outer_sensitive_layers_start", 0) > 0
or getattr(config, "outer_sensitive_layers_end", 0) > 0
):
outer_sensitive_backend = resolve_sensitive_backend(
getattr(config, "outer_sensitive_layer_precision", "bf16")
)
elif (
config.fp8 is not None
and HAVE_PRIMUS_TURBO_LOCAL
Expand All @@ -566,12 +578,22 @@ def get_flux_layer_spec(
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
outer_num_start = getattr(config, "outer_sensitive_layers_start", 0) if sensitive_enabled else 0
outer_num_end = getattr(config, "outer_sensitive_layers_end", 0) if sensitive_enabled else 0
total = config.num_joint_layers + config.num_single_layers

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))
Expand Down
6 changes: 6 additions & 0 deletions primus/backends/megatron/flux_pretrain_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -515,6 +515,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
),
Expand Down Expand Up @@ -707,6 +710,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",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,29 @@ def test_validation_axes_dim_positive_values(self):
config.validate()
self.assertIn("All axes_dim values must be positive", str(cm.exception))

def test_outer_sensitive_layers_require_sensitive_routing(self):
"""An outer override is meaningless without the enclosing boundary."""
with self.assertRaisesRegex(
ValueError, "outer sensitive layers require sensitive_layers_enabled=True"
):
FluxConfig.flux_12b(outer_sensitive_layers_start=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 layer counts must be non-negative"):
FluxConfig.flux_12b(outer_sensitive_layers_end=-1)


if __name__ == "__main__":
pytest.main([__file__, "-v"])
Original file line number Diff line number Diff line change
Expand Up @@ -77,3 +77,40 @@ 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):
"""Outermost BF16 overrides the enclosing FP8 boundary."""
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.flux_12b(
transformer_impl="local",
fp4="mxfp4",
fp4_recipe="mxfp4",
mxfp4_backward_precision="mxfp4",
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",
)

block_submodules = get_flux_layer_spec(config, backend=None)
qkv_classes = [
layer_spec.submodules.self_attention.submodules.linear_qkv
for layer_spec in block_submodules.layer_specs
]

assert len(qkv_classes) == 57
assert qkv_classes[0] is ColumnParallelLinear
assert all(qkv_classes[index] is Float8ColumnParallelLinear for index in [1, 2, 3, 53, 54, 55])
assert all(qkv_classes[index] is MXFP4ColumnParallelLinear for index in range(4, 53))
assert qkv_classes[56] is ColumnParallelLinear