@@ -248,6 +248,19 @@ def _is_mtp_megatron_param(param_name: str) -> bool:
248248 return param_name .startswith ("mtp." ) or ".mtp." in param_name
249249
250250
251+ def _grouped_expert_member_views (weight : torch .Tensor ) -> list [torch .Tensor ]:
252+ """Return cached TE grouped members without invoking unsupported indexing."""
253+ storage = weight .data if isinstance (weight , torch .nn .Parameter ) else weight
254+ splitter = getattr (storage , "split_into_quantized_tensors" , None )
255+ if callable (splitter ):
256+ members = getattr (storage , "quantized_tensors" , None )
257+ if members is None :
258+ members = splitter ()
259+ storage .quantized_tensors = members
260+ return list (members )
261+ return list (storage .unbind (0 ))
262+
263+
251264def _collect_mtp_hf_layer_names (conversion_tasks : Optional [list ]) -> set [str ]:
252265 """Return HF layer names whose weights originate from Megatron's MTP module.
253266
@@ -2474,6 +2487,44 @@ def _build_native_mxfp8_conversion_tasks(self) -> list[Any]:
24742487 )
24752488 self ._native_grouped_mxfp8_tasks = grouped_tasks
24762489 grouped_names = {task .global_param_name for task in grouped_tasks }
2490+ grouped_suffixes = (
2491+ ".mlp.experts.linear_fc1.weight" ,
2492+ ".mlp.experts.linear_fc2.weight" ,
2493+ )
2494+ num_experts = int (getattr (self .model .config , "num_moe_experts" , 0 ) or 0 )
2495+ ep_size = int (getattr (self .model .config , "expert_model_parallel_size" , 1 ) or 1 )
2496+ if num_experts and num_experts % ep_size :
2497+ raise ValueError (
2498+ f"num_moe_experts={ num_experts } must be divisible by "
2499+ f"expert_model_parallel_size={ ep_size } "
2500+ )
2501+ local_expert_count = num_experts // ep_size if num_experts else 0
2502+ grouped_misc_names : dict [str , list [str ]] = {}
2503+ expanded_global_names : list [str ] = []
2504+ for global_name in global_names :
2505+ if (
2506+ global_name .endswith (grouped_suffixes )
2507+ and global_name not in grouped_names
2508+ ):
2509+ if _is_mtp_megatron_param (global_name ):
2510+ raise ValueError (
2511+ "native MXFP8 refit does not yet support co-trained MTP "
2512+ "grouped experts"
2513+ )
2514+ if local_expert_count <= 0 :
2515+ raise ValueError (
2516+ f"Cannot expand grouped expert parameter { global_name !r} "
2517+ "without num_moe_experts"
2518+ )
2519+ expanded = [
2520+ f"{ global_name } { expert_id } "
2521+ for expert_id in range (local_expert_count )
2522+ ]
2523+ grouped_misc_names [global_name ] = expanded
2524+ expanded_global_names .extend (expanded )
2525+ else :
2526+ expanded_global_names .append (global_name )
2527+ global_names = expanded_global_names
24772528 remaining_names = [name for name in global_names if name not in grouped_names ]
24782529 if not remaining_names :
24792530 return grouped_tasks
@@ -2501,6 +2552,35 @@ def _build_native_mxfp8_conversion_tasks(self) -> list[Any]:
25012552 global_name = _megatron_local_name_to_global (
25022553 models , self .model .config , local_name , 0
25032554 )
2555+ expanded_names = grouped_misc_names .get (global_name )
2556+ if expanded_names is not None :
2557+ local_module , local_weight = get_module_and_param_from_name (
2558+ models , local_name , 0
2559+ )
2560+ members = (
2561+ []
2562+ if local_weight is None
2563+ else _grouped_expert_member_views (local_weight )
2564+ )
2565+ if len (members ) != len (expanded_names ):
2566+ raise ValueError (
2567+ f"Grouped expert parameter { global_name !r} has local shape "
2568+ f"{ getattr (local_weight , 'shape' , None )} , expected "
2569+ f"{ len (expanded_names )} local experts"
2570+ )
2571+ if local_module is not None and not hasattr (local_module , "config" ):
2572+ setattr (local_module , "config" , self .model .config )
2573+ for expert_id , expanded_name in enumerate (expanded_names ):
2574+ local_tasks [expanded_name ] = WeightConversionTask (
2575+ pp_rank = pp_rank ,
2576+ vp_stage = 0 ,
2577+ param_name = f"{ local_name } { expert_id } " ,
2578+ global_param_name = expanded_name ,
2579+ megatron_module = local_module ,
2580+ param_weight = members [expert_id ],
2581+ mapping = mappings [expanded_name ],
2582+ )
2583+ continue
25042584 if global_name not in remaining_set :
25052585 continue
25062586 local_module , local_weight = get_module_and_param_from_name (
@@ -2887,8 +2967,18 @@ def _task_uses_native_mxfp8_storage(self, task: Any, *, grouped: bool) -> bool:
28872967 members = get_grouped_quantized_members (
28882968 task .param_weight , create_if_missing = False
28892969 )
2890- except (RuntimeError , ValueError ):
2891- return False
2970+ except RuntimeError :
2971+ members = get_grouped_quantized_members (
2972+ task .param_weight , create_if_missing = True
2973+ )
2974+ except ValueError as error :
2975+ logical_name = self ._native_task_projections (task , grouped = True )[0 ][
2976+ 0
2977+ ]
2978+ raise ValueError (
2979+ f"Invalid grouped MXFP8 source { logical_name !r} role "
2980+ f"'weight': { error } "
2981+ ) from error
28922982 if not members :
28932983 logical_name = self ._native_task_projections (task , grouped = True )[0 ][
28942984 0
0 commit comments