Skip to content

Commit f620daf

Browse files
authored
Add SimplifiedLayerNormToRMSNorm surgery (microsoft#2348)
## Describe your changes ## Checklist before requesting a review - [x] Add unit tests for this change. - [x] Make sure all tests can pass. - [ ] Update documents if necessary. - [x] Lint and apply fixes to your code by running `lintrunner -a` - [ ] Is this a user-facing change? If yes, give a description of this change to be included in the release notes. ## (Optional) Issue link
1 parent 9ef9d0d commit f620daf

2 files changed

Lines changed: 368 additions & 0 deletions

File tree

olive/passes/onnx/graph_surgeries.py

Lines changed: 219 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
8531072
class SimplifiedLayerNormToL2Norm(ProtoSurgeon):
8541073
"""Replace Skip/SimplifiedLayerNormalization node with L2Norm subgraph.
8551074

test/passes/onnx/test_graph_surgeries.py

Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -588,6 +588,155 @@ def test_simplifiedlayernorm_to_l2norm_skip(tmp_path, all_ones, output_skip_sum)
588588
)
589589

590590

591+
def check_rmsnorm(
592+
original_model_path: str,
593+
modified_model_path: str,
594+
hidden_size: int,
595+
expected_num_nodes: int,
596+
has_skip: bool = False,
597+
):
598+
# check output values match
599+
input_session = InferenceSession(original_model_path)
600+
output_session = InferenceSession(modified_model_path)
601+
input_feed = {"x": np.random.randn(1, hidden_size).astype(np.float32)}
602+
if has_skip:
603+
input_feed["skip"] = np.random.randn(1, hidden_size).astype(np.float32)
604+
input_result = input_session.run(None, input_feed)
605+
output_result = output_session.run(None, input_feed)
606+
for i_r, o_r in zip(input_result, output_result):
607+
np.testing.assert_allclose(i_r, o_r, rtol=1e-3, atol=1e-3)
608+
609+
# count nodes and verify expected op types are present
610+
dag = OnnxDAG.from_model_path(modified_model_path)
611+
assert len(dag.nodes) == expected_num_nodes
612+
op_types = dag.get_node_op_types()
613+
assert "Pow" in op_types
614+
assert "ReduceMean" in op_types
615+
assert "Sqrt" in op_types
616+
assert "Div" in op_types
617+
assert "Mul" in op_types
618+
assert "SimplifiedLayerNormalization" not in op_types
619+
assert "SkipSimplifiedLayerNormalization" not in op_types
620+
621+
622+
@pytest.mark.parametrize("all_ones", [True, False])
623+
def test_simplifiedlayernorm_to_rmsnorm(tmp_path, all_ones):
624+
# setup
625+
hidden_size = 3
626+
inputs = [
627+
onnx.helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, hidden_size]),
628+
]
629+
outputs = [
630+
onnx.helper.make_tensor_value_info("y", TensorProto.FLOAT, [1, hidden_size]),
631+
]
632+
weight = (np.ones(hidden_size) if all_ones else np.random.randn(hidden_size)).astype(np.float32)
633+
initializers = [onnx.numpy_helper.from_array(weight, name="weight")]
634+
nodes = [
635+
onnx.helper.make_node(
636+
"SimplifiedLayerNormalization",
637+
inputs=["x", "weight"],
638+
outputs=["layernorm_output"],
639+
name="layernorm/LayerNorm",
640+
),
641+
onnx.helper.make_node("Identity", inputs=["layernorm_output"], outputs=["y"], name="Identity"),
642+
]
643+
graph = helper.make_graph(
644+
nodes=nodes,
645+
name="TestGraph",
646+
inputs=inputs,
647+
outputs=outputs,
648+
initializer=initializers,
649+
)
650+
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 20)])
651+
model.ir_version = 10
652+
onnx.save(model, str(tmp_path / "input_model.onnx"))
653+
input_model = ONNXModelHandler(model_path=str(tmp_path / "input_model.onnx"))
654+
655+
output_folder = str(tmp_path / "output")
656+
p = create_pass_from_dict(
657+
GraphSurgeries,
658+
{"surgeries": [{"surgeon": "SimplifiedLayerNormToRMSNorm"}]},
659+
disable_search=True,
660+
)
661+
662+
# execute
663+
onnx_model = p.run(input_model, output_folder)
664+
665+
# assert
666+
# Pow, ReduceMean, Add(eps), Sqrt, Div, Mul, Identity = 7 nodes
667+
check_rmsnorm(str(tmp_path / "input_model.onnx"), onnx_model.model_path, hidden_size, 7)
668+
669+
670+
@pytest.mark.parametrize("all_ones", [True, False])
671+
@pytest.mark.parametrize("output_skip_sum", [True, False])
672+
def test_simplifiedlayernorm_to_rmsnorm_skip(tmp_path, all_ones, output_skip_sum):
673+
# setup
674+
hidden_size = 3
675+
inputs = [
676+
onnx.helper.make_tensor_value_info("x", TensorProto.FLOAT, [1, hidden_size]),
677+
onnx.helper.make_tensor_value_info("skip", TensorProto.FLOAT, [1, hidden_size]),
678+
]
679+
outputs = [
680+
onnx.helper.make_tensor_value_info("y", TensorProto.FLOAT, [1, hidden_size]),
681+
]
682+
if output_skip_sum:
683+
outputs.append(
684+
onnx.helper.make_tensor_value_info("skip_sum", TensorProto.FLOAT, [1, hidden_size]),
685+
)
686+
initializers = [
687+
onnx.numpy_helper.from_array(
688+
(np.ones(hidden_size) if all_ones else np.random.randn(hidden_size)).astype(np.float32), name="weight"
689+
)
690+
]
691+
nodes = [
692+
onnx.helper.make_node(
693+
"SkipSimplifiedLayerNormalization",
694+
inputs=["x", "skip", "weight"],
695+
outputs=["layernorm_output"] if not output_skip_sum else ["layernorm_output", "", "", "layernorm_skip_sum"],
696+
name="layernorm/LayerNorm",
697+
domain=MSFT_DOMAIN,
698+
),
699+
onnx.helper.make_node("Identity", inputs=["layernorm_output"], outputs=["y"], name="Identity"),
700+
]
701+
if output_skip_sum:
702+
nodes.append(
703+
onnx.helper.make_node(
704+
"Identity", inputs=["layernorm_skip_sum"], outputs=["skip_sum"], name="Identity_skip_sum"
705+
)
706+
)
707+
graph = helper.make_graph(
708+
nodes=nodes,
709+
name="TestGraph",
710+
inputs=inputs,
711+
outputs=outputs,
712+
initializer=initializers,
713+
)
714+
model = helper.make_model(graph, opset_imports=[helper.make_opsetid("", 20)])
715+
model.ir_version = 10
716+
onnx.save(model, str(tmp_path / "input_model.onnx"))
717+
input_model = ONNXModelHandler(model_path=str(tmp_path / "input_model.onnx"))
718+
719+
output_folder = str(tmp_path / "output")
720+
p = create_pass_from_dict(
721+
GraphSurgeries,
722+
{"surgeries": [{"surgeon": "SimplifiedLayerNormToRMSNorm"}]},
723+
disable_search=True,
724+
)
725+
726+
# execute
727+
output_model = p.run(input_model, output_folder)
728+
729+
# assert
730+
# Add(skip), Pow, ReduceMean, Add(eps), Sqrt, Div, Mul, Identity[, Identity_skip_sum] = 8 or 9 nodes
731+
check_rmsnorm(
732+
str(tmp_path / "input_model.onnx"),
733+
output_model.model_path,
734+
hidden_size,
735+
8 + int(output_skip_sum),
736+
has_skip=True,
737+
)
738+
739+
591740
@pytest.mark.parametrize("use_large_cache", [True, False])
592741
def test_remove_rope_multi_cache(tmp_path, use_large_cache):
593742
# setup

0 commit comments

Comments
 (0)