Split symbolic dimension pass - #5123
shivadbhavsar wants to merge 71 commits into
Conversation
Regressions detected 🔴 * No develop baseline was found for this PR's branch point; compared against the latest available develop run instead. |
|
CharlieL7
left a comment
There was a problem hiding this comment.
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.
| if(has_finalize(op)) | ||
| return false; |
There was a problem hiding this comment.
What was finalize() doing before these symbolic changes? As in, are there operators with a finalize() that can be evaluated?
There was a problem hiding this comment.
gpu::literal is the only other one that has finalize and also is context free (ie. affected by this change). The other options are:
- Make eval_expr_from_shape not context free so it fails can_eval using that existing path
- update the operator interface to define something new for when an op is dependent on runtime symbol resolution
| return root_ins.name() == "@literal" and root_ins.get_literal() == literal; | ||
| }); | ||
| auto new_lit = | ||
| existing == root_module->end() ? root_module->add_literal(literal) : existing; |
There was a problem hiding this comment.
Is this so we don't have a literal at the end of a module? Why is that an issue?
There was a problem hiding this comment.
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
There was a problem hiding this comment.
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.
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>
…nto split_sym_dim
There was a problem hiding this comment.
Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit
format.py
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2297 in 17f7126
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2392 in 17f7126
There was a problem hiding this comment.
Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit
format.py
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2785 in fe76ee1
There was a problem hiding this comment.
Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit
format.py
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2217 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2279 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Lines 2297 to 2298 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Lines 2404 to 2405 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2407 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2445 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2616 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2647 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2728 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2744 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2786 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2819 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2871 in e37d2c8
[format.py] reported by reviewdog 🐶
AMDMIGraphX/test/split_sym_dim_test.cpp
Line 2966 in e37d2c8
|
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. |
| transform(output_dims, output_expressions.begin(), [](const auto& dim) { | ||
| return dim.sym_expr; | ||
| }); | ||
| const auto resolved_dims = info.add_instruction( |
There was a problem hiding this comment.
note: I had to change this to handle data-dependent symbols. Needed to figure out if the symbol comes from parameters or somewhere else.
| continue; | ||
|
|
||
| auto alias_ins = get_allocation(ins); | ||
| // A buffer holding one element of a tuple result cannot match the shape of the tuple |
There was a problem hiding this comment.
I don't understand this comment. AFAIK this is checking that the alias was also a tuple?
There was a problem hiding this comment.
its supposed to skip instructions that produce a tuple output. Will update to clarify
| return root_ins.name() == "@literal" and root_ins.get_literal() == literal; | ||
| }); | ||
| auto new_lit = | ||
| existing == root_module->end() ? root_module->add_literal(literal) : existing; |
There was a problem hiding this comment.
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.
|
|
||
| // 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) |
There was a problem hiding this comment.
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?
There was a problem hiding this comment.
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.
| EXPECT(p1.sort() == p2.sort()); | ||
| } | ||
|
|
||
| TEST_CASE(dot_add_dot_batched_n1) |
There was a problem hiding this comment.
Same question why doing GEG stuff here.
There was a problem hiding this comment.
model failure for actual LLM this needs to support. Probably can be pulled in separate PR, got left behind as an artifact
| end)); | ||
| } | ||
| auto ret = mm->add_return(outputs); | ||
| mm->add_debug_symbols(ret, {symbols.begin(), symbols.end()}); |
There was a problem hiding this comment.
to make sure replace_allocate preserves them for the tuple outputs
CharlieL7
left a comment
There was a problem hiding this comment.
I think the design is good. I have some questions about particulars above
Codecov Report❌ Patch coverage is
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
🚀 New features to boost your workflow:
|
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:
3. Collect symbolic roots
collect_roots()finds the independent symbols introduced by parameter dimensions.For each root, it:
4. Discover specialization blocks
discover_blocks()first creates one candidate block for every operation that requires padding or static rewriting.It determines:
5. Coalesce blocks
Coalescing combines candidate blocks so a connected region can use one
select_moduleinstead of dispatching every operation separately.Two blocks are merged only when:
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:
7. Materialize specialization clones
Materialization happens inside
specialize_blocks()throughbuild_clone().For every combination of root buckets, it:
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:select_modulereferencing the generated clones.Changelog Category
Add a
CHANGELOG.mdentry for any option other thanNot ApplicableFollow the LLVM AI Tool Use Policy for contributions using AI.