add pipeline stages synchronization
Validate Operations / validate-operations (push) Has been cancelled

full ops throughput validation now passes
This commit is contained in:
NiccoloN
2026-08-11 11:34:31 +02:00
parent c55d9f3dad
commit 45072ca743
20 changed files with 763 additions and 250 deletions
+6 -3
View File
@@ -695,13 +695,16 @@ void PimCodeGen::codeGenSendOp(pim::PimSendOp sendOp, const StaticValueKnowledge
void PimCodeGen::codeGenWaitOp(
pim::PimWaitOp waitOp, const StaticValueKnowledge& knowledge) const {
auto eventRegister = indexOf(waitOp.getEventRegister(), knowledge);
assert(succeeded(eventRegister)
&& "pim.wait event register must be statically resolvable during codegen");
auto waitValue = indexOf(waitOp.getWaitValue(), knowledge);
assert(succeeded(eventRegister) && succeeded(waitValue)
&& "pim.wait operands must be statically resolvable during codegen");
if (*waitValue == 0)
return;
pim_binary::InstructionRecord instruction;
instruction.opcode = pim_binary::Opcode::wait;
instruction.generic1 = pim::checkedI32OrCrash(
*eventRegister, "wait event register");
instruction.generic2 = waitOp.getWaitValue();
instruction.generic2 = pim::checkedI32OrCrash(*waitValue, "wait value");
emitInstruction(instruction);
}
+2
View File
@@ -12,6 +12,7 @@
#include <limits>
#include <tuple>
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/ONNXToSpatialOptions.hpp"
@@ -78,6 +79,7 @@ spatial::SchedulingTarget getDefaultPimSchedulingTarget() {
target.residentWeightCapacity = crossbarCountInCore.getValue();
target.matrixRows = crossbarSize.getValue();
target.matrixColumns = crossbarSize.getValue();
target.synchronizationRegisterCount = kPimEventRegisterCount;
setDefaultPimInterProcessorLatencies(target);
return target;
@@ -368,11 +368,14 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeOp(spatial::SpatScheduledCom
return failure();
PimWaitOp::create(
rewriter, receiveOp->getLoc(), hostWaitLoad.getEventRegister(),
rewriter.getI32IntegerAttr(1));
hostWaitLoad.getWaitValue());
received = PimMemCopyHostToDevOp::create(
rewriter, receiveOp->getLoc(), outputBuffer.getType(), zero,
hostWaitLoad.getHostOffset(), outputBuffer, *hostBuffer, *sizeAttr)
.getOutput();
PimSyncOp::create(
rewriter, receiveOp->getLoc(), hostWaitLoad.getSourceCoreId(),
hostWaitLoad.getAcknowledgementEventRegister());
} else {
received = PimReceiveOp::create(
rewriter, receiveOp->getLoc(), outputBuffer.getType(), outputBuffer,
@@ -147,15 +147,42 @@ struct HostWaitLoadLowering : OpRewritePattern<spatial::SpatHostWaitLoadOp> {
return failure();
auto wait = pim::PimWaitOp::create(
rewriter, op.getLoc(), op.getEventRegister(),
rewriter.getI32IntegerAttr(1));
op.getWaitValue());
copyRaptorDebugAttrs(op.getOperation(), wait.getOperation());
return pim::PimMemCopyHostToDevOp::create(
Value output = pim::PimMemCopyHostToDevOp::create(
rewriter, op.getLoc(), outputBuffer.getType(), zero,
op.getHostOffset(), outputBuffer, *hostBuffer, sizeAttr).getOutput();
auto sync = pim::PimSyncOp::create(
rewriter, op.getLoc(), op.getSourceCoreId(),
op.getAcknowledgementEventRegister());
copyRaptorDebugAttrs(op.getOperation(), sync.getOperation());
return output;
});
}
};
struct SyncLowering : OpRewritePattern<spatial::SpatSyncOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatSyncOp op,
PatternRewriter& rewriter) const override {
rewriter.replaceOpWithNewOp<pim::PimSyncOp>(
op, op.getTargetCoreId(), op.getEventRegister());
return success();
}
};
struct WaitLowering : OpRewritePattern<spatial::SpatWaitOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatWaitOp op,
PatternRewriter& rewriter) const override {
rewriter.replaceOpWithNewOp<pim::PimWaitOp>(
op, op.getEventRegister(), op.getWaitValue());
return success();
}
};
struct ExtractRowsLowering : OpRewritePattern<spatial::SpatExtractRowsOp> {
using OpRewritePattern::OpRewritePattern;
@@ -200,7 +227,8 @@ struct ConcatLowering : OpRewritePattern<spatial::SpatConcatOp> {
void populateChannelLoweringPatterns(RewritePatternSet& patterns) {
patterns.add<ChannelSendLowering, ChannelReceiveLowering,
HostStoreSyncLowering, HostWaitLoadLowering,
ExtractRowsLowering, ConcatLowering>(patterns.getContext());
SyncLowering, WaitLowering, ExtractRowsLowering,
ConcatLowering>(patterns.getContext());
}
} // namespace onnx_mlir
@@ -128,6 +128,8 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
spatial::SpatChannelSendOp,
spatial::SpatHostStoreSyncOp,
spatial::SpatHostWaitLoadOp,
spatial::SpatSyncOp,
spatial::SpatWaitOp,
spatial::SpatExtractRowsOp>();
RewritePatternSet initialPatterns(ctx);
@@ -223,6 +225,8 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
spatial::SpatChannelSendOp,
spatial::SpatHostStoreSyncOp,
spatial::SpatHostWaitLoadOp,
spatial::SpatSyncOp,
spatial::SpatWaitOp,
spatial::SpatExtractRowsOp>();
SmallVector<pim::PimCoreOp> coreOps;
@@ -274,6 +278,8 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
spatial::SpatChannelSendOp,
spatial::SpatHostStoreSyncOp,
spatial::SpatHostWaitLoadOp,
spatial::SpatSyncOp,
spatial::SpatWaitOp,
spatial::SpatExtractRowsOp>();
RewritePatternSet communicationPatterns(ctx);
@@ -302,12 +302,19 @@ static FailureOr<int64_t> getShapedByteSize(MemRefType type) {
return static_cast<int64_t>(*byteSize);
}
static FailureOr<SmallVector<int64_t>>
struct LogicalCopyShape {
SmallVector<int64_t> dimensions;
Type elementType;
};
static bool isPackedByteBuffer(MemRefType type) {
return type.getRank() == 1 && type.getElementType().isInteger(8);
}
static FailureOr<LogicalCopyShape>
inferLogicalCopyShape(MemRefType targetType, MemRefType sourceType, int64_t size) {
if (!targetType.hasStaticShape() || !sourceType.hasStaticShape())
return failure();
if (targetType.getElementType() != sourceType.getElementType() || targetType.getRank() != sourceType.getRank())
return failure();
auto targetBytes = getShapedByteSize(targetType);
auto sourceBytes = getShapedByteSize(sourceType);
@@ -316,18 +323,37 @@ inferLogicalCopyShape(MemRefType targetType, MemRefType sourceType, int64_t size
bool targetMatches = *targetBytes == size;
bool sourceMatches = *sourceBytes == size;
if (targetMatches && sourceMatches && targetType.getShape() != sourceType.getShape())
bool matchingTypes = targetType.getElementType() == sourceType.getElementType()
&& targetType.getRank() == sourceType.getRank();
if (matchingTypes) {
if (targetMatches && sourceMatches
&& targetType.getShape() != sourceType.getShape())
return failure();
MemRefType logicalType = targetMatches ? targetType : sourceType;
if (targetMatches || sourceMatches)
return LogicalCopyShape {
SmallVector<int64_t>(logicalType.getShape()),
logicalType.getElementType()};
return failure();
if (targetMatches)
return SmallVector<int64_t>(targetType.getShape().begin(), targetType.getShape().end());
if (sourceMatches)
return SmallVector<int64_t>(sourceType.getShape().begin(), sourceType.getShape().end());
}
if (targetMatches && isPackedByteBuffer(sourceType))
return LogicalCopyShape {
SmallVector<int64_t>(targetType.getShape()),
targetType.getElementType()};
if (sourceMatches && isPackedByteBuffer(targetType))
return LogicalCopyShape {
SmallVector<int64_t>(sourceType.getShape()),
sourceType.getElementType()};
return failure();
}
static FailureOr<int64_t> getContiguousSuffixRank(Value value, ArrayRef<int64_t> copyShape) {
static FailureOr<int64_t> getContiguousSuffixRank(
Value value, ArrayRef<int64_t> copyShape, Type elementType = {}) {
auto type = dyn_cast<MemRefType>(value.getType());
if (type && elementType && isPackedByteBuffer(type))
return copyShape.size();
if (!type || !type.hasStaticShape() || !hasByteSizedElementType(type.getElementType())
|| (elementType && type.getElementType() != elementType)
|| type.getRank() != static_cast<int64_t>(copyShape.size()))
return failure();
if (llvm::any_of(copyShape, [](int64_t dim) { return dim <= 0; }))
@@ -351,6 +377,30 @@ static FailureOr<int64_t> getContiguousSuffixRank(Value value, ArrayRef<int64_t>
return contiguousSuffixRank;
}
static FailureOr<SmallVector<int64_t>> getOuterByteStrides(
Value value, const LogicalCopyShape &copyShape, size_t outerRank) {
auto type = cast<MemRefType>(value.getType());
SmallVector<int64_t> strides;
if (isPackedByteBuffer(type))
strides = computeRowMajorStrides(copyShape.dimensions);
else {
auto proven = getProvenMemRefStrides(value);
if (failed(proven))
return failure();
strides = std::move(*proven);
}
int64_t elementByteWidth = static_cast<int64_t>(
getElementTypeSizeInBytes(copyShape.elementType));
SmallVector<int64_t> result;
for (int64_t stride : ArrayRef<int64_t>(strides).take_front(outerRank)) {
auto byteStride = checkedPositiveMul(stride, elementByteWidth);
if (failed(byteStride))
return failure();
result.push_back(*byteStride);
}
return result;
}
static FailureOr<CopyEndpointPlan> analyzeCopyEndpoint(Value value, Value initialByteOffset, MemRefType logicalType) {
if (!logicalType.hasStaticShape() || !hasByteSizedElementType(logicalType.getElementType()))
return failure();
@@ -448,8 +498,10 @@ analyzeCopyRewrite(Value target, Value source, Value targetOffset, Value sourceO
if (failed(logicalCopyShape))
return failure();
auto targetSuffixRank = getContiguousSuffixRank(target, *logicalCopyShape);
auto sourceSuffixRank = getContiguousSuffixRank(source, *logicalCopyShape);
auto targetSuffixRank = getContiguousSuffixRank(
target, logicalCopyShape->dimensions, logicalCopyShape->elementType);
auto sourceSuffixRank = getContiguousSuffixRank(
source, logicalCopyShape->dimensions, logicalCopyShape->elementType);
if (failed(targetSuffixRank) || failed(sourceSuffixRank))
return failure();
@@ -458,23 +510,24 @@ analyzeCopyRewrite(Value target, Value source, Value targetOffset, Value sourceO
plan.source = *sourcePlan;
int64_t contiguousSuffixRank = std::min(*targetSuffixRank, *sourceSuffixRank);
if (contiguousSuffixRank == static_cast<int64_t>(logicalCopyShape->size())) {
if (contiguousSuffixRank
== static_cast<int64_t>(logicalCopyShape->dimensions.size())) {
plan.kind = CopyRewritePlan::Kind::Direct;
plan.directBytes = size;
return plan;
}
auto targetStrides = getProvenMemRefStrides(target);
auto sourceStrides = getProvenMemRefStrides(source);
if (failed(targetStrides) || failed(sourceStrides))
return failure();
int64_t elementByteWidth = static_cast<int64_t>(getElementTypeSizeInBytes(targetType.getElementType()));
int64_t elementByteWidth = static_cast<int64_t>(
getElementTypeSizeInBytes(logicalCopyShape->elementType));
plan.kind = CopyRewritePlan::Kind::Loop;
plan.loop.targetBaseOffset = plan.target.offset;
plan.loop.sourceBaseOffset = plan.source.offset;
plan.loop.outerShape.assign(logicalCopyShape->begin(), logicalCopyShape->end() - contiguousSuffixRank);
SmallVector<int64_t> chunkShape(logicalCopyShape->end() - contiguousSuffixRank, logicalCopyShape->end());
plan.loop.outerShape.assign(
logicalCopyShape->dimensions.begin(),
logicalCopyShape->dimensions.end() - contiguousSuffixRank);
SmallVector<int64_t> chunkShape(
logicalCopyShape->dimensions.end() - contiguousSuffixRank,
logicalCopyShape->dimensions.end());
auto outerElements = checkedPositiveProduct(plan.loop.outerShape);
auto chunkElements = checkedPositiveProduct(chunkShape);
auto chunkBytes = failed(chunkElements)
@@ -484,18 +537,14 @@ analyzeCopyRewrite(Value target, Value source, Value targetOffset, Value sourceO
return failure();
plan.loop.outerElements = *outerElements;
plan.loop.chunkBytes = *chunkBytes;
for (int64_t stride : ArrayRef<int64_t>(*targetStrides).take_front(plan.loop.outerShape.size())) {
auto byteStride = checkedPositiveMul(stride, elementByteWidth);
if (failed(byteStride))
return failure();
plan.loop.targetOuterByteStrides.push_back(*byteStride);
}
for (int64_t stride : ArrayRef<int64_t>(*sourceStrides).take_front(plan.loop.outerShape.size())) {
auto byteStride = checkedPositiveMul(stride, elementByteWidth);
if (failed(byteStride))
return failure();
plan.loop.sourceOuterByteStrides.push_back(*byteStride);
}
auto targetStrides = getOuterByteStrides(
target, *logicalCopyShape, plan.loop.outerShape.size());
auto sourceStrides = getOuterByteStrides(
source, *logicalCopyShape, plan.loop.outerShape.size());
if (failed(targetStrides) || failed(sourceStrides))
return failure();
plan.loop.targetOuterByteStrides = std::move(*targetStrides);
plan.loop.sourceOuterByteStrides = std::move(*sourceStrides);
if (plan.loop.chunkBytes <= 0)
return failure();
return plan;
@@ -602,18 +602,30 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
PatternRewriter rewriter(ctx);
SmallVector<MemRefCopyWorkItem> copyWorklist;
SmallVector<PimMemCopyDevToHostOp> hostToHostCopies;
llvm::SmallPtrSet<Operation*, 16> seenCopyOps;
llvm::SmallPtrSet<Operation*, 4> seenHostToHostCopies;
auto addCopyOp = [&](memref::CopyOp copyOp, const StaticValueKnowledge& knowledge) {
if (seenCopyOps.insert(copyOp.getOperation()).second)
copyWorklist.push_back({copyOp, knowledge});
};
auto collectCopy = [&](Operation &op,
const StaticValueKnowledge &knowledge) {
if (auto copyOp = dyn_cast<memref::CopyOp>(&op))
addCopyOp(copyOp, knowledge);
if (auto copyOp = dyn_cast<PimMemCopyDevToHostOp>(&op);
copyOp
&& isHostBackedPimAddress(copyOp.getDeviceSource(), knowledge)
&& isHostBackedPimAddress(copyOp.getHostTarget(), knowledge)
&& seenHostToHostCopies.insert(copyOp).second)
hostToHostCopies.push_back(copyOp);
};
moduleOp.walk([&](pim::PimCoreOp coreOp) {
StaticValueKnowledge knowledge = seedCoreKnowledge(coreOp);
(void) walkPimCoreBlockStructurally(
coreOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) {
if (auto copyOp = dyn_cast<memref::CopyOp>(&op))
addCopyOp(copyOp, opKnowledge);
collectCopy(op, opKnowledge);
return success();
});
});
@@ -622,8 +634,7 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
StaticValueKnowledge knowledge = seedCoreBatchKnowledge(coreBatchOp, lane);
(void) walkPimCoreBlockStructurally(
coreBatchOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) {
if (auto copyOp = dyn_cast<memref::CopyOp>(&op))
addCopyOp(copyOp, opKnowledge);
collectCopy(op, opKnowledge);
return success();
});
}
@@ -631,6 +642,22 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
bool hasFailed = false;
Value zeroOffset = getOrCreateIndexConstant(rewriter, funcOp, 0);
for (PimMemCopyDevToHostOp copyOp : hostToHostCopies) {
rewriter.setInsertionPoint(copyOp);
auto scratchType = MemRefType::get(
{copyOp.getSize()}, rewriter.getI8Type());
Value scratch = memref::AllocOp::create(
rewriter, copyOp.getLoc(), scratchType);
auto load = PimMemCopyHostToDevOp::create(
rewriter, copyOp.getLoc(), scratchType, zeroOffset,
copyOp.getDeviceSourceOffset(), scratch, copyOp.getDeviceSource(),
copyOp.getSizeAttr());
auto store = PimMemCopyDevToHostOp::create(
rewriter, copyOp.getLoc(), copyOp.getHostTarget().getType(),
copyOp.getHostTargetOffset(), zeroOffset, copyOp.getHostTarget(),
load.getOutput(), copyOp.getSizeAttr());
rewriter.replaceOp(copyOp, store.getOutput());
}
for (const MemRefCopyWorkItem& workItem : copyWorklist) {
memref::CopyOp copyOp = workItem.copyOp;
rewriter.setInsertionPoint(copyOp);
+1 -1
View File
@@ -136,7 +136,7 @@ def PimWaitOp : PimOp<"wait", []> {
let arguments = (ins
Index:$eventRegister,
I32Attr:$waitValue
Index:$waitValue
);
let assemblyFormat = [{
@@ -242,11 +242,146 @@ static void appendReceive(BoundaryProgram &boundary,
target.collection, {slice}, {0, 1}, {target.position}, {lanes}, lanes});
}
struct HostTransferRef {
ExternalTransferFamily *family = nullptr;
size_t index = 0;
};
static unsigned getBarrierRoundCount(size_t coreCount) {
unsigned rounds = 0;
for (size_t distance = 1; distance < coreCount; distance *= 2)
++rounds;
return rounds;
}
static LogicalResult assignPipelineSynchronization(
DeferredTransferPlan &transfers,
ArrayRef<BoundaryProgram> boundaries,
size_t synchronizationRegisterCount) {
bool pipelined = false;
for (ScheduledInfo &scheduled : transfers.scheduled) {
if (scheduled.pipelineStages.empty())
continue;
pipelined = true;
llvm::append_range(transfers.downstreamCores, scheduled.cores);
for (auto [core, stage] :
llvm::zip_equal(scheduled.cores, scheduled.pipelineStages))
if (stage == 0)
transfers.stageZeroCores.push_back(core);
}
if (!pipelined)
return success();
transfers.synchronizationRegisterCount = synchronizationRegisterCount;
llvm::sort(transfers.stageZeroCores);
transfers.stageZeroCores.erase(
llvm::unique(transfers.stageZeroCores), transfers.stageZeroCores.end());
llvm::sort(transfers.downstreamCores);
transfers.downstreamCores.erase(
llvm::unique(transfers.downstreamCores),
transfers.downstreamCores.end());
llvm::erase_if(transfers.downstreamCores, [&](int64_t core) {
return llvm::is_contained(transfers.stageZeroCores, core);
});
DenseMap<int64_t, SmallVector<HostTransferRef>> incomingByCore;
DenseMap<ExternalTransferFamily *, SmallVector<int64_t>> eventRegisters;
DenseMap<ExternalTransferFamily *, SmallVector<int64_t>> waitValues;
DenseMap<ExternalTransferFamily *, SmallVector<int64_t>> acknowledgementRegisters;
auto initialize = [&](ExternalTransferFamily &family) {
size_t count = family.targetCores.size();
eventRegisters.try_emplace(&family, count, 0);
waitValues.try_emplace(&family, count, 0);
acknowledgementRegisters.try_emplace(&family, count, 0);
};
for (const BoundaryProgram &boundary : boundaries)
for (const BoundaryInstruction &instruction : boundary.instructions) {
auto *receive = std::get_if<EmitReceiveAssemblyRun>(&instruction);
if (!receive || receive->slices.empty()
|| !receive->slices.front().family->hostRouted)
continue;
for (const ScheduledTransferSlice &slice : receive->slices) {
ExternalTransferFamily &family = *slice.family;
initialize(family);
for (size_t offset = 0; offset < slice.transferCount; ++offset) {
size_t index = slice.familyOffset + offset;
int64_t source = family.sourceCores.valueAt(index);
int64_t target = family.targetCores.valueAt(index);
incomingByCore[target].push_back({&family, index});
++transfers.hostAcknowledgementCounts[source];
}
}
}
unsigned barrierRounds = getBarrierRoundCount(
transfers.stageZeroCores.size());
bool stageZeroNeedsAcknowledgements = llvm::any_of(
transfers.stageZeroCores, [&](int64_t core) {
return transfers.hostAcknowledgementCounts.contains(core);
});
for (auto &[target, incoming] : incomingByCore) {
bool needsAcknowledgementRegister =
transfers.hostAcknowledgementCounts.contains(target);
bool stageZero = llvm::is_contained(transfers.stageZeroCores, target);
size_t reserved = stageZero
? barrierRounds + (stageZeroNeedsAcknowledgements ? 1 : 0)
: 1 + (needsAcknowledgementRegister ? 1 : 0);
if (reserved >= synchronizationRegisterCount) {
incoming.front().family->requirement->exchange->deferred.emitOpError(
"pipeline synchronization leaves no event register for incoming host transfers");
return failure();
}
size_t groupCount = std::min(
incoming.size(), synchronizationRegisterCount - reserved);
// One wait consumes a complete consecutive group of producer signals.
SmallVector<size_t> groupSizes(groupCount);
for (size_t ordinal = 0; ordinal < incoming.size(); ++ordinal)
++groupSizes[ordinal * groupCount / incoming.size()];
SmallVector<bool> first(groupCount, true);
for (size_t ordinal = 0; ordinal < incoming.size(); ++ordinal) {
size_t group = ordinal * groupCount / incoming.size();
HostTransferRef transfer = incoming[ordinal];
eventRegisters[transfer.family][transfer.index] = group;
acknowledgementRegisters[transfer.family][transfer.index] =
synchronizationRegisterCount - 1;
if (first[group]) {
waitValues[transfer.family][transfer.index] = groupSizes[group];
first[group] = false;
}
}
}
for (auto &[family, values] : eventRegisters) {
family->eventRegisters = StaticIntSequence::fromValues(values);
family->waitValues = StaticIntSequence::fromValues(waitValues[family]);
family->acknowledgementEventRegisters =
StaticIntSequence::fromValues(acknowledgementRegisters[family]);
}
if (!transfers.stageZeroCores.empty()) {
size_t reserved = barrierRounds
+ (stageZeroNeedsAcknowledgements ? 1 : 0);
if (reserved > synchronizationRegisterCount)
return transfers.scheduled.front().op->emitOpError(
"pipeline stage-zero barrier requires more synchronization registers than the target provides");
}
if (!transfers.downstreamCores.empty()) {
bool needsAcknowledgements = llvm::any_of(
transfers.downstreamCores, [&](int64_t core) {
return transfers.hostAcknowledgementCounts.contains(core);
});
if (1 + (needsAcknowledgements ? 1 : 0)
> synchronizationRegisterCount)
return transfers.scheduled.front().op->emitOpError(
"pipeline stage-zero release requires more synchronization registers than the target provides");
}
return success();
}
} // namespace
FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan(
DeferredTransferPlan &transfers,
const ScheduledCommunicationPlan &schedule) {
const ScheduledCommunicationPlan &schedule,
size_t synchronizationRegisterCount) {
DeferredBoundaryPlan result;
SmallVector<BoundaryProgram> boundaries;
DenseMap<BoundaryKey, unsigned> indices;
@@ -373,6 +508,9 @@ FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan(
return std::tie(scheduledOrder[lhs.key.first], lhs.key.second)
< std::tie(scheduledOrder[rhs.key.first], rhs.key.second);
});
if (failed(assignPipelineSynchronization(
transfers, boundaries, synchronizationRegisterCount)))
return failure();
result.boundaries = std::move(boundaries);
return result;
}
@@ -53,6 +53,7 @@ struct DeferredBoundaryPlan {
};
mlir::FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan(DeferredTransferPlan& transfers,
const ScheduledCommunicationPlan& schedule);
const ScheduledCommunicationPlan& schedule,
size_t synchronizationRegisterCount);
} // namespace onnx_mlir::spatial
@@ -4,6 +4,7 @@
#include "DeferredBoundaryRealization.hpp"
#include "DeferredProjectionAnalysis.hpp"
#include "DeferredResultRealization.hpp"
#include "DeferredTransferPlanning.hpp"
#include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/StaticIntGrid.hpp"
#include "src/Accelerators/PIM/Common/IR/StaticIntSequence.hpp"
@@ -21,6 +22,8 @@ struct LogicalTransferMetadataView {
StaticIntSequenceChain targetCores;
StaticIntSequenceChain hostOffsets;
StaticIntSequenceChain eventRegisters;
StaticIntSequenceChain waitValues;
StaticIntSequenceChain acknowledgementEventRegisters;
StaticIntSequenceChain targetLanes;
StaticIntSequenceChain localOffsets;
SmallVector<StaticIntSequenceChain> projectionOffsets;
@@ -91,6 +94,9 @@ static void appendMetadata(const ScheduledTransferSlice &slice, LogicalTransferM
metadata.hostOffsets.append(family.hostOffsets, familyIndex, count);
metadata.eventRegisters.append(
family.eventRegisters, familyIndex, count);
metadata.waitValues.append(family.waitValues, familyIndex, count);
metadata.acknowledgementEventRegisters.append(
family.acknowledgementEventRegisters, familyIndex, count);
}
metadata.targetLanes.append(StaticIntSequence::affine(targetLane, 1, count));
if (family.requirement->producerLocalOffsets)
@@ -296,13 +302,21 @@ static FailureOr<Value> emitReceiveValue(ArrayRef<ScheduledTransferSlice> slices
if (failed(grids)) return failure();
std::optional<StaticIntGrid> hostOffsets;
std::optional<StaticIntGrid> eventRegisters;
std::optional<StaticIntGrid> waitValues;
std::optional<StaticIntGrid> acknowledgementEventRegisters;
if (slices.front().family->hostRouted) {
auto offsets = buildGrid(metadata.hostOffsets);
auto events = buildGrid(metadata.eventRegisters);
if (failed(offsets) || failed(events))
auto waits = buildGrid(metadata.waitValues);
auto acknowledgements = buildGrid(
metadata.acknowledgementEventRegisters);
if (failed(offsets) || failed(events) || failed(waits)
|| failed(acknowledgements))
return failure();
hostOffsets = std::move(*offsets);
eventRegisters = std::move(*events);
waitValues = std::move(*waits);
acknowledgementEventRegisters = std::move(*acknowledgements);
}
Value position = lane ? lane : context.constants.getIndex(0);
Value row = context.constants.getIndex(0);
@@ -319,6 +333,10 @@ static FailureOr<Value> emitReceiveValue(ArrayRef<ScheduledTransferSlice> slices
hostOffsets->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc()),
eventRegisters->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc()),
waitValues->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc()),
acknowledgementEventRegisters->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc()));
receive = op;
output = op.getOutput();
@@ -387,6 +405,8 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
std::optional<StaticIntGrid> positions;
std::optional<StaticIntGrid> hostOffsets;
std::optional<StaticIntGrid> eventRegisters;
std::optional<StaticIntGrid> waitValues;
std::optional<StaticIntGrid> acknowledgementEventRegisters;
bool hostRouted = run.slices.front().family->hostRouted;
auto metadataByEntry = buildRectangularReceiveMetadata(run, laneCount);
if (succeeded(metadataByEntry)) {
@@ -402,10 +422,17 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
&LogicalTransferMetadataView::hostOffsets);
auto events = buildRows(
&LogicalTransferMetadataView::eventRegisters);
if (failed(offsets) || failed(events))
auto waits = buildRows(
&LogicalTransferMetadataView::waitValues);
auto acknowledgements = buildRows(
&LogicalTransferMetadataView::acknowledgementEventRegisters);
if (failed(offsets) || failed(events) || failed(waits)
|| failed(acknowledgements))
return failure();
hostOffsets = std::move(*offsets);
eventRegisters = std::move(*events);
waitValues = std::move(*waits);
acknowledgementEventRegisters = std::move(*acknowledgements);
}
SmallVector<StaticIntSequence> positionRows;
for (unsigned position : run.positions)
@@ -456,10 +483,17 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
&LogicalTransferMetadataView::hostOffsets);
auto events = buildGrid(
&LogicalTransferMetadataView::eventRegisters);
if (failed(offsets) || failed(events))
auto waits = buildGrid(
&LogicalTransferMetadataView::waitValues);
auto acknowledgements = buildGrid(
&LogicalTransferMetadataView::acknowledgementEventRegisters);
if (failed(offsets) || failed(events) || failed(waits)
|| failed(acknowledgements))
return failure();
hostOffsets = std::move(*offsets);
eventRegisters = std::move(*events);
waitValues = std::move(*waits);
acknowledgementEventRegisters = std::move(*acknowledgements);
}
SmallVector<StaticIntSequence> positionColumns;
for (const StaticIntSequenceChain &values : positionsByLane)
@@ -491,6 +525,10 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
hostOffsets->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc),
eventRegisters->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc),
waitValues->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc),
acknowledgementEventRegisters->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc));
receive = op;
output = op.getOutput();
@@ -1158,9 +1196,199 @@ static LogicalResult emitBoundary(const BoundaryProgram &boundary, ArrayRef<Defe
return failed(values) ? failure() : replaceResults(exchanges, *values, replacements);
}
static unsigned getBarrierRoundCount(size_t coreCount) {
unsigned rounds = 0;
for (size_t distance = 1; distance < coreCount; distance *= 2)
++rounds;
return rounds;
}
static LogicalResult emitCompletionSynchronization(
DeferredTransferPlan &transfers, DeferredEmissionContext &context) {
if (transfers.synchronizationRegisterCount == 0)
return success();
size_t acknowledgementRegister =
transfers.synchronizationRegisterCount - 1;
unsigned barrierRounds = getBarrierRoundCount(
transfers.stageZeroCores.size());
bool stageZeroNeedsAcknowledgements = llvm::any_of(
transfers.stageZeroCores, [&](int64_t core) {
return transfers.hostAcknowledgementCounts.contains(core);
});
size_t firstBarrierRegister = acknowledgementRegister
- (stageZeroNeedsAcknowledgements ? 1 : 0);
DenseMap<int64_t, unsigned> stageZeroRank;
for (auto [rank, core] : llvm::enumerate(transfers.stageZeroCores))
stageZeroRank[core] = rank;
DenseMap<int64_t, unsigned> downstreamRank;
for (auto [rank, core] : llvm::enumerate(transfers.downstreamCores))
downstreamRank[core] = rank;
auto getReleaseRegister = [&](int64_t core) {
return acknowledgementRegister
- (transfers.hostAcknowledgementCounts.contains(core) ? 1 : 0);
};
for (ScheduledInfo &scheduled : transfers.scheduled) {
Block *block = scheduled.blocks.front();
context.rewriter.setInsertionPoint(block->getTerminator());
Location loc = scheduled.op->getLoc();
Value lane;
if (auto batch = dyn_cast<SpatScheduledComputeBatch>(scheduled.op))
lane = *batch.getLaneArgument();
SmallVector<int64_t> acknowledgementCounts, releaseRegisters;
SmallVector<int64_t> releaseWaitValues, leftTargets, leftRegisters;
SmallVector<int64_t> rightTargets, rightRegisters;
LaneSet barrierLanes, leaderLanes, leftLanes, rightLanes;
for (auto [index, core] : llvm::enumerate(scheduled.cores)) {
acknowledgementCounts.push_back(
transfers.hostAcknowledgementCounts.lookup(core));
if (stageZeroRank.contains(core))
barrierLanes = barrierLanes.unite(
LaneSet::range(index, index + 1));
if (!transfers.stageZeroCores.empty()
&& core == transfers.stageZeroCores.front())
leaderLanes = leaderLanes.unite(LaneSet::range(index, index + 1));
auto rank = downstreamRank.find(core);
if (rank == downstreamRank.end()) {
releaseRegisters.push_back(0);
releaseWaitValues.push_back(0);
leftTargets.push_back(core);
leftRegisters.push_back(0);
rightTargets.push_back(core);
rightRegisters.push_back(0);
continue;
}
releaseRegisters.push_back(getReleaseRegister(core));
releaseWaitValues.push_back(1);
size_t left = 2 * rank->second + 1;
size_t right = left + 1;
if (left < transfers.downstreamCores.size()) {
int64_t child = transfers.downstreamCores[left];
leftTargets.push_back(child);
leftRegisters.push_back(getReleaseRegister(child));
leftLanes = leftLanes.unite(LaneSet::range(index, index + 1));
} else {
leftTargets.push_back(core);
leftRegisters.push_back(0);
}
if (right < transfers.downstreamCores.size()) {
int64_t child = transfers.downstreamCores[right];
rightTargets.push_back(child);
rightRegisters.push_back(getReleaseRegister(child));
rightLanes = rightLanes.unite(LaneSet::range(index, index + 1));
} else {
rightTargets.push_back(core);
rightRegisters.push_back(0);
}
}
Value runtimeLane = lane ? lane : context.constants.getIndex(0);
auto emitForLanes = [&](const LaneSet &active, auto emit) -> LogicalResult {
if (active.empty())
return success();
if (!lane) {
if (active.contains(0))
emit();
return success();
}
auto condition = emitLaneCondition(
active, lane, scheduled.cores.size(), scheduled.op, context, loc);
if (failed(condition))
return failure();
auto conditional = scf::IfOp::create(
context.rewriter, loc, TypeRange {}, *condition, false);
OpBuilder::InsertionGuard guard(context.rewriter);
context.rewriter.setInsertionPoint(
conditional.getThenRegion().front().getTerminator());
emit();
return success();
};
Value acknowledgementCount = emitStaticIntLookup(
StaticIntSequence::fromValues(acknowledgementCounts),
runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
SpatWaitOp::create(
context.rewriter, loc,
context.constants.getIndex(acknowledgementRegister),
acknowledgementCount);
// Dissemination barrier: every round doubles the covered stage-zero peers.
auto emitBarrier = [&]() {
for (unsigned round = 0; round < barrierRounds; ++round) {
SmallVector<int64_t> targets;
targets.reserve(scheduled.cores.size());
size_t distance = size_t {1} << round;
for (int64_t core : scheduled.cores) {
auto rank = stageZeroRank.find(core);
targets.push_back(rank == stageZeroRank.end()
? core
: transfers.stageZeroCores[
(rank->second + distance)
% transfers.stageZeroCores.size()]);
}
Value target = emitStaticIntLookup(
StaticIntSequence::fromValues(targets),
runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
Value eventRegister = context.constants.getIndex(
firstBarrierRegister - round);
SpatSyncOp::create(
context.rewriter, loc, target, eventRegister);
SpatWaitOp::create(
context.rewriter, loc, eventRegister,
context.constants.getIndex(1));
}
};
if (barrierRounds > 0
&& failed(emitForLanes(barrierLanes, emitBarrier)))
return failure();
// Gate downstream restarts so no core advances the simulator input
// iteration ahead of stage zero.
if (!transfers.downstreamCores.empty()
&& failed(emitForLanes(leaderLanes, [&]() {
int64_t root = transfers.downstreamCores.front();
SpatSyncOp::create(
context.rewriter, loc, context.constants.getIndex(root),
context.constants.getIndex(getReleaseRegister(root)));
})))
return failure();
Value releaseRegister = emitStaticIntLookup(
StaticIntSequence::fromValues(releaseRegisters), runtimeLane,
scheduled.op, context.constants, context.rewriter, loc);
Value releaseWaitValue = emitStaticIntLookup(
StaticIntSequence::fromValues(releaseWaitValues), runtimeLane,
scheduled.op, context.constants, context.rewriter, loc);
SpatWaitOp::create(
context.rewriter, loc, releaseRegister, releaseWaitValue);
auto emitChild = [&](ArrayRef<int64_t> targets,
ArrayRef<int64_t> registers) {
Value target = emitStaticIntLookup(
StaticIntSequence::fromValues(targets), runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
Value eventRegister = emitStaticIntLookup(
StaticIntSequence::fromValues(registers), runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
SpatSyncOp::create(context.rewriter, loc, target, eventRegister);
};
if (failed(emitForLanes(leftLanes, [&]() {
emitChild(leftTargets, leftRegisters);
}))
|| failed(emitForLanes(rightLanes, [&]() {
emitChild(rightTargets, rightRegisters);
})))
return failure();
}
return success();
}
} // namespace
LogicalResult realizeDeferredBoundaries(ArrayRef<BoundaryProgram> boundaries, ArrayRef<DeferredResultPlan> results, DeferredEmissionContext &context,
LogicalResult realizeDeferredBoundaries(ArrayRef<BoundaryProgram> boundaries, ArrayRef<DeferredResultPlan> results,
DeferredTransferPlan &transfers, DeferredEmissionContext &context,
DeferredReplacementMap &replacements) {
ScheduledInfo *scheduled = nullptr;
for (const BoundaryProgram &boundary : boundaries) {
@@ -1171,7 +1399,7 @@ LogicalResult realizeDeferredBoundaries(ArrayRef<BoundaryProgram> boundaries, Ar
if (failed(emitBoundary(boundary, results, context, replacements)))
return boundary.key.first->op->emitOpError("phase 2 failed to realize a communication boundary");
}
return success();
return emitCompletionSynchronization(transfers, context);
}
} // namespace onnx_mlir::spatial
@@ -42,6 +42,7 @@ using DeferredReplacementMap =
mlir::LogicalResult realizeDeferredBoundaries(mlir::ArrayRef<BoundaryProgram> boundaries,
mlir::ArrayRef<DeferredResultPlan> results,
DeferredTransferPlan& transfers,
DeferredEmissionContext& context,
DeferredReplacementMap& replacements);
@@ -236,6 +236,9 @@ struct ExternalTransferFamily {
StaticIntSequence channelIds = StaticIntSequence::uniform(0, 1);
StaticIntSequence hostOffsets = StaticIntSequence::uniform(0, 1);
StaticIntSequence eventRegisters = StaticIntSequence::uniform(0, 1);
StaticIntSequence waitValues = StaticIntSequence::uniform(1, 1);
StaticIntSequence acknowledgementEventRegisters =
StaticIntSequence::uniform(0, 1);
bool hostRouted = false;
};
@@ -232,7 +232,8 @@ LogicalResult realizeDeferredCommunication(func::FuncOp funcOp,
auto schedule = scheduleDeferredCommunication(funcOp, *transfers);
if (failed(schedule) || failed(verifyPlannedCommunicationDeadlockFree(funcOp, transfers->stepCounts, *schedule)))
return funcOp.emitOpError("phase 2 failed to schedule symbolic communication");
auto boundaries = buildDeferredBoundaryPlan(*transfers, *schedule);
auto boundaries = buildDeferredBoundaryPlan(
*transfers, *schedule, target.synchronizationRegisterCount);
if (failed(boundaries))
return funcOp.emitOpError("phase 2 failed to build sparse boundary programs");
@@ -242,7 +243,9 @@ LogicalResult realizeDeferredCommunication(func::FuncOp funcOp,
ConstantPool constants(funcOp, rewriter);
DeferredEmissionContext context(rewriter, constants);
DeferredReplacementMap replacements;
if (failed(realizeDeferredBoundaries(boundaries->boundaries, boundaries->results, context, replacements)))
if (failed(realizeDeferredBoundaries(
boundaries->boundaries, boundaries->results, *transfers,
context, replacements)))
return failure();
for (auto [op, replacement] : replacements) {
if (op->getResult(0) == replacement)
@@ -322,8 +322,7 @@ static LogicalResult buildRequirementFamilies(DeferredTransferPlan& plan,
static LogicalResult buildAvailabilityFamilies(
DeferredTransferPlan &plan,
DeferredExchangePlan& exchange,
uint64_t& nextChannel,
DenseMap<int64_t, DenseMap<int64_t, unsigned>>& eventRegistersByTarget) {
uint64_t& nextChannel) {
enum class Availability { Local, Direct, Host };
for (RequirementFamily& requirement : exchange.requirements) {
for (LaneInterval interval : requirement.targetLanes.intervals()) {
@@ -357,18 +356,6 @@ static LogicalResult buildAvailabilityFamilies(
family.channelIds = StaticIntSequence::affine(nextChannel, 1, count);
family.hostRouted = runAvailability == Availability::Host;
if (family.hostRouted) {
SmallVector<int64_t> eventRegisters;
for (int64_t targetCore : targetCores) {
auto &registers = eventRegistersByTarget[targetCore];
auto it = registers.try_emplace(
requirement.producer->core, registers.size()).first;
if (it->second >= kPimEventRegisterCount)
return exchange.deferred.emitOpError(
"pipeline host transfer requires more event registers than the target core provides");
eventRegisters.push_back(it->second);
}
family.eventRegisters = StaticIntSequence::fromValues(
eventRegisters);
auto fragmentType = dyn_cast<ShapedType>(
requirement.publicationFragmentType);
auto fragmentBytes = fragmentType
@@ -431,7 +418,6 @@ static LogicalResult buildExchanges(func::FuncOp funcOp, DeferredTransferPlan& p
funcOp.walk([&](SpatDeferredCommunicationOp op) { deferredOps.push_back(op); });
GraphBatchPublicationCache publicationCache;
uint64_t nextChannel = 0;
DenseMap<int64_t, DenseMap<int64_t, unsigned>> eventRegistersByTarget;
for (SpatDeferredCommunicationOp deferred : deferredOps) {
Operation* targetOp = deferred->getParentOfType<SpatScheduledCompute>();
if (!targetOp)
@@ -451,8 +437,7 @@ static LogicalResult buildExchanges(func::FuncOp funcOp, DeferredTransferPlan& p
exchange->program = std::move(*program);
if (failed(buildRequirementFamilies(plan, *exchange, publicationCache)))
return failure();
if (failed(buildAvailabilityFamilies(
plan, *exchange, nextChannel, eventRegistersByTarget)))
if (failed(buildAvailabilityFamilies(plan, *exchange, nextChannel)))
return failure();
plan.exchanges.push_back(std::move(exchange));
}
@@ -13,6 +13,10 @@ struct DeferredTransferPlan {
llvm::DenseMap<int64_t, llvm::SmallVector<ProducedValue*>> producedByGraph;
llvm::SmallVector<std::unique_ptr<DeferredExchangePlan>> exchanges;
llvm::SmallVector<unsigned> stepCounts;
llvm::DenseMap<int64_t, unsigned> hostAcknowledgementCounts;
llvm::SmallVector<int64_t> stageZeroCores;
llvm::SmallVector<int64_t> downstreamCores;
size_t synchronizationRegisterCount = 0;
size_t pipelineHostBufferBytes = 0;
};
@@ -89,6 +89,8 @@ struct ScheduleAndRealizeSpatialPass final
return;
}
if (pipelineStages == 0 || target.processorCount % pipelineStages != 0
|| (pipelineStages > 1
&& target.synchronizationRegisterCount == 0)
|| target.residentWeightCapacity
> std::numeric_limits<size_t>::max() / pipelineStages) {
moduleOp.emitError("ScheduleAndRealizeSpatial requires valid pipeline stages and resource counts");
@@ -23,6 +23,7 @@ struct SchedulingTarget {
Cost transferWidthBytes = 8;
Cost vectorWidth = 16;
Cost vectorLatencyCycles = 4;
size_t synchronizationRegisterCount = 0;
Cost matrixRows = 128;
Cost matrixColumns = 128;
+32 -3
View File
@@ -592,13 +592,15 @@ def SpatHostStoreSyncOp : SpatOp<"host_store_sync", []> {
}
def SpatHostWaitLoadOp : SpatOp<"host_wait_load", []> {
let summary = "Wait for a producer and load its tensor from host memory";
let summary = "Wait for producers, load from host memory, and acknowledge consumption";
let arguments = (ins
Index:$sourceCoreId,
Index:$targetCoreId,
Index:$hostOffset,
Index:$eventRegister
Index:$eventRegister,
Index:$waitValue,
Index:$acknowledgementEventRegister
);
let results = (outs
@@ -607,7 +609,34 @@ def SpatHostWaitLoadOp : SpatOp<"host_wait_load", []> {
let assemblyFormat = [{
`from` $sourceCoreId `to` $targetCoreId
`host_offset` $hostOffset `event` $eventRegister attr-dict `:` type($output)
`host_offset` $hostOffset `event` $eventRegister `count` $waitValue
`ack` $acknowledgementEventRegister attr-dict `:` type($output)
}];
}
def SpatSyncOp : SpatOp<"sync", []> {
let summary = "Signal a synchronization register on another processor";
let arguments = (ins
Index:$targetCoreId,
Index:$eventRegister
);
let assemblyFormat = [{
$targetCoreId `event` $eventRegister attr-dict
}];
}
def SpatWaitOp : SpatOp<"wait", []> {
let summary = "Wait for a synchronization register value";
let arguments = (ins
Index:$eventRegister,
Index:$waitValue
);
let assemblyFormat = [{
$eventRegister `value` $waitValue attr-dict
}];
}