diff --git a/docs/design-docs/model-express-refit.md b/docs/design-docs/model-express-refit.md
new file mode 100644
index 0000000000..4153613c0c
--- /dev/null
+++ b/docs/design-docs/model-express-refit.md
@@ -0,0 +1,134 @@
+# ModelExpress Weight Refit
+
+## Problem
+
+NeMo RL periodically updates generation workers with policy weights produced by
+trainer workers. A non-colocated deployment needs to move those weights across
+different process groups and may use different parallel layouts for training
+and generation.
+
+ModelExpress (MX) provides source discovery and topology-aware transfer
+planning for this exchange. The NVIDIA Inference Xfer Library (NIXL) can move
+the resulting byte ranges directly between worker memory. The MX server stores
+metadata and lifecycle state; it does not carry model weights.
+
+## Recommended design
+
+Each trainer rank publishes the tensor shards it already owns. Generation
+workers discover a complete source set for one policy version, request the
+ranges required by their local layout, and install the received tensors through
+an inference-engine adapter.
+
+```mermaid
+flowchart TB
+ T["**1. Trainer ranks**
publish local tensor ownership"]
+ C["**2. MX catalog**
stores version and topology metadata"]
+ G["**3. Generation workers**
discover sources and build local plans"]
+ X["**4. NIXL data plane**
moves bytes directly between workers"]
+ I["**5. Inference adapter**
transforms and installs local weights"]
+
+ T --> C
+ C --> G
+ T --> X
+ G --> X
+ X --> I
+
+ classDef trainer fill:#C8E6C9,stroke:#2E7D32,stroke-width:2px,color:#000
+ classDef catalog fill:#BBDEFB,stroke:#1565C0,stroke-width:2px,color:#000
+ classDef target fill:#D1C4E9,stroke:#5E35B1,stroke-width:2px,color:#000
+ class T trainer
+ class C,X catalog
+ class G,I target
+```
+
+NeMo RL remains responsible for selecting the policy version, invoking the
+weight synchronizer, and deciding when the generation fleet is ready. MX owns
+source metadata, transfer planning, and the transport/installer boundary.
+
+## Trainer publication
+
+The trainer publisher describes:
+
+- model version and numerical format;
+- Tensor Parallelism (TP), Pipeline Parallelism (PP), and Expert Parallelism
+ (EP) coordinates;
+- each tensor's global shape and locally owned range;
+- expert ownership where applicable; and
+- the Megatron-to-Hugging Face name mapping required by receiver-side
+ translation.
+
+DTensor publication uses `DTensor.to_local()` and records its single sharded
+axis. Partial placements and tensors sharded across multiple mesh axes fail
+before publication. Megatron publication classifies native fused, replicated,
+and expert parameters without gathering a full model on rank zero.
+
+## NeMo RL integration boundary
+
+The shared `WeightSynchronizer` abstraction owns the complete refit lifecycle.
+A future `ModelExpressWeightSynchronizer` will call
+`publish_weights_for_model_express()` on the policy, coordinate generation-side
+discovery and apply, and report completion through the same interface used by
+the existing synchronizers.
+
+This change adds only the policy publication operation. It does not add a
+second synchronization lifecycle to algorithm code.
+
+## Current status
+
+**Partially aligned:** trainer-side rank-local publication and metadata
+construction are implemented for the existing DTensor worker and the Megatron
+worker. Focused unit tests cover DTensor shard ranges, unsupported placements,
+replicated-tensor ownership, fused QKV classification, and global expert IDs.
+
+Generation-side discovery, transfer, translation, installation, and
+end-to-end GPU validation are follow-up work. The AutoModel-based
+`DTensorPolicyWorkerV2` does not yet implement MX publication. Existing NeMo RL
+weight synchronizers remain unchanged.
+
+## Assumptions
+
+- Published tensor storage remains valid until the corresponding update
+ completes.
+- Trainer and generation adapters agree on global tensor names or provide an
+ explicit translation map.
+- The MX publisher exposes a public `reset_tensors()` lifecycle method.
+- A selected generation backend supports the requested source and target
+ layouts.
+
+## Tradeoffs
+
+- Rank-local publication avoids a trainer-side full-model gather, but requires
+ explicit ownership metadata.
+- Receiver-side planning supports different trainer and generation layouts,
+ but adds a metadata and planning stage before transfer.
+- Lazy ModelExpress imports keep the standard NeMo RL installation independent
+ of MX, but configuration errors appear when the MX backend is initialized.
+
+## Failure modes
+
+- Unsupported DTensor placement fails before metadata publication.
+- Missing ModelExpress APIs fail when the trainer publisher first initializes.
+- Incomplete source coverage must fail planning; a generation worker must not
+ install a partial model.
+- A failed trainer or generation rank causes the global NeMo RL weight update
+ to fail rather than advancing only part of the fleet.
+
+## Open questions
+
+- Final user-facing configuration for `ModelExpressWeightSynchronizer`.
+- Rank-local publication for `DTensorPolicyWorkerV2`.
+- The common install-plan contract between MX reshard planning and the
+ inference adapter.
+- Retention and read-lease behavior for trainer buffers during long updates.
+- Supported behavior for topology changes between policy versions.
+
+## Implementation references
+
+- Trainer helpers: `nemo_rl/distributed/mx_helpers.py`
+- Megatron tensor classification:
+ `nemo_rl/distributed/mx_megatron_helpers.py`
+- Policy interface: `nemo_rl/models/policy/interfaces.py`
+- DTensor worker:
+ `nemo_rl/models/policy/workers/dtensor_policy_worker.py`
+- Megatron worker:
+ `nemo_rl/models/policy/workers/megatron_policy_worker.py`
diff --git a/docs/index.md b/docs/index.md
index f57dc28510..2cc12a97e1 100644
--- a/docs/index.md
+++ b/docs/index.md
@@ -322,6 +322,7 @@ design-docs/uv.md
design-docs/dependency-management.md
design-docs/chat-datasets.md
design-docs/generation.md
+design-docs/model-express-refit.md
design-docs/checkpointing.md
design-docs/loss-functions.md
design-docs/fsdp2-parallel-plan.md
diff --git a/nemo_rl/distributed/mx_helpers.py b/nemo_rl/distributed/mx_helpers.py
new file mode 100644
index 0000000000..81fe73ac19
--- /dev/null
+++ b/nemo_rl/distributed/mx_helpers.py
@@ -0,0 +1,236 @@
+# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Trainer-side helpers for rank-local ModelExpress publication."""
+
+from __future__ import annotations
+
+import logging
+import os
+from collections.abc import Mapping
+from dataclasses import dataclass
+from typing import TYPE_CHECKING, Any, Protocol
+
+if TYPE_CHECKING:
+ import torch
+
+logger = logging.getLogger("nemo_rl.distributed.mx_helpers")
+
+
+@dataclass(frozen=True)
+class ModelExpressPublisherOptions:
+ """Internal settings supplied by a ModelExpress weight synchronizer.
+
+ Args:
+ mx_server_url: gRPC URL of the MX server.
+ nic_pin: NIC pinning strategy passed to ``pin_local_nic``:
+ ``"auto"`` | ``"off"`` | concrete ``"mlx5_"``.
+ megatron_role_overrides: Parameter-name substrings mapped to explicit
+ ModelExpress Megatron roles.
+
+ This is not a user-facing configuration schema. The future
+ ``ModelExpressWeightSynchronizer`` will validate user configuration and
+ construct these settings explicitly.
+ """
+
+ mx_server_url: str
+ nic_pin: str
+ megatron_role_overrides: Mapping[str, str]
+
+
+class ModelExpressPublisher(Protocol):
+ """Publisher operations used by NeMo RL trainer workers."""
+
+ def initialize(self, *, model_name: str, dtype: str) -> None: ...
+
+ def reset_tensors(self) -> None: ...
+
+ def set_megatron_sidecar(self, sidecar: dict[str, Any]) -> None: ...
+
+ def set_megatron_mesh_position(
+ self, *, tp_rank: int, pp_rank: int, ep_rank: int
+ ) -> None: ...
+
+ def add_tensor(
+ self,
+ *,
+ name: str,
+ tensor: "torch.Tensor",
+ is_expert: bool = False,
+ expert_axis: int = 0,
+ owned_expert_ids: tuple[int, ...] | set[int] | list[int] = (),
+ megatron_role: str | None = None,
+ megatron_extras: dict[str, str] | None = None,
+ shard_spec: Any | None = None,
+ ) -> None: ...
+
+ def publish(self, *, version: int) -> str: ...
+
+ def mark_ready(self) -> bool: ...
+
+
+@dataclass(frozen=True)
+class DTensorShardSpec:
+ """ModelExpress-compatible metadata for one materialized DTensor shard."""
+
+ global_shape: tuple[int, ...]
+ shard_axis: int
+ local_shard_range: tuple[int, int]
+
+
+def get_dtensor_local_shard(
+ tensor: Any,
+) -> tuple["torch.Tensor", DTensorShardSpec | None]:
+ """Materialize a DTensor's local buffer and describe its single shard axis.
+
+ ModelExpress currently represents one shard axis per tensor. Replicated
+ DTensors are supported, while partial or multi-axis sharding fails before
+ publication rather than falling back to an all-gather.
+ """
+ from torch.distributed.tensor import Partial, Replicate, Shard
+
+ local = tensor.to_local()
+ placements = tuple(tensor.placements)
+ if any(isinstance(placement, Partial) for placement in placements):
+ raise NotImplementedError("ModelExpress does not support partial DTensors")
+
+ sharded_mesh_dims = [
+ (mesh_dim, placement)
+ for mesh_dim, placement in enumerate(placements)
+ if isinstance(placement, Shard)
+ ]
+ if not sharded_mesh_dims:
+ if not all(isinstance(placement, Replicate) for placement in placements):
+ raise NotImplementedError(f"unsupported DTensor placements: {placements!r}")
+ return local, None
+ if len(sharded_mesh_dims) != 1:
+ raise NotImplementedError(
+ "ModelExpress supports one DTensor shard axis per tensor; "
+ f"got placements={placements!r}"
+ )
+
+ mesh_dim, placement = sharded_mesh_dims[0]
+ coordinate = tensor.device_mesh.get_coordinate()
+ if coordinate is None:
+ raise RuntimeError("current rank is not part of the DTensor device mesh")
+
+ axis = int(placement.dim)
+ global_shape = tuple(int(size) for size in tensor.shape)
+ world_size = int(tensor.device_mesh.size(mesh_dim))
+ rank = int(coordinate[mesh_dim])
+ chunk_size = (global_shape[axis] + world_size - 1) // world_size
+ start = min(rank * chunk_size, global_shape[axis])
+ local_extent = int(local.shape[axis])
+ end = start + local_extent
+ if end > global_shape[axis]:
+ raise ValueError(
+ f"local shard range ({start}, {end}) exceeds global axis "
+ f"size {global_shape[axis]}"
+ )
+
+ return local, DTensorShardSpec(
+ global_shape=global_shape,
+ shard_axis=axis,
+ local_shard_range=(start, end),
+ )
+
+
+def pin_local_nic(*, device_id: int, mode: str) -> None:
+ """Configure the requested NIC policy before NIXL initializes.
+
+ Automatic mode delegates topology selection to ModelExpress. A concrete
+ device name sets the UCX interface explicitly.
+ """
+ if mode == "off":
+ return
+
+ # Imported only when MX is selected so normal NeMo RL imports do not
+ # require the external modelexpress package.
+ from modelexpress.ucx_utils import apply_nic_pin_for_device
+
+ if mode == "auto":
+ apply_nic_pin_for_device(device_id=device_id)
+ logger.info("pinned NIC for device %d (auto)", device_id)
+ else:
+ os.environ["UCX_NET_DEVICES"] = mode
+ os.environ["MX_RDMA_NIC_PIN"] = "off"
+ logger.info("pinned NIC explicitly: %s", mode)
+
+
+def build_v2_publisher(
+ *,
+ rank: int,
+ device_id: int,
+ fsdp_world_size: int,
+ tp_world_size: int,
+ pp_world_size: int,
+ ep_world_size: int,
+ publisher_options: ModelExpressPublisherOptions,
+ agent_name: str | None = None,
+) -> ModelExpressPublisher:
+ """Construct a :class:`MxV2TrainingPublisher` and pin its NIC.
+
+ Returns a :class:`modelexpress.MxV2TrainingPublisher`. Caller must invoke
+ ``initialize(model_name=...)``, then ``add_tensor`` per tensor, then
+ ``publish(version=...)``, then ``mark_ready()``.
+ """
+ # Imported only when MX is selected so normal NeMo RL imports do not
+ # require the external modelexpress package.
+ from modelexpress import MxV2TrainingPublisher, TrainerWorldLayout
+
+ pin_local_nic(device_id=device_id, mode=publisher_options.nic_pin)
+
+ return MxV2TrainingPublisher(
+ agent_name=agent_name or f"nemo-rl-trainer-r{rank}",
+ device_id=device_id,
+ mx_server_url=publisher_options.mx_server_url,
+ worker_rank=rank,
+ world_layout=TrainerWorldLayout(
+ fsdp_world_size=fsdp_world_size,
+ tp_world_size=tp_world_size,
+ pp_world_size=pp_world_size,
+ ep_world_size=ep_world_size,
+ ),
+ heartbeat=True,
+ )
+
+
+def start_model_express_publication(publisher: ModelExpressPublisher) -> None:
+ """Clear per-version tensor registrations before adding the next version."""
+ publisher.reset_tensors()
+
+
+def finish_model_express_publication(
+ publisher: ModelExpressPublisher,
+ *,
+ version: int,
+ worker_rank: int,
+) -> str:
+ """Publish one complete version and require a successful READY transition."""
+ source_id = publisher.publish(version=version)
+ if not publisher.mark_ready():
+ raise RuntimeError(
+ f"ModelExpress failed to mark trainer rank {worker_rank} ready"
+ )
+ return source_id
+
+
+__all__ = [
+ "ModelExpressPublisher",
+ "ModelExpressPublisherOptions",
+ "build_v2_publisher",
+ "finish_model_express_publication",
+ "get_dtensor_local_shard",
+ "pin_local_nic",
+ "start_model_express_publication",
+]
diff --git a/nemo_rl/distributed/mx_megatron_helpers.py b/nemo_rl/distributed/mx_megatron_helpers.py
new file mode 100644
index 0000000000..2efc101a6c
--- /dev/null
+++ b/nemo_rl/distributed/mx_megatron_helpers.py
@@ -0,0 +1,478 @@
+# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+"""Classify Megatron-Core parameters for rank-local ModelExpress publication."""
+
+from __future__ import annotations
+
+import os
+from dataclasses import dataclass, field
+from typing import TYPE_CHECKING, Any, Iterator
+
+if TYPE_CHECKING:
+ import torch
+
+# Values consumed by ModelExpress's Megatron slice planner.
+ROLE_QKV_COLUMN = "qkv_column"
+ROLE_GATED_MLP_COLUMN = "gated_mlp_column"
+ROLE_COLUMN = "column"
+ROLE_ROW = "row"
+ROLE_VOCAB_PARALLEL = "vocab_parallel"
+ROLE_REPLICATED = "replicated"
+ROLE_EXPERT_COLUMN = "expert_column"
+ROLE_EXPERT_ROW = "expert_row"
+
+
+@dataclass
+class MegatronRoleSpec:
+ """Per-parameter classification result.
+
+ ``role`` is one of the role string constants. ``descriptor_extras`` is
+ the per-tensor ``extra_parameters`` payload the publisher will merge
+ into MX's ``identity.extra_parameters`` (and the v2 sidecar JSON).
+ Keys here MUST match the names ``modelexpress.nemo_rl_v2._extract_megatron_meta``
+ reads.
+ """
+
+ role: str
+ descriptor_extras: dict[str, str] = field(default_factory=dict)
+ is_expert: bool = False
+ expert_axis: int = 0
+ owned_expert_ids: set[int] = field(default_factory=set)
+
+
+# Heuristic name patterns for fused-QKV and fused-gate+up linears in
+# mainline Megatron-Core. Publisher options can override these roles for
+# deployments with different parameter names.
+_DEFAULT_FUSED_QKV_NAME_PATTERNS = ("linear_qkv", "qkv_proj", "fused_qkv")
+_DEFAULT_FUSED_GATED_MLP_PATTERNS = ("linear_fc1", "gate_up_proj")
+# Vocab / embedding name pattern.
+_DEFAULT_VOCAB_NAME_PATTERNS = (
+ "word_embeddings",
+ "embedding",
+ "lm_head",
+ "output_layer",
+)
+
+
+def _bridge_module_type_registry() -> dict[str, set[str]] | None:
+ """Return Bridge's authoritative module classifier registry, or None.
+
+ Bridge ships a curated dict of
+ ``{"column": {classes...}, "row": {...}, "replicated": {...}}`` covering
+ every TE / Inference / Quant variant. Importing it lazily avoids a hard
+ dependency: when Bridge is not in the import path (e.g. in unit tests on
+ a CPU-only env), the caller falls back to substring matching against
+ the base class names, which is correct for mainline Megatron-Core.
+ """
+ try:
+ from megatron.bridge.models.conversion.param_mapping import (
+ AutoMapping as _AM,
+ )
+
+ return dict(_AM._MODULE_TYPE_REGISTRY)
+ except (ImportError, AttributeError):
+ return None
+
+
+def _classify_module(module: "torch.nn.Module | None") -> str | None:
+ """Map a module to a Megatron-Bridge parallelism kind.
+
+ Returns one of ``"column"``, ``"row"``, ``"replicated"``, or ``None``
+ if no verified rule identifies its placement.
+ """
+ mod_class_name = _module_class_name(module)
+ if not mod_class_name:
+ return None
+ registry = _bridge_module_type_registry()
+ if registry is not None:
+ # Direct hit on Bridge's curated set (catches every TE / Inference /
+ # Quant variant by exact class name).
+ for kind, cls_set in registry.items():
+ if mod_class_name in cls_set:
+ return kind
+ # Bridge also has a special-case for the TE-fused
+ # LayerNormColumnParallelLinear: classify as column.
+ if "LayerNormColumnParallelLinear" in mod_class_name:
+ return "column"
+ # Fallback: substring match against the base names.
+ if "ColumnParallel" in mod_class_name or "VocabParallelEmbedding" in mod_class_name:
+ return "column"
+ if "RowParallel" in mod_class_name:
+ return "row"
+ if module is not None:
+ tensor_model_parallel = getattr(module, "tensor_model_parallel", None)
+ if tensor_model_parallel is False:
+ return "replicated"
+ if tensor_model_parallel is True:
+ partition_dim = getattr(module, "partition_dim", None)
+ if partition_dim == 0:
+ return "column"
+ if partition_dim == 1:
+ return "row"
+ if mod_class_name == "TELinear":
+ parallel_mode = getattr(module, "parallel_mode", None)
+ if parallel_mode in ("column", "row"):
+ return parallel_mode
+ if parallel_mode is None:
+ return "replicated"
+ if any(
+ needle in mod_class_name
+ for needle in (
+ "Norm",
+ "RMSNorm",
+ "L2Norm",
+ "TopKRouter",
+ "LinearForLastLayer",
+ "IdentityOp",
+ )
+ ):
+ return "replicated"
+ return None
+
+
+_PARAM_LEAF_NAMES = {"weight", "bias", "scale", "_extra_state"}
+
+
+def _is_param_leaf(name_part: str) -> bool:
+ """Return True for any trailing name that's a parameter rather than a child module.
+
+ Includes the standard ``weight``/``bias``/``scale``/``_extra_state``
+ and the grouped-MoE per-expert convention ``weight0``, ``weight1``,
+ ``weight127``, ``bias0``, etc. Megatron-Core's TE-grouped linears
+ expose one ``weight`` ``nn.Parameter`` per local expert.
+ """
+ if name_part in _PARAM_LEAF_NAMES:
+ return True
+ for base in ("weight", "bias", "scale"):
+ if name_part.startswith(base):
+ suffix = name_part[len(base) :]
+ if suffix and suffix.isdigit():
+ return True
+ return False
+
+
+def _expert_index_from_param(name_part: str) -> int | None:
+ """If ``name_part`` is ``weight``/``bias``/etc, return ``N``."""
+ for base in ("weight", "bias", "scale"):
+ if name_part.startswith(base):
+ suffix = name_part[len(base) :]
+ if suffix and suffix.isdigit():
+ return int(suffix)
+ return None
+
+
+def _expert_index_from_path(name: str) -> int | None:
+ """Return the local expert index from a module path when present."""
+ parts = name.split(".")
+ for marker in ("local_experts", "experts"):
+ for index, part in enumerate(parts[:-1]):
+ if part == marker and parts[index + 1].isdigit():
+ return int(parts[index + 1])
+ return None
+
+
+def _enclosing_module(name: str, model: "torch.nn.Module") -> "torch.nn.Module | None":
+ """Walk down model attributes to find the module that owns ``name``.
+
+ ``name`` is a parameter name like
+ ``decoder.layers.0.self_attention.linear_qkv.weight`` or
+ ``decoder.layers.0.mlp.experts.linear_fc1.weight0`` for grouped-MoE
+ per-expert parameters. Return the parent module of the final
+ parameter token.
+ """
+ parts = name.split(".")
+ if not parts or not _is_param_leaf(parts[-1]):
+ # Fall back to the deepest module — caller will get a leaf.
+ cur = model
+ for p in parts:
+ sub = getattr(cur, p, None)
+ if sub is None:
+ return None
+ cur = sub
+ return cur
+ cur: Any = model
+ for p in parts[:-1]:
+ sub = getattr(cur, p, None)
+ if sub is None:
+ return None
+ cur = sub
+ return cur
+
+
+def _module_class_name(mod: "torch.nn.Module | None") -> str:
+ if mod is None:
+ return ""
+ return type(mod).__name__
+
+
+def _is_fused_qkv_name(name: str) -> bool:
+ return any(p in name for p in _DEFAULT_FUSED_QKV_NAME_PATTERNS)
+
+
+def _is_fused_gated_mlp_name(name: str) -> bool:
+ return any(p in name for p in _DEFAULT_FUSED_GATED_MLP_PATTERNS)
+
+
+def _is_vocab_name(name: str) -> bool:
+ return any(p in name for p in _DEFAULT_VOCAB_NAME_PATTERNS)
+
+
+def _is_expert_name(name: str, *, expert_pattern: str) -> bool:
+ return expert_pattern in name
+
+
+def detect_megatron_role(
+ name: str,
+ param: "torch.Tensor",
+ *,
+ model: "torch.nn.Module",
+ tp_size: int,
+ ep_size: int,
+ ep_rank: int,
+ num_local_experts: int | None = None,
+ num_attention_heads: int | None = None,
+ num_kv_heads: int | None = None,
+ head_dim: int | None = None,
+ expert_pattern: str | None = None,
+ role_overrides: dict[str, str] | None = None,
+) -> MegatronRoleSpec:
+ """Classify a Megatron parameter into one of seven roles.
+
+ Returns the role and per-tensor metadata that the publisher attaches to
+ source metadata. Unknown placement at TP greater than one fails before
+ publication rather than silently treating a sharded tensor as replicated.
+
+ Args:
+ name: param name from ``model.named_parameters()`` (e.g.
+ ``decoder.layers.0.self_attention.linear_qkv.weight``).
+ param: the local shard tensor (Megatron stores native shards).
+ model: the root model module; used to walk attributes for the
+ enclosing module's class.
+ tp_size, ep_size, ep_rank: from ``parallel_state``.
+ num_attention_heads, num_kv_heads, head_dim: required for
+ ``qkv_column`` role; derived from the model config. Pass
+ ``None`` if unknown — the role still classifies but the
+ descriptor will be missing fields and the receiver will
+ fall back to its default un-interleave assumptions.
+ expert_pattern: Substring marker for MoE expert tensors.
+ role_overrides: optional ``{param_name_substring: role}`` dict
+ for forcing a role on a specific tensor (escape hatch for
+ non-mainline Megatron forks).
+ """
+ expert_pattern = expert_pattern or os.environ.get(
+ "NRL_MX_EXPERT_TENSOR_PATTERN", "experts"
+ )
+
+ # ---- 1. Explicit override wins. ----
+ if role_overrides:
+ for needle, role in role_overrides.items():
+ if needle in name:
+ return MegatronRoleSpec(role=role)
+
+ # Grouped-MoE parameters use either a ``weight`` leaf or an expert
+ # index in the module path. Both indices are local to the EP rank.
+ if _is_expert_name(name, expert_pattern=expert_pattern):
+ leaf = name.rsplit(".", 1)[-1] if "." in name else name
+ expert_idx = _expert_index_from_param(leaf)
+ if expert_idx is None:
+ expert_idx = _expert_index_from_path(name)
+ if expert_idx is not None:
+ global_idx = expert_idx
+ if num_local_experts:
+ global_idx = ep_rank * int(num_local_experts) + expert_idx
+ mod_class = _module_class_name(_enclosing_module(name, model))
+ sub_role = (
+ ROLE_EXPERT_ROW if "RowParallel" in mod_class else ROLE_EXPERT_COLUMN
+ )
+ return MegatronRoleSpec(
+ role=sub_role,
+ is_expert=True,
+ expert_axis=0,
+ owned_expert_ids={global_idx},
+ descriptor_extras={
+ "expert_axis": "0",
+ "expert_id": str(global_idx),
+ "local_expert_id": str(expert_idx),
+ "expert_layout": "grouped",
+ },
+ )
+
+ # Some grouped linears store all local experts in one leading-axis tensor.
+ # Only classify this layout when the module class and configured local
+ # expert count agree; ordinary expert linears also have a leading output
+ # dimension and must not be mistaken for grouped storage.
+ if (
+ _is_expert_name(name, expert_pattern=expert_pattern)
+ and ep_size > 1
+ and num_local_experts is not None
+ and param.ndim >= 2
+ ):
+ mod_class = _module_class_name(_enclosing_module(name, model))
+ if "Grouped" in mod_class and param.shape[0] == num_local_experts:
+ first_expert = ep_rank * num_local_experts
+ owned = set(range(first_expert, first_expert + num_local_experts))
+ sub_role = ROLE_EXPERT_COLUMN
+ if "RowParallel" in mod_class:
+ sub_role = ROLE_EXPERT_ROW
+ return MegatronRoleSpec(
+ role=sub_role,
+ is_expert=True,
+ expert_axis=0,
+ owned_expert_ids=owned,
+ descriptor_extras={
+ "expert_axis": "0",
+ "expert_layout": "leading_axis",
+ },
+ )
+
+ # ---- 3. Walk to the enclosing module + classify against Bridge's
+ # AutoMapping._MODULE_TYPE_REGISTRY (or fall back to substring match). ----
+ mod = _enclosing_module(name, model)
+ mod_class = _module_class_name(mod)
+ parallelism = _classify_module(mod)
+
+ # ---- 4. VocabParallelEmbedding / lm_head sharded along rows. ----
+ if mod_class == "VocabParallelEmbedding" or (
+ _is_vocab_name(name)
+ and tp_size > 1
+ and param.ndim >= 2
+ and parallelism == "column"
+ ):
+ return MegatronRoleSpec(role=ROLE_VOCAB_PARALLEL)
+
+ # ---- 5. Column-parallel linears (incl. all TE / Inference / Quant variants). ----
+ if parallelism == "column":
+ if _is_fused_qkv_name(name):
+ extras: dict[str, str] = {"qkv_interleave": "by_head"}
+ if num_attention_heads is not None and tp_size > 0:
+ extras["num_heads_local"] = str(num_attention_heads // tp_size)
+ if num_kv_heads is not None and tp_size > 0:
+ extras["num_kv_heads_local"] = str(num_kv_heads // tp_size)
+ if head_dim is not None:
+ extras["head_dim"] = str(head_dim)
+ return MegatronRoleSpec(role=ROLE_QKV_COLUMN, descriptor_extras=extras)
+ if _is_fused_gated_mlp_name(name):
+ return MegatronRoleSpec(
+ role=ROLE_GATED_MLP_COLUMN,
+ descriptor_extras={"gated_mlp_order": "gate_then_up"},
+ )
+ return MegatronRoleSpec(role=ROLE_COLUMN)
+
+ # ---- 6. Row-parallel linears. ----
+ if parallelism == "row":
+ return MegatronRoleSpec(role=ROLE_ROW)
+
+ # ---- 7. Replicated (LayerNorms, biases, scalars, routers, etc.). ----
+ if parallelism == "replicated" or tp_size <= 1:
+ return MegatronRoleSpec(role=ROLE_REPLICATED)
+
+ raise ValueError(
+ "cannot determine Megatron tensor placement for "
+ f"parameter {name!r} owned by module {mod_class!r}; "
+ "register the module with Megatron Bridge or provide an explicit "
+ "megatron_role_overrides entry"
+ )
+
+
+def collect_megatron_publish_set(
+ model: "torch.nn.Module",
+ *,
+ tp_size: int,
+ ep_size: int,
+ ep_rank: int,
+ tp_rank: int,
+ num_local_experts: int | None = None,
+ num_attention_heads: int | None = None,
+ num_kv_heads: int | None = None,
+ head_dim: int | None = None,
+ expert_pattern: str | None = None,
+ role_overrides: dict[str, str] | None = None,
+ target_dtype: "torch.dtype | None" = None,
+) -> Iterator[tuple[str, "torch.Tensor", MegatronRoleSpec]]:
+ """Yield ``(name, local_shard, role_spec)`` for the publisher.
+
+ For each parameter:
+
+ * Skips replicated tensors when ``tp_rank != 0``. The MX Megatron receiver
+ handles rank-0 replicated model tensors specially; publishing local
+ copies from non-zero TP ranks can make vLLM's rank-local loader treat
+ them as global tensors and slice past the end.
+ * Returns the parameter as-is — Megatron stores native shards, so
+ the param tensor IS the local shard. No allgather, no Bridge call.
+ Caller is responsible for invoking ``add_tensor`` and
+ ``publish(version=...)`` on the publisher.
+ """
+ for raw_name, param in model.named_parameters():
+ if not param.is_floating_point():
+ # Skip non-float buffers (rotary inv_freq, etc.); they aren't
+ # weight-refit material.
+ continue
+
+ # `model.named_parameters()` returns names with a `module.` prefix
+ # when the model is wrapped (DDP-style). Two distinct uses of the
+ # name:
+ #
+ # 1. The model-walking classifier needs the ORIGINAL prefixed
+ # name to descend through `model.module.decoder.layers...` —
+ # stripping the prefix breaks `_enclosing_module` and every
+ # non-expert tensor falls to ROLE_REPLICATED.
+ # 2. The PUBLISHED name on the catalog has to match Bridge's
+ # name_map (which uses unprefixed names from
+ # `get_conversion_tasks`) so the receiver's name-map lookup
+ # finds the HF target names.
+ #
+ # Classify with `raw_name`; publish with the stripped form.
+ name = (
+ raw_name[len("module.") :] if raw_name.startswith("module.") else raw_name
+ )
+
+ spec = detect_megatron_role(
+ raw_name,
+ param,
+ model=model,
+ tp_size=tp_size,
+ ep_size=ep_size,
+ ep_rank=ep_rank,
+ num_local_experts=num_local_experts,
+ num_attention_heads=num_attention_heads,
+ num_kv_heads=num_kv_heads,
+ head_dim=head_dim,
+ expert_pattern=expert_pattern,
+ role_overrides=role_overrides,
+ )
+
+ if spec.role == ROLE_REPLICATED and tp_rank != 0:
+ continue
+
+ local = param.detach()
+ if target_dtype is not None and local.dtype != target_dtype:
+ local = local.to(target_dtype, non_blocking=True)
+ local = local.contiguous()
+
+ yield name, local, spec
+
+
+__all__ = [
+ "MegatronRoleSpec",
+ "ROLE_COLUMN",
+ "ROLE_EXPERT_COLUMN",
+ "ROLE_EXPERT_ROW",
+ "ROLE_GATED_MLP_COLUMN",
+ "ROLE_QKV_COLUMN",
+ "ROLE_REPLICATED",
+ "ROLE_ROW",
+ "ROLE_VOCAB_PARALLEL",
+ "collect_megatron_publish_set",
+ "detect_megatron_role",
+]
diff --git a/nemo_rl/models/policy/interfaces.py b/nemo_rl/models/policy/interfaces.py
index f0c1ad6bb8..51a26297c4 100644
--- a/nemo_rl/models/policy/interfaces.py
+++ b/nemo_rl/models/policy/interfaces.py
@@ -19,6 +19,7 @@
from nemo_rl.algorithms.loss.interfaces import LossFunction
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
+from nemo_rl.distributed.mx_helpers import ModelExpressPublisherOptions
from nemo_rl.models.generation.interfaces import GenerationDatumSpec
from nemo_rl.utils.timer import Timer
@@ -222,6 +223,28 @@ def broadcast_weights_for_collective(
) -> list[ray.ObjectRef]:
pass
+ def publish_weights_for_model_express(
+ self,
+ *,
+ version: int,
+ publisher_options: ModelExpressPublisherOptions,
+ kv_scales: Optional[dict[str, float]] = None,
+ ) -> list[ray.ObjectRef]:
+ """Publish this policy's rank-local weights through ModelExpress.
+
+ Args:
+ version: Monotonically increasing model version.
+ publisher_options: Internal publisher settings supplied by the
+ ModelExpress weight synchronizer.
+ kv_scales: Optional named FP8 Q/K/V scale values.
+
+ Returns:
+ One Ray object reference per trainer worker.
+ """
+ raise NotImplementedError(
+ "publish_weights_for_model_express is not implemented for this policy worker"
+ )
+
@abstractmethod
def prepare_for_lp_inference(self) -> None:
pass
diff --git a/nemo_rl/models/policy/lm_policy.py b/nemo_rl/models/policy/lm_policy.py
index 397b4e086b..67e3a4d724 100644
--- a/nemo_rl/models/policy/lm_policy.py
+++ b/nemo_rl/models/policy/lm_policy.py
@@ -30,6 +30,7 @@
SequencePackingArgs,
SlicedDataDict,
)
+from nemo_rl.distributed.mx_helpers import ModelExpressPublisherOptions
from nemo_rl.distributed.named_sharding import NamedSharding
from nemo_rl.distributed.virtual_cluster import RayVirtualCluster
from nemo_rl.distributed.worker_groups import RayWorkerBuilder, RayWorkerGroup
@@ -1074,6 +1075,21 @@ def broadcast_weights_for_collective(
# this function should co-work with vllm, so we should wait for all futures to complete outside
return futures
+ def publish_weights_for_model_express(
+ self,
+ *,
+ version: int,
+ publisher_options: ModelExpressPublisherOptions,
+ kv_scales: Optional[dict[str, float]] = None,
+ ) -> list[ray.ObjectRef]:
+ """Publish rank-local weights through ModelExpress."""
+ return self.worker_group.run_all_workers_single_data(
+ "publish_weights_for_model_express",
+ version=int(version),
+ publisher_options=publisher_options,
+ kv_scales=kv_scales,
+ )
+
def offload_before_refit(self) -> None:
"""Offload the optimizer and buffers to the CPU."""
futures = self.worker_group.run_all_workers_single_data("offload_before_refit")
diff --git a/nemo_rl/models/policy/workers/dtensor_policy_worker.py b/nemo_rl/models/policy/workers/dtensor_policy_worker.py
index 0083e5301e..14acbb0c59 100644
--- a/nemo_rl/models/policy/workers/dtensor_policy_worker.py
+++ b/nemo_rl/models/policy/workers/dtensor_policy_worker.py
@@ -65,6 +65,14 @@
distributed_vocab_topk,
get_logprobs_from_vocab_parallel_logits,
)
+from nemo_rl.distributed.mx_helpers import (
+ ModelExpressPublisher,
+ ModelExpressPublisherOptions,
+ build_v2_publisher,
+ finish_model_express_publication,
+ get_dtensor_local_shard,
+ start_model_express_publication,
+)
from nemo_rl.models.dtensor.parallelize import (
_parallelize_model,
clip_grad_by_total_norm_,
@@ -239,6 +247,7 @@ def __init__(
configure_dynamo_cache()
self.cfg = config
+ self._model_express_publisher: ModelExpressPublisher | None = None
# torch distributed init. Envars for rank, world_size, and master_addr and master_port are set from the ray remote call
torch.distributed.init_process_group(backend="nccl")
self.rank = torch.distributed.get_rank()
@@ -1883,6 +1892,69 @@ def dtensor_params_generator():
worker_name=str(self),
)
+ @torch.no_grad()
+ @wrap_with_nvtx_name("dtensor_policy_worker/publish_weights_for_model_express")
+ def publish_weights_for_model_express(
+ self,
+ *,
+ version: int,
+ publisher_options: ModelExpressPublisherOptions,
+ kv_scales: Optional[dict[str, float]] = None,
+ ) -> None:
+ """Publish this worker's local DTensor shards through ModelExpress."""
+ if kv_scales is not None:
+ raise NotImplementedError(
+ "FP8 kvcache scales are only supported on the Megatron MX path"
+ )
+
+ if self.cpu_offload:
+ self.model = self.move_to_cuda(self.model)
+
+ try:
+ if self._model_express_publisher is None:
+ tp_size = self.tp_size or 1
+ self._model_express_publisher = build_v2_publisher(
+ rank=self.rank,
+ # Ray exposes one GPU to each worker, so NIXL must use the
+ # process-local CUDA index rather than the global rank.
+ device_id=torch.cuda.current_device(),
+ fsdp_world_size=self.dp_size,
+ tp_world_size=tp_size,
+ pp_world_size=1,
+ ep_world_size=1,
+ publisher_options=publisher_options,
+ )
+ self._model_express_publisher.initialize(
+ model_name=self.cfg["model_name"],
+ dtype=str(self.dtype).removeprefix("torch."),
+ )
+
+ start_model_express_publication(self._model_express_publisher)
+ for name, tensor in self.model.state_dict().items():
+ shard_spec = None
+ if isinstance(tensor, DTensor):
+ local, shard_spec = get_dtensor_local_shard(tensor)
+ else:
+ local = tensor
+ if local.is_floating_point() and local.dtype != self.dtype:
+ local = local.to(self.dtype, non_blocking=True)
+ local = local.contiguous()
+
+ self._model_express_publisher.add_tensor(
+ name=name,
+ tensor=local,
+ shard_spec=shard_spec,
+ )
+
+ finish_model_express_publication(
+ self._model_express_publisher,
+ version=int(version),
+ worker_rank=self.rank,
+ )
+ finally:
+ if self.cpu_offload:
+ self.model = self.move_to_cpu(self.model)
+
@torch.no_grad()
def broadcast_weights_for_collective(
self, kv_scales: Optional[dict[str, float]] = None
diff --git a/nemo_rl/models/policy/workers/megatron_policy_worker.py b/nemo_rl/models/policy/workers/megatron_policy_worker.py
index 6eea8bc4f8..5a4a009152 100644
--- a/nemo_rl/models/policy/workers/megatron_policy_worker.py
+++ b/nemo_rl/models/policy/workers/megatron_policy_worker.py
@@ -50,6 +50,18 @@
from nemo_rl.algorithms.loss.interfaces import LossFunction
from nemo_rl.data_plane.worker_mixin import TQWorkerMixin
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
+from nemo_rl.distributed.mx_helpers import (
+ ModelExpressPublisher,
+ ModelExpressPublisherOptions,
+ build_v2_publisher,
+ finish_model_express_publication,
+ start_model_express_publication,
+)
+from nemo_rl.distributed.mx_megatron_helpers import (
+ ROLE_COLUMN,
+ ROLE_REPLICATED,
+ collect_megatron_publish_set,
+)
from nemo_rl.distributed.named_sharding import NamedSharding
from nemo_rl.models.generation.interfaces import GenerationDatumSpec
from nemo_rl.models.generation.megatron.megatron_worker import (
@@ -292,6 +304,8 @@ def __init__(
bind_to_gpu_numa(local_rank)
self.cfg = config
+ self._model_express_publisher: ModelExpressPublisher | None = None
+ self._model_express_megatron_sidecar: dict[str, Any] | None = None
self._router_replay_enabled = router_replay_enabled(config)
# Set rank for non-collocated to check which ranks to broadcast from
@@ -1953,6 +1967,247 @@ def broadcast_weights_for_collective(
def _use_real_quant_refit(self) -> bool:
return False
+ @torch.no_grad()
+ @wrap_with_nvtx_name("megatron_policy_worker/publish_weights_for_model_express")
+ def publish_weights_for_model_express(
+ self,
+ *,
+ version: int,
+ publisher_options: ModelExpressPublisherOptions,
+ kv_scales: Optional[dict[str, float]] = None,
+ ) -> None:
+ """Publish this worker's native Megatron shards through ModelExpress."""
+ from megatron.core import parallel_state
+
+ tp_size = parallel_state.get_tensor_model_parallel_world_size()
+ tp_rank = parallel_state.get_tensor_model_parallel_rank()
+ pp_size = parallel_state.get_pipeline_model_parallel_world_size()
+ pp_rank = parallel_state.get_pipeline_model_parallel_rank()
+ ep_size = parallel_state.get_expert_model_parallel_world_size()
+ ep_rank = parallel_state.get_expert_model_parallel_rank()
+
+ # ---- Lazy-init the publisher (once per worker lifetime). ----
+ if self._model_express_publisher is None:
+ mx_device_id = torch.cuda.current_device()
+ self._model_express_publisher = build_v2_publisher(
+ rank=self.rank,
+ device_id=mx_device_id,
+ fsdp_world_size=self.dp_size,
+ tp_world_size=tp_size,
+ pp_world_size=pp_size,
+ ep_world_size=ep_size,
+ publisher_options=publisher_options,
+ agent_name=f"nemo-rl-megatron-trainer-r{self.rank}",
+ )
+ self._model_express_publisher.initialize(
+ model_name=self.cfg["model_name"],
+ dtype=str(self.dtype).removeprefix("torch."),
+ )
+ self._model_express_megatron_sidecar = self._build_megatron_sidecar()
+ self._model_express_publisher.set_megatron_sidecar(
+ self._model_express_megatron_sidecar
+ )
+
+ # Resolve attention-head metadata for fused-QKV descriptors.
+ tcfg = getattr(self.megatron_bridge, "transformer_config", None)
+ num_attention_heads = (
+ getattr(tcfg, "num_attention_heads", None) if tcfg else None
+ )
+ num_kv_heads = (
+ getattr(tcfg, "num_query_groups", None) if tcfg else None
+ ) or num_attention_heads
+ num_moe_experts = getattr(tcfg, "num_moe_experts", None) if tcfg else None
+ num_local_experts = (
+ int(num_moe_experts) // ep_size
+ if num_moe_experts is not None and int(num_moe_experts) % ep_size == 0
+ else None
+ )
+ head_dim = (getattr(tcfg, "kv_channels", None) if tcfg else None) or (
+ num_attention_heads
+ and getattr(tcfg, "hidden_size", 0) // num_attention_heads
+ if tcfg
+ else None
+ )
+
+ role_overrides = self._mx_megatron_role_overrides_from_sidecar(
+ role=ROLE_COLUMN,
+ )
+ role_overrides.update(publisher_options.megatron_role_overrides)
+
+ start_model_express_publication(self._model_express_publisher)
+
+ # Stamp the publisher's mesh position into source metadata.
+ self._model_express_publisher.set_megatron_mesh_position(
+ tp_rank=tp_rank,
+ pp_rank=pp_rank,
+ ep_rank=ep_rank,
+ )
+
+ for name, local, spec in collect_megatron_publish_set(
+ self.model,
+ tp_size=tp_size,
+ ep_size=ep_size,
+ ep_rank=ep_rank,
+ tp_rank=tp_rank,
+ num_local_experts=num_local_experts,
+ num_attention_heads=num_attention_heads,
+ num_kv_heads=num_kv_heads,
+ head_dim=head_dim,
+ role_overrides=role_overrides,
+ target_dtype=self.dtype,
+ ):
+ self._model_express_publisher.add_tensor(
+ name=name,
+ tensor=local,
+ is_expert=spec.is_expert,
+ expert_axis=spec.expert_axis,
+ owned_expert_ids=spec.owned_expert_ids,
+ megatron_role=spec.role,
+ megatron_extras=spec.descriptor_extras,
+ )
+
+ # KV scales are small replicated values. Publish them from every TP
+ # rank so each rank-local source set is independently complete.
+ if kv_scales:
+ for name, scale_value in sorted(kv_scales.items()):
+ scale_tensor = torch.tensor(
+ float(scale_value),
+ dtype=torch.float32,
+ device="cuda",
+ ).reshape(1)
+ self._model_express_publisher.add_tensor(
+ name=name,
+ tensor=scale_tensor,
+ is_expert=False,
+ expert_axis=0,
+ owned_expert_ids=(),
+ megatron_role=ROLE_REPLICATED,
+ megatron_extras={"fp8_kv_scale": "1"},
+ )
+
+ finish_model_express_publication(
+ self._model_express_publisher,
+ version=int(version),
+ worker_rank=self.rank,
+ )
+
+ def _mx_megatron_role_overrides_from_sidecar(self, *, role: str) -> dict[str, str]:
+ """Derive publish role overrides from Bridge's Megatron-to-HF name map."""
+ sidecar = self._model_express_megatron_sidecar or {}
+ name_map = sidecar.get("megatron_hf_name_map", [])
+ role_overrides: dict[str, str] = {}
+ for entry in name_map:
+ if not isinstance(entry, (list, tuple)) or len(entry) != 2:
+ continue
+ megatron_name = str(entry[0])
+ if ".experts." in megatron_name:
+ continue
+ if not any(
+ marker in megatron_name for marker in ("linear_fc1", "gate_up_proj")
+ ):
+ continue
+ hf_names = [str(hf_name) for hf_name in entry[1]]
+ if len(hf_names) != 1:
+ continue
+ hf_name = hf_names[0]
+ if "up_proj" in hf_name and "gate_proj" not in hf_name:
+ role_overrides[megatron_name] = role
+ return role_overrides
+
+ def _build_megatron_sidecar(self) -> dict[str, Any]:
+ """Serialize Megatron-Bridge introspection results at trainer init.
+
+ Two pieces:
+ 1. ``megatron_transformer_config`` — head counts + dims read
+ from the trainer's :class:`TransformerConfig`.
+ 2. ``megatron_hf_name_map`` — list of
+ ``[megatron_param_name, [hf_name_1, hf_name_2, ...]]`` derived
+ from a Bridge introspection pass over the local model
+ (no weight transfer; just iterates the conversion-task list).
+ """
+ sidecar: dict[str, Any] = {}
+
+ # --- transformer_config ---
+ tcfg = getattr(self.megatron_bridge, "transformer_config", None)
+ if tcfg is not None:
+ num_heads = getattr(tcfg, "num_attention_heads", None)
+ kv_groups = getattr(tcfg, "num_query_groups", None) or num_heads
+ kv_channels = getattr(tcfg, "kv_channels", None)
+ if kv_channels is None and num_heads:
+ kv_channels = getattr(tcfg, "hidden_size", 0) // num_heads
+ sidecar["megatron_transformer_config"] = {
+ "num_attention_heads": num_heads,
+ "num_query_groups": kv_groups,
+ "kv_channels": kv_channels,
+ "hidden_size": getattr(tcfg, "hidden_size", None),
+ }
+
+ # --- hf_name_map ---
+ # Walk the Bridge mapping registry to derive
+ # (megatron_local_name, [hf_name_1, hf_name_2, ...]). We use the
+ # registry's resolved tasks rather than calling
+ # ``export_hf_weights`` so we don't pay the gather/broadcast cost
+ # at startup. The receiver only needs the name pairings; head
+ # counts come from transformer_config.
+ tasks = self.megatron_bridge.get_conversion_tasks([self.model])
+ name_map: defaultdict[str, list[str]] = defaultdict(list)
+
+ def _ordered_hf_names(hf_names: list[str]) -> list[str]:
+ unique: list[str] = []
+ for hf_name in hf_names:
+ if hf_name not in unique:
+ unique.append(hf_name)
+
+ def _priority(name: str, markers: tuple[str, ...]) -> int:
+ for index, marker in enumerate(markers):
+ if marker in name:
+ return index
+ return len(markers)
+
+ if any(
+ marker in name
+ for name in unique
+ for marker in ("q_proj", "k_proj", "v_proj")
+ ):
+ return sorted(
+ unique,
+ key=lambda name: _priority(name, ("q_proj", "k_proj", "v_proj")),
+ )
+ if any(
+ marker in name for name in unique for marker in ("gate_proj", "up_proj")
+ ):
+ return sorted(
+ unique,
+ key=lambda name: _priority(name, ("gate_proj", "up_proj")),
+ )
+ return unique
+
+ for task in tasks:
+ if task is None:
+ continue
+ m_name = getattr(task, "global_param_name", None) or task.param_name
+ # Each task's mapping declares one or more HF names. Resolve
+ # via the mapping's hf_param attribute (str or dict).
+ hf_attr = getattr(task.mapping, "hf_param", None)
+ if isinstance(hf_attr, str):
+ hf_names = [hf_attr]
+ elif isinstance(hf_attr, dict):
+ # Order matters for QKV: q, k, v. Bridge's QKVMapping
+ # uses keys "q", "k", "v" — preserve that ordering.
+ if set(hf_attr.keys()) == {"q", "k", "v"}:
+ hf_names = [hf_attr["q"], hf_attr["k"], hf_attr["v"]]
+ else:
+ hf_names = list(hf_attr.values())
+ else:
+ continue
+ name_map[m_name].extend(str(hf_name) for hf_name in hf_names)
+ sidecar["megatron_hf_name_map"] = [
+ (m_name, _ordered_hf_names(hf_names))
+ for m_name, hf_names in name_map.items()
+ ]
+
+ return sidecar
+
def prepare_for_lp_inference(self):
self.model = self.move_model(self.model, "cuda", move_grads=False)
self.model.eval()
diff --git a/pyrefly.toml b/pyrefly.toml
index 117ec1fdcd..856dfe6376 100644
--- a/pyrefly.toml
+++ b/pyrefly.toml
@@ -117,6 +117,8 @@ project-includes = [
"nemo_rl/data_plane/worker_mixin.py",
"nemo_rl/distributed/__init__.py",
"nemo_rl/distributed/collectives.py",
+ "nemo_rl/distributed/mx_helpers.py",
+ "nemo_rl/distributed/mx_megatron_helpers.py",
"nemo_rl/distributed/named_sharding.py",
"nemo_rl/distributed/numa_utils.py",
"nemo_rl/distributed/ray_actor_environment_registry.py",
diff --git a/tests/unit/distributed/test_mx_helpers.py b/tests/unit/distributed/test_mx_helpers.py
new file mode 100644
index 0000000000..2a68f3e877
--- /dev/null
+++ b/tests/unit/distributed/test_mx_helpers.py
@@ -0,0 +1,167 @@
+import pytest
+import torch
+from torch.distributed.tensor import Partial, Replicate, Shard
+
+from nemo_rl.distributed.mx_helpers import (
+ ModelExpressPublisherOptions,
+ finish_model_express_publication,
+ get_dtensor_local_shard,
+ start_model_express_publication,
+)
+
+
+class FakeMesh:
+ def __init__(self, coordinate: tuple[int, ...] | None, sizes: tuple[int, ...]):
+ self._coordinate = coordinate
+ self._sizes = sizes
+
+ def get_coordinate(self) -> list[int] | None:
+ return list(self._coordinate) if self._coordinate is not None else None
+
+ def size(self, mesh_dim: int) -> int:
+ return self._sizes[mesh_dim]
+
+
+class FakeDTensor:
+ def __init__(
+ self,
+ local: torch.Tensor,
+ *,
+ global_shape: tuple[int, ...],
+ placements: tuple[object, ...],
+ coordinate: tuple[int, ...] | None,
+ mesh_sizes: tuple[int, ...],
+ ):
+ self._local = local
+ self.shape = global_shape
+ self.placements = placements
+ self.device_mesh = FakeMesh(coordinate, mesh_sizes)
+
+ def to_local(self) -> torch.Tensor:
+ return self._local
+
+
+class FakePublisher:
+ def __init__(self, *, ready: bool = True):
+ self.ready = ready
+ self.reset_count = 0
+ self.published_versions: list[int] = []
+
+ def reset_tensors(self) -> None:
+ self.reset_count += 1
+
+ def publish(self, *, version: int) -> str:
+ self.published_versions.append(version)
+ return f"source-{version}"
+
+ def mark_ready(self) -> bool:
+ return self.ready
+
+
+def test_publisher_options_preserve_megatron_role_overrides():
+ options = ModelExpressPublisherOptions(
+ mx_server_url="modelexpress-server:8001",
+ nic_pin="auto",
+ megatron_role_overrides={"linear_fc1": "column"},
+ )
+
+ assert options.megatron_role_overrides == {"linear_fc1": "column"}
+
+
+def test_model_express_publication_lifecycle_repeats_per_version():
+ publisher = FakePublisher()
+
+ for version in (7, 8):
+ start_model_express_publication(publisher)
+ source_id = finish_model_express_publication(
+ publisher,
+ version=version,
+ worker_rank=3,
+ )
+ assert source_id == f"source-{version}"
+
+ assert publisher.reset_count == 2
+ assert publisher.published_versions == [7, 8]
+
+
+def test_model_express_publication_requires_ready_transition():
+ publisher = FakePublisher(ready=False)
+
+ with pytest.raises(RuntimeError, match="trainer rank 3 ready"):
+ finish_model_express_publication(
+ publisher,
+ version=7,
+ worker_rank=3,
+ )
+
+
+def test_get_dtensor_local_shard_describes_uneven_last_shard():
+ tensor = FakeDTensor(
+ torch.ones(2, 4),
+ global_shape=(5, 4),
+ placements=(Shard(0),),
+ coordinate=(1,),
+ mesh_sizes=(2,),
+ )
+
+ local, shard_spec = get_dtensor_local_shard(tensor)
+
+ assert local is tensor._local
+ assert shard_spec is not None
+ assert shard_spec.global_shape == (5, 4)
+ assert shard_spec.shard_axis == 0
+ assert shard_spec.local_shard_range == (3, 5)
+
+
+def test_get_dtensor_local_shard_accepts_replicated_tensor():
+ tensor = FakeDTensor(
+ torch.ones(2, 4),
+ global_shape=(2, 4),
+ placements=(Replicate(),),
+ coordinate=(0,),
+ mesh_sizes=(1,),
+ )
+
+ local, shard_spec = get_dtensor_local_shard(tensor)
+
+ assert local is tensor._local
+ assert shard_spec is None
+
+
+def test_get_dtensor_local_shard_rejects_multiple_shard_axes():
+ tensor = FakeDTensor(
+ torch.ones(2, 2),
+ global_shape=(4, 4),
+ placements=(Shard(0), Shard(1)),
+ coordinate=(0, 0),
+ mesh_sizes=(2, 2),
+ )
+
+ with pytest.raises(NotImplementedError, match="one DTensor shard axis"):
+ get_dtensor_local_shard(tensor)
+
+
+def test_get_dtensor_local_shard_rejects_partial_placement():
+ tensor = FakeDTensor(
+ torch.ones(2, 4),
+ global_shape=(2, 4),
+ placements=(Partial(),),
+ coordinate=(0,),
+ mesh_sizes=(1,),
+ )
+
+ with pytest.raises(NotImplementedError, match="partial DTensors"):
+ get_dtensor_local_shard(tensor)
+
+
+def test_get_dtensor_local_shard_rejects_rank_outside_mesh():
+ tensor = FakeDTensor(
+ torch.ones(2, 4),
+ global_shape=(4, 4),
+ placements=(Shard(0),),
+ coordinate=None,
+ mesh_sizes=(2,),
+ )
+
+ with pytest.raises(RuntimeError, match="not part of the DTensor device mesh"):
+ get_dtensor_local_shard(tensor)
diff --git a/tests/unit/distributed/test_mx_megatron_helpers.py b/tests/unit/distributed/test_mx_megatron_helpers.py
new file mode 100644
index 0000000000..0a7abc2d76
--- /dev/null
+++ b/tests/unit/distributed/test_mx_megatron_helpers.py
@@ -0,0 +1,147 @@
+import torch
+import pytest
+
+from nemo_rl.distributed.mx_megatron_helpers import (
+ ROLE_EXPERT_COLUMN,
+ ROLE_QKV_COLUMN,
+ ROLE_REPLICATED,
+ ROLE_ROW,
+ collect_megatron_publish_set,
+ detect_megatron_role,
+)
+
+
+class ReplicatedOnlyModule(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.weight = torch.nn.Parameter(torch.ones(2))
+
+
+class ColumnParallelLinear(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.weight = torch.nn.Parameter(torch.ones(6, 2))
+
+
+class TEColumnParallelGroupedLinear(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.weight0 = torch.nn.Parameter(torch.ones(4, 2))
+
+
+class AttentionModule(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.linear_qkv = ColumnParallelLinear()
+
+
+class ExpertModule(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.experts = TEColumnParallelGroupedLinear()
+
+
+class TELinear(torch.nn.Module):
+ def __init__(self, parallel_mode: str):
+ super().__init__()
+ self.parallel_mode = parallel_mode
+ self.weight = torch.nn.Parameter(torch.ones(2, 2))
+
+
+class UnknownParallelModule(torch.nn.Module):
+ def __init__(self):
+ super().__init__()
+ self.weight = torch.nn.Parameter(torch.ones(2, 2))
+
+
+def _published_names(*, tp_rank: int) -> list[str]:
+ model = ReplicatedOnlyModule()
+ published = collect_megatron_publish_set(
+ model,
+ tp_size=2,
+ ep_size=1,
+ ep_rank=0,
+ tp_rank=tp_rank,
+ role_overrides={"weight": ROLE_REPLICATED},
+ )
+ return [name for name, _, _ in published]
+
+
+def test_collect_megatron_publish_set_skips_replicated_on_nonzero_tp_rank():
+ assert _published_names(tp_rank=1) == []
+
+
+def test_collect_megatron_publish_set_publishes_replicated_on_zero_tp_rank():
+ assert _published_names(tp_rank=0) == ["weight"]
+
+
+def test_collect_megatron_publish_set_classifies_fused_qkv():
+ published = list(
+ collect_megatron_publish_set(
+ AttentionModule(),
+ tp_size=2,
+ ep_size=1,
+ ep_rank=0,
+ tp_rank=0,
+ num_attention_heads=8,
+ num_kv_heads=4,
+ head_dim=16,
+ )
+ )
+
+ name, _, spec = published[0]
+ assert name == "linear_qkv.weight"
+ assert spec.role == ROLE_QKV_COLUMN
+ assert spec.descriptor_extras == {
+ "qkv_interleave": "by_head",
+ "num_heads_local": "4",
+ "num_kv_heads_local": "2",
+ "head_dim": "16",
+ }
+
+
+def test_collect_megatron_publish_set_uses_global_grouped_expert_id():
+ published = list(
+ collect_megatron_publish_set(
+ ExpertModule(),
+ tp_size=1,
+ ep_size=2,
+ ep_rank=1,
+ tp_rank=0,
+ num_local_experts=4,
+ )
+ )
+
+ name, _, spec = published[0]
+ assert name == "experts.weight0"
+ assert spec.role == ROLE_EXPERT_COLUMN
+ assert spec.owned_expert_ids == {4}
+
+
+def test_detect_megatron_role_uses_te_parallel_mode():
+ model = TELinear(parallel_mode="row")
+
+ spec = detect_megatron_role(
+ "weight",
+ model.weight,
+ model=model,
+ tp_size=2,
+ ep_size=1,
+ ep_rank=0,
+ )
+
+ assert spec.role == ROLE_ROW
+
+
+def test_detect_megatron_role_rejects_unknown_tp_module():
+ model = UnknownParallelModule()
+
+ with pytest.raises(ValueError, match="cannot determine Megatron tensor placement"):
+ detect_megatron_role(
+ "weight",
+ model.weight,
+ model=model,
+ tp_size=2,
+ ep_size=1,
+ ep_rank=0,
+ )