Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
6c00c53
fix single weight - first draft
zhongbozhu Jun 23, 2026
9b52af5
update unit test
zhongbozhu Jun 23, 2026
acc9f92
fix for gradient_accumulation_fusion
zhongbozhu Jun 24, 2026
da3368f
checks all ranks
zhongbozhu Jun 24, 2026
9035b18
increase UT coverage
zhongbozhu Jun 24, 2026
b57aeae
resolve comments
zhongbozhu Jun 24, 2026
97d09c3
resolve comments and refactor param remapping logic for better clarity
zhongbozhu Jun 24, 2026
54a9768
add mcore warning about use_transformer_engine_op_fuser and single we…
zhongbozhu Jun 24, 2026
f721718
linter
zhongbozhu Jun 24, 2026
fb8d700
another linter
zhongbozhu Jun 27, 2026
e02cc75
fix a no_grad bug in E2E traning, add repro to unit test
zhongbozhu Jun 27, 2026
35802db
fix unit test
zhongbozhu Jun 27, 2026
c1f4450
continue improve UT
zhongbozhu Jun 27, 2026
43569d3
improve UT, fix grad norm spike after eval
zhongbozhu Jun 28, 2026
6f70e0f
run UT in CI
zhongbozhu Jun 28, 2026
d2bc40e
lint
zhongbozhu Jun 28, 2026
d9f678b
include checkpointing to the unit test
zhongbozhu Jun 29, 2026
e7c21b2
reapply https://github.com/NVIDIA/Megatron-LM/pull/4994
zhongbozhu Jun 29, 2026
b15f248
fix grouped tensor remap bug, improve UT
zhongbozhu Jun 30, 2026
0f4244b
resolve comments, fix CI failed test
zhongbozhu Jul 6, 2026
2e8da4d
continue to resolve comments, fix CI test cases
zhongbozhu Jul 6, 2026
8c5db51
Use portable model setup calls in single-weight tests
zhongbozhu Jul 7, 2026
babf2b2
linter
zhongbozhu Jul 7, 2026
1053622
enforce single weight to be used with te op fuser, bug fixes
zhongbozhu Jul 9, 2026
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
18 changes: 17 additions & 1 deletion megatron/core/distributed/distributed_data_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -408,7 +408,10 @@ def disable_forward_pre_hook(self, param_sync: bool = True):

# Force synchronize parameters.
if param_sync:
self.start_param_sync(force_sync=True)
# Hook-disable paths (eval/checkpointing/shutdown) synchronize params as an
# explicit state update, not as differentiable forward compute.
with torch.no_grad():
Comment thread
kunlunl marked this conversation as resolved.
self.start_param_sync(force_sync=True)

def _make_forward_pre_hook(self):
"""
Expand Down Expand Up @@ -529,6 +532,19 @@ def start_param_sync(self, *unused, force_sync: bool = False, force_dispatch: bo
for bucket_group in self.bucket_groups + self.expert_parallel_bucket_groups:
self._start_bucket_group_param_sync(bucket_group, force_sync=force_sync)

def reset_param_sync_dispatch_state(self):
"""Mark DDP param all-gathers as not dispatched for the next forward pre-hook."""
for bucket_group in self.bucket_groups + self.expert_parallel_bucket_groups:
# A non-None handle means the previous all-gather is still in flight. Resetting only
# the dispatch flag would create the invalid state
# `param_gather_dispatched=False, param_gather_handle!=None` and could dispatch a
# second all-gather into the same parameter buffer.
assert bucket_group.param_gather_handle is None, (
"Cannot reset parameter all-gather dispatch state while an asynchronous "
"parameter all-gather is still in flight."
)
bucket_group.param_gather_dispatched = False

def start_grad_sync(self, *unused):
"""
Initiates grad sync (all-reduce or reduce-scatter) communication operations
Expand Down
222 changes: 164 additions & 58 deletions megatron/core/distributed/param_and_grad_buffer.py

Large diffs are not rendered by default.

49 changes: 48 additions & 1 deletion megatron/core/fp4_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,11 @@
import torch

from megatron.core.enums import Fp4Recipe
from megatron.core.fp8_utils import _get_custom_recipe
from megatron.core.fp8_utils import (
_get_custom_recipe,
_get_grouped_quantized_recipe,
_unwrap_parameter_data,
)
from megatron.core.transformer.transformer_config import TransformerConfig
from megatron.core.utils import is_te_min_version

Expand Down Expand Up @@ -55,6 +59,14 @@ def is_nvfp4tensor(tensor: torch.Tensor) -> bool:
return HAVE_TE_FP4_TENSOR_CLASS and isinstance(tensor, FP4_TENSOR_CLASS)


def is_grouped_nvfp4tensor(tensor: torch.Tensor) -> bool:
"""Check if a TE GroupedTensor stores NVFP4 member tensors."""
if not HAVE_TE_FP4_TENSOR_CLASS:
return False
recipe = _get_grouped_quantized_recipe(tensor)
return recipe is not None and hasattr(recipe, "nvfp4") and recipe.nvfp4()


def get_nvfp4_rowwise_packed_shape(shape: torch.Size) -> torch.Size:
"""Return packed byte shape for NVFP4 rowwise storage (last dim // 2)."""
if len(shape) == 0:
Expand Down Expand Up @@ -85,6 +97,41 @@ def modify_nvfp4_rowwise_storage(fp4_tensor: torch.Tensor, new_rowwise_data: tor
del old_rowwise


def modify_grouped_nvfp4_rowwise_storage(
grouped_tensor: torch.Tensor, new_rowwise_data: torch.Tensor
) -> None:
"""Replace grouped NVFP4 rowwise data with a new uint8 storage view.

The name intentionally mirrors `modify_nvfp4_rowwise_storage`: only the
packed rowwise byte buffer is remapped into the DDP buffer. The grouped
scale, amax, and columnwise buffers remain owned by the original tensor.
"""
tensor = _unwrap_parameter_data(grouped_tensor)
if not is_grouped_nvfp4tensor(tensor):
raise ValueError("modify_grouped_nvfp4_rowwise_storage expects grouped NVFP4 storage")

old_rowwise = getattr(tensor, "rowwise_data", None)
if old_rowwise is None:
raise RuntimeError("Grouped NVFP4 tensor is missing rowwise data to replace")

new_rowwise_data = new_rowwise_data.view(-1)
if old_rowwise.numel() != new_rowwise_data.numel():
raise ValueError(
"Grouped NVFP4 rowwise storage size mismatch: "
f"old numel={old_rowwise.numel()}, new numel={new_rowwise_data.numel()}"
)
assert (
old_rowwise.dtype == new_rowwise_data.dtype == torch.uint8
), "Grouped NVFP4 rowwise storage must be uint8"

new_rowwise_data.detach().copy_(old_rowwise.view(-1))
tensor.rowwise_data = new_rowwise_data
# Member views capture data pointers. Refresh them after swapping rowwise storage while
# preserving the existing scale/amax/columnwise grouped buffers.
tensor.quantized_tensors = tensor.split_into_quantized_tensors()
del old_rowwise
Comment thread
kunlunl marked this conversation as resolved.


def quantize_nvfp4_param_shard(
model_params, main_params, start_offsets, data_parallel_group, fsdp_shard_model_params=None
):
Expand Down
166 changes: 163 additions & 3 deletions megatron/core/fp8_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,14 @@
# MXFP8Tensor not found
HAVE_TE_MXFP8TENSOR = False

try:
from transformer_engine.pytorch.tensor.grouped_tensor import GroupedTensor

HAVE_TE_GROUPED_TENSOR_CLASS = True
except (ImportError, ModuleNotFoundError):
GroupedTensor = None
HAVE_TE_GROUPED_TENSOR_CLASS = False

if HAVE_TE:
from megatron.core.extensions.transformer_engine import (
TEColumnParallelLinear,
Expand Down Expand Up @@ -93,6 +101,24 @@
te_post_all_gather_processing = None


def _unwrap_parameter_data(tensor: torch.Tensor) -> torch.Tensor:
"""Return underlying tensor data when PyTorch wraps a tensor subclass as a Parameter."""
if HAVE_TE_GROUPED_TENSOR_CLASS and isinstance(tensor, GroupedTensor):
# TE GroupedTensor stores its real payload in Python-side metadata fields
# such as rowwise_data/scale_inv. PyTorch marks tensor-subclass parameters
# as Parameters, so tensor.data would create a detached wrapper copy. Return
# the live wrapper so storage metadata mutations update the module parameter.
return tensor
return tensor.data if isinstance(tensor, torch.nn.Parameter) else tensor


def _is_instance_or_param_data(tensor: torch.Tensor, tensor_class: type) -> bool:
"""Check a tensor subclass, including when wrapped by torch.nn.Parameter."""
return isinstance(tensor, tensor_class) or isinstance(
_unwrap_parameter_data(tensor), tensor_class
)


def is_float8tensor(tensor: torch.Tensor) -> bool:
"""Check if a tensor is a Transformer Engine Float8Tensor.

Expand All @@ -102,12 +128,136 @@ def is_float8tensor(tensor: torch.Tensor) -> bool:
are both inherited from QuantizedTensor. So, for TE1.x, FP8_TENSOR_CLASS is Float8Tensor,
and for TE2.x, FP8_TENSOR_CLASS is QuantizedTensor.
"""
return HAVE_TE_FP8_TENSOR_CLASS and isinstance(tensor, FP8_TENSOR_CLASS)
return HAVE_TE_FP8_TENSOR_CLASS and _is_instance_or_param_data(tensor, FP8_TENSOR_CLASS)


def is_mxfp8tensor(tensor: torch.Tensor) -> bool:
"""Check if a tensor is a Transformer Engine MXFP8Tensor"""
return HAVE_TE_MXFP8TENSOR and isinstance(tensor, MXFP8Tensor)
return HAVE_TE_MXFP8TENSOR and _is_instance_or_param_data(tensor, MXFP8Tensor)


def is_grouped_tensor(tensor: torch.Tensor) -> bool:
"""Check if a tensor is a Transformer Engine GroupedTensor."""
return HAVE_TE_GROUPED_TENSOR_CLASS and _is_instance_or_param_data(tensor, GroupedTensor)


def is_grouped_tensor_with_quantized_storage(tensor: torch.Tensor) -> bool:
"""Check if a Transformer Engine GroupedTensor owns quantized primary storage."""
tensor = _unwrap_parameter_data(tensor)
if not is_grouped_tensor(tensor):
return False
rowwise_data = getattr(tensor, "rowwise_data", None)
return rowwise_data is not None and rowwise_data.dtype == torch.uint8


def _get_grouped_quantized_recipe(tensor: torch.Tensor):
"""Return TE recipe for grouped quantized storage, or None if unavailable."""
tensor = _unwrap_parameter_data(tensor)
if not is_grouped_tensor_with_quantized_storage(tensor):
return None

quantizer = getattr(tensor, "quantizer", None)
if quantizer is None or not hasattr(quantizer, "_get_compatible_recipe"):
return None
return quantizer._get_compatible_recipe()


def is_grouped_mxfp8tensor(tensor: torch.Tensor) -> bool:
"""Check if a TE GroupedTensor stores MXFP8 member tensors."""
if not HAVE_TE_MXFP8TENSOR:
return False
recipe = _get_grouped_quantized_recipe(tensor)
return recipe is not None and hasattr(recipe, "mxfp8") and recipe.mxfp8()


def get_grouped_quantized_members(
tensor: torch.Tensor, *, create_if_missing: bool = False
) -> List[torch.Tensor]:
"""Return cached per-member views for a grouped quantized tensor."""
grouped_tensor = _unwrap_parameter_data(tensor)
if not is_grouped_tensor_with_quantized_storage(grouped_tensor):
raise ValueError("get_grouped_quantized_members expects grouped quantized storage.")

quantized_members = getattr(grouped_tensor, "quantized_tensors", None)
if quantized_members is None:
if not create_if_missing:
raise RuntimeError(
"Grouped quantized parameter is missing cached member tensors. "
"Create them outside the training critical path."
)
quantized_members = grouped_tensor.split_into_quantized_tensors()
grouped_tensor.quantized_tensors = quantized_members
return quantized_members


def copy_tensor_to_quantized_param(param: torch.Tensor, src: torch.Tensor) -> None:
"""Copy high-precision values into TE quantized parameter storage."""
dst = _unwrap_parameter_data(param)

if is_grouped_tensor_with_quantized_storage(dst):
if src.numel() != dst.numel():
raise ValueError(
"Grouped quantized parameter copy size mismatch: "
f"src numel={src.numel()}, dst numel={dst.numel()}"
)
if not dst.all_same_shape():
raise NotImplementedError(
"Copying into grouped quantized parameters requires uniform member shapes."
)

# Grouped quantized tensors cannot use GroupedTensor.copy_ here because
# the generic grouped path can rebuild member tensors through
# split_into_quantized_tensors(), which is not graph safe. Update cached
# member tensors in place instead.
quantized_members = get_grouped_quantized_members(dst)
src_members = src.view(dst.shape).unbind(dim=0)
if len(src_members) != len(quantized_members):
raise RuntimeError(
"Grouped quantized parameter member count mismatch: "
f"src members={len(src_members)}, dst members={len(quantized_members)}"
)

for src_member, dst_member in zip(src_members, quantized_members):
dst.quantizer.update_quantized(src_member, dst_member)
return

# Plain TE quantized tensors override copy_ to requantize into their
# backing storage.
dst.copy_(src.view(dst.shape))
Comment thread
kunlunl marked this conversation as resolved.


def modify_grouped_tensor_rowwise_storage(tensor: torch.Tensor, new_storage: torch.Tensor) -> None:
"""Replace a high-precision Transformer Engine GroupedTensor's rowwise storage."""
tensor = _unwrap_parameter_data(tensor)
Comment thread
kunlunl marked this conversation as resolved.
if not is_grouped_tensor(tensor):
raise ValueError("modify_grouped_tensor_rowwise_storage expects a GroupedTensor.")
if is_grouped_tensor_with_quantized_storage(tensor):
raise ValueError(
"modify_grouped_tensor_rowwise_storage only supports high-precision GroupedTensor "
"storage. Quantized grouped storage also owns scale buffers."
)

old_rowwise_data = getattr(tensor, "rowwise_data", None)
if old_rowwise_data is None:
raise RuntimeError("GroupedTensor is missing rowwise_data.")

new_storage = new_storage.view(-1)
if old_rowwise_data.numel() != new_storage.numel():
raise ValueError(
"GroupedTensor backing storage size mismatch: "
f"old numel={old_rowwise_data.numel()}, new numel={new_storage.numel()}"
)
if old_rowwise_data.dtype != new_storage.dtype:
raise ValueError(
"GroupedTensor backing storage dtype mismatch: "
f"old dtype={old_rowwise_data.dtype}, new dtype={new_storage.dtype}"
)

new_storage.detach().copy_(old_rowwise_data)
tensor.rowwise_data = new_storage
tensor.columnwise_data = None
tensor.quantized_tensors = None
del old_rowwise_data


def dequantize_fp8_tensor(fp8_tensor: torch.Tensor) -> torch.Tensor:
Expand Down Expand Up @@ -502,8 +652,18 @@ def post_all_gather_processing(model_params):
- tensorwise: may need to create a transposed view to match backend GEMM.
- blockwise: create column-wise storage.
"""
if not isinstance(model_params, list):
model_params = [model_params]

expanded_model_params = []
for param in model_params:
if is_grouped_tensor_with_quantized_storage(param):
expanded_model_params.extend(get_grouped_quantized_members(param))
else:
expanded_model_params.append(param)

if te_post_all_gather_processing is not None:
te_post_all_gather_processing(model_params)
te_post_all_gather_processing(expanded_model_params)
else:
# If the TE version is old and does not have post_all_gather_processing function, this is
# a no-op, and the transpose/columnwise data will be created in the next forward pass.
Expand Down
Loading
Loading