Resnet is fast

This commit is contained in:
ilgeco
2026-07-20 11:34:59 +02:00
parent 5f42da36ae
commit 6bad9a8008
16 changed files with 883 additions and 138 deletions
@@ -29,6 +29,11 @@ static bool isUsedOnlyAsExplicitHostOperand(Value value) {
});
}
static bool isUsedOnlyByExtractSlices(Value value) {
return !value.use_empty()
&& llvm::all_of(value.getUsers(), [](Operation* user) { return isa<tensor::ExtractSliceOp>(user); });
}
static FailureOr<unsigned> getDirectReturnOperandIndex(OpResult result) {
if (!result.hasOneUse())
return failure();
@@ -357,6 +362,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
rewriter.createBlock(&coreBatchOp.getBody(), coreBatchOp.getBody().end(), TypeRange(blockArgTypes), blockArgLocs);
IRMapping mapper;
SmallPtrSet<Value, 4> hostResidentTensors;
rewriter.setInsertionPointToStart(newBlock);
auto oldLaneArg = computeBatchOp.getLaneArgument();
if (!oldLaneArg)
@@ -523,8 +529,10 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
if (isa_and_present<memref::GetGlobalOp>(toTensorOp.getBuffer().getDefiningOp())) {
Operation* cloned = rewriter.clone(op, mapper);
auto clonedTensor = cloned->getResult(0);
if (isUsedOnlyAsExplicitHostOperand(toTensorOp.getResult())) {
if (isUsedOnlyAsExplicitHostOperand(toTensorOp.getResult())
|| isUsedOnlyByExtractSlices(toTensorOp.getResult())) {
mapper.map(toTensorOp.getResult(), clonedTensor);
hostResidentTensors.insert(toTensorOp.getResult());
continue;
}
auto clonedType = cast<ShapedType>(clonedTensor.getType());
@@ -542,6 +550,28 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
}
}
if (auto extractSlice = dyn_cast<tensor::ExtractSliceOp>(op);
extractSlice && hostResidentTensors.contains(extractSlice.getSource())) {
Operation* cloned = rewriter.clone(op, mapper);
Value hostSlice = cloned->getResult(0);
auto outputBuffer = createEmptyTensorFromShaped(rewriter, loc, cast<ShapedType>(hostSlice.getType()));
Value zeroOffset = getOrCreateIndexConstant(rewriter, coreBatchOp.getOperation(), 0);
auto sizeAttr = getTensorSizeInBytesAttr(rewriter, coreBatchOp.getOperation(), hostSlice);
if (failed(sizeAttr))
return failure();
auto copied = pim::PimMemCopyHostToDevOp::create(rewriter,
loc,
outputBuffer.getType(),
zeroOffset,
zeroOffset,
outputBuffer,
hostSlice,
*sizeAttr)
.getOutput();
mapper.map(extractSlice.getResult(), copied);
continue;
}
for (auto [operandIndex, operand] : llvm::enumerate(op.getOperands())) {
if (!isa<TensorType>(operand.getType()) || mapper.contains(operand))
continue;