Skip to content

Commit 5e583a9

Browse files
committed
[backend](fix) preserve addresses across pointer bitcasts
1 parent d011fad commit 5e583a9

20 files changed

Lines changed: 4627 additions & 591 deletions

third_party/ascend/include/TritonToLinalg/TritonOpConverter.h

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -100,14 +100,18 @@ class SelectCanonicalizer : public OpRewritePattern<arith::SelectOp> {
100100
};
101101

102102
/*
103-
* Move tt.bitcast to a previous location if tt.bitcast is not directly applied
104-
* on function arguments
103+
* Preserve different-width pointer bitcasts as exact address boundaries and
104+
* canonicalize only same-width pointer shape operations.
105105
*/
106106
class BitcastCanonicalizer : public OpRewritePattern<triton::BitcastOp> {
107107
public:
108-
using OpRewritePattern<triton::BitcastOp>::OpRewritePattern;
108+
BitcastCanonicalizer(MLIRContext *context, bool &hadError)
109+
: OpRewritePattern<triton::BitcastOp>(context), hadError(hadError) {}
109110
LogicalResult matchAndRewrite(triton::BitcastOp bitcastOp,
110111
PatternRewriter &rewriter) const override;
112+
113+
private:
114+
bool &hadError;
111115
};
112116

113117
template <typename MathOp>

third_party/ascend/include/TritonToLinalg/TritonToLinalgPass.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -83,7 +83,8 @@ class TritonToLinalgPass : public TritonToLinalgBase<TritonToLinalgPass> {
8383
TritonTypeConverter &tritonTypeConverter);
8484

8585
void
86-
populateTritonToLinalgCanonicalizationPatterns(RewritePatternSet &patterns);
86+
populateTritonToLinalgCanonicalizationPatterns(RewritePatternSet &patterns,
87+
bool &hadError);
8788

8889
void populateTritonToLinalgConversionPatterns(TypeConverter &typeConverter,
8990
RewritePatternSet &patterns,

third_party/ascend/include/TritonToUnstructure/OffsetAnalysis.h

Lines changed: 124 additions & 97 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929
#include "mlir/IR/OpDefinition.h"
3030
#include "mlir/IR/PatternMatch.h"
3131
#include "mlir/IR/Value.h"
32+
#include "mlir/Support/LogicalResult.h"
3233
#include "mlir/Transforms/DialectConversion.h"
3334
#include "triton/Dialect/Triton/IR/Dialect.h"
3435
#include "llvm/ADT/DenseMap.h"
@@ -95,6 +96,7 @@ struct PtrOffsetInfo {
9596
SmallVector<Value> getOffsets() const;
9697
SmallVector<Value> &getOffsetsRef();
9798
bool isScalarLike() const;
99+
bool isByteAddressed() const;
98100
SmallVector<AxisInfo> &getStructuredRef();
99101
const SmallVector<AxisInfo> &getStructured() const;
100102
int getRank() const;
@@ -110,6 +112,7 @@ struct PtrOffsetInfo {
110112
void setStructured(ArrayRef<AxisInfo> structured);
111113
void setStructured(const PtrOffsetInfo &other);
112114
void setScalarLike(bool scalarLike);
115+
void setByteAddressed(bool byteAddressed = true);
113116

114117
bool isStructured(int dim) const;
115118
bool isStructured() const;
@@ -120,157 +123,181 @@ struct PtrOffsetInfo {
120123

121124
private:
122125
Value ptr;
126+
// The offset normally uses the pointee element as its unit, matching
127+
// tt.addptr. After a different-width pointer bitcast, combining offsets in
128+
// either the source or destination element unit can lose address bits. In
129+
// that case byteAddressed is set and this same field stores an exact signed
130+
// byte offset from ptr. Every later AddPtr contributes
131+
// offset * sizeof(current pointee), while Bitcast itself contributes zero.
132+
// Consumers must inspect byteAddressed before interpreting offset.
123133
Value offset;
124134
SmallVector<Value> tptOffsets;
125135

126136
bool scalarLike = false;
137+
bool byteAddressed = false;
127138

128139
SmallVector<AxisInfo> structured;
129140
};
130141

131142
PtrOffsetInfo combineInfo(const PtrOffsetInfo &lhs, const PtrOffsetInfo &rhs);
132143

144+
// Compatibility entry point for legacy argument reconstruction. New analysis
145+
// code must use parseChecked so diagnostics stop the enclosing pass.
133146
void parse(Value operand, const Location &loc, RewriterBase &rewriter,
134147
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
135148

136-
void parseLoopRegionIterArg(LoopLikeOpInterface loopOp, const Location &loc,
137-
RewriterBase &rewriter,
138-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap,
139-
BlockArgument regionIterArg);
149+
LogicalResult parseChecked(Value operand, const Location &loc,
150+
RewriterBase &rewriter,
151+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
140152

141-
void parseArithOp(Operation *arithOp, const Location &loc,
142-
RewriterBase &rewriter,
143-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
153+
LogicalResult
154+
parseLoopRegionIterArg(LoopLikeOpInterface loopOp, const Location &loc,
155+
RewriterBase &rewriter,
156+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap,
157+
BlockArgument regionIterArg);
144158

145-
void parseTritonOp(Operation *tritonOp, const Location &loc,
146-
RewriterBase &rewriter,
147-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
159+
LogicalResult parseArithOp(Operation *arithOp, const Location &loc,
160+
RewriterBase &rewriter,
161+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
148162

149-
void parseTritonOp(Operation *tritonOp, const Location &loc,
150-
RewriterBase &rewriter,
151-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
163+
LogicalResult parseTritonOp(Operation *tritonOp, const Location &loc,
164+
RewriterBase &rewriter,
165+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
152166

153-
void parseAddPtr(triton::AddPtrOp op, const Location &loc,
154-
RewriterBase &rewriter,
155-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
167+
LogicalResult parseAddPtr(triton::AddPtrOp op, const Location &loc,
168+
RewriterBase &rewriter,
169+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
156170

157-
void parseSplat(triton::SplatOp op, const Location &loc, RewriterBase &rewriter,
158-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
171+
LogicalResult parseSplat(triton::SplatOp op, const Location &loc,
172+
RewriterBase &rewriter,
173+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
159174

160175
template <typename BinOpTy>
161-
void parseBinaryOp(BinOpTy op, const Location &loc, RewriterBase &rewriter,
162-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
176+
LogicalResult parseBinaryOp(BinOpTy op, const Location &loc,
177+
RewriterBase &rewriter,
178+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
163179

164-
void parseAddI(arith::AddIOp op, const Location &loc, RewriterBase &rewriter,
165-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
180+
LogicalResult parseAddI(arith::AddIOp op, const Location &loc,
181+
RewriterBase &rewriter,
182+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
166183

167-
void parseSubI(arith::SubIOp op, const Location &loc, RewriterBase &rewriter,
168-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
184+
LogicalResult parseSubI(arith::SubIOp op, const Location &loc,
185+
RewriterBase &rewriter,
186+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
169187

170-
void parseIndexCast(arith::IndexCastOp op, const Location &loc,
171-
RewriterBase &rewriter,
172-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
188+
LogicalResult parseIndexCast(arith::IndexCastOp op, const Location &loc,
189+
RewriterBase &rewriter,
190+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
173191

174192
template <typename ConstOpTy>
175-
void parseConstantOp(ConstOpTy dst, const Location &loc, RewriterBase &rewriter,
176-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
193+
LogicalResult parseConstantOp(ConstOpTy dst, const Location &loc,
194+
RewriterBase &rewriter,
195+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
177196

178-
void parseMakeRange(triton::MakeRangeOp op, const Location &loc,
179-
RewriterBase &rewriter,
180-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
197+
LogicalResult parseMakeRange(triton::MakeRangeOp op, const Location &loc,
198+
RewriterBase &rewriter,
199+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
181200

182-
void parseExtSI(arith::ExtSIOp op, const Location &loc, RewriterBase &rewriter,
183-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
201+
LogicalResult parseExtSI(arith::ExtSIOp op, const Location &loc,
202+
RewriterBase &rewriter,
203+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
184204

185-
void parseBitcast(triton::BitcastOp op, const Location &loc,
186-
RewriterBase &rewriter,
187-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
205+
LogicalResult parseBitcast(triton::BitcastOp op, const Location &loc,
206+
RewriterBase &rewriter,
207+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
188208

189-
void parseLoad(triton::LoadOp op, const Location &loc, RewriterBase &rewriter,
190-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
209+
LogicalResult parseLoad(triton::LoadOp op, const Location &loc,
210+
RewriterBase &rewriter,
211+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
191212

192-
void parseMulI(arith::MulIOp op, const Location &loc, RewriterBase &rewriter,
193-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
213+
LogicalResult parseMulI(arith::MulIOp op, const Location &loc,
214+
RewriterBase &rewriter,
215+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
194216

195-
void parseBroadcast(triton::BroadcastOp op, const Location &loc,
196-
RewriterBase &rewriter,
197-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
217+
LogicalResult parseBroadcast(triton::BroadcastOp op, const Location &loc,
218+
RewriterBase &rewriter,
219+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
198220

199-
void parseExpandDims(triton::ExpandDimsOp op, const Location &loc,
200-
RewriterBase &rewriter,
201-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
221+
LogicalResult parseExpandDims(triton::ExpandDimsOp op, const Location &loc,
222+
RewriterBase &rewriter,
223+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
202224

203-
void parseClampF(triton::ClampFOp op, const Location &loc,
204-
RewriterBase &rewriter,
205-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
225+
LogicalResult parseClampF(triton::ClampFOp op, const Location &loc,
226+
RewriterBase &rewriter,
227+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
206228

207-
void parseSelect(arith::SelectOp op, const Location &loc,
208-
RewriterBase &rewriter,
209-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
229+
LogicalResult parseSelect(arith::SelectOp op, const Location &loc,
230+
RewriterBase &rewriter,
231+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
210232

211-
void parseFPToSI(arith::FPToSIOp op, const Location &loc,
212-
RewriterBase &rewriter,
213-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
233+
LogicalResult parseFPToSI(arith::FPToSIOp op, const Location &loc,
234+
RewriterBase &rewriter,
235+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
214236

215-
void parseSIToFP(arith::SIToFPOp op, const Location &loc,
216-
RewriterBase &rewriter,
217-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
237+
LogicalResult parseSIToFP(arith::SIToFPOp op, const Location &loc,
238+
RewriterBase &rewriter,
239+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
218240

219241
// FIXME:Z|wait triton version upgrade to 3.4
220242
// void parseMakeTensorDesc(triton::MakeTensorDescOp op, const Location &loc,
221243
// RewriterBase &rewriter,
222244
// llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
223245

224-
void parseMakeTensorPtr(triton::MakeTensorPtrOp op, const Location &loc,
225-
RewriterBase &rewriter,
226-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
227-
228-
void parseAdvance(triton::AdvanceOp op, const Location &loc,
229-
RewriterBase &rewriter,
230-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
246+
LogicalResult
247+
parseMakeTensorPtr(triton::MakeTensorPtrOp op, const Location &loc,
248+
RewriterBase &rewriter,
249+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
231250

232-
void parseReduce(triton::ReduceOp op, const Location &loc,
233-
RewriterBase &rewriter,
234-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
251+
LogicalResult parseAdvance(triton::AdvanceOp op, const Location &loc,
252+
RewriterBase &rewriter,
253+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
235254

236-
void parseReduceReturn(triton::ReduceReturnOp op, const Location &loc,
237-
RewriterBase &rewriter,
238-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
255+
LogicalResult parseReduce(triton::ReduceOp op, const Location &loc,
256+
RewriterBase &rewriter,
257+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
239258

240-
void parseIf(scf::IfOp op, const Location &loc, RewriterBase &rewriter,
241-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap, Value dst);
259+
LogicalResult
260+
parseReduceReturn(triton::ReduceReturnOp op, const Location &loc,
261+
RewriterBase &rewriter,
262+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
242263

243-
void parseYield(scf::YieldOp op, const Location &loc, RewriterBase &rewriter,
244-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
264+
LogicalResult parseIf(scf::IfOp op, const Location &loc, RewriterBase &rewriter,
265+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap,
266+
Value dst);
245267

246-
void parseLoopOp(LoopLikeOpInterface op, const Location &loc,
247-
RewriterBase &rewriter,
248-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap, Value dst);
268+
LogicalResult parseYield(scf::YieldOp op, const Location &loc,
269+
RewriterBase &rewriter,
270+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
249271

250-
void parseExtractSlice(tensor::ExtractSliceOp op, const Location &loc,
251-
RewriterBase &rewriter,
252-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
272+
LogicalResult parseLoopOp(LoopLikeOpInterface op, const Location &loc,
273+
RewriterBase &rewriter,
274+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap,
275+
Value dst);
253276

254-
void parseInsertSlice(tensor::InsertSliceOp op, const Location &loc,
255-
RewriterBase &rewriter,
256-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
257-
258-
void parseExtract(tensor::ExtractOp op, const Location &loc,
277+
LogicalResult
278+
parseExtractSlice(tensor::ExtractSliceOp op, const Location &loc,
259279
RewriterBase &rewriter,
260280
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
261281

262-
void parseInsert(tensor::InsertOp op, const Location &loc,
263-
RewriterBase &rewriter,
264-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
282+
LogicalResult parseInsertSlice(tensor::InsertSliceOp op, const Location &loc,
283+
RewriterBase &rewriter,
284+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
265285

266-
void parseIntToPtr(triton::IntToPtrOp op, const Location &loc,
267-
RewriterBase &rewriter,
268-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
286+
LogicalResult parseExtract(tensor::ExtractOp op, const Location &loc,
287+
RewriterBase &rewriter,
288+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
269289

270-
void parseStructuredCustomOp(Operation *op, const Location &loc,
271-
RewriterBase &rewriter,
272-
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap,
273-
unsigned resultIdx);
290+
LogicalResult parseInsert(tensor::InsertOp op, const Location &loc,
291+
RewriterBase &rewriter,
292+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
293+
294+
LogicalResult parseIntToPtr(triton::IntToPtrOp op, const Location &loc,
295+
RewriterBase &rewriter,
296+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap);
297+
298+
LogicalResult parseStructuredCustomOp(
299+
Operation *op, const Location &loc, RewriterBase &rewriter,
300+
llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap, unsigned resultIdx);
274301
} // namespace triton
275302

276303
} // namespace mlir

third_party/ascend/include/TritonToUnstructure/UnstructureConversionPass.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -114,7 +114,8 @@ class UnstructuredMemAccessConverter : public OpRewritePattern<MemAccOpTy> {
114114

115115
template <typename... Args>
116116
MemAccOpTy createMemAccOp(MemAccOpTy op, Value ptrToAccess, Location loc,
117-
PatternRewriter &rewriter, Args &&...args) const;
117+
PatternRewriter &rewriter, bool preserveLoadMask,
118+
Args &&...args) const;
118119

119120
const llvm::DenseMap<Value, PtrOffsetInfo> &offsetMap;
120121
const llvm::SmallDenseMap<Value, bool> &fromTensorArg;

third_party/ascend/include/Utils/Utils.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,8 @@ const std::string discreteMaskAttrName = "DiscreteMask";
5050
const std::string discreteAttrName = "DiscreteMemAccess";
5151
const std::string continuousAttrName = "ContinuousMemAccess";
5252
const std::string customSrcPtrIndexAttrName = "SrcPtrIndex";
53+
inline constexpr llvm::StringLiteral pointerBitcastPointerCastAttrName =
54+
"tt.pointer_bitcast_pointer_cast";
5355

5456
bool isaPermutedMemRefType(MemRefType);
5557

0 commit comments

Comments
 (0)