Resnet is fast
This commit is contained in:
@@ -780,7 +780,7 @@ ComputeGraph buildComputeGraph(Operation* entryOp) {
|
||||
if (auto batch = dyn_cast<SpatComputeBatch>(&op)) {
|
||||
if (isUsedAsWeightOnly(batch.getOperation()))
|
||||
continue;
|
||||
size_t chunkCount = getBatchChunkTargetCount(batch.getLaneCount());
|
||||
size_t chunkCount = getBatchChunkTargetCount(batch);
|
||||
for (size_t chunkIndex = 0; chunkIndex < chunkCount; ++chunkIndex) {
|
||||
ComputeInstance instance = getBatchChunkForIndex(batch, chunkIndex);
|
||||
size_t index = graph.nodes.size();
|
||||
|
||||
+51
-13
@@ -1,9 +1,11 @@
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
|
||||
#include "ComputeGraph.hpp"
|
||||
#include "ComputeInstanceUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
|
||||
|
||||
@@ -12,20 +14,16 @@ using namespace mlir;
|
||||
namespace onnx_mlir {
|
||||
namespace spatial {
|
||||
|
||||
static constexpr llvm::StringLiteral kMergeChunkCountAttr = "spat.merge_chunk_count";
|
||||
|
||||
size_t getSchedulingCpuBudget() {
|
||||
if (coresCount.getValue() > 0)
|
||||
return static_cast<size_t>(coresCount.getValue());
|
||||
return std::numeric_limits<size_t>::max();
|
||||
}
|
||||
|
||||
size_t getBatchChunkTargetCount(int32_t laneCount) {
|
||||
static BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkCount, size_t chunkIndex) {
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
return std::min(static_cast<size_t>(laneCount), getSchedulingCpuBudget());
|
||||
}
|
||||
|
||||
BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex) {
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
size_t chunkCount = getBatchChunkTargetCount(laneCount);
|
||||
assert(chunkIndex < chunkCount && "chunkIndex out of range");
|
||||
|
||||
size_t laneCountSize = static_cast<size_t>(laneCount);
|
||||
@@ -38,11 +36,51 @@ BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex) {
|
||||
return {static_cast<uint32_t>(start), static_cast<uint32_t>(count)};
|
||||
}
|
||||
|
||||
size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane) {
|
||||
static bool batchChunksFit(SpatComputeBatch batch, size_t chunkCount, size_t crossbarCapacity) {
|
||||
for (size_t chunkIndex = 0; chunkIndex < chunkCount; ++chunkIndex) {
|
||||
BatchChunkRange chunk = getBatchChunkRange(batch.getLaneCount(), chunkCount, chunkIndex);
|
||||
ComputeInstance instance {batch.getOperation(), chunk.laneStart, chunk.laneCount};
|
||||
if (getComputeInstanceCrossbarUsage(instance).size() > crossbarCapacity)
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
size_t getBatchChunkTargetCount(SpatComputeBatch batch) {
|
||||
if (auto chunkCount = batch->getAttrOfType<IntegerAttr>(kMergeChunkCountAttr))
|
||||
return static_cast<size_t>(chunkCount.getInt());
|
||||
|
||||
int32_t laneCount = batch.getLaneCount();
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
size_t maxChunkCount = std::min(static_cast<size_t>(laneCount), getSchedulingCpuBudget());
|
||||
size_t crossbarCapacity = crossbarCountInCore.getValue();
|
||||
CrossbarUsage fullUsage = collectDistinctCrossbarWeights(batch.getOperation());
|
||||
if (fullUsage.empty() || crossbarCapacity == 0) {
|
||||
batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), maxChunkCount));
|
||||
return maxChunkCount;
|
||||
}
|
||||
|
||||
size_t chunkCount = std::max<size_t>(1, (fullUsage.size() + crossbarCapacity - 1) / crossbarCapacity);
|
||||
for (; chunkCount <= maxChunkCount; ++chunkCount) {
|
||||
if (batchChunksFit(batch, chunkCount, crossbarCapacity)) {
|
||||
batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), chunkCount));
|
||||
return chunkCount;
|
||||
}
|
||||
}
|
||||
batch->setAttr(kMergeChunkCountAttr, IntegerAttr::get(IndexType::get(batch.getContext()), maxChunkCount));
|
||||
return maxChunkCount;
|
||||
}
|
||||
|
||||
BatchChunkRange getBatchChunkRange(SpatComputeBatch batch, size_t chunkIndex) {
|
||||
return getBatchChunkRange(batch.getLaneCount(), getBatchChunkTargetCount(batch), chunkIndex);
|
||||
}
|
||||
|
||||
size_t getBatchChunkIndexForLane(SpatComputeBatch batch, uint32_t lane) {
|
||||
int32_t laneCount = batch.getLaneCount();
|
||||
assert(laneCount > 0 && "laneCount must be positive");
|
||||
assert(lane < static_cast<uint32_t>(laneCount) && "lane out of range");
|
||||
|
||||
size_t chunkCount = getBatchChunkTargetCount(laneCount);
|
||||
size_t chunkCount = getBatchChunkTargetCount(batch);
|
||||
size_t laneCountSize = static_cast<size_t>(laneCount);
|
||||
size_t baseChunkSize = laneCountSize / chunkCount;
|
||||
size_t remainder = laneCountSize % chunkCount;
|
||||
@@ -56,12 +94,12 @@ size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane) {
|
||||
}
|
||||
|
||||
ComputeInstance getBatchChunkForIndex(SpatComputeBatch batch, size_t chunkIndex) {
|
||||
BatchChunkRange chunk = getBatchChunkRange(batch.getLaneCount(), chunkIndex);
|
||||
BatchChunkRange chunk = getBatchChunkRange(batch, chunkIndex);
|
||||
return {batch.getOperation(), chunk.laneStart, chunk.laneCount};
|
||||
}
|
||||
|
||||
ComputeInstance getBatchChunkForLane(SpatComputeBatch batch, uint32_t lane) {
|
||||
return getBatchChunkForIndex(batch, getBatchChunkIndexForLane(batch.getLaneCount(), lane));
|
||||
return getBatchChunkForIndex(batch, getBatchChunkIndexForLane(batch, lane));
|
||||
}
|
||||
|
||||
llvm::SmallVector<ComputeInstance, 4>
|
||||
@@ -74,8 +112,8 @@ getBatchChunksForRange(SpatComputeBatch batch, uint32_t laneStart, uint32_t lane
|
||||
assert(laneEnd >= laneStart && "lane range overflow");
|
||||
assert(laneEnd <= static_cast<uint32_t>(batch.getLaneCount()) && "lane range out of bounds");
|
||||
|
||||
size_t firstChunk = getBatchChunkIndexForLane(batch.getLaneCount(), laneStart);
|
||||
size_t lastChunk = getBatchChunkIndexForLane(batch.getLaneCount(), laneEnd - 1);
|
||||
size_t firstChunk = getBatchChunkIndexForLane(batch, laneStart);
|
||||
size_t lastChunk = getBatchChunkIndexForLane(batch, laneEnd - 1);
|
||||
chunks.reserve(lastChunk - firstChunk + 1);
|
||||
for (size_t chunkIndex = firstChunk; chunkIndex <= lastChunk; ++chunkIndex)
|
||||
chunks.push_back(getBatchChunkForIndex(batch, chunkIndex));
|
||||
|
||||
+3
-3
@@ -27,9 +27,9 @@ struct BatchChunkRange {
|
||||
};
|
||||
|
||||
size_t getSchedulingCpuBudget();
|
||||
size_t getBatchChunkTargetCount(int32_t laneCount);
|
||||
BatchChunkRange getBatchChunkRange(int32_t laneCount, size_t chunkIndex);
|
||||
size_t getBatchChunkIndexForLane(int32_t laneCount, uint32_t lane);
|
||||
size_t getBatchChunkTargetCount(SpatComputeBatch batch);
|
||||
BatchChunkRange getBatchChunkRange(SpatComputeBatch batch, size_t chunkIndex);
|
||||
size_t getBatchChunkIndexForLane(SpatComputeBatch batch, uint32_t lane);
|
||||
ComputeInstance getBatchChunkForIndex(SpatComputeBatch batch, size_t chunkIndex);
|
||||
ComputeInstance getBatchChunkForLane(SpatComputeBatch batch, uint32_t lane);
|
||||
llvm::SmallVector<ComputeInstance, 4>
|
||||
|
||||
@@ -6,7 +6,9 @@
|
||||
|
||||
#include <cmath>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
#include <queue>
|
||||
#include <tuple>
|
||||
#include <vector>
|
||||
|
||||
#include "PeftScheduler.hpp"
|
||||
@@ -133,6 +135,55 @@ void verifyOctTableSize(size_t nodeCount, size_t processorCount) {
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<CrossbarUsage> planCrossbarResidency(const ComputeGraph& graph,
|
||||
size_t processorCount,
|
||||
size_t crossbarCapacity,
|
||||
const MeshModel& mesh) {
|
||||
std::vector<size_t> weightedTasks;
|
||||
for (size_t task = 0; task < graph.nodes.size(); ++task)
|
||||
if (!graph.nodes[task].crossbarUsage.empty())
|
||||
weightedTasks.push_back(task);
|
||||
llvm::sort(weightedTasks, [&](size_t lhs, size_t rhs) {
|
||||
if (graph.nodes[lhs].crossbarUsage.size() != graph.nodes[rhs].crossbarUsage.size())
|
||||
return graph.nodes[lhs].crossbarUsage.size() > graph.nodes[rhs].crossbarUsage.size();
|
||||
return graph.nodes[lhs].originalOrder < graph.nodes[rhs].originalOrder;
|
||||
});
|
||||
|
||||
std::vector<CrossbarUsage> residency(processorCount);
|
||||
for (size_t task : weightedTasks) {
|
||||
size_t bestProcessor = std::numeric_limits<size_t>::max();
|
||||
using ResidencyScore = std::tuple<size_t, size_t, size_t, size_t>;
|
||||
std::optional<ResidencyScore> bestScore;
|
||||
for (size_t processor = 0; processor < processorCount; ++processor) {
|
||||
size_t crossbarUnion = getCrossbarUnionSize(residency[processor], graph.nodes[task].crossbarUsage);
|
||||
if (crossbarUnion > crossbarCapacity)
|
||||
continue;
|
||||
size_t addedCrossbars = crossbarUnion - residency[processor].size();
|
||||
ResidencyScore score {addedCrossbars,
|
||||
crossbarCapacity - crossbarUnion,
|
||||
mesh.getCenterDistance(processor),
|
||||
processor};
|
||||
if (!bestScore || score < *bestScore) {
|
||||
bestProcessor = processor;
|
||||
bestScore = score;
|
||||
}
|
||||
}
|
||||
if (bestProcessor == std::numeric_limits<size_t>::max()) {
|
||||
std::string message =
|
||||
llvm::formatv("PEFT residency planner: cannot place task {0} with {1} distinct weights in {2} "
|
||||
"processors of capacity {3}",
|
||||
graph.nodes[task].originalOrder,
|
||||
graph.nodes[task].crossbarUsage.size(),
|
||||
processorCount,
|
||||
crossbarCapacity)
|
||||
.str();
|
||||
llvm::report_fatal_error(llvm::StringRef(message));
|
||||
}
|
||||
insertCrossbarWeights(residency[bestProcessor], graph.nodes[task].crossbarUsage);
|
||||
}
|
||||
return residency;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
Time getPeftTransferTime(Time transferCost, size_t sourceProcessor, size_t targetProcessor, size_t processorCount) {
|
||||
@@ -145,6 +196,8 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
|
||||
if (processorCount == 0)
|
||||
llvm::report_fatal_error("PEFT scheduler: processor count must be positive");
|
||||
MeshModel mesh = MeshModel::infer(processorCount);
|
||||
std::vector<CrossbarUsage> plannedResidency =
|
||||
planCrossbarResidency(graph, processorCount, options.crossbarCapacity, mesh);
|
||||
|
||||
verifyOctTableSize(nodeCount, processorCount);
|
||||
std::vector<std::vector<size_t>> reverseLevels = buildReverseLevels(graph);
|
||||
@@ -237,20 +290,23 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
|
||||
size_t bestProcessor = std::numeric_limits<size_t>::max();
|
||||
Time bestEst = 0;
|
||||
Time bestEft = 0;
|
||||
Time bestOeft = std::numeric_limits<Time>::max();
|
||||
unsigned int bestOverlapCount = 0;
|
||||
size_t bestCenterDistance = std::numeric_limits<size_t>::max();
|
||||
using CandidateScore = std::tuple<size_t, size_t, Time, Time, Time, size_t, unsigned int>;
|
||||
std::optional<CandidateScore> bestScore;
|
||||
size_t smallestCrossbarUnion = std::numeric_limits<size_t>::max();
|
||||
bool crossbarRejected = false;
|
||||
|
||||
for (size_t processor = 0; processor < processorCount; ++processor) {
|
||||
unsigned int overlapCount = countCrossbarOverlap(processorCrossbars[processor], graph.nodes[task].crossbarUsage);
|
||||
if (!graph.nodes[task].crossbarUsage.empty()
|
||||
&& getCrossbarUnionSize(processorCrossbars[processor], graph.nodes[task].crossbarUsage)
|
||||
> options.crossbarCapacity) {
|
||||
&& countCrossbarOverlap(plannedResidency[processor], graph.nodes[task].crossbarUsage)
|
||||
!= graph.nodes[task].crossbarUsage.size())
|
||||
continue;
|
||||
unsigned int overlapCount = countCrossbarOverlap(processorCrossbars[processor], graph.nodes[task].crossbarUsage);
|
||||
size_t crossbarUnion = getCrossbarUnionSize(processorCrossbars[processor], graph.nodes[task].crossbarUsage);
|
||||
smallestCrossbarUnion = std::min(smallestCrossbarUnion, crossbarUnion);
|
||||
if (!graph.nodes[task].crossbarUsage.empty() && crossbarUnion > options.crossbarCapacity) {
|
||||
crossbarRejected = true;
|
||||
continue;
|
||||
}
|
||||
|
||||
Time dataReady = 0;
|
||||
for (const auto& [pred, comm] : graph.predecessors[task]) {
|
||||
const ScheduledTask& predSchedule = schedules[pred];
|
||||
@@ -282,41 +338,32 @@ MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftSchedu
|
||||
Time eft = addOrMax(est, computeCost);
|
||||
Time oeft = addOrMax(eft, oct[task * processorCount + processor]);
|
||||
size_t centerDistance = mesh.getCenterDistance(processor);
|
||||
|
||||
if (oeft < bestOeft || (oeft == bestOeft && eft < bestEft)
|
||||
|| (oeft == bestOeft && eft == bestEft && est < bestEst)) {
|
||||
CandidateScore score {0,
|
||||
0,
|
||||
oeft,
|
||||
eft,
|
||||
est,
|
||||
centerDistance,
|
||||
overlapCount};
|
||||
if (!bestScore || score < *bestScore) {
|
||||
bestProcessor = processor;
|
||||
bestEst = est;
|
||||
bestEft = eft;
|
||||
bestOeft = oeft;
|
||||
bestOverlapCount = overlapCount;
|
||||
bestCenterDistance = centerDistance;
|
||||
}
|
||||
else if (oeft == bestOeft && eft == bestEft && est == bestEst
|
||||
&& centerDistance < bestCenterDistance) {
|
||||
bestProcessor = processor;
|
||||
bestEst = est;
|
||||
bestEft = eft;
|
||||
bestOeft = oeft;
|
||||
bestOverlapCount = overlapCount;
|
||||
bestCenterDistance = centerDistance;
|
||||
}
|
||||
else if (oeft == bestOeft && eft == bestEft && est == bestEst
|
||||
&& centerDistance == bestCenterDistance && overlapCount < bestOverlapCount) {
|
||||
bestProcessor = processor;
|
||||
bestEst = est;
|
||||
bestEft = eft;
|
||||
bestOeft = oeft;
|
||||
bestOverlapCount = overlapCount;
|
||||
bestCenterDistance = centerDistance;
|
||||
bestScore = score;
|
||||
}
|
||||
}
|
||||
|
||||
if (bestProcessor == std::numeric_limits<size_t>::max()) {
|
||||
if (crossbarRejected) {
|
||||
const ComputeInstance& instance = graph.nodes[task].instance;
|
||||
std::string message =
|
||||
llvm::formatv("PEFT scheduler: no valid processor for task {0}; crossbar capacity {1} is exhausted",
|
||||
llvm::formatv("PEFT scheduler: no valid processor for task {0} (lanes {1}..{2}, {3} distinct weights); "
|
||||
"smallest processor union is {4}, exceeding crossbar capacity {5}",
|
||||
graph.nodes[task].originalOrder,
|
||||
instance.laneStart,
|
||||
instance.laneStart + instance.laneCount,
|
||||
graph.nodes[task].crossbarUsage.size(),
|
||||
smallestCrossbarUnion,
|
||||
options.crossbarCapacity)
|
||||
.str();
|
||||
llvm::report_fatal_error(llvm::StringRef(message));
|
||||
|
||||
Reference in New Issue
Block a user