add pipeline stages synchronization
Validate Operations / validate-operations (push) Has been cancelled
Validate Operations / validate-operations (push) Has been cancelled
full ops throughput validation now passes
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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 ©Shape, 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);
|
||||
|
||||
@@ -136,7 +136,7 @@ def PimWaitOp : PimOp<"wait", []> {
|
||||
|
||||
let arguments = (ins
|
||||
Index:$eventRegister,
|
||||
I32Attr:$waitValue
|
||||
Index:$waitValue
|
||||
);
|
||||
|
||||
let assemblyFormat = [{
|
||||
|
||||
+139
-1
@@ -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;
|
||||
}
|
||||
|
||||
+2
-1
@@ -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
|
||||
|
||||
+233
-5
@@ -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
|
||||
|
||||
+1
@@ -42,6 +42,7 @@ using DeferredReplacementMap =
|
||||
|
||||
mlir::LogicalResult realizeDeferredBoundaries(mlir::ArrayRef<BoundaryProgram> boundaries,
|
||||
mlir::ArrayRef<DeferredResultPlan> results,
|
||||
DeferredTransferPlan& transfers,
|
||||
DeferredEmissionContext& context,
|
||||
DeferredReplacementMap& replacements);
|
||||
|
||||
|
||||
+3
@@ -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;
|
||||
};
|
||||
|
||||
|
||||
+5
-2
@@ -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)
|
||||
|
||||
+2
-17
@@ -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 ®isters = 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));
|
||||
}
|
||||
|
||||
+4
@@ -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");
|
||||
|
||||
+1
@@ -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;
|
||||
|
||||
@@ -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
|
||||
}];
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user