@@ -850,6 +850,225 @@ def get_rmsnorm_nodes(pow_node: str, dag: OnnxDAG) -> list[str] | None:
850850 return rmsnorm_nodes if len (rmsnorm_nodes ) >= (len (pattern ) - 1 ) else []
851851
852852
853+ class SimplifiedLayerNormToRMSNorm (ProtoSurgeon ):
854+ """Replace SimplifiedLayerNormalization or SkipSimplifiedLayerNormalization with an RMSNorm subgraph built from elementwise ops.
855+
856+ RMS(x) = sqrt(mean(x^2, axis=-1, keepdims=1) + eps)
857+ y = (x / RMS(x)) * gamma
858+
859+ For SkipSimplifiedLayerNormalization, we first do:
860+ s = input + skip
861+ and use 's' as x for RMSNorm. If the original node exposes a second output
862+ (residual sum), we rewire its consumers to 's' to preserve graph behavior.
863+
864+ IMPORTANT: ReduceMean schema change across opsets:
865+ - opset < 18: axes is an ATTRIBUTE
866+ - opset >=18: axes is an INPUT tensor (int64), keepdims remains an attribute.
867+ """
868+
869+ def __call__ (self , model : onnx .ModelProto ):
870+ from onnx import numpy_helper
871+ from onnx .helper import tensor_dtype_to_np_dtype
872+
873+ dag = OnnxDAG (model )
874+
875+ # Determine the default ONNX opset for the main domain ("", "ai.onnx").
876+ # We'll use this to decide how to build ReduceMean.
877+ default_opset = None
878+ for imp in model .opset_import :
879+ if imp .domain in ("" , "ai.onnx" ):
880+ default_opset = imp .version
881+ break
882+ if default_opset is None :
883+ # Fall back defensively; most models have a default import.
884+ default_opset = 13
885+
886+ use_axes_input_for_reduce_mean = default_opset >= 18
887+
888+ modified = 0
889+
890+ for node_name in dag .get_node_names ():
891+ op_type = dag .get_node_op_type (node_name )
892+ if op_type not in {"SimplifiedLayerNormalization" , "SkipSimplifiedLayerNormalization" }:
893+ continue
894+
895+ graph_idx = dag .get_graph_idx (node_name )
896+ inputs = dag .get_node_inputs (node_name , True )
897+ outputs = dag .get_node_outputs (node_name , True )
898+
899+ # ---------------------------
900+ # Build the input to be normalized: ln_input
901+ # ---------------------------
902+ if op_type == "SkipSimplifiedLayerNormalization" :
903+ # Expect inputs: [input, skip, gamma]
904+ if len (inputs ) != 3 :
905+ continue
906+ root1 , root2 , gamma = inputs
907+
908+ # Add(input, skip) => skip_add_out
909+ skip_add_name = self .create_new_name (node_name , op_type , "Add" )
910+ skip_add_out = f"{ skip_add_name } _out"
911+ skip_add_node = onnx .helper .make_node (
912+ "Add" ,
913+ inputs = [root1 , root2 ],
914+ outputs = [skip_add_out ],
915+ name = skip_add_name ,
916+ )
917+ dag .add_node (skip_add_node , graph_idx )
918+
919+ ln_input = skip_add_out
920+ else :
921+ # SimplifiedLayerNormalization: inputs = [x, gamma]
922+ if len (inputs ) != 2 :
923+ continue
924+ ln_input , gamma = inputs
925+
926+ # The original primary output (normalized tensor)
927+ ln_output = outputs [0 ]
928+
929+ ln_elem_type = dag .get_io_elem_type (inputs [0 ]) or onnx .TensorProto .FLOAT
930+ ln_np_dtype = tensor_dtype_to_np_dtype (ln_elem_type )
931+
932+ # ---------------------------
933+ # Step 1: Pow(x, 2)
934+ # ---------------------------
935+ pow_name = self .create_new_name (node_name , op_type , "Pow" )
936+ pow_out = f"{ pow_name } _out"
937+ pow_const = numpy_helper .from_array (np .array ([2.0 ], dtype = ln_np_dtype ), name = f"{ pow_name } _const" )
938+ dag .add_initializer (pow_const , graph_idx )
939+ pow_node = onnx .helper .make_node (
940+ "Pow" ,
941+ inputs = [ln_input , pow_const .name ],
942+ outputs = [pow_out ],
943+ name = pow_name ,
944+ )
945+ dag .add_node (pow_node , graph_idx )
946+
947+ # ---------------------------
948+ # Step 2: ReduceMean over last dim, keepdims=1
949+ # - opset < 18 : axes is an attribute
950+ # - opset >= 18: axes is an input tensor (INT64)
951+ # ---------------------------
952+ mean_name = self .create_new_name (node_name , op_type , "ReduceMean" )
953+ mean_out = f"{ mean_name } _out"
954+
955+ if use_axes_input_for_reduce_mean :
956+ axes_init = numpy_helper .from_array (np .array ([- 1 ], dtype = np .int64 ), name = f"{ mean_name } _axes" )
957+ dag .add_initializer (axes_init , graph_idx )
958+
959+ mean_node = onnx .helper .make_node (
960+ "ReduceMean" ,
961+ inputs = [pow_out , axes_init .name ],
962+ outputs = [mean_out ],
963+ name = mean_name ,
964+ keepdims = 1 ,
965+ )
966+ else :
967+ # Older schema: axes is an attribute
968+ mean_node = onnx .helper .make_node (
969+ "ReduceMean" ,
970+ inputs = [pow_out ],
971+ outputs = [mean_out ],
972+ name = mean_name ,
973+ axes = [- 1 ],
974+ keepdims = 1 ,
975+ )
976+ dag .add_node (mean_node , graph_idx )
977+
978+ # ---------------------------
979+ # Step 3: Add epsilon
980+ # ---------------------------
981+ eps_value = 1e-06
982+ add_eps_name = self .create_new_name (node_name , op_type , "AddEps" )
983+ add_eps_out = f"{ add_eps_name } _out"
984+
985+ eps_const = numpy_helper .from_array (np .array ([eps_value ], dtype = ln_np_dtype ), name = f"{ add_eps_name } _const" )
986+ dag .add_initializer (eps_const , graph_idx )
987+
988+ add_eps_node = onnx .helper .make_node (
989+ "Add" ,
990+ inputs = [mean_out , eps_const .name ],
991+ outputs = [add_eps_out ],
992+ name = add_eps_name ,
993+ )
994+ dag .add_node (add_eps_node , graph_idx )
995+
996+ # ---------------------------
997+ # Step 4: Sqrt
998+ # ---------------------------
999+ sqrt_name = self .create_new_name (node_name , op_type , "Sqrt" )
1000+ sqrt_out = f"{ sqrt_name } _out"
1001+ sqrt_node = onnx .helper .make_node (
1002+ "Sqrt" ,
1003+ inputs = [add_eps_out ],
1004+ outputs = [sqrt_out ],
1005+ name = sqrt_name ,
1006+ )
1007+ dag .add_node (sqrt_node , graph_idx )
1008+
1009+ # ---------------------------
1010+ # Step 5: Div (x / sqrt(...))
1011+ # ---------------------------
1012+ div_name = self .create_new_name (node_name , op_type , "Div" )
1013+ div_out = f"{ div_name } _out"
1014+ div_node = onnx .helper .make_node (
1015+ "Div" ,
1016+ inputs = [ln_input , sqrt_out ],
1017+ outputs = [div_out ],
1018+ name = div_name ,
1019+ )
1020+ dag .add_node (div_node , graph_idx )
1021+
1022+ # ---------------------------
1023+ # Step 6: Mul with gamma
1024+ # ---------------------------
1025+ mul_name = self .create_new_name (node_name , op_type , "Mul" )
1026+ mul_out = f"{ mul_name } _out"
1027+ mul_node = onnx .helper .make_node (
1028+ "Mul" ,
1029+ inputs = [div_out , gamma ],
1030+ outputs = [mul_out ],
1031+ name = mul_name ,
1032+ )
1033+ dag .add_node (mul_node , graph_idx )
1034+
1035+ # ---------------------------
1036+ # Rewire consumers of the original main output
1037+ # ---------------------------
1038+ for consumer in dag .get_consumers (ln_output ):
1039+ dag .replace_node_input (consumer , ln_output , mul_out )
1040+
1041+ # ---------------------------
1042+ # For SkipSimplifiedLayerNormalization that had two outputs:
1043+ # - Output 1 is typically residual sum (input_skip_bias_sum)
1044+ # - Redirect its consumers to the skip-sum Add output
1045+ # ---------------------------
1046+ if op_type == "SkipSimplifiedLayerNormalization" and len (outputs ) == 2 :
1047+ second_output = outputs [1 ]
1048+
1049+ second_vi = dag .get_value_info_proto (second_output )
1050+ if second_vi is not None :
1051+ new_vi = onnx .ValueInfoProto ()
1052+ new_vi .CopyFrom (second_vi )
1053+ new_vi .name = skip_add_out
1054+ dag .add_value_info (new_vi , graph_idx )
1055+
1056+ # Redirect all consumers of the second output
1057+ for consumer in dag .get_consumers (second_output ):
1058+ dag .replace_node_input (consumer , second_output , skip_add_out )
1059+
1060+ dag .remove_node (node_name )
1061+ modified += 1
1062+
1063+ if modified > 0 :
1064+ logger .debug (
1065+ "Replaced %d Simplified/SkipSimplifiedLayerNormalization nodes with RMSNorm subgraphs" , modified
1066+ )
1067+
1068+ dag .update ()
1069+ return dag .model
1070+
1071+
8531072class SimplifiedLayerNormToL2Norm (ProtoSurgeon ):
8541073 """Replace Skip/SimplifiedLayerNormalization node with L2Norm subgraph.
8551074
0 commit comments