Skip to content

Split symbolic dimension pass - #5123

Open
shivadbhavsar wants to merge 71 commits into
developfrom
split_sym_dim
Open

shivadbhavsar wants to merge 71 commits into
developfrom
split_sym_dim

Conversation

@shivadbhavsar

@shivadbhavsar shivadbhavsar commented Aug 7, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

Currently we only handle the dynamic batch case for using select modules, this creates the framework for generalizing for all axes.

Technical Details

The padding and slicing semantics are defined in terms of op-families where each family of ops (ie. convolution, pointwise, gemm, reduce, etc.) defines the characteristics of its input and output axes. These characteristics are then use to materialize pad-slice wrappings for each operator and then used to coalesce chains of operations that have compatible padding/masking characteristics.

1. Normalize runtime shape dependencies

Shape expressions are rewritten to read symbolic values directly from module parameters.

2. Analyze symbolic operations

Each symbolic instruction is classified to determine:

  • Which axes depend on symbolic dimensions.
  • Which inputs require padding.
  • Which values require masking.
  • Which shape-only inputs can be removed from clones.
  • Whether the operator needs rewriting with fixed attributes.
  • Whether the operation can be safely specialized.

3. Collect symbolic roots

collect_roots() finds the independent symbols introduced by parameter dimensions.
For each root, it:

  • Validates that every occurrence has the same interval and optimal values.
  • Divides the interval into specialization buckets.
  • Assigns each bucket a fixed target extent.

4. Discover specialization blocks

discover_blocks() first creates one candidate block for every operation that requires padding or static rewriting.
It determines:

  • Which symbolic roots the operation depends on.
  • How many clones those roots would produce.

5. Coalesce blocks

Coalescing combines candidate blocks so a connected region can use one select_module instead of dispatching every operation separately.
Two blocks are merged only when:

  • Their combined dynamic dependencies are closed.
  • Their required root specializations are compatible.
  • Their Cartesian product does not exceed max_clones.
    Connected producer-consumer blocks are merged first. Independent blocks are then merged when they can safely share the same dispatch.

6. Prepare clone boundaries

prepare_clone_infos() determines how values cross block boundaries.
It records:

  • Which inputs must be sliced.
  • Which inputs must be padded.
  • Which masks must be inserted.
  • Which internal edges already have the target extent.
  • The target-substituted output shape used by dispatch.

7. Materialize specialization clones

Materialization happens inside specialize_blocks() through build_clone().
For every combination of root buckets, it:

  • Creates a clone module.
  • Gives parameters the runtime range accepted by that clone.
  • Freezes symbolic target extents to the bucket target.
  • Inserts required padding and masks.
  • Rewrites shape-dependent operators to fixed forms.
  • Verifies that emitted operation shapes are static.
    The runtime shape inputs remain responsible for selecting the compatible clone and evaluating any required runtime masks.

8. Insert dispatch and reconnect the graph

specialize_blocks() then:

  • Inserts a select_module referencing the generated clones.
  • Passes the original runtime arguments into the selector.
  • Extracts the selected outputs.
  • Slices padded outputs back to their actual runtime extents.
  • Replaces uses of the original symbolic block.
  • Rewrites the module return and removes dead instructions.

Changelog Category

Add a CHANGELOG.md entry for any option other than Not Applicable

    • Added: New functionality.
    • Changed: Changes to existing functionality.
    • Removed: Functionality or support that has been removed. (Compared to a previous release)
    • Optimized: Component performance that has been optimized or improved.
    • Resolved Issues: Known issues from a previous version that have been resolved.
    • Not Applicable: This PR is not to be included in the changelog.

Follow the LLVM AI Tool Use Policy for contributions using AI.

@gh-app-migraphx-bot-pr-write

gh-app-migraphx-bot-pr-write Bot commented Aug 8, 2026 •

Copy link
Copy Markdown
Test Batch New Rate (4ed7b4) Old Rate (3a503c)* Diff Status
torchvision-resnet50 64 3,300.06 3,264.92 1.08% ✅
torchvision-resnet50_fp16 64 7,881.82 7,548.67 4.41% ✅
torchvision-densenet121 32 2,491.36 2,483.99 0.30% ✅
torchvision-densenet121_fp16 32 5,012.77 5,004.24 0.17% ✅
torchvision-inceptionv3 32 2,060.36 2,058.51 0.09% ✅
torchvision-inceptionv3_fp16 32 4,449.52 4,416.99 0.74% ✅
cadene-inceptionv4 16 817.23 820.61 -0.41% ✅
cadene-resnext64x4 16 784.97 782.78 0.28% ✅
slim-mobilenet 64 8,373.77 8,386.36 -0.15% ✅
slim-nasnetalarge 64 229.27 228.86 0.18% ✅
slim-resnet50v2 64 3,230.32 3,180.91 1.55% ✅
bert-mrpc-onnx 8 1,170.53 1,168.84 0.14% ✅
bert-mrpc-tf 1 502.47 498.63 0.77% ✅
pytorch-examples-wlang-gru 1 451.99 473.35 -4.51% ✅
pytorch-examples-wlang-lstm 1 594.17 384.83 54.40% 🔆
torchvision-resnet50_1 1 1,045.75 1,046.63 -0.08% ✅
cadene-dpn92_1 1 449.47 437.32 2.78% ✅
cadene-resnext101_1 1 363.77 365.89 -0.58% ✅
onnx-taau-downsample 1 845.42 844.09 0.16% ✅
dlrm-criteoterabyte 1 32.23 32.42 -0.58% ✅
dlrm-criteoterabyte_fp16 1 51.47 51.80 -0.64% ✅
agentmodel 1 14,460.27 9,209.12 57.02% 🔆
unet_fp16 2 58.27 58.80 -0.90% ✅
resnet50v1_fp16 1 1,412.38 1,366.11 3.39% ✅
resnet50v1_int8 1 1,765.85 1,883.96 -6.27% 🔴
bert_base_cased_fp16 64 1,098.68 1,098.16 0.05% ✅
bert_large_uncased_fp16 32 347.27 345.59 0.49% ✅
bert_large_fp16 1 207.04 206.59 0.22% ✅
distilgpt2_fp16 16 2,100.39 2,092.89 0.36% ✅
yolov5s 1 567.72 558.33 1.68% ✅
tinyllama 1 45.83 45.83 0.01% ✅
vicuna-fastchat 1 44.23 44.20 0.06% ✅
whisper-tiny-encoder 1 413.06 411.87 0.29% ✅
whisper-tiny-decoder 1 409.05 408.48 0.14% ✅
llama2_7b 1 20.85 20.84 0.07% ✅
qwen1.5-7b 1 21.81 23.58 -7.51% 🔴
phi3-3.8b 1 28.36 26.72 6.15% 🔆
llama3-8b 1 22.65 21.80 3.89% ✅
whisper-large-encoder 1 10.14 10.18 -0.41% ✅
whisper-large-decoder 1 105.99 105.30 0.65% ✅
mistral-7b 1 23.54 23.78 -0.98% ✅
FLUX.1-schnell 1 781.81 755.22 3.52% ✅

Regressions detected 🔴

* No develop baseline was found for this PR's branch point; compared against the latest available develop run instead.

@gh-app-migraphx-bot-pr-write

gh-app-migraphx-bot-pr-write Bot commented Aug 8, 2026 •

Copy link
Copy Markdown
Test Status Result
bert-mrpc-onnx ✅ PASSED: MIGraphX meets tolerance
bert-mrpc-tf ❌ ERROR - check error output
traceback
Traceback (most recent call last):
File "/src/AMDMIGraphX/tools/accuracy/accuracy_checker.py", line 377, in
main()
File "/src/AMDMIGraphX/tools/accuracy/accuracy_checker.py", line 313, in main
import tensorflow as tf
File "/usr/local/lib/python3.12/dist-packages/tensorflow/init.py", line 40, in
from tensorflow.python import pywrap_tensorflow as _pywrap_tensorflow # pylint: disable=unused-import
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/usr/local/lib/python3.12/dist-packages/tensorflow/python/pywrap_tensorflow.py", line 37, in
self_check.preload_check()
File "/usr/local/lib/python3.12/dist-packages/tensorflow/python/platform/self_check.py", line 63, in preload_check
from tensorflow.python.platform import _pywrap_cpu_feature_guard
ImportError: libnuma.so.1: cannot open shared object file: No such file or directory
pytorch-examples-wlang-gru 🔴 FAILED: MIGraphX is not within tolerance - check verbose output
pytorch-examples-wlang-lstm 🔴 FAILED: MIGraphX is not within tolerance - check verbose output
dlrm-criteoterabyte ✅ PASSED: MIGraphX meets tolerance
agentmodel ✅ PASSED: MIGraphX meets tolerance
unet ✅ PASSED: MIGraphX meets tolerance
resnet50v1 ✅ PASSED: MIGraphX meets tolerance
bert_base_cased_fp16 ✅ PASSED: MIGraphX meets tolerance
bert_large_uncased_fp16 🔴 FAILED: MIGraphX is not within tolerance - check verbose output
bert_large ✅ PASSED: MIGraphX meets tolerance
yolov5s ✅ PASSED: MIGraphX meets tolerance
tinyllama ✅ PASSED: MIGraphX meets tolerance
vicuna-fastchat ✅ PASSED: MIGraphX meets tolerance
whisper-tiny-encoder ✅ PASSED: MIGraphX meets tolerance
whisper-tiny-decoder ✅ PASSED: MIGraphX meets tolerance
distilgpt2_fp16 🔴 FAILED: MIGraphX is not within tolerance - check verbose output
llama2_7b ✅ PASSED: MIGraphX meets tolerance
qwen1.5-7b ✅ PASSED: MIGraphX meets tolerance
phi3-3.8b ✅ PASSED: MIGraphX meets tolerance
llama3-8b ✅ PASSED: MIGraphX meets tolerance
whisper-large-encoder ✅ PASSED: MIGraphX meets tolerance
whisper-large-decoder ✅ PASSED: MIGraphX meets tolerance
mistral-7b ✅ PASSED: MIGraphX meets tolerance
FLUX.1-schnell ✅ PASSED: MIGraphX meets tolerance

@CharlieL7 CharlieL7 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I would like to see an explanation of the steps in the split_sym_dim pass. Especially what the collect_roots() , discover_blocks(), materialize(), coalesce() and specialize_blocks() functions are trying to do.

I also see that this is using broadcast and multibroadcast with symbolic output_dyn_dims but without 2 inputs. Which is why they then have to be materialized to using the out_lens attribute. Our original idea was to use a symbolic map to evaluate at runtime, but that plan changed.

I don't think the multibroadcast or broadcast with symbolic output_dyn_dim with only 1 input would be able to run compute(). So it's an op that has to be rewritten. That's the same issue as the TopK redesign I had. To avoid that, multibroadcast and broadcast should be handled like dyn_slice.

Comment thread src/targets/gpu/eliminate_data_type_for_gpu.cpp Outdated
Comment thread src/targets/gpu/target.cpp
Comment thread src/eliminate_data_type.cpp Outdated
Comment thread src/instruction.cpp
Comment on lines +382 to +383
if(has_finalize(op))
return false;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

What was finalize() doing before these symbolic changes? As in, are there operators with a finalize() that can be evaluated?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

gpu::literal is the only other one that has finalize and also is context free (ie. affected by this change). The other options are:

  1. Make eval_expr_from_shape not context free so it fails can_eval using that existing path
  2. update the operator interface to define something new for when an op is dependent on runtime symbol resolution

Comment thread src/promote_literals.cpp Outdated
return root_ins.name() == "@literal" and root_ins.get_literal() == literal;
});
auto new_lit =
existing == root_module->end() ? root_module->add_literal(literal) : existing;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is this so we don't have a literal at the end of a module? Why is that an issue?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

no its just checking if the root module already has this literal and it not it inserts it. Its preventing it from duplicating literals in each select_module clone

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The submodule literals should be pulled out by promote_literals pass. They will be de-duplicated by CSE as Paul mentioned. Not sure if it's worth it to try to do that here.

Comment thread src/split_sym_dim.cpp Outdated
Comment thread src/split_sym_dim.cpp Outdated
Comment thread src/split_sym_dim.cpp
Comment thread src/split_sym_dim.cpp
Comment thread src/split_sym_dim.cpp
shivadbhavsar and others added 11 commits August 15, 2026 22:03
Bring the PR #5112-based parser prerequisite branch onto the current development baseline before applying its remaining changes.

Co-authored-by: Cursor <cursoragent@cursor.com>
Keep the merged resolver regression compatible with the dynamic-slice interface introduced by PR #5112.

Co-authored-by: Cursor <cursoragent@cursor.com>
Track exact integral shape values through parser operations so dynamic consumers retain symbolic output relationships without changing runtime dataflow.

Co-authored-by: Cursor <cursoragent@cursor.com>
Keep the symbolic-value change focused by restoring existing resolver diagnostics and simplifying the signed Gather size declaration.

Co-authored-by: Cursor <cursoragent@cursor.com>

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit

format.py

[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit

format.py

[format.py] reported by reviewdog 🐶

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit

format.py

[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶

auto& m = *p.get_main_module();
auto ids = m.add_parameter("input_ids",


[format.py] reported by reviewdog 🐶

auto select = add_select_module(expected_main,
{expected_ids, expected_mask_buffer, expected_update_buffer},


[format.py] reported by reviewdog 🐶

{symbolic_shape({lit(1), target_sequence, lit(4)})});


[format.py] reported by reviewdog 🐶

{{"expressions", migraphx::to_value(std::vector<se>{lit(4)})}}),


[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶

{{"padding", {0}}, {"stride", {1}}, {"dilation", {1}}}),


[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶

{{"padding", {0}}, {"stride", {1}}, {"dilation", {1}}}),


[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶


[format.py] reported by reviewdog 🐶

@CharlieL7

Copy link
Copy Markdown
Collaborator

Note that I've needed to change both step 1 and step 3 to support data-dependent shapes. Step 1 because the source of symbols can come from ops. Step 3 because the same for roots. That and adding more analyze structs for split_sym_dim. The base design is the same, just extending functionality.

Comment thread src/onnx/parse_reshape.cpp Outdated
transform(output_dims, output_expressions.begin(), [](const auto& dim) {
return dim.sym_expr;
});
const auto resolved_dims = info.add_instruction(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

note: I had to change this to handle data-dependent symbols. Needed to figure out if the symbol comes from parameters or somewhere else.

Comment thread src/targets/gpu/fuse_mlir.cpp
Comment thread src/targets/gpu/device/fixed_pad.cpp
Comment thread src/adjust_allocation.cpp Outdated
continue;

auto alias_ins = get_allocation(ins);
// A buffer holding one element of a tuple result cannot match the shape of the tuple

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I don't understand this comment. AFAIK this is checking that the alias was also a tuple?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

its supposed to skip instructions that produce a tuple output. Will update to clarify

Comment thread src/promote_literals.cpp Outdated
return root_ins.name() == "@literal" and root_ins.get_literal() == literal;
});
auto new_lit =
existing == root_module->end() ? root_module->add_literal(literal) : existing;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The submodule literals should be pulled out by promote_literals pass. They will be de-duplicated by CSE as Paul mentioned. Not sure if it's worth it to try to do that here.

Comment thread src/replace_allocate.cpp

// A lowered select_module writes each output into its own trailing buffer, which its operator
// alias cannot distinguish because the alias covers every buffer at once.
optional<instruction_ref> get_select_module_buffer(instruction_ref ins)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I don't follow here. Shouldn't the output buffer for select_module be a tuple? Or are saying that we want a separate allocate for each actual output buffer of select_module?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The latter. There was an issue with replace_allocate being able to properly generate output buffers properly and it was also doing copies. Went down separate allocate path to prevent that.

Comment thread test/gpu/fuse_mlir.cpp
EXPECT(p1.sort() == p2.sort());
}

TEST_CASE(dot_add_dot_batched_n1)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same question why doing GEG stuff here.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

model failure for actual LLM this needs to support. Probably can be pulled in separate PR, got left behind as an artifact

Comment thread test/op_shape_test.cpp
Comment thread test/replace_allocate.cpp
end));
}
auto ret = mm->add_return(outputs);
mm->add_debug_symbols(ret, {symbols.begin(), symbols.end()});

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why debug symbols?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

to make sure replace_allocate preserves them for the tuple outputs

Comment thread test/replace_allocate.cpp

@CharlieL7 CharlieL7 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think the design is good. I have some questions about particulars above

@codecov

codecov Bot commented Oct 1, 2026

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 87.74510% with 25 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
src/include/migraphx/op/select_module.hpp 48.39% 16 Missing ⚠️
src/shape.cpp 86.36% 3 Missing ⚠️
src/common.cpp 90.91% 2 Missing ⚠️
...clude/migraphx/op/builder/broadcast_dimensions.hpp 90.48% 2 Missing ⚠️
src/adjust_allocation.cpp 66.67% 1 Missing ⚠️
src/replace_allocate.cpp 95.65% 1 Missing ⚠️
Additional details and impacted files
@@             Coverage Diff             @@
##           develop    #5123      +/-   ##
===========================================
+ Coverage    92.70%   92.78%   +0.09%     
===========================================
  Files          633      635       +2     
  Lines        35759    37303    +1544     
===========================================
+ Hits         33147    34610    +1463     
- Misses        2612     2693      +81     
Files with missing lines Coverage Δ
src/eliminate_common_subexpression.cpp 100.00% <100.00%> (ø)
src/eliminate_data_type.cpp 64.44% <ø> (ø)
src/include/migraphx/op/broadcast_with_dims.hpp 96.67% <100.00%> (ø)
src/include/migraphx/op/concat_past_present.hpp 98.25% <100.00%> (ø)
src/include/migraphx/op/eval_expr_from_shape.hpp 88.89% <100.00%> (+2.84%) ⬆️
src/include/migraphx/op/fixed_pad.hpp 100.00% <100.00%> (ø)
src/include/migraphx/op/multibroadcast.hpp 96.15% <100.00%> (+0.70%) ⬆️
src/include/migraphx/split_sym_dim.hpp 100.00% <100.00%> (ø)
src/onnx/onnx_parser.cpp 88.99% <ø> (ø)
src/onnx/parse_instancenorm.cpp 83.08% <100.00%> (ø)
... and 13 more

... and 2 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.
  • 📦 JS Bundle Analysis: Save yourself from yourself by tracking and limiting bundle sizes in JS merges.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants