Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 15 additions & 2 deletions docs/designs/ptoas-multi-buffer-explicit-design.md
Original file line number Diff line number Diff line change
Expand Up @@ -227,7 +227,7 @@ struct SlotInfo {
| Const(a) ↔ Dyn(%k) | – | ⚠️ 保守 alias,全同步 | dyn event id(仅 %k==a 时同步) |
| Dyn(%j) ↔ Dyn(%k), 表达式相同 | – | ⚠️ 保守 alias | 真冲突 + dyn event id |
| Dyn(%j) ↔ Dyn(%k), 可证 disjoint | – | ⚠️ 保守 alias | 同 iter 不冲突,跨 iter dyn event id |
| Dyn ↔ Dyn 不可证 | – | ⚠️ 保守 alias | N 个 dyn event id |
| Dyn ↔ Dyn 不可证 | – | ⚠️ 保守 alias | 保留依赖,使用静态 event id |

dyn event id 分配 + `set_flag_dyn` / `wait_flag_dyn` 生成留作 follow-up(需要扩展 `SyncEventIdAllocation` 和 `SyncCodegen`)。

Expand Down Expand Up @@ -355,6 +355,19 @@ ptoas 自动行为:
- 同步分析:producer slot 表达式 `(iv+1)%2`,consumer slot `iv%2` → 同 iter disjoint / 跨 iter 冲突 → 2 个 dyn event id;
- emit `set_flag_dyn` / `wait_flag_dyn`,event id value 与 slot 表达式同源。

循环边界必须按槽位配平事件,不能对所有槽位统一 prime / drain:

- V→MTE2:循环前只初始化 slot 1 的事件。slot 0 的第一次循环内写入必须等待第 0 次计算完成。
- MTE2→V:循环前只初始化 slot 0 的循环事件;预加载本身另有完整的 MTE2→V 同步。slot 1 必须等待实际预取完成。
- 执行 `n` 次后,V→MTE2 只 drain slot `(n + 1) % 2`,MTE2→V 只 drain slot `n % 2`。`n = 0` 时也遵循同一规则。

对步长为 1、非负常量下界、槽位形如 `(iv + 非负常量) remui N` 的循环,
按每个槽位第一次 set / wait 的先后关系生成边界同步,并在循环结束时按最终轮转位置计算剩余事件。
生产者和消费者必须直接位于同一个循环体内。嵌套循环在所属循环的每次入口和出口完成配平,
外层循环不再为这对操作添加重复的跨迭代事件。若两个槽位表达式不同且无法证明这样的轮转,
则保留同次迭代的依赖,并使用单个静态事件进行保守同步;不假定未知表达式会遍历所有槽位。
已证明轮转的事件组在资源不足时使用 `PIPE_ALL` 保守同步,不将多个槽位折叠成单个事件。

### 7.3 例 3:N=4 同表达式轮转

```mlir
Expand Down Expand Up @@ -476,7 +489,7 @@ lit test/lit/pto/multi_tile_prefetch_insert_sync.pto

### 当前限制

- **affine 分析仅覆盖核心几种形态**:`compareSlotSSA` 当前能证 `iv % N` / `(iv ± c) % N` / 同 SSA / 纯常量;不能证 `(iv * c) % N`、跨函数 / 跨循环的 SSA 等价、非 `arith.remui` 包装的 slot 表达式。命中不到时退回 kUnknown / 保守 N dyn event id。
- **affine 分析仅覆盖核心几种形态**:`compareSlotSSA` 当前能证 `iv % N` / `(iv ± c) % N` / 同 SSA / 纯常量;不能证 `(iv * c) % N`、跨函数 / 跨循环的 SSA 等价、非 `arith.remui` 包装的 slot 表达式。不同表达式的轮转无法证明时,使用静态事件保守同步(见 §7.2)。
- **PlanMemory N>2 不复用 Stage1**:N>2 的兄弟 slot 不走 SPEC_LEVEL_1 "ping/pong 相邻摆放"优化,用更多内存。N=2 路径不变。
- 初版仅支持 `loc=vec` / `loc=mat` local memory。
- function argument / return 上的 `multi_tile_buf` 不支持(多 buffer 所有权限定在 ptoas 内)。
Expand Down
14 changes: 13 additions & 1 deletion include/PTO/Transforms/InsertSync/SyncCommon.h
Original file line number Diff line number Diff line change
Expand Up @@ -166,8 +166,17 @@ struct BaseMemInfo {

using DepBaseMemInfoPairVec =
SmallVector<std::pair<const BaseMemInfo *, const BaseMemInfo *>>;

// 表示一个具体的同步指令 (Set, Wait, Barrier)
// A unit-step loop visits slot `(iv + offset) % count` once per rotation.
// Boundary flags compare each lane's first set and wait in that rotation.
struct SlotEventSchedule {
Operation* loop{nullptr};
uint32_t producerOffset{0};
uint32_t consumerOffset{0};
bool producerBeforeConsumer{false};
};

class SyncOperation {
public:
enum class TYPE {
Expand Down Expand Up @@ -198,6 +207,9 @@ class SyncOperation {
// hardware event-id index. Empty when this sync is single-buffer.
Value slotSSAExpr;
uint32_t slotCount{1};
std::optional<SlotEventSchedule> slotSchedule;
// Present only on the synthetic loop prime/drain for this event lane.
std::optional<uint32_t> boundarySlot;
Value lowestCommonAncestorBuffer{nullptr};
int reuseCntForWiden{0};
bool reallocatedLoopHeadTailSync{false};
Expand Down
8 changes: 8 additions & 0 deletions include/PTO/Transforms/SlotAffineAnalysis.h
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@
// provably disjoint modulo N, or indeterminate. The result lets sync
// shrink event-id count or skip same-iter forward syncs entirely when
// producer and consumer touch different slots in every iteration.
// Rotation offsets also determine which slot events are initially available
// and which releases remain outstanding at the loop boundary.
//
//===----------------------------------------------------------------------===//

Expand All @@ -22,6 +24,7 @@

#include "mlir/IR/Value.h"
#include <cstdint>
#include <optional>

namespace mlir {
namespace pto {
Expand Down Expand Up @@ -53,6 +56,11 @@ mlir::Value findMultiTileSlotExpr(mlir::Value v);
/// compareSlotSSA(arith.constant 0, arith.constant 1) -> kDisjoint
SlotRelation compareSlotSSA(mlir::Value a, mlir::Value b, uint32_t N);

/// Recognize `(iv + nonnegative constant) remui N` for boundary event
/// accounting. The returned offset is reduced modulo N. Other expressions
/// require conservative synchronization rather than assumed slot rotation.
std::optional<uint32_t> getSlotRotationOffset(mlir::Value slot, mlir::Value inductionVar, uint32_t count);

} // namespace pto
} // namespace mlir

Expand Down
212 changes: 135 additions & 77 deletions lib/PTO/Transforms/InsertSync/InsertSyncAnalysis.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -477,30 +477,89 @@ void InsertSyncAnalysis::InsertSync(
MemAnalyze(nowCompound, frontCompound, syncRecordList, forEndIndex);
}

static std::optional<std::pair<Value, Value>> getDependencySlots(const DepBaseMemInfoPairVec& dependencies)
{
Value producerSlot;
Value consumerSlot;
for (const auto& pair : dependencies) {
if (!pair.first || !pair.second) {
return std::nullopt;
}
Value producer = findMultiTileSlotExpr(pair.second->baseBuffer);
Value consumer = findMultiTileSlotExpr(pair.first->baseBuffer);
if (!producer || !consumer) {
return std::nullopt;
}
bool mismatchedProducer = producerSlot && producerSlot != producer;
bool mismatchedConsumer = consumerSlot && consumerSlot != consumer;
if (mismatchedProducer || mismatchedConsumer) {
return std::nullopt;
}
producerSlot = producer;
consumerSlot = consumer;
}
if (!producerSlot || !consumerSlot) {
return std::nullopt;
}
return std::make_pair(producerSlot, consumerSlot);
}

static std::optional<SlotEventSchedule> getSlotEventSchedule(
Value producerSlot, Value consumerSlot, Operation* producer, Operation* consumer, uint32_t count)
{
Block* producerBlock = producer->getBlock();
Block* consumerBlock = consumer->getBlock();
if (producerBlock != consumerBlock) {
return std::nullopt;
}
auto loop = dyn_cast<scf::ForOp>(producer->getParentOp());
if (!loop) {
return std::nullopt;
}
IntegerAttr lowerBound;
bool unitStep = matchPattern(loop.getStep(), m_One());
bool constantLowerBound = matchPattern(loop.getLowerBound(), m_Constant(&lowerBound));
if (!unitStep || !constantLowerBound) {
return std::nullopt;
}
if (lowerBound.getValue().isNegative()) {
return std::nullopt;
}
auto producerOffset = getSlotRotationOffset(producerSlot, loop.getInductionVar(), count);
auto consumerOffset = getSlotRotationOffset(consumerSlot, loop.getInductionVar(), count);
if (!producerOffset || !consumerOffset) {
return std::nullopt;
}
return SlotEventSchedule{loop, *producerOffset, *consumerOffset, producer->isBeforeInBlock(consumer)};
}

// Returns true if a *same-iter* multi-buffer dep pair can be dropped
// because the producer's and consumer's slot SSA expressions are provably
// disjoint modulo N. Only applied to forward (non-back-edge) deps -- the
// back-edge path still needs to sync per-slot via dyn event id (the
// prefetch idiom). When the analysis is inconclusive (kUnknown / kEqual)
// the dep is kept and the existing conservative path runs.
static bool isForwardDepDroppableBySlotAffine(const BaseMemInfo *a,
const BaseMemInfo *b) {
if (!a || !b) {
return false;
}
size_t aN = a->baseAddresses.size();
size_t bN = b->baseAddresses.size();
size_t n = std::max(aN, bN);
if (n < kPtoMultiBufferMinNum) {
return false;
}
Value slotA = findMultiTileSlotExpr(a->baseBuffer);
Value slotB = findMultiTileSlotExpr(b->baseBuffer);
if (!slotA || !slotB) {
return false;
}
return compareSlotSSA(slotA, slotB, static_cast<uint32_t>(n)) ==
SlotRelation::kDisjoint;
static bool isForwardDepDroppableBySlotAffine(
const BaseMemInfo* a, const BaseMemInfo* b, Operation* consumer, Operation* producer)
{
if (!a || !b) {
return false;
}
size_t aN = a->baseAddresses.size();
size_t bN = b->baseAddresses.size();
size_t n = std::max(aN, bN);
if (n < kPtoMultiBufferMinNum) {
return false;
}
Value slotA = findMultiTileSlotExpr(a->baseBuffer);
Value slotB = findMultiTileSlotExpr(b->baseBuffer);
if (!slotA || !slotB) {
return false;
}
// Removing the forward edge requires a balanced slot rotation on the back
// edge. Unknown schedules retain static program-order synchronization.
return getSlotEventSchedule(slotB, slotA, producer, consumer, static_cast<uint32_t>(n)).has_value() &&
compareSlotSSA(slotA, slotB, static_cast<uint32_t>(n)) == SlotRelation::kDisjoint;
}

void InsertSyncAnalysis::MemAnalyze(
Expand All @@ -518,18 +577,31 @@ void InsertSyncAnalysis::MemAnalyze(

// Same-iter (forward) deps: drop pairs that the affine analysis proves
// touch disjoint slots in every iteration of the multi-buffer loop.
// Back-edge deps stay untouched -- they still need per-slot syncing
// through the dyn-event-id pipeline.
if (!forEndIndex.has_value()) {
auto isDroppable = [](const std::pair<const BaseMemInfo *,
const BaseMemInfo *> &pair) {
return isForwardDepDroppableBySlotAffine(pair.first, pair.second);
};
depVec.erase(std::remove_if(depVec.begin(), depVec.end(), isDroppable),
depVec.end());
if (depVec.empty()) {
return;
}
// The owning loop still needs per-slot back-edge synchronization; its
// boundaries also close the dependency across enclosing iterations.
auto dependencySlots = getDependencySlots(depVec);
int slotCount = GetEventIdNum(depVec);
bool uniformSlots = slotCount > 1 && dependencySlots.has_value();
if (forEndIndex && uniformSlots) {
auto schedule = getSlotEventSchedule(
dependencySlots->first, dependencySlots->second, frontCompound->elementOp, nowCompound->elementOp, slotCount);
Operation* scope = syncIR_[*forEndIndex]->elementOp;
if (schedule && scope->isProperAncestor(schedule->loop)) {
// The owning loop drains its slot events on every exit. An enclosing
// loop must not add another carried token to these same operations.
return;
}
}
bool forwardDependency = !forEndIndex.has_value();
if (forwardDependency && uniformSlots) {
auto isDroppable = [nowCompound, frontCompound](const std::pair<const BaseMemInfo*, const BaseMemInfo*>& pair) {
return isForwardDepDroppableBySlotAffine(
pair.first, pair.second, nowCompound->elementOp, frontCompound->elementOp);
};
depVec.erase(std::remove_if(depVec.begin(), depVec.end(), isDroppable), depVec.end());
if (depVec.empty()) {
return;
}
}

if (CanPrunePipeVBarrier(nowCompound, frontCompound, depVec, forEndIndex)) {
Expand Down Expand Up @@ -673,50 +745,36 @@ void InsertSyncAnalysis::InsertPipeBarrierSync(

// Resolve one unambiguous producer/consumer slot SSA pair for the whole
// dependency group and configure a dynamic set/wait pair with it. Returns the
// effective event-id count: when the group decomposes into a single slot
// expression per side, slotSSAExpr/slotCount are filled and `eventIdNum` is
// kept; otherwise a single static event id is used for the whole group.
// effective event-id count. Keep per-slot events for a balanced rotation or
// equal slot expressions; missing, ambiguous, or unproven distinct expressions
// use one static event for the whole group.
static int configureDynEventSlots(
SyncOperation *setOp, SyncOperation *waitOp,
const DepBaseMemInfoPairVec &depBaseMemInfosVec, int eventIdNum) {
if (eventIdNum <= 1) {
SyncOperation* setOp, SyncOperation* waitOp, const DepBaseMemInfoPairVec& depBaseMemInfosVec, int eventIdNum,
Operation* producer, Operation* consumer, Operation* loop)
{
if (eventIdNum <= 1) {
return eventIdNum;
}
auto slots = getDependencySlots(depBaseMemInfosVec);
if (!slots) {
return 1;
}
auto [producerSlot, consumerSlot] = *slots;
auto schedule = getSlotEventSchedule(producerSlot, consumerSlot, producer, consumer, eventIdNum);
if (schedule && schedule->loop != loop) {
schedule.reset();
}
bool sameSlot = compareSlotSSA(producerSlot, consumerSlot, eventIdNum) == SlotRelation::kEqual;
if (!schedule && !sameSlot) {
return 1;
}
setOp->slotSchedule = schedule;
waitOp->slotSchedule = schedule;
setOp->slotSSAExpr = producerSlot;
setOp->slotCount = static_cast<uint32_t>(eventIdNum);
waitOp->slotSSAExpr = consumerSlot;
waitOp->slotCount = static_cast<uint32_t>(eventIdNum);
return eventIdNum;
}
Value producerSlot;
Value consumerSlot;
bool hasAmbiguousSlot = false;
for (auto &pair : depBaseMemInfosVec) {
Value pairProducerSlot;
Value pairConsumerSlot;
if (pair.second && pair.second->baseBuffer) {
pairProducerSlot = findMultiTileSlotExpr(pair.second->baseBuffer);
}
if (pair.first && pair.first->baseBuffer) {
pairConsumerSlot = findMultiTileSlotExpr(pair.first->baseBuffer);
}
if (!pairProducerSlot || !pairConsumerSlot) {
hasAmbiguousSlot = true;
break;
}
if ((producerSlot && producerSlot != pairProducerSlot) ||
(consumerSlot && consumerSlot != pairConsumerSlot)) {
hasAmbiguousSlot = true;
break;
}
producerSlot = pairProducerSlot;
consumerSlot = pairConsumerSlot;
}
if (hasAmbiguousSlot || !producerSlot || !consumerSlot) {
// Missing or ambiguous slot SSA -- fall back to a single event id. This
// also keeps non-multi-buffer codepaths untouched if their baseAddresses
// have multiple entries for another reason.
return 1;
}
setOp->slotSSAExpr = producerSlot;
setOp->slotCount = static_cast<uint32_t>(eventIdNum);
waitOp->slotSSAExpr = consumerSlot;
waitOp->slotCount = static_cast<uint32_t>(eventIdNum);
return eventIdNum;
}

void InsertSyncAnalysis::InsertCrossPipeEventSync(
Expand All @@ -740,11 +798,11 @@ void InsertSyncAnalysis::InsertCrossPipeEventSync(
// dyn event IDs are warranted, also plumb the per-side slot SSA so
// codegen can lower into `pto.set_flag_dyn` / `pto.wait_flag_dyn`.
if (forEndIndex.has_value()) {
int eventIdNum = configureDynEventSlots(
setOp.get(), waitOp.get(), depBaseMemInfosVec,
GetEventIdNum(depBaseMemInfosVec));
setOp->eventIdNum = eventIdNum;
waitOp->eventIdNum = eventIdNum;
int eventIdNum = configureDynEventSlots(
setOp.get(), waitOp.get(), depBaseMemInfosVec, GetEventIdNum(depBaseMemInfosVec), frontCompound->elementOp,
nowCompound->elementOp, syncIR_[*forEndIndex]->elementOp);
setOp->eventIdNum = eventIdNum;
waitOp->eventIdNum = eventIdNum;
}

syncIR_[insertSetId]->pipeAfter.push_back(setOp.get());
Expand Down
42 changes: 42 additions & 0 deletions lib/PTO/Transforms/InsertSync/SyncCodegen.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,44 @@ static void createSetOrWaitFlagOp(IRRewriter &rewriter, Operation *op,
rewriter.create<pto::SetFlagOp>(op->getLoc(), srcPipe, dstPipe, eventId);
}

// The first visit to a lane occurs at `(lane - phase - offset) mod N`.
// Work with residues so the boundary calculation cannot overflow the IV.
static Value getFirstSlotVisit(
IRRewriter& rewriter, Location loc, Value phase, uint32_t lane, uint32_t offset, uint32_t count)
{
Value modulus = rewriter.create<arith::ConstantIndexOp>(loc, count);
Value phaseMod = rewriter.create<arith::RemUIOp>(loc, phase, modulus);
uint32_t residue = (lane + count - offset) % count;
Value shiftedLane = rewriter.create<arith::ConstantIndexOp>(loc, residue + count);
Value distance = rewriter.create<arith::SubIOp>(loc, shiftedLane, phaseMod);
return rewriter.create<arith::RemUIOp>(loc, distance, modulus);
}

static void createSlotBoundaryFlag(
IRRewriter& rewriter, Operation* op, SyncOperation* sync, pto::PipeAttr srcPipe, pto::PipeAttr dstPipe,
pto::EventAttr eventId)
{
const auto& schedule = *sync->slotSchedule;
auto loop = cast<scf::ForOp>(schedule.loop);
Location loc = op->getLoc();
Value phase = loop.getLowerBound();
if (sync->isSyncWaitType()) {
// For a unit-step loop, max(lb, ub) is the next IV even for zero trips.
Value entered = rewriter.create<arith::CmpIOp>(loc, arith::CmpIPredicate::slt, phase, loop.getUpperBound());
phase = rewriter.create<arith::SelectOp>(loc, entered, loop.getUpperBound(), phase);
}
Value firstSet =
getFirstSlotVisit(rewriter, loc, phase, *sync->boundarySlot, schedule.producerOffset, sync->slotCount);
Value firstWait =
getFirstSlotVisit(rewriter, loc, phase, *sync->boundarySlot, schedule.consumerOffset, sync->slotCount);
auto predicate = schedule.producerBeforeConsumer ? arith::CmpIPredicate::ult : arith::CmpIPredicate::ule;
Value pending = rewriter.create<arith::CmpIOp>(loc, predicate, firstWait, firstSet);
auto guard = rewriter.create<scf::IfOp>(loc, pending, false);
OpBuilder::InsertionGuard insertionGuard(rewriter);
rewriter.setInsertionPointToStart(guard.thenBlock());
createSetOrWaitFlagOp(rewriter, op, sync, srcPipe, dstPipe, eventId);
}

// ==============================================================================
// 2. SyncCodegen Implementation
// ==============================================================================
Expand Down Expand Up @@ -367,6 +405,10 @@ void SyncCodegen::CreateSetWaitOpForSingleBuffer(IRRewriter &rewriter,
auto srcPipe = getPipeAttr(rewriter, sync->GetActualSrcPipe());
auto dstPipe = getPipeAttr(rewriter, sync->GetActualDstPipe());
auto eventId = getEventAttr(rewriter, sync->eventIds[0]);
if (sync->boundarySlot) {
createSlotBoundaryFlag(rewriter, op, sync, srcPipe, dstPipe, eventId);
return;
}
createSetOrWaitFlagOp(rewriter, op, sync, srcPipe, dstPipe, eventId);
}

Expand Down
Loading
Loading