Skip to content
Open
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
105 changes: 96 additions & 9 deletions megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,8 @@
HAVE_EMERGING_OPTIMIZERS,
_create_emerging_optimizer,
_get_qkv_split_shapes,
_localize_qkv_split_shapes,
_qkv_split_groups_are_complete,
)
from .fully_sharded_optimizer import FullyShardedOptimizer
from .grad_scaler import ConstantGradScaler, DynamicGradScaler
Expand Down Expand Up @@ -771,6 +773,8 @@ def _get_megatron_emerging_optimizer(
raise ValueError(f"Unsupported emerging optimizer: {eopt_name}")
if config.fp16:
raise ValueError('emerging optimizer with fp16 is not supported.')
if config.muon_split_qkv_per_head and not config.muon_split_qkv:
raise ValueError("muon_split_qkv_per_head requires muon_split_qkv=True")

if pg_collection is None:
pg_collection = ProcessGroupCollection.use_mpu_process_groups()
Expand All @@ -779,25 +783,108 @@ def _get_megatron_emerging_optimizer(

# Tag parameters with optimizer-specific attributes (expert_tp, is_qkv).
for model_chunk in model_chunks:
qkv_split_shapes = None
for name, param in model_chunk.named_parameters():
if not param.requires_grad:
continue
if 'experts' in name and 'shared' not in name:
param.expert_tp = True
# TODO(deyuf): support MLA
if 'linear_qkv.weight' in name and len(param.shape) == 2:
if qkv_split_shapes is None:
qkv_split_shapes = _get_qkv_split_shapes(model_chunk.config)
if param.shape[0] % sum(qkv_split_shapes) == 0:
param.is_qkv = True
param.qkv_split_shapes = qkv_split_shapes
qkv_layout = getattr(param, 'qkv_layout', None)
if (qkv_layout is not None or 'linear_qkv.weight' in name) and len(param.shape) == 2:
if qkv_layout is not None:
qkv_split_shapes = _get_qkv_split_shapes(
qkv_layout, split_qkv_per_head=config.muon_split_qkv_per_head
)
logical_split_shapes = (
qkv_split_shapes
if config.muon_split_qkv_per_head
else qkv_split_shapes * qkv_layout.num_groups
)
else:
# Backward compatibility for custom QKV modules that do not annotate
# their weight with the owning attention layer's logical layout.
qkv_split_shapes = _get_qkv_split_shapes(
model_chunk.config, split_qkv_per_head=config.muon_split_qkv_per_head
)
logical_split_shapes = (
qkv_split_shapes
if config.muon_split_qkv_per_head
else qkv_split_shapes * model_chunk.config.num_query_groups
)

tp_group = (
pg_collection.expt_tp
if getattr(param, 'expert_tp', False)
else pg_collection.tp
)
tp_size = get_pg_size(tp_group)
tp_rank = get_pg_rank(tp_group)
gtp_remat_group = (
(
pg_collection.expt_gtp_remat
if getattr(param, 'expert_tp', False)
else pg_collection.gtp_remat
)
if getattr(param, 'is_gtp_weight_remat', False)
else None
)
gtp_size = get_pg_size(gtp_remat_group)
gtp_rank = get_pg_rank(gtp_remat_group)

qkv_gtp_pad_length = (
int(getattr(param, 'pad_length', 0))
if getattr(param, 'is_gtp_weight_remat', False)
else 0
)
physical_tp_local_rows = param.shape[0] * gtp_size
if not 0 <= qkv_gtp_pad_length < physical_tp_local_rows:
raise RuntimeError(
f"Invalid Muon QKV GTP padding for {name}: "
f"pad_length={qkv_gtp_pad_length}, "
f"physical_tp_local_rows={physical_tp_local_rows}"
)
logical_tp_local_rows = physical_tp_local_rows - qkv_gtp_pad_length
expected_logical_rows = logical_tp_local_rows * tp_size
if expected_logical_rows != sum(logical_split_shapes):
log_single_rank(
logger,
logging.DEBUG,
f"Emerging optimizer QKV split skipped for {name}: "
f"shape={tuple(param.shape)}, split_shapes={qkv_split_shapes}",
f"logical_rows={sum(logical_split_shapes)}, "
f"local_rows={param.shape[0]}, tp_size={tp_size}, "
f"gtp_remat_size={gtp_size}, "
f"gtp_pad_length={qkv_gtp_pad_length}",
)
param.is_qkv = False
param.qkv_split_shapes = None
param.qkv_split_shapes_global = None
param.qkv_gtp_pad_length = 0
param.qkv_split_groups_are_complete = False
param.qkv_split_heads_are_complete = False
continue

param.is_qkv = True
param.qkv_split_shapes_global = logical_split_shapes
param.qkv_gtp_pad_length = qkv_gtp_pad_length
local_start = tp_rank * logical_tp_local_rows + gtp_rank * param.shape[0]
if config.muon_split_qkv_per_head:
if qkv_gtp_pad_length > 0:
# A padded GTP shard does not have a purely logical local layout.
# Force the per-head path to use qkv_split_shapes_global instead.
param.qkv_split_shapes = None
param.qkv_split_heads_are_complete = False
Comment thread
philipcmonk marked this conversation as resolved.
else:
param.qkv_split_shapes, param.qkv_split_heads_are_complete = (
_localize_qkv_split_shapes(
qkv_split_shapes, local_start=local_start, local_rows=param.shape[0]
)
)
else:
param.qkv_split_shapes = qkv_split_shapes
param.qkv_split_groups_are_complete = (
qkv_gtp_pad_length == 0
and _qkv_split_groups_are_complete(
qkv_split_shapes, local_start=local_start, local_rows=param.shape[0]
)
)

# Apply optimizer-specific default param overrides (e.g. muon: non-linear -> adam).
Expand Down
Loading