From b491ff77b13940b309286784022bd548b19ec41f Mon Sep 17 00:00:00 2001 From: NiccoloN Date: Wed, 22 Jul 2026 15:43:27 +0200 Subject: [PATCH] fix trivial merge minor validate.py fix --- .../TrivialGraphComputeMergePass.cpp | 22 +++++++++++++++---- validation/validate.py | 6 ++--- 2 files changed, 21 insertions(+), 7 deletions(-) diff --git a/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp b/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp index efcb9a4..9bd691e 100644 --- a/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/TrivialGraphComputeMergePass.cpp @@ -42,6 +42,17 @@ static bool hasCapacityFor(Operation* producer, Operation* consumer) { return getCrossbarUnionSize(producerWeights, consumerWeights) <= static_cast(crossbarCountInCore.getValue()); } +template +static bool isUniqueGraphComputePredecessor(Operation *candidate, ConsumerOp consumer) { + llvm::SmallSetVector predecessors; + for (Value input : consumer.getInputs()) { + Operation *producer = input.getDefiningOp(); + if (producer && isGraphComputeLike(producer)) + predecessors.insert(producer); + } + return predecessors.size() == 1 && predecessors.front() == candidate; +} + struct TrivialGraphMergeStats { size_t scalarBefore = 0; size_t batchBefore = 0; @@ -153,7 +164,8 @@ struct MergeTrivialScalarComputes : OpRewritePattern { for (Value input : consumer.getInputs()) { auto candidate = input.getDefiningOp(); if (candidate && candidate->getBlock() == consumer->getBlock() && hasOnlyStructuralAttrs(candidate) - && hasOnlyStructuralAttrs(consumer) && isExclusivelyConsumedBy(candidate, consumer) + && hasOnlyStructuralAttrs(consumer) && isUniqueGraphComputePredecessor(candidate, consumer) + && isExclusivelyConsumedBy(candidate, consumer) && hasCapacityFor(candidate, consumer) && hasNoNestedArgumentCaptures(candidate) && hasNoNestedArgumentCaptures(consumer)) { producer = candidate; @@ -319,7 +331,8 @@ struct FoldBatchLeadingUnitNormalization : OpRewritePattern { ? SpatGraphComputeBatch() : consumer.getInputs().front().getDefiningOp(); if (!producer || producer->getBlock() != consumer->getBlock() || !hasOnlyStructuralAttrs(producer) - || !hasOnlyStructuralAttrs(consumer) || !isExclusivelyConsumedBy(producer, consumer)) + || !hasOnlyStructuralAttrs(consumer) || !isUniqueGraphComputePredecessor(producer, consumer) + || !isExclusivelyConsumedBy(producer, consumer)) return failure(); auto fragments = collectPublishedFragments(producer); if (!matchLeadingUnitNormalization(producer, consumer) || failed(fragments)) @@ -394,7 +407,8 @@ struct MergeTrivialBatchComputes : OpRewritePattern { auto candidate = input.getDefiningOp(); if (candidate && candidate->getBlock() == consumer->getBlock() && candidate.getLaneCount() == consumer.getLaneCount() && hasOnlyStructuralAttrs(candidate) - && hasOnlyStructuralAttrs(consumer) && isExclusivelyConsumedBy(candidate, consumer) + && hasOnlyStructuralAttrs(consumer) && isUniqueGraphComputePredecessor(candidate, consumer) + && isExclusivelyConsumedBy(candidate, consumer) && hasCapacityFor(candidate, consumer) && hasDirectLaneConsumers(candidate, consumer) && succeeded(fragments = collectPublishedFragments(candidate))) { producer = candidate; @@ -462,7 +476,7 @@ struct TrivialGraphComputeMergePass final : PassWrapper