Skip to content

Commit 4cc8a1f

Browse files
dpathikondaDatta Nagraj Pathikonda
authored andcommitted
Fix SIGABRT for 1D circular-buffer split pointer (ptr[i % N])
Extend rewriteSplitPtr to handle rank-1 split pointers produced by PtrAnalysis::visitOperandRem. Previously the assert required rank == 2. The failing kernel was triton_poi_fused_mul_4, the RMSNorm weight-broadcast multiply (weight[i % 2048] * input[i]), when compiled via torch.compile/inductor with --enable-triton: @triton_heuristics.pointwise( size_hints={'x': 65536}, filename=__file__, triton_meta={'signature': {'in_out_ptr0': '*fp32', 'in_ptr0': '*fp32', 'xnumel': 'i32', 'XBLOCK': 'constexpr'}, 'device': DeviceProperties(type='qaic', index=0, multi_processor_count=16, cc=None, major=None, regs_per_multiprocessor=None, max_threads_per_multi_processor=None, max_threads_per_block=1024, warp_size=32), 'constants': {}, 'native_matmul': False, 'configs': [{(0,): [['tt.divisibility', 16]], (2,): [['tt.divisibility', 16]]}], 'enable_fp_fusion': True}, inductor_meta={'grid_type': 'Grid1D', 'autotune_hints': set(), 'kernel_name': 'triton_poi_fused_mul_4', 'mutated_arg_names': ['in_out_ptr0'], 'optimize_mem': True, 'no_x_dim': False, ...}, min_elem_per_thread=0 ) @triton.jit def triton_poi_fused_mul_4(in_out_ptr0, in_ptr0, xnumel, XBLOCK : tl.constexpr): xoffset = tl.program_id(0) * XBLOCK xindex = xoffset + tl.arange(0, XBLOCK)[:] xmask = xindex < xnumel x0 = (xindex % 2048) x2 = xindex tmp0 = tl.load(in_ptr0 + (x0), xmask, eviction_policy='evict_last') tmp1 = tl.load(in_out_ptr0 + (x2), xmask) tmp2 = tmp0 * tmp1 tl.store(in_out_ptr0 + (x2), tmp2, xmask) The load of in_ptr0 + (x0) — where x0 = xindex % 2048 — is a purely 1D circular-buffer access. PtrAnalysis::visitOperandRem sets state.shape[0] = N (here 2048), making isSplitPtr() return true with parentShape.size() == 1. rewriteSplitPtr previously asserted rank == 2, causing the SIGABRT. Add create1DCastOps (two ReinterpretCastOps for the two contiguous circular-buffer segments), create1DCopies, get1DSubviews, and a WRAP_1D store path to StoreConverter. Add wraparound_1d.mlir FileCheck test.
1 parent bce5578 commit 4cc8a1f

2 files changed

Lines changed: 361 additions & 21 deletions

File tree

lib/Conversion/StructuredToMemref/StructuredToMemref.cpp

Lines changed: 221 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ using namespace mlir;
4444

4545
static const std::string WRAP_SIDE_BY_SIDE = "wrap_side_by_side";
4646
static const std::string WRAP_STACKED = "wrap_stacked";
47+
static const std::string WRAP_1D = "wrap_1d";
4748

4849
static memref::SubViewOp getSubview(int rank, ArrayRef<OpFoldResult> dims,
4950
Value source, Location loc, OpBuilder &b) {
@@ -456,33 +457,108 @@ struct MakeTensorPtrConverter
456457
return {cast1, cast2};
457458
}
458459

460+
// Handles 1D circular-buffer wrap-around: ptr[i % N] where the tile
461+
// [startOffset, startOffset+XBLOCK) may straddle the modulo boundary N.
462+
//
463+
// x = startOffset % N (position within circular buffer)
464+
// d1 = min(x + XBLOCK, N) - x (elements before wrap)
465+
// d2 = XBLOCK - d1 (elements after wrap, from 0)
466+
//
467+
// chunk1: base[startOffset .. startOffset + d1) stride 1
468+
// chunk2: base[startOffset - x .. startOffset - x + d2) stride 1
469+
// (i.e. from the start of the row in flat memory)
470+
//
471+
// When d2 == 0 (no wrap, common case) chunk2 is a zero-size memref; the
472+
// downstream CopyOp over it is a no-op.
473+
std::pair<memref::ReinterpretCastOp, memref::ReinterpretCastOp>
474+
create1DCastOps(tts::MakeTensorPtrOp op, OpAdaptor adaptor,
475+
ConversionPatternRewriter &rewriter) const {
476+
auto loc = op->getLoc();
477+
478+
auto resultType = getResultMemrefType(
479+
op, /* offset */ ShapedType::kDynamic,
480+
/* staticStrides */ SmallVector<int64_t>{ShapedType::kDynamic},
481+
/* resultShape */ SmallVector<int64_t>{ShapedType::kDynamic});
482+
483+
auto targetOffset = ofrToIndexValue(
484+
accumulateTargetOffset(op.getLoc(), op.getMixedOffsets(), rewriter),
485+
loc, rewriter);
486+
487+
// N = modulo bound stored in shape[0] by PtrAnalysis::visitOperandRem
488+
Value modN = ofrToIndexValue(op.getMixedShape()[0], loc, rewriter);
489+
490+
Value blockSz = arith::ConstantOp::create(
491+
rewriter, loc, rewriter.getIndexAttr(op.getSizes()[0]));
492+
493+
Value stride1 = arith::ConstantOp::create(
494+
rewriter, loc, rewriter.getIndexAttr(1));
495+
496+
// x = targetOffset % N
497+
Value x = arith::RemSIOp::create(rewriter, loc, targetOffset, modN);
498+
499+
// wrappedBase = targetOffset - x (start of the circular-buffer row)
500+
Value wrappedBase = arith::SubIOp::create(rewriter, loc, targetOffset, x);
501+
502+
// d1 = min(x + blockSz, N) - x
503+
Value nextOff = arith::AddIOp::create(rewriter, loc, x, blockSz);
504+
Value clampedOff = arith::MinSIOp::create(rewriter, loc, nextOff, modN);
505+
Value d1 = arith::SubIOp::create(rewriter, loc, clampedOff, x);
506+
507+
// chunk1: [targetOffset, targetOffset + d1)
508+
SmallVector<Value> sizes1{d1};
509+
SmallVector<Value> strides1{stride1};
510+
auto cast1 = memref::ReinterpretCastOp::create(
511+
rewriter, loc, resultType, adaptor.getBase(),
512+
targetOffset, sizes1, strides1);
513+
514+
// d2 = blockSz - d1 (wrap-around part, may be zero)
515+
Value d2 = arith::SubIOp::create(rewriter, loc, blockSz, d1);
516+
517+
// chunk2: [wrappedBase, wrappedBase + d2)
518+
SmallVector<Value> sizes2{d2};
519+
auto cast2 = memref::ReinterpretCastOp::create(
520+
rewriter, loc, resultType, adaptor.getBase(),
521+
wrappedBase, sizes2, strides1);
522+
523+
return {cast1, cast2};
524+
}
525+
459526
LogicalResult rewriteSplitPtr(tts::MakeTensorPtrOp op, OpAdaptor adaptor,
460527
ConversionPatternRewriter &rewriter) const {
461528
auto parentShape = op.getStaticShape();
462-
assert(parentShape.size() == 2 &&
463-
"Only support split pointer for 2D tensors only");
464529
SmallVector<Value> casts;
465530
StringRef wrapType;
466531

467-
// For split pointers, a split dimension is either a dynamic or a non-zero
468-
// value. The other dimension must be zero.
469-
auto isSplitDimension = [](int64_t dim) {
470-
return dim == ShapedType::kDynamic || dim != 0;
471-
};
472-
473-
if (isSplitDimension(parentShape[0])) {
474-
// Stacked case
475-
assert(parentShape[1] == 0);
476-
auto [cast1, cast2] = createStackedCastOps(op, adaptor, rewriter);
477-
casts = {cast1.getResult(), cast2.getResult()};
478-
wrapType = WRAP_STACKED;
479-
} else if (isSplitDimension(parentShape[1])) {
480-
assert(parentShape[0] == 0);
481-
auto [cast1, cast2] = createSideBySideCastOps(op, adaptor, rewriter);
532+
if (parentShape.size() == 1) {
533+
// 1D circular-buffer wrap-around: ptr[i % N]
534+
// shape[0] carries N (set by PtrAnalysis::visitOperandRem for rank-1).
535+
auto [cast1, cast2] = create1DCastOps(op, adaptor, rewriter);
482536
casts = {cast1.getResult(), cast2.getResult()};
483-
wrapType = WRAP_SIDE_BY_SIDE;
537+
wrapType = WRAP_1D;
484538
} else {
485-
llvm_unreachable("Unexpected split pointer shape");
539+
assert(parentShape.size() == 2 &&
540+
"Only support split pointer for 1D and 2D tensors");
541+
542+
// For split pointers, a split dimension is either a dynamic or a non-zero
543+
// value. The other dimension must be zero.
544+
auto isSplitDimension = [](int64_t dim) {
545+
return dim == ShapedType::kDynamic || dim != 0;
546+
};
547+
548+
if (isSplitDimension(parentShape[0])) {
549+
// Stacked case
550+
assert(parentShape[1] == 0);
551+
auto [cast1, cast2] = createStackedCastOps(op, adaptor, rewriter);
552+
casts = {cast1.getResult(), cast2.getResult()};
553+
wrapType = WRAP_STACKED;
554+
} else if (isSplitDimension(parentShape[1])) {
555+
assert(parentShape[0] == 0);
556+
auto [cast1, cast2] = createSideBySideCastOps(op, adaptor, rewriter);
557+
casts = {cast1.getResult(), cast2.getResult()};
558+
wrapType = WRAP_SIDE_BY_SIDE;
559+
} else {
560+
llvm_unreachable("Unexpected split pointer shape");
561+
}
486562
}
487563

488564
auto combinedCast = UnrealizedConversionCastOp::create(
@@ -643,6 +719,34 @@ struct LoadConverter : public OpConversionPattern<tts::LoadOp> {
643719
memref::CopyOp::create(rewriter, loc, block2, block2Dst);
644720
}
645721

722+
// 1D wrap copy: block1 = [startOffset, startOffset+d1), fills dst[0..d1).
723+
// block2 = wrapped part [wrappedBase, wrappedBase+d2), fills dst[d1..XBLOCK).
724+
// When d2 == 0 (no wrap) block2 has size 0 and the second CopyOp is a no-op.
725+
void create1DCopies(Value block1, Value block2, Value dst, Location loc,
726+
ConversionPatternRewriter &rewriter) const {
727+
auto zero =
728+
arith::ConstantOp::create(rewriter, loc, rewriter.getIndexAttr(0));
729+
auto one =
730+
arith::ConstantOp::create(rewriter, loc, rewriter.getIndexAttr(1));
731+
732+
Value d1 = memref::DimOp::create(rewriter, loc, block1, 0);
733+
Value d2 = memref::DimOp::create(rewriter, loc, block2, 0);
734+
735+
// dst[0 : d1] <- block1
736+
auto block1Dst = memref::SubViewOp::create(rewriter, loc, dst,
737+
ValueRange{zero},
738+
ValueRange{d1},
739+
ValueRange{one});
740+
// dst[d1 : d1 + d2] <- block2
741+
auto block2Dst = memref::SubViewOp::create(rewriter, loc, dst,
742+
ValueRange{d1},
743+
ValueRange{d2},
744+
ValueRange{one});
745+
746+
memref::CopyOp::create(rewriter, loc, block1, block1Dst);
747+
memref::CopyOp::create(rewriter, loc, block2, block2Dst);
748+
}
749+
646750
memref::SubViewOp createSubview(Value src, ArrayRef<OpFoldResult> offsets,
647751
ArrayRef<OpFoldResult> sizes,
648752
ArrayRef<OpFoldResult> strides, Location loc,
@@ -695,6 +799,25 @@ struct LoadConverter : public OpConversionPattern<tts::LoadOp> {
695799
return {sv1, sv2};
696800
}
697801

802+
// 1D masked subviews: each chunk is already sized to d1/d2 by the cast ops,
803+
// so we read the actual dim and clip to the mask dimension.
804+
std::pair<memref::SubViewOp, memref::SubViewOp>
805+
get1DSubviews(ArrayRef<OpFoldResult> dims, Value block1, Value block2,
806+
Location loc, ConversionPatternRewriter &rewriter) const {
807+
// dims[0] is the mask size for the whole 1D block.
808+
// chunk1 contributes min(d1, maskSize) elements; chunk2 the remainder.
809+
// We use the actual sizes already baked into the ReinterpretCastOps (d1/d2)
810+
// directly, clipped by the mask dimension via SubViewOp.
811+
OpFoldResult d1 = memref::DimOp::create(rewriter, loc, block1, 0).getResult();
812+
OpFoldResult d2 = memref::DimOp::create(rewriter, loc, block2, 0).getResult();
813+
814+
SmallVector<OpFoldResult> offsets{rewriter.getIndexAttr(0)};
815+
SmallVector<OpFoldResult> strides{rewriter.getIndexAttr(1)};
816+
auto sv1 = createSubview(block1, offsets, {d1}, strides, loc, rewriter);
817+
auto sv2 = createSubview(block2, offsets, {d2}, strides, loc, rewriter);
818+
return {sv1, sv2};
819+
}
820+
698821
LogicalResult
699822
rewriteStructuredLoad(tts::LoadOp op, OpAdaptor adaptor,
700823
ConversionPatternRewriter &rewriter) const {
@@ -715,7 +838,8 @@ struct LoadConverter : public OpConversionPattern<tts::LoadOp> {
715838

716839
auto ptrDefiningOp = ptr.getDefiningOp();
717840
if (ptrDefiningOp->hasAttr(WRAP_SIDE_BY_SIDE) ||
718-
ptrDefiningOp->hasAttr(WRAP_STACKED)) {
841+
ptrDefiningOp->hasAttr(WRAP_STACKED) ||
842+
ptrDefiningOp->hasAttr(WRAP_1D)) {
719843

720844
auto unrealizedCast = cast<UnrealizedConversionCastOp>(ptrDefiningOp);
721845
auto memrefs = unrealizedCast.getOperands();
@@ -727,6 +851,8 @@ struct LoadConverter : public OpConversionPattern<tts::LoadOp> {
727851
createSideBySideCopies(block1, block2, alloc, loc, rewriter);
728852
} else if (unrealizedCast->hasAttr(WRAP_STACKED)) {
729853
createStackedCopies(block1, block2, alloc, loc, rewriter);
854+
} else if (unrealizedCast->hasAttr(WRAP_1D)) {
855+
create1DCopies(block1, block2, alloc, loc, rewriter);
730856
} else {
731857
llvm_unreachable("unexpected wraparound type");
732858
}
@@ -765,7 +891,8 @@ struct LoadConverter : public OpConversionPattern<tts::LoadOp> {
765891

766892
auto ptrDefiningOp = ptr.getDefiningOp();
767893
if (ptrDefiningOp->hasAttr(WRAP_SIDE_BY_SIDE) ||
768-
ptrDefiningOp->hasAttr(WRAP_STACKED)) {
894+
ptrDefiningOp->hasAttr(WRAP_STACKED) ||
895+
ptrDefiningOp->hasAttr(WRAP_1D)) {
769896

770897
auto unrealizedCast = cast<UnrealizedConversionCastOp>(ptrDefiningOp);
771898

@@ -782,6 +909,10 @@ struct LoadConverter : public OpConversionPattern<tts::LoadOp> {
782909
auto [subview1, subview2] =
783910
getStackedSubviews(mixedDims, block1, block2, loc, rewriter);
784911
createStackedCopies(subview1, subview2, alloc, loc, rewriter);
912+
} else if (unrealizedCast->hasAttr(WRAP_1D)) {
913+
auto [subview1, subview2] =
914+
get1DSubviews(mixedDims, block1, block2, loc, rewriter);
915+
create1DCopies(subview1, subview2, alloc, loc, rewriter);
785916
} else {
786917
llvm_unreachable("unexpected wraparound type");
787918
}
@@ -1123,6 +1254,75 @@ struct StoreConverter : public OpConversionPattern<tts::StoreOp> {
11231254
auto storeValue = op.getValue();
11241255
auto rank = cast<RankedTensorType>(storeValue.getType()).getRank();
11251256

1257+
// Handle 1D circular-buffer wrap-around store: store into two contiguous
1258+
// segments that together span the block.
1259+
auto ptrDefiningOp = ptr.getDefiningOp();
1260+
if (ptrDefiningOp && ptrDefiningOp->hasAttr(WRAP_1D)) {
1261+
auto unrealizedCast = cast<UnrealizedConversionCastOp>(ptrDefiningOp);
1262+
auto memrefs = unrealizedCast.getOperands();
1263+
assert(memrefs.size() == 2);
1264+
Value block1 = memrefs[0];
1265+
Value block2 = memrefs[1];
1266+
1267+
Value d1 = memref::DimOp::create(rewriter, loc, block1, 0);
1268+
Value d2 = memref::DimOp::create(rewriter, loc, block2, 0);
1269+
Value zero =
1270+
arith::ConstantOp::create(rewriter, loc, rewriter.getIndexAttr(0));
1271+
Value one =
1272+
arith::ConstantOp::create(rewriter, loc, rewriter.getIndexAttr(1));
1273+
1274+
if (op.hasMask()) {
1275+
auto mixedDims = op.getMixedMaskDims();
1276+
// Clip store slices to the mask size (dims[0] for the 1D case).
1277+
// Use d1/d2 directly from the cast ops since they are already sized.
1278+
auto slice1 = tensor::ExtractSliceOp::create(
1279+
rewriter, loc, storeValue,
1280+
SmallVector<OpFoldResult>{rewriter.getIndexAttr(0)},
1281+
SmallVector<OpFoldResult>{OpFoldResult(d1)},
1282+
SmallVector<OpFoldResult>{rewriter.getIndexAttr(1)});
1283+
auto dst1Subview = memref::SubViewOp::create(
1284+
rewriter, loc, block1, ValueRange{zero}, ValueRange{d1},
1285+
ValueRange{one});
1286+
auto store1 = bufferization::MaterializeInDestinationOp::create(
1287+
rewriter, loc, slice1, dst1Subview);
1288+
store1.setWritable(true);
1289+
1290+
auto slice2 = tensor::ExtractSliceOp::create(
1291+
rewriter, loc, storeValue,
1292+
SmallVector<OpFoldResult>{OpFoldResult(d1)},
1293+
SmallVector<OpFoldResult>{OpFoldResult(d2)},
1294+
SmallVector<OpFoldResult>{rewriter.getIndexAttr(1)});
1295+
auto dst2Subview = memref::SubViewOp::create(
1296+
rewriter, loc, block2, ValueRange{zero}, ValueRange{d2},
1297+
ValueRange{one});
1298+
auto store2 = bufferization::MaterializeInDestinationOp::create(
1299+
rewriter, loc, slice2, dst2Subview);
1300+
store2.setWritable(true);
1301+
} else {
1302+
// Unmasked: store first d1 elements to chunk1, next d2 to chunk2.
1303+
auto slice1 = tensor::ExtractSliceOp::create(
1304+
rewriter, loc, storeValue,
1305+
SmallVector<OpFoldResult>{rewriter.getIndexAttr(0)},
1306+
SmallVector<OpFoldResult>{OpFoldResult(d1)},
1307+
SmallVector<OpFoldResult>{rewriter.getIndexAttr(1)});
1308+
auto store1 = bufferization::MaterializeInDestinationOp::create(
1309+
rewriter, loc, slice1, block1);
1310+
store1.setWritable(true);
1311+
1312+
auto slice2 = tensor::ExtractSliceOp::create(
1313+
rewriter, loc, storeValue,
1314+
SmallVector<OpFoldResult>{OpFoldResult(d1)},
1315+
SmallVector<OpFoldResult>{OpFoldResult(d2)},
1316+
SmallVector<OpFoldResult>{rewriter.getIndexAttr(1)});
1317+
auto store2 = bufferization::MaterializeInDestinationOp::create(
1318+
rewriter, loc, slice2, block2);
1319+
store2.setWritable(true);
1320+
}
1321+
1322+
rewriter.eraseOp(op);
1323+
return success();
1324+
}
1325+
11261326
if (op.hasMask()) {
11271327
auto mixedDims = op.getMixedMaskDims();
11281328

0 commit comments

Comments
 (0)