#include #include #include #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 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 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 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 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 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 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 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 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; }