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
9 changes: 7 additions & 2 deletions miles/backends/experimental/fsdp_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
from miles.utils.tracking_utils import init_tracking

from ....utils.profile_utils import TrainProfiler
from ...training_utils.ci_utils import check_grad_norm
from ...training_utils.ci_utils import assert_rollout_engine_weight_versions, check_grad_norm
from ...training_utils.data import DataIterator, get_batch, get_data_iterator, get_rollout_data
from ...training_utils.log_utils import (
aggregate_forward_results,
Expand Down Expand Up @@ -569,7 +569,12 @@ def update_weights(self) -> None: # type: ignore[override]

self.weight_updater.update_weights()

if self.args.ci_test and len(rollout_engines) > 0:
if getattr(self.args, "check_all_engine_weight_versions", False):
assert_rollout_engine_weight_versions(
rollout_engines,
self.weight_updater.weight_version,
)
elif self.args.ci_test and len(rollout_engines) > 0:
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
if str(engine_version) != str(self.weight_updater.weight_version):
Expand Down
19 changes: 13 additions & 6 deletions miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from ...utils.profile_utils import TrainProfiler
from ...utils.tensor_backper import TensorBackuper
from ..training_utils.cp_utils import slice_with_cp
from ..training_utils.ci_utils import assert_rollout_engine_weight_versions
from ..training_utils.data import DataIterator, get_data_iterator, get_rollout_data, sync_actor_critic_data
from ..training_utils.log_utils import log_cpu_memory, log_perf_data, log_rollout_data
from ..training_utils.loss import compute_advantages_and_returns, get_log_probs_and_entropy, get_values
Expand Down Expand Up @@ -1013,13 +1014,19 @@ def update_weights(self) -> None:
self.weight_updater.update_weights()
print_memory("after update_weights")

if self.args.ci_test and len(rollout_engines) > 0 and not is_lora_enabled(self.args):
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
if str(engine_version) != str(self.weight_updater.weight_version):
raise RuntimeError(
f"Weight version mismatch! Engine: {engine_version}, Updater: {self.weight_updater.weight_version}"
if not is_lora_enabled(self.args):
if getattr(self.args, "check_all_engine_weight_versions", False):
assert_rollout_engine_weight_versions(
rollout_engines,
self.weight_updater.weight_version,
)
elif self.args.ci_test and len(rollout_engines) > 0:
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
if str(engine_version) != str(self.weight_updater.weight_version):
raise RuntimeError(
f"Weight version mismatch! Engine: {engine_version}, Updater: {self.weight_updater.weight_version}"
)

if getattr(self.args, "keep_old_actor", False):
if self.args.update_weights_interval == 1:
Expand Down
198 changes: 183 additions & 15 deletions miles/backends/megatron_utils/megatron_to_hf/xllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,20 +3,142 @@
import torch


def _is_mova(args) -> bool:
return getattr(args, "mova_num_value_experts", 0) > 0


def _attention_geometry(args) -> tuple[int, int, int, int]:
hidden_size = args.hidden_size
num_attention_heads = args.num_attention_heads
num_query_groups = args.num_query_groups
kv_channels = getattr(args, "kv_channels", None)
head_dim = kv_channels if kv_channels is not None else hidden_size // num_attention_heads

if num_attention_heads % num_query_groups:
raise ValueError(
f"num_attention_heads={num_attention_heads} must be divisible by "
f"num_query_groups={num_query_groups}"
)
if head_dim <= 0 or head_dim % 2:
raise ValueError(f"xLLM MoVA requires an even positive head dimension, got {head_dim}")
return hidden_size, num_attention_heads, num_query_groups, head_dim


def _permute_qk_to_hf(
weight: torch.Tensor,
*,
num_heads: int,
head_dim: int,
hidden_size: int,
name: str,
) -> torch.Tensor:
"""Convert MCore's adjacent-complex-pair Q/K rows to xLLM HF layout."""

expected_shape = (num_heads * head_dim, hidden_size)
if tuple(weight.shape) != expected_shape:
raise ValueError(f"Invalid {name} shape: got {tuple(weight.shape)}, expected {expected_shape}")
return (
weight.reshape(num_heads, head_dim // 2, 2, hidden_size)
.transpose(1, 2)
.reshape(expected_shape)
.contiguous()
)


def _unpack_grouped_attention_projection(
args,
name: str,
param: torch.Tensor,
*,
include_value: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]:
"""Unpack MCore's per-GQA-group Q/G/K[/V] projection."""

hidden_size, num_attention_heads, num_query_groups, head_dim = _attention_geometry(args)
query_heads_per_group = num_attention_heads // num_query_groups
segment_heads = [query_heads_per_group, query_heads_per_group, 1]
if include_value:
segment_heads.append(1)

expected_rows = num_query_groups * sum(segment_heads) * head_dim
if tuple(param.shape) != (expected_rows, hidden_size):
raise ValueError(
f"Invalid {name} shape: got {tuple(param.shape)}, expected "
f"{(expected_rows, hidden_size)}"
)

packed = param.reshape(num_query_groups, sum(segment_heads), head_dim, hidden_size)
chunks = torch.split(packed, segment_heads, dim=1)
query = chunks[0].reshape(num_attention_heads * head_dim, hidden_size)
gate = chunks[1].reshape(num_attention_heads * head_dim, hidden_size)
key = chunks[2].reshape(num_query_groups * head_dim, hidden_size)
value = chunks[3].reshape(num_query_groups * head_dim, hidden_size) if include_value else None

query = _permute_qk_to_hf(
query,
num_heads=num_attention_heads,
head_dim=head_dim,
hidden_size=hidden_size,
name=f"{name}.query",
)
key = _permute_qk_to_hf(
key,
num_heads=num_query_groups,
head_dim=head_dim,
hidden_size=hidden_size,
name=f"{name}.key",
)
return query, gate, key, value


def _convert_grouped_value_experts(
args, layer_idx: str, name: str, param: torch.Tensor
) -> list[tuple[str, torch.Tensor]]:
"""Convert gathered MCore Wv [expert, hidden, value] to HF [value, hidden]."""

if param.ndim != 3:
raise ValueError(
f"Invalid {name} shape: got {tuple(param.shape)}, expected "
"[num_value_experts, hidden_size, value_width]"
)
expected_experts = getattr(args, "mova_num_value_experts", 0)
if expected_experts and param.shape[0] != expected_experts:
raise ValueError(
f"Invalid {name} expert count: got {param.shape[0]}, expected {expected_experts}"
)
if param.shape[1] != args.hidden_size:
raise ValueError(
f"Invalid {name} hidden dimension: got {param.shape[1]}, expected {args.hidden_size}. "
"The grouped MoVA weight must be gathered across regular attention TP before conversion."
)

return [
(
f"model.layers.{layer_idx}.self_attn.v_experts.{expert_idx}.weight",
expert_weight.transpose(0, 1).contiguous(),
)
for expert_idx, expert_weight in enumerate(param.unbind(dim=0))
]


def convert_xllm_to_hf(args, name, param):
"""Convert Megatron parameter names/tensors to HuggingFace xLLM format."""
"""Convert Megatron parameter names/tensors to HuggingFace xLLM format.

MoVA uses MCore's interleaved Q/G/K[/V] projections and a grouped value
weight stored as ``[expert, hidden / TP, value_width]``. The caller first
gathers those tensors over regular attention TP; this function then emits
canonical xLLM HF names consumed by both SGLang broadcast and P2P loading.
"""

if name == "module.module.embedding.word_embeddings.weight":
return [("model.embed_tokens.weight", param)]
if name == "module.module.output_layer.weight":
return [("lm_head.weight", param)]
if name == "module.module.decoder.final_layernorm.weight":
return [("model.norm.weight", param)]

try:
head_dim = args.kv_channels if args.kv_channels is not None else args.hidden_size // args.num_attention_heads
except AttributeError:
head_dim = args.hidden_size // args.num_attention_heads
value_num_per_group = args.num_attention_heads // args.num_query_groups
hidden_size, num_attention_heads, num_query_groups, head_dim = _attention_geometry(args)
query_heads_per_group = num_attention_heads // num_query_groups

decoder_layers_pattern = r"module\.module\.decoder\.layers\.(\d+)\.(.+)"
match = re.match(decoder_layers_pattern, name)
Expand Down Expand Up @@ -53,18 +175,60 @@ def convert_xllm_to_hf(args, name, param):

if rest == "self_attention.linear_proj.weight":
return [(f"model.layers.{layer_idx}.self_attn.o_proj.weight", param)]

if rest == "self_attention.linear_qkv.weight":
param = param.view(args.num_query_groups, -1, head_dim, args.hidden_size)
q_param, k_param, v_param = torch.split(param, [value_num_per_group, 1, 1], dim=1)
q_param = q_param.reshape(-1, args.hidden_size)
k_param = k_param.reshape(-1, args.hidden_size)
v_param = v_param.reshape(-1, args.hidden_size)
if _is_mova(args):
query, gate, key, value = _unpack_grouped_attention_projection(
args, name, param, include_value=True
)
assert value is not None
return [
(f"model.layers.{layer_idx}.self_attn.q_proj.weight", query),
(f"model.layers.{layer_idx}.self_attn.attn_gate_proj.weight", gate),
(f"model.layers.{layer_idx}.self_attn.k_proj.weight", key),
(f"model.layers.{layer_idx}.self_attn.v_proj.weight", value),
]

# Preserve the legacy xLLM Q/K/V contract when MoVA is disabled.
packed = param.view(num_query_groups, -1, head_dim, hidden_size)
query, key, value = torch.split(packed, [query_heads_per_group, 1, 1], dim=1)
return [
(f"model.layers.{layer_idx}.self_attn.q_proj.weight", q_param),
(f"model.layers.{layer_idx}.self_attn.k_proj.weight", k_param),
(f"model.layers.{layer_idx}.self_attn.v_proj.weight", v_param),
(f"model.layers.{layer_idx}.self_attn.q_proj.weight", query.reshape(-1, hidden_size)),
(f"model.layers.{layer_idx}.self_attn.k_proj.weight", key.reshape(-1, hidden_size)),
(f"model.layers.{layer_idx}.self_attn.v_proj.weight", value.reshape(-1, hidden_size)),
]

if rest == "self_attention.linear_qkg.weight":
if not _is_mova(args):
raise ValueError(f"Found MoVA Q/K/gate projection while MoVA is disabled: {name}")
query, gate, key, value = _unpack_grouped_attention_projection(
args, name, param, include_value=False
)
assert value is None
return [
(f"model.layers.{layer_idx}.self_attn.q_proj.weight", query),
(f"model.layers.{layer_idx}.self_attn.attn_gate_proj.weight", gate),
(f"model.layers.{layer_idx}.self_attn.k_proj.weight", key),
]

if rest == "self_attention.value_projection.experts.weight":
if not _is_mova(args):
raise ValueError(f"Found grouped MoVA value experts while MoVA is disabled: {name}")
return _convert_grouped_value_experts(args, layer_idx, name, param)

sequential_value_expert_pattern = (
r"self_attention\.value_projection\.experts\.experts\.(\d+)\.weight"
)
match = re.match(sequential_value_expert_pattern, rest)
if match:
expert_idx = match.group(1)
return [(f"model.layers.{layer_idx}.self_attn.v_experts.{expert_idx}.weight", param)]

if rest == "self_attention.value_projection.router.weight":
return [(f"model.layers.{layer_idx}.self_attn.v_router.weight", param)]
if rest == "self_attention.value_projection.router.expert_bias":
return [(f"model.layers.{layer_idx}.self_attn.v_router.bias", param)]

if rest == "mlp.linear_fc1.weight":
gate_weight, up_weight = param.chunk(2, dim=0)
return [
Expand All @@ -74,7 +238,11 @@ def convert_xllm_to_hf(args, name, param):
if rest == "mlp.linear_fc2.weight":
return [(f"model.layers.{layer_idx}.mlp.down_proj.weight", param)]

if rest in ("self_attention.linear_qkv.layer_norm_weight", "input_layernorm.weight"):
if rest in (
"self_attention.linear_qkv.layer_norm_weight",
"self_attention.linear_qkg.layer_norm_weight",
"input_layernorm.weight",
):
return [(f"model.layers.{layer_idx}.input_layernorm.weight", param)]
if rest in ("mlp.linear_fc1.layer_norm_weight", "pre_mlp_layernorm.weight"):
return [(f"model.layers.{layer_idx}.post_attention_layernorm.weight", param)]
Expand Down
44 changes: 41 additions & 3 deletions miles/backends/megatron_utils/model_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,24 @@
logger = logging.getLogger(__name__)


def _get_mova_model_components():
"""Import MoVA components only when the architecture is requested.

Keeping this import lazy preserves compatibility with ordinary Megatron
installations that have not yet taken the native MoVA extension.
"""

try:
from megatron.core.models.gpt.mova_layer_specs import get_mova_gpt_decoder_block_spec
from megatron.core.transformer.mova import MoVATransformerConfig
except ImportError as error:
raise RuntimeError(
"Native MoVA was requested, but this Megatron installation does not "
"provide MoVATransformerConfig/get_mova_gpt_decoder_block_spec"
) from error
return MoVATransformerConfig, get_mova_gpt_decoder_block_spec


# Adapt from https://github.com/volcengine/verl/blob/c3b20575d2bc815fcccd84bddb4c0401fc4b632b/verl/models/llama/megatron/layers/parallel_linear.py#L82
class LinearForLastLayer(torch.nn.Linear):
def __init__(
Expand Down Expand Up @@ -59,6 +77,14 @@ def get_model_provider_func(
args: argparse.Namespace,
role: Literal["actor", "critic"] = "actor",
):
is_mova = getattr(args, "mova_num_value_experts", 0) > 0
if is_mova:
# Validate before any provider branch so MoVA can never silently build a
# custom/ordinary provider or select an unsupported converter.
from miles.utils.arguments import validate_mova_args

validate_mova_args(args)

# Support custom model provider path (similar to --custom-rm-path for reward models)
if getattr(args, "custom_model_provider_path", None):

Expand Down Expand Up @@ -147,9 +173,21 @@ def model_provider(

# Experimental loading arguments from yaml
assert config is None, "miles builds the config from args, so it expects config to be None"
config = core_transformer_config_from_args(args)

if args.spec is not None:
if is_mova:
mova_config_class, get_mova_block_spec = _get_mova_model_components()
config = core_transformer_config_from_args(args, mova_config_class)
else:
config = core_transformer_config_from_args(args)

if is_mova:
transformer_layer_spec = get_mova_block_spec(
config,
use_transformer_engine=use_te,
moe_grouped_gemm=args.moe_grouped_gemm,
moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm,
vp_stage=vp_stage,
)
elif args.spec is not None:
transformer_layer_spec = import_module(args.spec)
# Allow the spec to be a function so that user can use customized Megatron easier.
if callable(transformer_layer_spec):
Expand Down
Loading
Loading