Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
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
132 changes: 125 additions & 7 deletions primus/backends/megatron/core/models/diffusion/common/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand All @@ -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(
Expand All @@ -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
Expand Down
137 changes: 120 additions & 17 deletions primus/backends/megatron/core/models/diffusion/flux/layer_spec.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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))
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 @@ -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
),
Expand Down Expand Up @@ -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",
Expand Down
Loading