Skip to content

Commit ba2b83e

Browse files
justinchubyCopilot
andauthored
Migrate op_multi_out to __getattr__ dispatch with _outputs (#294)
## Problem `op.op_multi_out()` is being removed from onnxscript/onnx_ir. 9 call sites across 5 rewrite rule files use this API. ## Fix Replace all `op.op_multi_out()` calls with the standard `__getattr__` dispatch pattern using `_outputs=N`: ```python # Before outputs = op.op_multi_out( 'GroupQueryAttention', inputs=[q, k, v, past_k, past_v, seqlens, total_seq, cos, sin], domain='com.microsoft', attributes=gqa_attrs, num_outputs=3, ) # After outputs = op.GroupQueryAttention( q, k, v, past_k, past_v, seqlens, total_seq, cos, sin, _domain='com.microsoft', _outputs=3, **gqa_attrs, ) ``` ## Files Changed (5 files, 9 calls) - `_group_query_attention.py`: 4 calls - `_skip_layer_norm.py`: 2 calls - `_skip_norm.py`: 1 call - `_unpack_qkv.py`: 1 call - `_separate_rope.py`: 1 call ## Testing - ✅ All 2718 tests pass - ✅ Verified GQA (24 ops), SkipNorm (48 ops), RoPE separation all produce correct output on Qwen2.5-0.5B with CUDA EP --------- Signed-off-by: Justin Chu <justinchu@microsoft.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent 8f7c007 commit ba2b83e

5 files changed

Lines changed: 78 additions & 74 deletions

File tree

src/mobius/rewrite_rules/_group_query_attention.py

Lines changed: 40 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -244,22 +244,19 @@ def rewrite(
244244
}
245245
if softcap:
246246
gqa_attrs["softcap"] = softcap
247-
outputs = op.op_multi_out(
248-
"GroupQueryAttention",
249-
inputs=[
250-
q_pre,
251-
k_pre,
252-
v,
253-
past_key,
254-
past_value,
255-
self._seqlens_k,
256-
self._total_seq_len,
257-
self._cos_cache,
258-
self._sin_cache,
259-
],
260-
domain="com.microsoft",
261-
attributes=gqa_attrs,
262-
num_outputs=3,
247+
outputs = op.GroupQueryAttention(
248+
q_pre,
249+
k_pre,
250+
v,
251+
past_key,
252+
past_value,
253+
self._seqlens_k,
254+
self._total_seq_len,
255+
self._cos_cache,
256+
self._sin_cache,
257+
_domain="com.microsoft",
258+
_outputs=3,
259+
**gqa_attrs,
263260
)
264261

265262
return outputs[0], outputs[1], outputs[2]
@@ -377,17 +374,14 @@ def rewrite(
377374
gqa_node = gqa_out.producer()
378375
attrs = {key: gqa_node.attributes[key].value for key in gqa_node.attributes}
379376

380-
outputs = op.op_multi_out(
381-
"GroupQueryAttention",
382-
inputs=[
383-
packed_qkv,
384-
None,
385-
None,
386-
*gqa_node.inputs[3:],
387-
],
388-
domain="com.microsoft",
389-
attributes=attrs,
390-
num_outputs=3,
377+
outputs = op.GroupQueryAttention(
378+
packed_qkv,
379+
None,
380+
None,
381+
*gqa_node.inputs[3:],
382+
_domain="com.microsoft",
383+
_outputs=3,
384+
**attrs,
391385
)
392386

393387
return outputs[0], outputs[1], outputs[2]
@@ -507,17 +501,14 @@ def rewrite(
507501
gqa_node = gqa_out.producer()
508502
attrs = {key: gqa_node.attributes[key].value for key in gqa_node.attributes}
509503

510-
outputs = op.op_multi_out(
511-
"GroupQueryAttention",
512-
inputs=[
513-
packed_qkv,
514-
None,
515-
None,
516-
*gqa_node.inputs[3:],
517-
],
518-
domain="com.microsoft",
519-
attributes=attrs,
520-
num_outputs=3,
504+
outputs = op.GroupQueryAttention(
505+
packed_qkv,
506+
None,
507+
None,
508+
*gqa_node.inputs[3:],
509+
_domain="com.microsoft",
510+
_outputs=3,
511+
**attrs,
521512
)
522513

523514
return outputs[0], outputs[1], outputs[2]
@@ -655,12 +646,17 @@ def rewrite(
655646
if softcap:
656647
gqa_attrs["softcap"] = softcap
657648

658-
outputs = op.op_multi_out(
659-
"GroupQueryAttention",
660-
inputs=[q, k, v, past_key, past_value, self._seqlens_k, self._total_seq_len],
661-
domain="com.microsoft",
662-
attributes=gqa_attrs,
663-
num_outputs=3,
649+
outputs = op.GroupQueryAttention(
650+
q,
651+
k,
652+
v,
653+
past_key,
654+
past_value,
655+
self._seqlens_k,
656+
self._total_seq_len,
657+
_domain="com.microsoft",
658+
_outputs=3,
659+
**gqa_attrs,
664660
)
665661
return outputs[0], outputs[1], outputs[2]
666662

src/mobius/rewrite_rules/_separate_rope.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -133,12 +133,14 @@ def rewrite(self, op, q, k, v, gqa_out, present_key, present_value, **_):
133133
# past_key, past_value, seqlens_k, total_sequence_length.
134134
remaining = list(gqa_node.inputs[3:7])
135135

136-
outputs = op.op_multi_out(
137-
"GroupQueryAttention",
138-
inputs=[q_rot, k_rot, v, *remaining],
139-
domain="com.microsoft",
140-
attributes=attrs,
141-
num_outputs=3,
136+
outputs = op.GroupQueryAttention(
137+
q_rot,
138+
k_rot,
139+
v,
140+
*remaining,
141+
_domain="com.microsoft",
142+
_outputs=3,
143+
**attrs,
142144
)
143145
return outputs[0], outputs[1], outputs[2]
144146

src/mobius/rewrite_rules/_skip_layer_norm.py

Lines changed: 15 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -114,12 +114,14 @@ def rewrite(self, op, add_out, weight, bias, norm_out, **_):
114114
input_a = add_node.inputs[0]
115115
input_b = add_node.inputs[1]
116116

117-
outputs = op.op_multi_out(
118-
"SkipLayerNormalization",
119-
inputs=[input_a, input_b, weight, bias],
120-
domain="com.microsoft",
121-
attributes={"epsilon": epsilon},
122-
num_outputs=4,
117+
outputs = op.SkipLayerNormalization(
118+
input_a,
119+
input_b,
120+
weight,
121+
bias,
122+
_domain="com.microsoft",
123+
epsilon=epsilon,
124+
_outputs=4,
123125
)
124126
new_norm_out = outputs[0]
125127
skip_out = outputs[3]
@@ -200,12 +202,13 @@ def rewrite(self, op, add_out, weight, norm_out, **_):
200202
input_b = add_node.inputs[1]
201203

202204
# SkipLayerNormalization with gamma only (no beta)
203-
outputs = op.op_multi_out(
204-
"SkipLayerNormalization",
205-
inputs=[input_a, input_b, weight],
206-
domain="com.microsoft",
207-
attributes={"epsilon": epsilon},
208-
num_outputs=4,
205+
outputs = op.SkipLayerNormalization(
206+
input_a,
207+
input_b,
208+
weight,
209+
_domain="com.microsoft",
210+
epsilon=epsilon,
211+
_outputs=4,
209212
)
210213
new_norm_out = outputs[0]
211214
skip_out = outputs[3]

src/mobius/rewrite_rules/_skip_norm.py

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -102,12 +102,13 @@ def rewrite(self, op, add_out, weight, norm_out, **_):
102102
input_a = add_node.inputs[0]
103103
input_b = add_node.inputs[1]
104104

105-
outputs = op.op_multi_out(
106-
"SkipSimplifiedLayerNormalization",
107-
inputs=[input_a, input_b, weight],
108-
domain="com.microsoft",
109-
attributes={"epsilon": epsilon},
110-
num_outputs=4,
105+
outputs = op.SkipSimplifiedLayerNormalization(
106+
input_a,
107+
input_b,
108+
weight,
109+
_domain="com.microsoft",
110+
epsilon=epsilon,
111+
_outputs=4,
111112
)
112113
new_norm_out = outputs[0]
113114
skip_out = outputs[3]

src/mobius/rewrite_rules/_unpack_qkv.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -285,12 +285,14 @@ def _proj(w: np.ndarray, name: str) -> ir.Value:
285285
attrs = {key: gqa_node.attributes[key].value for key in gqa_node.attributes}
286286
remaining = list(gqa_node.inputs[3:]) # everything after the packed slot
287287

288-
outputs = op.op_multi_out(
289-
"GroupQueryAttention",
290-
inputs=[q, k, v, *remaining],
291-
domain="com.microsoft",
292-
attributes=attrs,
293-
num_outputs=3,
288+
outputs = op.GroupQueryAttention(
289+
q,
290+
k,
291+
v,
292+
*remaining,
293+
_domain="com.microsoft",
294+
_outputs=3,
295+
**attrs,
294296
)
295297
return outputs[0], outputs[1], outputs[2]
296298

0 commit comments

Comments
 (0)