diff --git a/src/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/Scheduling/PipelineScheduling.cpp b/src/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/Scheduling/PipelineScheduling.cpp index ab035cf..1639001 100644 --- a/src/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/Scheduling/PipelineScheduling.cpp +++ b/src/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/Scheduling/PipelineScheduling.cpp @@ -841,7 +841,7 @@ static FailureOr assignPipelineStages( std::string &error) { if (graph.nodes.empty()) return PipelineStageAssignment { - {}, std::vector(layout.getStageCount(), 1)}; + {}, std::vector(layout.getStageSizes())}; std::vector tasksByOrder(graph.nodes.size()); std::iota(tasksByOrder.begin(), tasksByOrder.end(), 0); llvm::sort(tasksByOrder, [&](size_t lhs, size_t rhs) { diff --git a/test/PIM/SpatialSchedulingTargetTest.cpp b/test/PIM/SpatialSchedulingTargetTest.cpp index c6c937c..06f6166 100644 --- a/test/PIM/SpatialSchedulingTargetTest.cpp +++ b/test/PIM/SpatialSchedulingTargetTest.cpp @@ -130,6 +130,14 @@ int main() { 3, 3, 3, 0, }; std::string pipelineError; + ComputeGraph emptyGraph; + MergeScheduleResult emptySchedule; + emptySchedule.processorCount = 2; + assert(mlir::succeeded(applyPipelineScheduling( + emptyGraph, emptySchedule, 2, physical, pipelineError))); + assert(emptySchedule.processorCount == physical.processorCount); + assert(emptySchedule.processorStages == std::vector({0, 0, 1, 1})); + MergeScheduleResult pipelineSchedule = logicalSchedule; assert(mlir::succeeded(applyPipelineScheduling( graph, pipelineSchedule, 2, physical, pipelineError)));