Skip to content
Merged
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
3 changes: 0 additions & 3 deletions .gitmodules
Original file line number Diff line number Diff line change
@@ -1,3 +0,0 @@
[submodule "triton"]
path = triton
url = https://github.com/triton-lang/triton.git
6 changes: 4 additions & 2 deletions backend/driver.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,8 +54,10 @@ def _ty_to_cpp(ty):
"u16": "uint16_t",
"u32": "uint32_t",
"u64": "uint64_t",
"fp16": "float",
Comment thread
lechenyu marked this conversation as resolved.
"bf16": "float",
# Proper support for bfloat16 and float16 is not yet handled.
# https://github.com/microsoft/triton-shared/issues/348
# "fp16": "TODO",
# "bf16": "TODO",
"fp32": "float",
"f32": "float",
"fp64": "double",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -360,7 +360,7 @@ struct LoadConverter : public OpConversionPattern<triton::LoadOp> {
loc, rewriter);
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());
auto loadOp = rewriter.create<affine::AffineLoadOp>(
op.getLoc(), sMemRef, zeroMap, std::nullopt);
op.getLoc(), sMemRef, zeroMap, ValueRange{});
rewriter.replaceOp(op, loadOp.getResult());
return success();
}
Expand Down Expand Up @@ -520,7 +520,7 @@ struct StoreConverter : public OpConversionPattern<triton::StoreOp> {
PtrAnalysis::getScalarMemRef(op.getPtr(), ptr, loc, rewriter);
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());
rewriter.create<affine::AffineStoreOp>(loc, val, sMemRef, zeroMap,
std::nullopt);
ValueRange{});
rewriter.eraseOp(op);
return success();
}
Expand Down Expand Up @@ -649,6 +649,28 @@ struct SplatConverter : public OpConversionPattern<triton::SplatOp> {
}
};

struct UnsplatConverter : public OpConversionPattern<triton::UnsplatOp> {
using OpConversionPattern::OpConversionPattern;

LogicalResult
matchAndRewrite(triton::UnsplatOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto tensorType = op.getSrc().getType();

// Only generate indices for non-zero rank tensors.
SmallVector<Value, 1> indices(tensorType.getRank());
if (indices.size() > 0) {
auto zeroIdx =
rewriter.createOrFold<arith::ConstantIndexOp>(op.getLoc(), 0);
llvm::fill(indices, zeroIdx);
}

rewriter.replaceOpWithNewOp<tensor::ExtractOp>(op, adaptor.getSrc(),
indices);
return success();
}
};

struct BroadcastConverter : public OpConversionPattern<triton::BroadcastOp> {
private:
using OpConversionPattern<triton::BroadcastOp>::OpConversionPattern;
Expand Down Expand Up @@ -1397,24 +1419,6 @@ struct ReduceConverter : public OpConversionPattern<triton::ReduceOp> {
return success();
}

LogicalResult
convertToTensorExtract(triton::ReduceOp op,
typename triton::ReduceOp::Adaptor adaptor,
ConversionPatternRewriter &rewriter) const {
assert(llvm::hasSingleElement(op.getSrcs()));

auto returnOp = cast<triton::ReduceReturnOp>(*op.getOps().begin());
assert(llvm::hasSingleElement(returnOp.getResult()));
assert(cast<BlockArgument>(returnOp.getResult().front()).getArgNumber() ==
0);

auto source = op.getSrcs().front();
auto zeroIdx =
rewriter.createOrFold<arith::ConstantIndexOp>(op.getLoc(), 0);
rewriter.replaceOpWithNewOp<tensor::ExtractOp>(op, source, zeroIdx);
return success();
}

public:
LogicalResult
matchAndRewrite(triton::ReduceOp op,
Expand All @@ -1431,14 +1435,6 @@ struct ReduceConverter : public OpConversionPattern<triton::ReduceOp> {
"axis is within "
"operand's rank");

// Unsplat is implemented as a single element, rank 1 reduction where
// single element is yielded immediately. This can be simplified into
// a single element extract.
if (llvm::hasSingleElement(op.getOps()) && sourceType.getRank() == 1 &&
sourceType.getShape()[0] == 1) {
return convertToTensorExtract(op, adaptor, rewriter);
}

return convertToLinalgReduce(op, adaptor, rewriter);
}
};
Expand Down
3 changes: 0 additions & 3 deletions include/triton-shared/Dialect/TPtr/IR/TPtrDialect.td
Original file line number Diff line number Diff line change
Expand Up @@ -109,9 +109,6 @@ def TPTR_TypeOffsetOp : TPTR_Op<"type_offset", [ConstantLike, Pure]> {

let arguments = (ins TypeAttr:$baseType);
let results = (outs AnySignlessIntegerOrIndex:$result);
let builders = [
OpBuilder<(ins "TypeAttr":$baseType, CArg<"Type", "nullptr">:$resultTy)>
];
let assemblyFormat = [{
attr-dict $baseType custom<IntType>(type($result))
}];
Expand Down
15 changes: 6 additions & 9 deletions lib/Conversion/StructuredToMemref/StructuredToMemref.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -577,21 +577,18 @@ struct MakeTensorPtrConverter

struct MakeGatherScatterTensorPtrConverter
: public OpConversionPattern<tts::MakeGatherScatterTensorPtrOp> {
private:
using OpConversionPattern<tts::MakeGatherScatterTensorPtrOp>::OpConversionPattern;

public:
MakeGatherScatterTensorPtrConverter(const TypeConverter &typeConverter,
MLIRContext *context)
: OpConversionPattern<tts::MakeGatherScatterTensorPtrOp>(typeConverter, context) {}
using OpConversionPattern::OpConversionPattern;

LogicalResult
matchAndRewrite(tts::MakeGatherScatterTensorPtrOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
// The gatherScatterPtr is rewritten as separate rows during load/store
// operations. Therefore, no action is needed here except saving
// adaptor.getBase().
rewriter.replaceOp(op, adaptor.getBase());
// adaptor.getBase(). DialectConversion will ignore pure type conversion if
// we were to simply replace the op with adaptor.getBase(). To circumvent
// this we create an identity cast.
rewriter.replaceOpWithNewOp<UnrealizedConversionCastOp>(
op, adaptor.getBase().getType(), adaptor.getBase());
return success();
}
};
Expand Down
1 change: 1 addition & 0 deletions lib/Conversion/TritonArithToLinalg/TritonArithToLinalg.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ void mlir::triton::populateTritonArithToLinalgConversionPatterns(
patterns.add<ClampConverter>(patterns.getContext());
patterns.add<MatmulConverter>(patterns.getContext());
patterns.add<SplatConverter>(patterns.getContext());
patterns.add<UnsplatConverter>(patterns.getContext());
patterns.add<DenseConstantConverter>(patterns.getContext());
patterns.add<CumSumConverter>(patterns.getContext());
patterns.add<ReshapeConverter>(patterns.getContext());
Expand Down
1 change: 1 addition & 0 deletions lib/Conversion/TritonToLinalg/TritonToLinalg.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ void mlir::triton::populateTritonToLinalgConversionPatterns(
patterns.add<AssertConverter>(patterns.getContext());
patterns.add<MatmulConverter>(patterns.getContext());
patterns.add<SplatConverter>(patterns.getContext());
patterns.add<UnsplatConverter>(patterns.getContext());
patterns.add<DenseConstantConverter>(patterns.getContext());
patterns.add<UnrealizedCastConverter>(patterns.getContext());
patterns.add<CumSumConverter>(patterns.getContext());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,7 @@ struct ScalarLoadConverter : public OpConversionPattern<tts::GatherOp> {
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());

auto scalarLoadOp = rewriter.create<affine::AffineLoadOp>(
loc, memref, zeroMap, std::nullopt);
loc, memref, zeroMap, ValueRange{});

rewriter.replaceOp(gatherOp, scalarLoadOp.getResult());

Expand Down Expand Up @@ -150,7 +150,7 @@ struct ScalarStoreConverter : public OpConversionPattern<tts::ScatterOp> {
auto zeroMap = AffineMap::getConstantMap(0, rewriter.getContext());

rewriter.create<affine::AffineStoreOp>(loc, storeVal, memref, zeroMap,
std::nullopt);
ValueRange{});
rewriter.eraseOp(scatterOp);

return success();
Expand Down
26 changes: 24 additions & 2 deletions python/examples/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,19 @@ def empty_decorator(func):
def device(request):
return "cpu"


# this fixture is used for test_enable_fp_fusion
@pytest.fixture
def fresh_knobs():
from triton._internal_testing import _fresh_knobs_impl

fresh_function, reset_function = _fresh_knobs_impl()
try:
yield fresh_function()
finally:
reset_function()


# this fixture is used for test_trans_4d && test_trans_reshape
@pytest.fixture
def with_allocator():
Expand All @@ -32,7 +45,7 @@ def with_allocator():
triton.set_allocator(NullAllocator())


tests_supported = {
core_tests_supported = {
"test_store_eviction_policy",
"test_unary_op",
"test_umulhi",
Expand Down Expand Up @@ -77,6 +90,11 @@ def with_allocator():
"test_arange",
}

annotations_tests_supported = {
"test_int_annotation",
"test_unknown_annotation",
}


def pytest_collection_modifyitems(config, items):
skip_marker = pytest.mark.skip(reason="CPU backend does not support it yet")
Expand All @@ -89,7 +107,11 @@ def pytest_collection_modifyitems(config, items):
test_func_name = item.originalname if item.originalname else item.name

test_file = str(item.fspath)
if test_file.endswith("test_core.py") and test_func_name not in tests_supported:
if test_file.endswith("test_core.py") and test_func_name not in core_tests_supported:
item.add_marker(skip_marker)
continue

if test_file.endswith("test_annotations.py") and test_func_name not in annotations_tests_supported:
item.add_marker(skip_marker)
continue

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,7 @@ module {
%0 = tt.splat %arg0 : !tt.ptr<i32> -> tensor<1x!tt.ptr<i32>>
%1 = tt.load %0 : tensor<1x!tt.ptr<i32>>
%2 = arith.cmpi sgt, %1, %cst : tensor<1xi32>
%3 = "tt.reduce"(%2) <{axis = 0 : i32}> ({
^bb0(%arg1: i1, %arg2: i1):
tt.reduce.return %arg1 : i1
}) : (tensor<1xi1>) -> i1
%3 = tt.unsplat %2 : tensor<1xi1>
scf.if %3 {
tt.store %arg0, %c42_i32 : !tt.ptr<i32>
}
Expand Down
1 change: 1 addition & 0 deletions tools/triton-shared-opt/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ target_link_libraries(triton-shared-opt PRIVATE
# MLIR core
MLIROptLib
MLIRPass
MLIRRegisterAllPasses
MLIRTransforms
)

Expand Down
6 changes: 3 additions & 3 deletions tools/triton-shared-opt/RegisterTritonSharedDialects.h
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,7 @@ inline void registerTritonSharedDialects(mlir::DialectRegistry &registry) {
mlir::ttx::TritonTilingExtDialect, mlir::tts::TritonStructuredDialect,
mlir::triton::TritonDialect, mlir::cf::ControlFlowDialect,
mlir::math::MathDialect, mlir::arith::ArithDialect, mlir::scf::SCFDialect,
mlir::gpu::GPUDialect, mlir::linalg::LinalgDialect,
mlir::func::FuncDialect, mlir::tensor::TensorDialect,
mlir::memref::MemRefDialect, mlir::bufferization::BufferizationDialect>();
mlir::linalg::LinalgDialect, mlir::func::FuncDialect,
mlir::tensor::TensorDialect, mlir::memref::MemRefDialect,
mlir::bufferization::BufferizationDialect>();
}
2 changes: 1 addition & 1 deletion triton-hash.txt
Original file line number Diff line number Diff line change
@@ -1 +1 @@
ec8cb09329cf25ac241a7dee1eea5a5d94daef8a
e44bd1c83c1c3e8deac7c4f02683cfb3cc395c8b