Skip to content

Commit 2230e8e

Browse files
committed
add callback for fp8 sonicmoe
1 parent 463d28f commit 2230e8e

4 files changed

Lines changed: 52 additions & 6 deletions

File tree

paddleformers/trainer/trainer.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -173,6 +173,7 @@
173173
InterleaveGateUpCallback,
174174
PrinterCallback,
175175
ProgressCallback,
176+
SonicMoELayoutSwitchCallback,
176177
SPGradSyncCallback,
177178
TrainerCallback,
178179
TrainerControl,
@@ -1648,8 +1649,10 @@ def train(
16481649
self.add_non_zcc_ema_callback(resume_from_checkpoint)
16491650

16501651
if self.args.using_sonic_moe:
1651-
callback = InterleaveGateUpCallback(self.model, resume_from_checkpoint, self.args.output_dir)
1652-
self.add_callback(callback)
1652+
# callback = InterleaveGateUpCallback(self.model, resume_from_checkpoint, self.args.output_dir)
1653+
# self.add_callback(callback)
1654+
print("==== add sonicmoe callback ====")
1655+
self.add_callback(SonicMoELayoutSwitchCallback())
16531656

16541657
self.log_trainable_numel(model)
16551658

paddleformers/trainer/trainer_callback.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,13 +42,17 @@
4242
# Conditionally import paddlefleet modules
4343
if is_paddlefleet_available():
4444
from paddlefleet.models.gpt import GPTModel
45+
from paddlefleet.transformer.moe.moe_expert import SonicMoEExpert
4546
from paddlefleet.transformer.moe.moe_layer import MoELayer
4647
from paddlefleet.transformer.moe.moe_router import StandardMoERouter
4748
else:
4849

4950
class GPTModel:
5051
pass
5152

53+
class SonicMoEExpert:
54+
pass
55+
5256
class MoELayer:
5357
pass
5458

@@ -80,6 +84,7 @@ class StandardMoERouter:
8084
"MoEGateSpGradSyncCallBack",
8185
"SPGradSyncCallback",
8286
"EMAStateAssemblerCallback",
87+
"SonicMoELayoutSwitchCallback",
8388
]
8489

8590

@@ -715,6 +720,9 @@ def on_step_begin(self, args, state, control, **kwargs):
715720
"""
716721
Quantize expert weights to FP8 before each training step
717722
"""
723+
if args.using_sonic_moe:
724+
# sonicmoe cannot support offline quant now.
725+
return
718726
model = kwargs["model"]
719727
optimizer = kwargs["optimizer"]
720728
global skip_count
@@ -754,6 +762,9 @@ def on_optimizer_begin(self, args, state, control, **kwargs):
754762
"""
755763
Reload weights before optimizer step
756764
"""
765+
if args.using_sonic_moe:
766+
# sonicmoe cannot support offline quant now.
767+
return
757768
model = kwargs["model"]
758769
optimizer = kwargs["optimizer"]
759770
global skip_count
@@ -932,6 +943,32 @@ def on_step_end(self, args, state, control, **kwargs):
932943
logger.info(f"[EMAStateAssembler] Assembling EMA state took {duration:.3f} seconds.")
933944

934945

946+
class SonicMoELayoutSwitchCallback(TrainerCallback):
947+
def _apply_to_sonic_moe_experts(self, model, fn_name):
948+
def apply_layout_switch(layer):
949+
if isinstance(layer, SonicMoEExpert):
950+
getattr(layer, fn_name)()
951+
952+
model.apply(apply_layout_switch)
953+
954+
def _prepare_sonic_moe_fp8_weights(self, model):
955+
def prepare_fp8_weights(layer):
956+
if isinstance(layer, SonicMoEExpert):
957+
layer.convert_weights_to_sonic_layout()
958+
layer.quant_weight()
959+
960+
model.apply(prepare_fp8_weights)
961+
962+
def on_step_begin(self, args, state, control, **kwargs):
963+
if args.using_sonic_moe:
964+
self._prepare_sonic_moe_fp8_weights(kwargs["model"])
965+
# kwargs["optimizer"].clear_param_storage("moe_expert")
966+
967+
def on_optimizer_begin(self, args, state, control, **kwargs):
968+
if args.using_sonic_moe:
969+
self._apply_to_sonic_moe_experts(kwargs["model"], "convert_weights_to_grouped_layout")
970+
971+
935972
class InterleaveGateUpCallback(TrainerCallback):
936973
def __init__(self, model, resume_from_checkpoint=None, output_dir=None):
937974
self.model = model

paddleformers/transformers/configuration_utils.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -444,6 +444,12 @@ class LlmMetaConfig:
444444
True,
445445
"Whether to use FP8 for gradient storage during training (only effective if `fp8=True`). Further reduces memory footprint but may introduce minor numerical error. Defaults to False.",
446446
),
447+
(
448+
"use_ue8m0",
449+
bool,
450+
False,
451+
"Whether to use UE8M0 packed scaling factors for FP8 on Blackwell GPUs (SM100+). Enables deep_gemm backend for weight gradient computation. Defaults to False.",
452+
),
447453
]
448454

449455
model_conf = [

paddleformers/transformers/qwen3_moe/modeling.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -820,8 +820,8 @@ def _gen_aoa_config(cls, config: Qwen3MoeConfig):
820820
group_gemm1 = ",".join(ep_weight1)
821821
group_gemm2 = ",".join(ep_weight2)
822822
aoa_config["aoa_statements"] += [
823-
f"{group_gemm1} -> {tgt_prefix}.mlp.grouped_gemm_experts.weight1, axis=0"
824-
f"{group_gemm2} -> {tgt_prefix}.mlp.grouped_gemm_experts.weight2, axis=0"
823+
f"{group_gemm1} -> {tgt_prefix}.mlp.grouped_gemm_experts.weight1, axis=0",
824+
f"{group_gemm2} -> {tgt_prefix}.mlp.grouped_gemm_experts.weight2, axis=0",
825825
]
826826
else:
827827
if config.get("fd_fallback", False):
@@ -903,8 +903,8 @@ def _gen_inv_aoa_config(cls, config: Qwen3MoeConfig):
903903
group_gemm1 = ",".join(ep_weight1)
904904
group_gemm2 = ",".join(ep_weight2)
905905
aoa_statements += [
906-
f"{model_prefix}layers.{layer_id}.mlp.grouped_gemm_experts.weight1 -> {group_gemm1}, axis=0"
907-
f"{model_prefix}layers.{layer_id}.mlp.grouped_gemm_experts.weight2 -> {group_gemm2}, axis=0"
906+
f"{model_prefix}layers.{layer_id}.mlp.grouped_gemm_experts.weight1 -> {group_gemm1}, axis=0",
907+
f"{model_prefix}layers.{layer_id}.mlp.grouped_gemm_experts.weight2 -> {group_gemm2}, axis=0",
908908
]
909909
else:
910910
if config.get("fd_fallback", False):

0 commit comments

Comments
 (0)