@@ -44,6 +44,7 @@ using namespace mlir;
4444
4545static const std::string WRAP_SIDE_BY_SIDE = " wrap_side_by_side" ;
4646static const std::string WRAP_STACKED = " wrap_stacked" ;
47+ static const std::string WRAP_1D = " wrap_1d" ;
4748
4849static 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