Skip to content

Canonicalize XLA wrapper inputs - #3350

Merged
wsmoses merged 5 commits into
mainfrom
vim/remove-unused-xla-wrapper-inputs
Oct 6, 2026
Merged

wsmoses merged 5 commits into
mainfrom
vim/remove-unused-xla-wrapper-inputs

Conversation

@vimarsh6739

@vimarsh6739 vimarsh6739 commented Oct 4, 2026 •

Copy link
Copy Markdown
Member

XLA wrappers can pass unused inputs or pass the same buffer more than once. Add two rewrites to standard XLAWrapperOp canonicalization.

Remove a buffer input and its matching result when the function only returns that input unchanged. For example, double_data(data, spare) returns (data + data, spare). The reduced function takes only data and returns data + data. Also remove specialized scalar inputs that the function does not use. Keep the remaining scalar inputs at the end and update num_specialized.

For unused-input removal, change the original function in place when it is private, all symbol uses are known, and every use is a supported wrapper call in the same symbol scope. Update all those callers together. Otherwise, keep the original signature and let compatible wrappers share one reduced clone. Keep calls with no remaining inputs so that their side effects still execute.

Merge duplicate buffer inputs when their returned values and attributes agree. For example, a function computes sum = x + y and returns (sum, sum). A wrapper that passes (a, a) can use a reduced function that computes a + a and returns one result. Also merge specialized scalar inputs when the wrapper passes the same SSA value and their types and argument attributes agree. For example, advance(data, n, n, scale) becomes advance(data, n, scale), with both uses of n retained in the body. Reduce num_specialized without removing a buffer result. Use a clone so that other callers can still pass distinct inputs. Keep the remaining specialized scalar inputs in their original order. Keep duplicate buffers when their results or attributes differ.

Remove argument and result attributes with their slots. Keep the attributes of each remaining wrapper input, function argument, and buffer result. This also permits callers with different attributes to share the reduced function after unused-input removal.

Buffer identity uses enzyme::oputils::getBaseObject with offsetAllowed=false. The pointer conversion interface support is already on main through #3377. This PR targets main, including the merged wrapper fusion from #2744.

Validation: the optimized build passes. The full MLIR LIT suite has 1,424 passes and 9 expected failures, with no unexpected failures. Formatting and whitespace checks pass. The tests cover private functions with one or several callers, shared clones, unknown symbol uses, buffer writes, side effects, specialized scalar removal, metadata removal and preservation, duplicate buffers with scalar inputs, specialized scalar deduplication, callers with distinct scalar values, combined buffer and scalar deduplication, conflicting attributes, and symbol scope. The new scalar deduplication test fails with the old optimizer and passes with this change.

@vimarsh6739
vimarsh6739 force-pushed the vim/remove-unused-xla-wrapper-inputs branch 3 times, most recently from ff13883 to 6ddd149 Compare October 4, 2026 18:30
@vimarsh6739 vimarsh6739 changed the title Remove unused XLA wrapper inputs Canonicalize XLA wrapper inputs Oct 4, 2026
@vimarsh6739
vimarsh6739 marked this pull request as ready for review October 4, 2026 19:40
@vimarsh6739
vimarsh6739 requested a review from wsmoses October 4, 2026 19:40
Comment thread src/enzyme_ad/jax/Dialect/Ops.cpp Outdated
Comment thread src/enzyme_ad/jax/Dialect/Ops.cpp
@codecov

codecov Bot commented Oct 4, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 29.61%. Comparing base (58fe151) to head (59a440e).
⚠️ Report is 16 commits behind head on main.

Additional details and impacted files
@@            Coverage Diff             @@
##             main    #3350      +/-   ##
==========================================
- Coverage   29.61%   29.61%   -0.01%     
==========================================
  Files         239      240       +1     
  Lines       48403    48459      +56     
==========================================
+ Hits        14336    14350      +14     
- Misses      34067    34109      +42     

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@vimarsh6739
vimarsh6739 force-pushed the vim/remove-unused-xla-wrapper-inputs branch 2 times, most recently from 259a07c to ce1bcf6 Compare October 4, 2026 23:29
Comment thread src/enzyme_ad/jax/Dialect/Ops.cpp Outdated
};

/// Remove casts that preserve the buffer address. Keep views that change it.
static Value getXLAWrapperBufferIdentity(Value value) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

also this should use getBaseObject, no?

Comment thread src/enzyme_ad/jax/Dialect/Ops.cpp Outdated
: public OpRewritePattern<XLAWrapperOp> {
static bool canRewrite(XLAWrapperOp wrapper) {
return !wrapper.getArgAttrsAttr() && !wrapper.getResAttrsAttr() &&
!wrapper.getNumSpecialized();

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

you should support specialized

Comment thread src/enzyme_ad/jax/Dialect/Ops.cpp Outdated
: public OpRewritePattern<XLAWrapperOp> {
static bool canRewrite(XLAWrapperOp wrapper) {
return !wrapper.getArgAttrsAttr() && !wrapper.getResAttrsAttr() &&
!wrapper.getNumSpecialized();

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

you also should support argattrs and resattrs, there's a helper fn in llvm for removing from them

Comment thread src/enzyme_ad/jax/Dialect/Ops.cpp Outdated

auto function = dyn_cast_or_null<FunctionOpInterface>(
SymbolTable::lookupNearestSymbolFrom(op, op.getFnAttr()));
auto hasMetadata = [](ArrayAttr attributes) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

avoid lambda fns

@vimarsh6739
vimarsh6739 force-pushed the vim/remove-unused-xla-wrapper-inputs branch from 042956c to 64098fb Compare October 6, 2026 03:04
Remove unused input/result pairs and merge duplicate buffer inputs when their results agree. Share one reduced function across compatible callers when removing unused inputs. Preserve callers that need the original signature.

Register both rewrites for ordinary canonicalization and add LIT coverage.

Assisted-By: OpenAI Codex
Mark the pointer conversion operations with NoOffsetViewInterface and
require a zero-offset result layout for Pointer2MemrefOp. Add the interface
headers and TableGen dependency.

Use getBaseObject with offsets disabled to identify duplicate wrapper
inputs. Remove the local helper and pin Enzyme to 95acfd4c from
EnzymeAD/Enzyme#3420.

Assisted-By: OpenAI Codex
Use OffsetViewInterface for pointer2memref and memref2pointer. Keep the
pointer2memref zero-offset verifier. The Enzyme dependency and shared pointer
utility cleanup now come from main through Enzyme-JAX#3373.

Update wrapper metadata in the LIT tests to use the current property syntax.

Assisted-By: OpenAI Codex
Remove unused specialized scalars and update their count. Remove argument and result attributes with their slots, and preserve attributes on retained inputs and results. Keep duplicate buffers separate when their attributes differ. Use static helpers instead of closures.

Assisted-By: OpenAI Codex
@vimarsh6739
vimarsh6739 force-pushed the vim/remove-unused-xla-wrapper-inputs branch from 64098fb to bd8cad5 Compare October 6, 2026 03:11
for (Operation *caller : callers) {
auto wrapper = cast<XLAWrapperOp>(caller);
rewriter.startOpModification(wrapper);
removeWrapperInputs(wrapper, unusedArguments, unusedResults);

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

do you have a test case also doing a removal of a specialized input?

if (canonicalResult(returnOp->getOperand(index), body, representatives) !=
canonicalResult(returnOp->getOperand(representative), body,
representatives))
return failure();

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

same comment of do you have a deduplicate specialized test?

Merge repeated specialized scalar arguments when their types and attributes agree. Keep buffer results and update num_specialized. Keep the original function for callers that pass distinct values.

Add direct tests for specialized input removal and deduplication. Check scalar order, distinct callers, combined buffer and scalar deduplication, and the result after wrapper fusion.

Validation: optimized build passed; MLIR LIT suite had 1,424 passes and 9 expected failures.

Assisted-By: OpenAI Codex
@vimarsh6739
vimarsh6739 requested a review from wsmoses October 6, 2026 05:03
@wsmoses
wsmoses merged commit f98db73 into main Oct 6, 2026
20 of 33 checks passed
@wsmoses
wsmoses deleted the vim/remove-unused-xla-wrapper-inputs branch October 6, 2026 20:00
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.

2 participants