feat: add EP support in tuning (#12012)
This commit is contained in:
@@ -421,69 +421,88 @@ def main(args: argparse.Namespace):
|
|||||||
|
|
||||||
config = AutoConfig.from_pretrained(args.model, trust_remote_code=True)
|
config = AutoConfig.from_pretrained(args.model, trust_remote_code=True)
|
||||||
if config.architectures[0] == "DbrxForCausalLM":
|
if config.architectures[0] == "DbrxForCausalLM":
|
||||||
E = config.ffn_config.moe_num_experts
|
E = config.ffn_config.moe_num_experts // args.ep_size
|
||||||
topk = config.ffn_config.moe_top_k
|
topk = config.ffn_config.moe_top_k
|
||||||
intermediate_size = config.ffn_config.ffn_hidden_size
|
intermediate_size = config.ffn_config.ffn_hidden_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
elif config.architectures[0] == "JambaForCausalLM":
|
elif config.architectures[0] == "JambaForCausalLM":
|
||||||
E = config.num_experts
|
E = config.num_experts // args.ep_size
|
||||||
topk = config.num_experts_per_tok
|
topk = config.num_experts_per_tok
|
||||||
intermediate_size = config.intermediate_size
|
intermediate_size = config.intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
elif config.architectures[0] in [
|
elif config.architectures[0] in [
|
||||||
"Qwen2MoeForCausalLM",
|
"Qwen2MoeForCausalLM",
|
||||||
"Qwen3MoeForCausalLM",
|
"Qwen3MoeForCausalLM",
|
||||||
"Qwen3NextForCausalLM",
|
"Qwen3NextForCausalLM",
|
||||||
]:
|
]:
|
||||||
E = config.num_experts
|
E = config.num_experts // args.ep_size
|
||||||
topk = config.num_experts_per_tok
|
topk = config.num_experts_per_tok
|
||||||
intermediate_size = config.moe_intermediate_size
|
intermediate_size = config.moe_intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
elif config.architectures[0] in ["DeepseekV2ForCausalLM", "DeepseekV3ForCausalLM"]:
|
elif config.architectures[0] in ["DeepseekV2ForCausalLM", "DeepseekV3ForCausalLM"]:
|
||||||
E = (
|
E = (config.n_routed_experts // args.ep_size) + (
|
||||||
config.n_routed_experts + (0 if args.disable_shared_experts_fusion else 1)
|
0
|
||||||
if config.architectures[0] in ["DeepseekV3ForCausalLM"]
|
if args.disable_shared_experts_fusion
|
||||||
else config.n_routed_experts
|
or config.architectures[0] not in ["DeepseekV3ForCausalLM"]
|
||||||
|
else 1
|
||||||
)
|
)
|
||||||
topk = config.num_experts_per_tok
|
topk = config.num_experts_per_tok
|
||||||
intermediate_size = config.moe_intermediate_size
|
intermediate_size = config.moe_intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
elif config.architectures[0] == "Llama4ForConditionalGeneration":
|
elif config.architectures[0] == "Llama4ForConditionalGeneration":
|
||||||
E = config.text_config.num_local_experts + (
|
E = config.text_config.num_local_experts // args.ep_size + (
|
||||||
0 if args.disable_shared_experts_fusion else 1
|
0 if args.disable_shared_experts_fusion else 1
|
||||||
)
|
)
|
||||||
topk = config.text_config.num_experts_per_tok
|
topk = config.text_config.num_experts_per_tok
|
||||||
intermediate_size = config.text_config.intermediate_size
|
intermediate_size = config.text_config.intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
elif config.architectures[0] in [
|
elif config.architectures[0] in [
|
||||||
"Grok1ForCausalLM",
|
"Grok1ForCausalLM",
|
||||||
"Grok1ImgGen",
|
"Grok1ImgGen",
|
||||||
"Grok1AForCausalLM",
|
"Grok1AForCausalLM",
|
||||||
]:
|
]:
|
||||||
E = config.num_local_experts
|
E = config.num_local_experts // args.ep_size
|
||||||
topk = config.num_experts_per_tok
|
topk = config.num_experts_per_tok
|
||||||
intermediate_size = config.moe_intermediate_size
|
intermediate_size = config.moe_intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
elif config.architectures[0] in [
|
elif config.architectures[0] in [
|
||||||
"BailingMoEForCausalLM",
|
"BailingMoEForCausalLM",
|
||||||
"BailingMoeForCausalLM",
|
"BailingMoeForCausalLM",
|
||||||
"BailingMoeV2ForCausalLM",
|
"BailingMoeV2ForCausalLM",
|
||||||
]:
|
]:
|
||||||
E = config.num_experts
|
E = config.num_experts // args.ep_size
|
||||||
topk = config.num_experts_per_tok
|
topk = config.num_experts_per_tok
|
||||||
intermediate_size = config.moe_intermediate_size
|
intermediate_size = config.moe_intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
elif config.architectures[0] in ["Glm4MoeForCausalLM"]:
|
elif config.architectures[0] in ["Glm4MoeForCausalLM"]:
|
||||||
E = config.n_routed_experts
|
E = config.n_routed_experts // args.ep_size
|
||||||
topk = config.num_experts_per_tok
|
topk = config.num_experts_per_tok
|
||||||
intermediate_size = config.moe_intermediate_size
|
intermediate_size = config.moe_intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# Default: Mixtral
|
# Default: Mixtral
|
||||||
E = config.num_local_experts
|
E = config.num_local_experts // args.ep_size
|
||||||
topk = config.num_experts_per_tok
|
topk = config.num_experts_per_tok
|
||||||
intermediate_size = config.intermediate_size
|
intermediate_size = config.intermediate_size
|
||||||
shard_intermediate_size = 2 * intermediate_size // args.tp_size
|
shard_intermediate_size = (
|
||||||
|
2 * intermediate_size // (args.tp_size // args.ep_size)
|
||||||
|
)
|
||||||
|
|
||||||
hidden_size = getattr(config, "hidden_size", None) or config.text_config.hidden_size
|
hidden_size = getattr(config, "hidden_size", None) or config.text_config.hidden_size
|
||||||
dtype = config.torch_dtype
|
dtype = config.torch_dtype
|
||||||
@@ -626,6 +645,7 @@ if __name__ == "__main__":
|
|||||||
"--model", type=str, default="mistralai/Mixtral-8x7B-Instruct-v0.1"
|
"--model", type=str, default="mistralai/Mixtral-8x7B-Instruct-v0.1"
|
||||||
)
|
)
|
||||||
parser.add_argument("--tp-size", "--tp", type=int, default=2)
|
parser.add_argument("--tp-size", "--tp", type=int, default=2)
|
||||||
|
parser.add_argument("--ep-size", "--ep", type=int, default=1)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--dtype",
|
"--dtype",
|
||||||
type=str,
|
type=str,
|
||||||
|
|||||||
Reference in New Issue
Block a user