@@ -585,3 +585,110 @@ def test_no_shared_expert_is_none(self) -> None:
585585 wrapper .m = moe
586586 replace_moe_with_quantized_op (wrapper , group_size = 32 , weight_nbit = 4 )
587587 self .assertIsNone (wrapper .m .shared_expert )
588+
589+
590+ class TestExportPipelineWiring (unittest .TestCase ):
591+ """The export pipeline correctly includes the MoE transform."""
592+
593+ def test_get_source_transforms_includes_moe_when_enabled (self ) -> None :
594+ from functools import partial
595+
596+ from executorch .examples .models .llama .export_llama_lib import (
597+ _get_source_transforms ,
598+ )
599+
600+ transforms = _get_source_transforms (
601+ dtype_override = torch .float32 , use_moe_quantized_op = True
602+ )
603+ moe_transforms = [
604+ t
605+ for t in transforms
606+ if isinstance (t , partial )
607+ and t .func .__name__ == "replace_moe_with_quantized_op"
608+ ]
609+ self .assertEqual (len (moe_transforms ), 1 )
610+ self .assertEqual (moe_transforms [0 ].keywords ["group_size" ], 32 )
611+ self .assertEqual (moe_transforms [0 ].keywords ["weight_nbit" ], 4 )
612+
613+ def test_get_source_transforms_excludes_moe_when_disabled (self ) -> None :
614+ from functools import partial
615+
616+ from executorch .examples .models .llama .export_llama_lib import (
617+ _get_source_transforms ,
618+ )
619+
620+ transforms = _get_source_transforms (
621+ dtype_override = torch .float32 , use_moe_quantized_op = False
622+ )
623+ moe_transforms = [
624+ t
625+ for t in transforms
626+ if isinstance (t , partial )
627+ and hasattr (t .func , "__name__" )
628+ and t .func .__name__ == "replace_moe_with_quantized_op"
629+ ]
630+ self .assertEqual (len (moe_transforms ), 0 )
631+
632+ def test_get_source_transforms_passes_custom_group_size (self ) -> None :
633+ from functools import partial
634+
635+ from executorch .examples .models .llama .export_llama_lib import (
636+ _get_source_transforms ,
637+ )
638+
639+ transforms = _get_source_transforms (
640+ dtype_override = torch .float32 ,
641+ use_moe_quantized_op = True ,
642+ group_size = 64 ,
643+ )
644+ moe_transforms = [
645+ t
646+ for t in transforms
647+ if isinstance (t , partial )
648+ and t .func .__name__ == "replace_moe_with_quantized_op"
649+ ]
650+ self .assertEqual (len (moe_transforms ), 1 )
651+ self .assertEqual (moe_transforms [0 ].keywords ["group_size" ], 64 )
652+
653+ def test_sentinel_op_is_registered (self ) -> None :
654+ self .assertTrue (hasattr (torch .ops .llama , "_quantized_moe_ffn_active" ))
655+ self .assertTrue (torch .ops .llama ._quantized_moe_ffn_active ())
656+
657+
658+ class TestLlmConfigMoeFlag (unittest .TestCase ):
659+ """llm_config wires the --use_moe_quantized_op flag correctly."""
660+
661+ def test_from_args_sets_flag (self ) -> None :
662+ import argparse
663+
664+ from executorch .extension .llm .export .config .llm_config import LlmConfig
665+
666+ args = argparse .Namespace (use_moe_quantized_op = True )
667+ config = LlmConfig .from_args (args )
668+ self .assertTrue (config .model .use_moe_quantized_op )
669+
670+ def test_default_is_false (self ) -> None :
671+ from executorch .extension .llm .export .config .llm_config import ModelConfig
672+
673+ self .assertFalse (ModelConfig ().use_moe_quantized_op )
674+
675+ def test_from_args_missing_field_defaults_false (self ) -> None :
676+ import argparse
677+
678+ from executorch .extension .llm .export .config .llm_config import LlmConfig
679+
680+ args = argparse .Namespace ()
681+ config = LlmConfig .from_args (args )
682+ self .assertFalse (config .model .use_moe_quantized_op )
683+
684+ def test_argparser_flag_true (self ) -> None :
685+ from executorch .examples .models .llama .export_llama_lib import build_args_parser
686+
687+ args = build_args_parser ().parse_args (["--use_moe_quantized_op" ])
688+ self .assertTrue (args .use_moe_quantized_op )
689+
690+ def test_argparser_flag_default_false (self ) -> None :
691+ from executorch .examples .models .llama .export_llama_lib import build_args_parser
692+
693+ args = build_args_parser ().parse_args ([])
694+ self .assertFalse (args .use_moe_quantized_op )
0 commit comments