f47fcebb83
Validate Operations / validate-operations (push) Has been cancelled
better reports cleanups
121 lines
3.9 KiB
C++
121 lines
3.9 KiB
C++
#include <cassert>
|
|
#include <cstdlib>
|
|
#include <iostream>
|
|
|
|
#include "src/Accelerators/PIM/Dialect/Pim/Transforms/LocalMemoryPlanning/LocalMemoryPlanning.hpp"
|
|
|
|
using onnx_mlir::LocalMemoryPlacement;
|
|
using onnx_mlir::pim::LocalMemoryInterval;
|
|
using onnx_mlir::planLocalMemoryPlacements;
|
|
|
|
namespace {
|
|
|
|
LocalMemoryInterval makeInterval(size_t size, uint64_t start, uint64_t end) {
|
|
LocalMemoryInterval interval;
|
|
interval.size = size;
|
|
interval.start = start;
|
|
interval.end = end;
|
|
return interval;
|
|
}
|
|
|
|
size_t getPeak(llvm::ArrayRef<LocalMemoryPlacement> placements) {
|
|
size_t peak = 0;
|
|
for (const auto& placement : placements)
|
|
peak = std::max(peak, placement.address + placement.size);
|
|
return peak;
|
|
}
|
|
|
|
void assertSinglePlacementCase(LocalMemoryInterval a, LocalMemoryInterval b, size_t expectedSize) {
|
|
llvm::SmallVector<LocalMemoryInterval, 4> intervals = {a, b};
|
|
auto placements = planLocalMemoryPlacements(intervals, 1024);
|
|
assert(mlir::succeeded(placements));
|
|
assert(placements->size() == 1);
|
|
assert(placements->front().size == expectedSize);
|
|
assert(placements->front().intervalIndices.size() == 2);
|
|
}
|
|
|
|
int testSameSizeNonOverlap() {
|
|
std::cout << "testSameSizeNonOverlap:" << std::endl;
|
|
assertSinglePlacementCase(makeInterval(64, 0, 10), makeInterval(64, 11, 20), 64);
|
|
return 0;
|
|
}
|
|
|
|
int testLargerFirst() {
|
|
std::cout << "testLargerFirst:" << std::endl;
|
|
llvm::SmallVector<LocalMemoryInterval, 4> intervals = {
|
|
makeInterval(100, 0, 10), makeInterval(40, 11, 20)};
|
|
auto placements = planLocalMemoryPlacements(intervals, 1024);
|
|
assert(mlir::succeeded(placements) && placements->size() == 2);
|
|
assert(getPeak(*placements) == 100);
|
|
return 0;
|
|
}
|
|
|
|
int testSmallerFirst() {
|
|
std::cout << "testSmallerFirst:" << std::endl;
|
|
llvm::SmallVector<LocalMemoryInterval, 4> intervals = {
|
|
makeInterval(40, 0, 10), makeInterval(100, 11, 20)};
|
|
auto placements = planLocalMemoryPlacements(intervals, 1024);
|
|
assert(mlir::succeeded(placements) && placements->size() == 2);
|
|
assert(getPeak(*placements) == 100);
|
|
return 0;
|
|
}
|
|
|
|
int testOverlapNeedsTwoSlots() {
|
|
std::cout << "testOverlapNeedsTwoSlots:" << std::endl;
|
|
llvm::SmallVector<LocalMemoryInterval, 4> intervals = {
|
|
makeInterval(100, 0, 20), makeInterval(40, 10, 30)};
|
|
auto placements = planLocalMemoryPlacements(intervals, 1024);
|
|
assert(mlir::succeeded(placements) && placements->size() == 2);
|
|
assert((*placements)[0].address != (*placements)[1].address);
|
|
return 0;
|
|
}
|
|
|
|
int testReuseChain() {
|
|
std::cout << "testReuseChain:" << std::endl;
|
|
llvm::SmallVector<LocalMemoryInterval, 4> intervals = {
|
|
makeInterval(40, 0, 10), makeInterval(100, 11, 20), makeInterval(20, 21, 30)};
|
|
auto placements = planLocalMemoryPlacements(intervals, 1024);
|
|
assert(mlir::succeeded(placements) && placements->size() == 3);
|
|
assert(getPeak(*placements) == 100);
|
|
return 0;
|
|
}
|
|
|
|
int testPartialAddressReuse() {
|
|
std::cout << "testPartialAddressReuse:" << std::endl;
|
|
llvm::SmallVector<LocalMemoryInterval, 4> intervals = {
|
|
makeInterval(100, 0, 10), makeInterval(60, 11, 20), makeInterval(40, 11, 20)};
|
|
auto placements = planLocalMemoryPlacements(intervals, 1024);
|
|
assert(mlir::succeeded(placements));
|
|
assert(getPeak(*placements) == 100);
|
|
return 0;
|
|
}
|
|
|
|
int testAddressLimit() {
|
|
std::cout << "testAddressLimit:" << std::endl;
|
|
llvm::SmallVector<LocalMemoryInterval, 4> intervals = {
|
|
makeInterval(100, 0, 20), makeInterval(40, 10, 30)};
|
|
assert(mlir::failed(planLocalMemoryPlacements(intervals, 128)));
|
|
return 0;
|
|
}
|
|
|
|
} // namespace
|
|
|
|
int main(int argc, char *argv[]) {
|
|
(void) argc;
|
|
(void) argv;
|
|
|
|
int failures = 0;
|
|
failures += testSameSizeNonOverlap();
|
|
failures += testLargerFirst();
|
|
failures += testSmallerFirst();
|
|
failures += testOverlapNeedsTwoSlots();
|
|
failures += testReuseChain();
|
|
failures += testPartialAddressReuse();
|
|
failures += testAddressLimit();
|
|
if (failures != 0) {
|
|
std::cerr << failures << " test failures\n";
|
|
return EXIT_FAILURE;
|
|
}
|
|
return EXIT_SUCCESS;
|
|
}
|