From 5f42da36aeef07731cce2213b18a6838ea8bc042 Mon Sep 17 00:00:00 2001 From: ilgeco Date: Fri, 17 Jul 2026 11:08:11 +0200 Subject: [PATCH] blazingly fast --- .../DeferredBoundaryPlanning.cpp | 55 ++++ .../DeferredBoundaryPlanning.hpp | 5 +- .../DeferredBoundaryRealization.cpp | 266 +++++++++++++++++- .../DeferredBoundaryRealization.hpp | 7 +- .../DeferredCommunicationModel.hpp | 8 + .../DeferredCommunicationRealization.cpp | 10 +- .../DeferredCommunicationScheduling.cpp | 2 + .../DeferredCommunicationScheduling.hpp | 2 + .../DeferredProjectionAnalysis.cpp | 25 +- .../DeferredResultRealization.cpp | 16 +- .../DeferredResultRealization.hpp | 6 +- .../DeferredTransferPlanning.cpp | 59 +++- 12 files changed, 426 insertions(+), 35 deletions(-) diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.cpp index 701c898..b305a94 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.cpp @@ -1,3 +1,5 @@ +#include "mlir/Dialect/SCF/IR/SCF.h" + #include "llvm/ADT/MapVector.h" #include "DeferredBoundaryPlanning.hpp" @@ -99,6 +101,7 @@ static size_t hashSemanticKey(const SemanticKey& key) { key.emission.payload.getAsOpaquePointer(), key.emission.fragmentType.getAsOpaquePointer(), key.emission.hasGraphLane, + key.emission.hasProducerProjection, key.emission.sourceIsBatch); if (key.kind == SemanticKind::Result) return llvm::hash_combine(key.kind, key.exchange); @@ -343,6 +346,44 @@ static bool matchesProjectionAssembly(ArrayRef actions, return true; } +static size_t getReceiveBundleLength(ArrayRef actions, + size_t begin, + const LaneSet& lanes) { + if (begin >= actions.size()) + return 0; + const CanonicalAction& first = actions[begin]; + if (first.key.kind != SemanticKind::Availability) + return 0; + Type fragmentType = first.key.fragmentType; + auto rankedFragment = dyn_cast(fragmentType); + if (!rankedFragment || !rankedFragment.hasStaticShape()) + return 0; + Value firstOutput = first.key.exchange->deferred.getOutput(); + if (!firstOutput.hasOneUse()) + return 0; + Operation* firstUser = *firstOutput.getUsers().begin(); + auto selection = firstUser->getParentOfType(); + if (!selection) + return 0; + size_t end = begin; + while (end < actions.size()) { + const CanonicalAction& action = actions[end]; + Value output = action.key.exchange + ? action.key.exchange->deferred.getOutput() + : Value(); + Operation* user = output && output.hasOneUse() + ? *output.getUsers().begin() + : nullptr; + if (action.key.kind != SemanticKind::Availability || !action.locals.empty() + || action.slices.empty() || !(action.receiveLanes == lanes) + || action.key.fragmentType != fragmentType + || !user || user->getParentOfType() != selection) + break; + ++end; + } + return end - begin >= 2 ? end - begin : 0; +} + static BoundaryInstructionList materializeInstructions(ArrayRef actions, const LaneSet& lanes) { BoundaryInstructionList result; @@ -393,6 +434,20 @@ static BoundaryInstructionList materializeInstructions(ArrayRef index += projectionPositions.size(); continue; } + if (size_t bundleLength = getReceiveBundleLength(actions, index, lanes)) { + EmitReceiveBundle bundle; + for (size_t offset = 0; offset < bundleLength; ++offset) { + const CanonicalAction& entry = actions[index + offset]; + EmitReceiveRun receive; + receive.slices = entry.slices; + receive.entryOffsets = {0, receive.slices.size()}; + receive.lanes = entry.receiveLanes; + bundle.entries.push_back(std::move(receive)); + } + instructions.push_back(std::move(bundle)); + index += bundleLength; + continue; + } if (action.key.kind == SemanticKind::Availability) { ResolveAvailability availability; availability.exchange = action.key.exchange; diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.hpp index 773567d..46506cb 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryPlanning.hpp @@ -26,6 +26,9 @@ struct EmitReceiveRun { llvm::SmallVector entryOffsets; LaneSet lanes; }; +struct EmitReceiveBundle { + llvm::SmallVector entries; +}; struct EmitReceiveAssemblyRun { llvm::SmallVector slices; llvm::SmallVector entryOffsets; @@ -57,7 +60,7 @@ struct ProduceDeferredResult { struct BoundaryInstructionList; struct LaneDispatch; using BoundaryInstruction = std::variant>; struct BoundaryInstructionList { llvm::SmallVector instructions; diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp index 236ccdc..e9b787c 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.cpp @@ -21,6 +21,9 @@ struct LogicalTransferMetadataView { StaticIntSequenceChain targetCores; StaticIntSequenceChain targetLanes; StaticIntSequenceChain localOffsets; + SmallVector projectionOffsets; + SmallVector projectionSizes; + SmallVector projectionStrides; size_t size() const { return channels.size(); } }; @@ -48,6 +51,25 @@ static void appendMetadata(const ScheduledTransferSlice& slice, targetLane - requirementLanes.begin, count); else metadata.localOffsets.append(StaticIntSequence::uniform(0, count)); + if (family.requirement->producerProjection) { + const DeferredStaticSliceGeometry& geometry = + *family.requirement->producerProjection; + if (metadata.projectionOffsets.empty()) { + metadata.projectionOffsets.resize(geometry.offsets.size()); + metadata.projectionSizes.resize(geometry.sizes.size()); + metadata.projectionStrides.resize(geometry.strides.size()); + } + size_t geometryIndex = targetLane - requirementLanes.begin; + for (auto [target, source] : + llvm::zip_equal(metadata.projectionOffsets, geometry.offsets)) + target.append(source, geometryIndex, count); + for (auto [target, source] : + llvm::zip_equal(metadata.projectionSizes, geometry.sizes)) + target.append(source, geometryIndex, count); + for (auto [target, source] : + llvm::zip_equal(metadata.projectionStrides, geometry.strides)) + target.append(source, geometryIndex, count); + } } static LogicalTransferMetadataView @@ -106,9 +128,18 @@ static OpFoldResult lookupGeometry( static FailureOr materializeSendPayload(const RequirementFamily& requirement, Value localOffset, + const MixedSliceGeometry* producerProjection, DeferredEmissionContext& context, Location loc) { Value payload = requirement.producer->payload; + if (producerProjection) { + auto fragmentType = dyn_cast( + requirement.publicationFragmentType); + if (!fragmentType) + return failure(); + return extractMixedSliceOrIdentity( + context.rewriter, loc, payload, fragmentType, *producerProjection); + } if (payload.getType() == requirement.publicationFragmentType) return payload; auto payloadType = dyn_cast(payload.getType()); @@ -148,6 +179,20 @@ emitSendRun(const EmitSendRun& run, Value lane, unsigned laneCount, DeferredEmis actionCount * laneCount, logical.targetCores.valueAt(0)); SmallVector offsetTable( actionCount * laneCount, logical.localOffsets.valueAt(0)); + SmallVector> projectionOffsetTables; + SmallVector> projectionSizeTables; + SmallVector> projectionStrideTables; + auto initializeGeometryTables = [&](ArrayRef source, + SmallVectorImpl>& target) { + for (const StaticIntSequenceChain& sequence : source) + target.emplace_back(actionCount * laneCount, sequence.valueAt(0)); + }; + initializeGeometryTables(logical.projectionOffsets, + projectionOffsetTables); + initializeGeometryTables(logical.projectionSizes, + projectionSizeTables); + initializeGeometryTables(logical.projectionStrides, + projectionStrideTables); SmallVector counts(laneCount); for (unsigned sourceLane = 0; sourceLane < laneCount; ++sourceLane) { const LogicalTransferMetadataView& source = metadataByLane[sourceLane]; @@ -158,6 +203,18 @@ emitSendRun(const EmitSendRun& run, Value lane, unsigned laneCount, DeferredEmis sourceTable[index] = source.sourceCores.valueAt(action); targetTable[index] = source.targetCores.valueAt(action); offsetTable[index] = source.localOffsets.valueAt(action); + for (auto [table, sequence] : + llvm::zip_equal(projectionOffsetTables, + source.projectionOffsets)) + table[index] = sequence.valueAt(action); + for (auto [table, sequence] : + llvm::zip_equal(projectionSizeTables, + source.projectionSizes)) + table[index] = sequence.valueAt(action); + for (auto [table, sequence] : + llvm::zip_equal(projectionStrideTables, + source.projectionStrides)) + table[index] = sequence.valueAt(action); } } ExternalTransferFamily& firstFamily = *run.slices.front().family; @@ -166,7 +223,20 @@ emitSendRun(const EmitSendRun& run, Value lane, unsigned laneCount, DeferredEmis Location loc = requirement.exchange->deferred.getLoc(); auto emitOne = [&](Value position) -> LogicalResult { Value localOffset = lookup(offsetTable, position, anchor, context, loc); - auto payload = materializeSendPayload(requirement, localOffset, context, loc); + MixedSliceGeometry projection; + for (ArrayRef table : projectionOffsetTables) + projection.offsets.push_back( + lookupGeometry(table, position, anchor, context, loc)); + for (ArrayRef table : projectionSizeTables) + projection.sizes.push_back( + lookupGeometry(table, position, anchor, context, loc)); + for (ArrayRef table : projectionStrideTables) + projection.strides.push_back( + lookupGeometry(table, position, anchor, context, loc)); + auto payload = materializeSendPayload( + requirement, localOffset, + projectionOffsetTables.empty() ? nullptr : &projection, + context, loc); if (failed(payload)) return failure(); auto send = SpatChannelSendOp::create(context.rewriter, @@ -231,6 +301,125 @@ emitReceiveValue(const EmitReceiveRun& run, Value lane, unsigned laneCount, Defe return receive.getOutput(); } +static LogicalResult emitReceiveBundle(const EmitReceiveBundle& bundle, + Value lane, + unsigned laneCount, + DeferredEmissionContext& context) { + if (bundle.entries.size() < 2) + return failure(); + RequirementFamily& reference = + *bundle.entries.front().slices.front().family->requirement; + auto fragmentType = dyn_cast( + reference.publicationFragmentType); + if (!fragmentType || !fragmentType.hasStaticShape()) + return failure(); + for (const EmitReceiveRun& entry : bundle.entries) { + if (entry.slices.empty() + || entry.slices.front().family->requirement->publicationFragmentType + != fragmentType) + return failure(); + } + + SmallVector slices; + for (const EmitReceiveRun& entry : bundle.entries) + llvm::append_range(slices, entry.slices); + LogicalTransferMetadataView metadata = buildMetadataView(slices); + size_t entryCount = bundle.entries.size(); + SmallVector channelTable( + entryCount * laneCount, metadata.channels.valueAt(0)); + SmallVector sourceTable( + entryCount * laneCount, metadata.sourceCores.valueAt(0)); + SmallVector targetTable( + entryCount * laneCount, metadata.targetCores.valueAt(0)); + for (auto [entryIndex, entry] : llvm::enumerate(bundle.entries)) { + LogicalTransferMetadataView item = buildMetadataView(entry.slices); + for (size_t index = 0; index < item.size(); ++index) { + size_t tableIndex = + entryIndex * laneCount + item.targetLanes.valueAt(index); + channelTable[tableIndex] = item.channels.valueAt(index); + sourceTable[tableIndex] = item.sourceCores.valueAt(index); + targetTable[tableIndex] = item.targetCores.valueAt(index); + } + } + + DeferredExchangePlan* exchange = reference.exchange; + Location loc = exchange->deferred.getLoc(); + SmallVector bundleShape {static_cast(entryCount)}; + llvm::append_range(bundleShape, fragmentType.getShape()); + auto bundleType = RankedTensorType::get( + bundleShape, fragmentType.getElementType()); + Value initial = tensor::EmptyOp::create( + context.rewriter, loc, bundleShape, fragmentType.getElementType()); + auto loop = buildNormalizedScfFor( + context.rewriter, + loc, + context.constants.getIndex(0), + context.constants.getIndex(entryCount), + context.constants.getIndex(1), + ValueRange {initial}, + [&](OpBuilder&, Location, Value entry, ValueRange iterArgs, + SmallVectorImpl& yielded) -> LogicalResult { + Value tableIndex = entry; + if (lane) { + Value base = affineMulConst( + context.rewriter, loc, entry, laneCount, exchange->deferred); + tableIndex = arith::AddIOp::create( + context.rewriter, loc, base, lane); + } + auto receive = SpatChannelReceiveOp::create( + context.rewriter, + loc, + fragmentType, + lookup(channelTable, tableIndex, exchange->deferred, context, loc), + lookup(sourceTable, tableIndex, exchange->deferred, context, loc), + lookup(targetTable, tableIndex, exchange->deferred, context, loc)); + setLogicalTransferMetadata(receive, metadata); + auto source = addLeadingUnitTensorDimension( + context.rewriter, loc, receive.getOutput()); + if (failed(source)) + return failure(); + MixedSliceGeometry geometry; + geometry.offsets.assign(bundleType.getRank(), + context.rewriter.getIndexAttr(0)); + geometry.offsets.front() = entry; + for (int64_t dimension : cast(source->getType()).getShape()) + geometry.sizes.push_back( + context.rewriter.getIndexAttr(dimension)); + geometry.strides.assign(bundleType.getRank(), + context.rewriter.getIndexAttr(1)); + yielded.push_back(insertMixedSlice( + context.rewriter, loc, *source, iterArgs.front(), geometry)); + return success(); + }); + if (failed(loop)) + return failure(); + + SmallVector unitShape {1}; + llvm::append_range(unitShape, fragmentType.getShape()); + auto unitType = RankedTensorType::get( + unitShape, fragmentType.getElementType()); + for (auto [entryIndex, entry] : llvm::enumerate(bundle.entries)) { + MixedSliceGeometry geometry; + geometry.offsets.assign(bundleType.getRank(), + context.rewriter.getIndexAttr(0)); + geometry.offsets.front() = context.rewriter.getIndexAttr(entryIndex); + for (int64_t dimension : unitShape) + geometry.sizes.push_back( + context.rewriter.getIndexAttr(dimension)); + geometry.strides.assign(bundleType.getRank(), + context.rewriter.getIndexAttr(1)); + Value unit = extractMixedSliceOrIdentity( + context.rewriter, loc, loop->results.front(), unitType, geometry); + auto fragment = removeLeadingUnitTensorDimension( + context.rewriter, loc, unit, fragmentType); + if (failed(fragment)) + return failure(); + for (const ScheduledTransferSlice& slice : entry.slices) + context.receives[slice.family->requirement] = *fragment; + } + return success(); +} + static LogicalResult emitConditionalSendRun(const EmitSendRun& run, Value lane, unsigned laneCount, DeferredEmissionContext& context) { if (run.lanes.size() == laneCount) @@ -259,19 +448,67 @@ static FailureOr materializeLocalValue(const MaterializeLocalFamily& loca DeferredEmissionContext& context) { RequirementFamily& reference = *local.families.front()->requirement; Value fragment = reference.producer->payload; - if (fragment.getType() != reference.publicationFragmentType) { + if (reference.producerProjection + || fragment.getType() != reference.publicationFragmentType) { SmallVector offsets(laneCount); + SmallVector> projectionOffsets; + SmallVector> projectionSizes; + SmallVector> projectionStrides; + auto initializeGeometry = [&](ArrayRef source, + SmallVectorImpl>& target) { + for (const StaticIntSequence& sequence : source) + target.emplace_back(laneCount, sequence.valueAt(0)); + }; + if (reference.producerProjection) { + initializeGeometry(reference.producerProjection->offsets, + projectionOffsets); + initializeGeometry(reference.producerProjection->sizes, + projectionSizes); + initializeGeometry(reference.producerProjection->strides, + projectionStrides); + } for (LocalAvailabilityFamily* family : local.families) { RequirementFamily& requirement = *family->requirement; LaneInterval requirementLanes = requirement.targetLanes.intervals().front(); for (LaneInterval interval : family->targetLanes.intervals()) - for (unsigned targetLane = interval.begin; targetLane < interval.end; ++targetLane) - offsets[targetLane] = requirement.producerLocalOffsets->valueAt(targetLane - requirementLanes.begin); + for (unsigned targetLane = interval.begin; + targetLane < interval.end; ++targetLane) { + size_t position = targetLane - requirementLanes.begin; + if (requirement.producerLocalOffsets) + offsets[targetLane] = + requirement.producerLocalOffsets->valueAt(position); + if (requirement.producerProjection) { + for (auto [table, sequence] : + llvm::zip_equal(projectionOffsets, + requirement.producerProjection->offsets)) + table[targetLane] = sequence.valueAt(position); + for (auto [table, sequence] : + llvm::zip_equal(projectionSizes, + requirement.producerProjection->sizes)) + table[targetLane] = sequence.valueAt(position); + for (auto [table, sequence] : + llvm::zip_equal(projectionStrides, + requirement.producerProjection->strides)) + table[targetLane] = sequence.valueAt(position); + } + } } Location loc = reference.exchange->deferred.getLoc(); Value position = lane ? lane : context.constants.getIndex(0); + MixedSliceGeometry projection; + for (ArrayRef table : projectionOffsets) + projection.offsets.push_back(lookupGeometry( + table, position, reference.exchange->deferred, context, loc)); + for (ArrayRef table : projectionSizes) + projection.sizes.push_back(lookupGeometry( + table, position, reference.exchange->deferred, context, loc)); + for (ArrayRef table : projectionStrides) + projection.strides.push_back(lookupGeometry( + table, position, reference.exchange->deferred, context, loc)); auto materialized = materializeSendPayload( - reference, lookup(offsets, position, reference.exchange->deferred, context, loc), context, loc); + reference, + lookup(offsets, position, reference.exchange->deferred, context, loc), + projectionOffsets.empty() ? nullptr : &projection, context, loc); if (failed(materialized)) return failure(); fragment = *materialized; @@ -721,6 +958,11 @@ static FailureOr> emitInstructions( return failure(); continue; } + if (auto bundle = std::get_if(&instruction)) { + if (failed(emitReceiveBundle(*bundle, lane, laneCount, context))) + return failure(); + continue; + } if (auto assembly = std::get_if(&instruction)) { auto value = assembly->projectionLeaf ? emitProjectionAssemblyRun( @@ -774,12 +1016,12 @@ static void setInsertionAtBoundary(IRRewriter& rewriter, const BoundaryKey& key) } static LogicalResult -replaceResults(ArrayRef exchanges, ValueRange replacements, DeferredEraseSet& erase) { +replaceResults(ArrayRef exchanges, ValueRange replacements, + DeferredReplacementMap& deferredReplacements) { if (exchanges.size() != replacements.size()) return failure(); for (auto [exchange, replacement] : llvm::zip_equal(exchanges, replacements)) { - exchange->deferred.getOutput().replaceAllUsesWith(replacement); - erase.insert(exchange->deferred); + deferredReplacements.insert({exchange->deferred, replacement}); } return success(); } @@ -787,7 +1029,7 @@ replaceResults(ArrayRef exchanges, ValueRange replacement static LogicalResult emitBoundary(const BoundaryProgram& boundary, ArrayRef results, DeferredEmissionContext& context, - DeferredEraseSet& erase) { + DeferredReplacementMap& replacements) { setInsertionAtBoundary(context.rewriter, boundary.key); unsigned laneCount = boundary.key.scheduled->cores.size(); Value lane; @@ -798,7 +1040,7 @@ static LogicalResult emitBoundary(const BoundaryProgram& boundary, auto values = emitInstructions( boundary.root, lane, laneCount, results, context); return failed(values) ? failure() - : replaceResults(exchanges, *values, erase); + : replaceResults(exchanges, *values, replacements); } } // namespace @@ -806,7 +1048,7 @@ static LogicalResult emitBoundary(const BoundaryProgram& boundary, LogicalResult realizeDeferredBoundaries(ArrayRef boundaries, ArrayRef results, DeferredEmissionContext& context, - DeferredEraseSet& erase) { + DeferredReplacementMap& replacements) { ScheduledInfo* scheduled = nullptr; for (const BoundaryProgram& boundary : boundaries) { if (scheduled != boundary.key.scheduled) { @@ -815,7 +1057,7 @@ LogicalResult realizeDeferredBoundaries(ArrayRef boundaries, context.projectionAssemblies.clear(); scheduled = boundary.key.scheduled; } - if (failed(emitBoundary(boundary, results, context, erase))) + if (failed(emitBoundary(boundary, results, context, replacements))) return boundary.key.scheduled->op->emitOpError("phase 2 failed to realize a communication boundary"); } return success(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.hpp index d48e64d..b4f4e8f 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredBoundaryRealization.hpp @@ -1,5 +1,7 @@ #pragma once +#include "llvm/ADT/MapVector.h" + #include "mlir/IR/PatternMatch.h" #include "DeferredBoundaryPlanning.hpp" @@ -20,11 +22,12 @@ struct DeferredEmissionContext { projectionAssemblies; }; -using DeferredEraseSet = llvm::SetVector; +using DeferredReplacementMap = + llvm::MapVector; mlir::LogicalResult realizeDeferredBoundaries(mlir::ArrayRef boundaries, mlir::ArrayRef results, DeferredEmissionContext& context, - DeferredEraseSet& erase); + DeferredReplacementMap& replacements); } // namespace onnx_mlir::spatial diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp index e94fd91..5794662 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationModel.hpp @@ -114,6 +114,7 @@ struct RequirementCoordinate { enum class DeferredLeafForm { DirectSource, + ScalarProjection, GraphBatchProjection }; enum class DeferredAssemblySourceTransform { @@ -128,6 +129,12 @@ struct DeferredSliceTemplate { llvm::SmallVector strides; }; +struct DeferredStaticSliceGeometry { + llvm::SmallVector offsets; + llvm::SmallVector sizes; + llvm::SmallVector strides; +}; + struct DeferredProjectionLeafTemplate { DeferredLeafForm form = DeferredLeafForm::DirectSource; mlir::Value sourceRoot; @@ -201,6 +208,7 @@ struct RequirementFamily { mlir::Type publicationFragmentType; std::optional graphLanes; std::optional producerLocalOffsets; + std::optional producerProjection; }; struct LocalAvailabilityFamily { diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationRealization.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationRealization.cpp index 35f8b97..6def91b 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationRealization.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationRealization.cpp @@ -244,11 +244,15 @@ LogicalResult realizeDeferredCommunication(func::FuncOp funcOp) { return failure(); ConstantPool constants(funcOp, rewriter); DeferredEmissionContext context(rewriter, constants); - DeferredEraseSet erase; + DeferredReplacementMap replacements; if (failed(realizeDeferredBoundaries( - boundaries->boundaries, boundaries->results, context, erase))) + boundaries->boundaries, boundaries->results, context, replacements))) return failure(); - for (Operation *op : erase) { + for (auto [op, replacement] : replacements) { + if (op->getResult(0) == replacement) + return op->emitOpError( + "phase 2 cannot replace deferred communication with itself"); + op->getResult(0).replaceAllUsesWith(replacement); if (!op->use_empty()) return op->emitOpError( "phase 2 cannot erase deferred communication with live uses"); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.cpp index 77c6509..3045528 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.cpp @@ -39,6 +39,7 @@ static size_t hashSignature(const TransferEmissionSignature& signature) { signature.payload.getAsOpaquePointer(), signature.fragmentType.getAsOpaquePointer(), signature.hasGraphLane, + signature.hasProducerProjection, signature.sourceIsBatch); } @@ -218,6 +219,7 @@ TransferEmissionSignature getTransferEmissionSignature(const ExternalTransferFam producer->payload, family.requirement->publicationFragmentType, family.requirement->graphLanes.has_value(), + family.requirement->producerProjection.has_value(), producer->scheduled->isBatch()}; } diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.hpp index 13f4e64..3bfc109 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredCommunicationScheduling.hpp @@ -26,12 +26,14 @@ struct TransferEmissionSignature { mlir::Value payload; mlir::Type fragmentType; bool hasGraphLane = false; + bool hasProducerProjection = false; bool sourceIsBatch = false; bool operator==(const TransferEmissionSignature& other) const { return scheduled == other.scheduled && payload == other.payload && fragmentType == other.fragmentType && hasGraphLane == other.hasGraphLane + && hasProducerProjection == other.hasProducerProjection && sourceIsBatch == other.sourceIsBatch; } }; diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp index 42c836e..d15cc5a 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredProjectionAnalysis.cpp @@ -518,9 +518,15 @@ FailureOr analyzeDeferredProgramTemplate( auto result = dyn_cast(deferred.getSources()[index]); return result && isa(result.getOwner()); }); - if (graphProjection) { + bool scalarProjection = succeeded(sources) + && llvm::all_of(*sources, [&](unsigned index) { + auto result = dyn_cast(deferred.getSources()[index]); + return result && isa(result.getOwner()); + }); + if (graphProjection || scalarProjection) { DeferredProjectionLeafTemplate leaf; - leaf.form = DeferredLeafForm::GraphBatchProjection; + leaf.form = graphProjection ? DeferredLeafForm::GraphBatchProjection + : DeferredLeafForm::ScalarProjection; leaf.sourceRoot = slice.getSource(); leaf.replacementRoot = value; leaf.leadingProjection = slice; @@ -528,13 +534,14 @@ FailureOr analyzeDeferredProgramTemplate( SmallVector(slice.getMixedOffsets()), SmallVector(slice.getMixedSizes()), SmallVector(slice.getMixedStrides())}; - leaf.innerGeometry = { - SmallVector( - ArrayRef(slice.getMixedOffsets()).drop_front()), - SmallVector( - ArrayRef(slice.getMixedSizes()).drop_front()), - SmallVector( - ArrayRef(slice.getMixedStrides()).drop_front())}; + if (graphProjection) + leaf.innerGeometry = { + SmallVector( + ArrayRef(slice.getMixedOffsets()).drop_front()), + SmallVector( + ArrayRef(slice.getMixedSizes()).drop_front()), + SmallVector( + ArrayRef(slice.getMixedStrides()).drop_front())}; leaf.reconstructedType = cast(value.getType()); program.leaves.push_back(std::move(leaf)); return success(); diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp index e51b4f5..12166d7 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.cpp @@ -235,6 +235,20 @@ FailureOr materializeDeferredRequirement(RequirementFamily& requirement, return received; ProducedValue& producer = *requirement.producer; Value payload = producer.payload; + Location loc = requirement.exchange->deferred.getLoc(); + Value position = getSequencePosition( + requirement.targetLanes, lane, requirement.exchange->deferred, context, loc); + if (requirement.producerProjection) { + auto fragmentType = dyn_cast( + requirement.publicationFragmentType); + if (!fragmentType) + return failure(); + MixedSliceGeometry geometry = materializeGeometry( + *requirement.producerProjection, position, + requirement.exchange->deferred, context, loc); + return extractMixedSliceOrIdentity( + context.rewriter, loc, payload, fragmentType, geometry); + } if (payload.getType() == requirement.publicationFragmentType) return payload; auto payloadType = dyn_cast(payload.getType()); @@ -242,8 +256,6 @@ FailureOr materializeDeferredRequirement(RequirementFamily& requirement, if (!payloadType || !fragmentType || !requirement.producerLocalOffsets || payloadType.getRank() != fragmentType.getRank() + 1) return failure(); - Location loc = requirement.exchange->deferred.getLoc(); - Value position = getSequencePosition(requirement.targetLanes, lane, requirement.exchange->deferred, context, loc); Value offset = emitStaticIntLookup(*requirement.producerLocalOffsets, position, requirement.exchange->deferred, diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp index 8946cf2..c6c738a 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredResultRealization.hpp @@ -9,11 +9,7 @@ struct DeferredEmissionContext; struct DeferredResultPlan { DeferredExchangePlan* exchange = nullptr; llvm::SmallVector requirements; - struct SliceGeometry { - llvm::SmallVector offsets; - llvm::SmallVector sizes; - llvm::SmallVector strides; - }; + using SliceGeometry = DeferredStaticSliceGeometry; llvm::SmallVector innerGeometry; llvm::SmallVector assemblyGeometry; llvm::DenseMap residualValues; diff --git a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp index fc59f26..d55da71 100644 --- a/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp +++ b/src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/DeferredTransferPlanning.cpp @@ -294,16 +294,38 @@ static FailureOr findProducer(DeferredTransferPlan& plan, } struct RequirementPoint { + struct SliceGeometry { + SmallVector offsets; + SmallVector sizes; + SmallVector strides; + }; + ProducedValue* producer = nullptr; Type fragmentType; std::optional graphLane; std::optional localOffset; + std::optional producerProjection; bool sameFamily(const RequirementPoint& other) const { return producer == other.producer && fragmentType == other.fragmentType; } }; +static FailureOr> evaluateGeometryValues( + ArrayRef values, + DeferredLaneValueEvaluator& evaluator, + unsigned lane) { + SmallVector result; + result.reserve(values.size()); + for (OpFoldResult value : values) { + auto sequence = evaluator.evaluate(value); + if (failed(sequence)) + return failure(); + result.push_back(sequence->valueAt(lane)); + } + return result; +} + static FailureOr> resolveRequirementPoint(DeferredTransferPlan& plan, DeferredExchangePlan& exchange, @@ -355,11 +377,25 @@ resolveRequirementPoint(DeferredTransferPlan& plan, else { if (position != 0 || !isa(source.getOwner())) return std::optional(); - point.fragmentType = source.getType(); + point.fragmentType = leaf.form == DeferredLeafForm::ScalarProjection + ? Type(leaf.reconstructedType) + : source.getType(); auto producer = findProducer(plan, exchange.deferred, graphId.getInt(), source.getResultNumber(), std::nullopt); if (failed(producer)) return failure(); point.producer = *producer; + if (leaf.form == DeferredLeafForm::ScalarProjection) { + RequirementPoint::SliceGeometry geometry; + auto offsets = evaluateGeometryValues(leaf.leadingGeometry.offsets, evaluator, lane); + auto sizes = evaluateGeometryValues(leaf.leadingGeometry.sizes, evaluator, lane); + auto strides = evaluateGeometryValues(leaf.leadingGeometry.strides, evaluator, lane); + if (failed(offsets) || failed(sizes) || failed(strides)) + return failure(); + geometry.offsets = std::move(*offsets); + geometry.sizes = std::move(*sizes); + geometry.strides = std::move(*strides); + point.producerProjection = std::move(geometry); + } } return std::optional(point); } @@ -384,6 +420,27 @@ static void appendRequirementFamily(DeferredExchangePlan& exchange, }; family.graphLanes = sequence(&RequirementPoint::graphLane); family.producerLocalOffsets = sequence(&RequirementPoint::localOffset); + if (points.front().producerProjection) { + family.producerProjection.emplace(); + auto appendGeometry = [&](auto member, + SmallVectorImpl& target) { + for (size_t dimension = 0; + dimension < ((*points.front().producerProjection).*member).size(); + ++dimension) { + SmallVector values; + values.reserve(points.size()); + for (const RequirementPoint& point : points) + values.push_back(((*point.producerProjection).*member)[dimension]); + target.push_back(StaticIntSequence::fromValues(values)); + } + }; + appendGeometry(&RequirementPoint::SliceGeometry::offsets, + family.producerProjection->offsets); + appendGeometry(&RequirementPoint::SliceGeometry::sizes, + family.producerProjection->sizes); + appendGeometry(&RequirementPoint::SliceGeometry::strides, + family.producerProjection->strides); + } exchange.requirements.push_back(std::move(family)); }