Skip to content
Draft
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
166 changes: 166 additions & 0 deletions src/enzyme_ad/jax/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -577,6 +577,105 @@ td_library(
],
)

td_library(
name = "AxisDialectFiles",
srcs = [
"Dialect/Axis/Dialect.td",
"Dialect/Axis/Interfaces.td",
"Dialect/Axis/Ops.td",
"Dialect/Axis/Types.td",
],
includes = ["."],
deps = [
"@llvm-project//mlir:BuiltinDialectTdFiles",
"@llvm-project//mlir:InferTypeOpInterfaceTdFiles",
"@llvm-project//mlir:OpBaseTdFiles",
],
)

gentbl_cc_library(
name = "AxisDialectIncGen",
tbl_outs = [
(
[
"-gen-dialect-decls",
"-dialect=axis",
],
"Dialect/Axis/AxisDialect.h.inc",
),
(
[
"-gen-dialect-defs",
"-dialect=axis",
],
"Dialect/Axis/AxisDialect.cpp.inc",
),
],
tblgen = "@llvm-project//mlir:mlir-tblgen",
td_file = "Dialect/Axis/Dialect.td",
deps = [
":AxisDialectFiles",
],
)

gentbl_cc_library(
name = "AxisOpsIncGen",
tbl_outs = [
(
["-gen-op-decls"],
"Dialect/Axis/AxisOps.h.inc",
),
(
["-gen-op-defs"],
"Dialect/Axis/AxisOps.cpp.inc",
),
],
tblgen = "@llvm-project//mlir:mlir-tblgen",
td_file = "Dialect/Axis/Ops.td",
deps = [
":AxisDialectFiles",
"@llvm-project//mlir:InferTypeOpInterface",
],
)

gentbl_cc_library(
name = "AxisTypesIncGen",
tbl_outs = [
(
["--gen-typedef-decls"],
"Dialect/Axis/AxisTypes.h.inc",
),
(
["--gen-typedef-defs"],
"Dialect/Axis/AxisTypes.cpp.inc",
),
],
tblgen = "@llvm-project//mlir:mlir-tblgen",
td_file = "Dialect/Axis/Types.td",
deps = [
":AxisDialectFiles",
],
)

gentbl_cc_library(
name = "AxisTypeInterfacesIncGen",
tbl_outs = [
(
["--gen-type-interface-decls"],
"Dialect/Axis/AxisTypeInterfaces.h.inc",
),
(
["--gen-type-interface-defs"],
"Dialect/Axis/AxisTypeInterfaces.cpp.inc",
),
],
tblgen = "@llvm-project//mlir:mlir-tblgen",
td_file = "Dialect/Axis/Interfaces.td",
deps = [
":AxisDialectFiles",
],
)

td_library(
name = "DistributedDialectFiles",
srcs = [
Expand All @@ -585,9 +684,17 @@ td_library(
"Dialect/Distributed/Ops.td",
"Dialect/Distributed/Types.td",
],
includes = [
".",
"../../../external/shardy",
],
deps = [
":AxisDialectFiles",
"@llvm-project//mlir:BuiltinDialectTdFiles",
"@llvm-project//mlir:InferTypeOpInterfaceTdFiles",
"@llvm-project//mlir:OpBaseTdFiles",
"@shardy//shardy/dialect/sdy/ir:sdy_td_files",
"@stablehlo//:stablehlo_ops_td_files",
],
)

Expand Down Expand Up @@ -632,6 +739,8 @@ gentbl_cc_library(
td_file = "Dialect/Distributed/Ops.td",
deps = [
":DistributedDialectFiles",
"@llvm-project//mlir:InferTypeOpInterface",
"@shardy//shardy/dialect/sdy/ir:sdy_td_files",
],
)

Expand Down Expand Up @@ -673,6 +782,25 @@ gentbl_cc_library(
],
)

gentbl_cc_library(
name = "DistributedTypeInterfacesIncGen",
tbl_outs = [
(
["--gen-type-interface-decls"],
"Dialect/Distributed/DistributedTypeInterfaces.h.inc",
),
(
["--gen-type-interface-defs"],
"Dialect/Distributed/DistributedTypeInterfaces.cpp.inc",
),
],
tblgen = "@llvm-project//mlir:mlir-tblgen",
td_file = "Dialect/Distributed/Interfaces.td",
deps = [
":DistributedDialectFiles",
],
)

td_library(
name = "TesseraDialectTdFiles",
srcs = [
Expand Down Expand Up @@ -824,6 +952,33 @@ gentbl_cc_library(
],
)

td_library(
name = "DistributedPassesTdFiles",
srcs = [],
deps = [
"@llvm-project//mlir:PassBaseTdFiles",
],
)

gentbl_cc_library(
name = "DistributedPassesIncGen",
tbl_outs = [
(
[
"-gen-pass-decls",
"-name=distributed",
],
"Passes/Distributed/Passes.h.inc",
),
],
tblgen = "@llvm-project//mlir:mlir-tblgen",
td_file = "Passes/Distributed/Passes.td",
deps = [
":DistributedPassesTdFiles",
"@llvm-project//mlir:FunctionInterfacesTdFiles",
],
)

cc_library(
name = "CheckedRewrite",
hdrs = ["CheckedRewrite.h"],
Expand Down Expand Up @@ -1195,8 +1350,10 @@ cc_library(
"Analysis/*.cpp",
"Implementations/*.cpp",
"Passes/*.cpp",
"Passes/Distributed/*.cpp",
"Passes/Tessera/*.cpp",
"Dialect/*.cpp",
"Dialect/Axis/*.cpp",
"Dialect/Distributed/*.cpp",
"Dialect/Tessera/*.cpp",
"Dialect/Perfify/*.cpp",
Expand All @@ -1213,8 +1370,10 @@ cc_library(
"Analysis/*.h",
"Implementations/*.h",
"Passes/*.h",
"Passes/Distributed/*.h",
"Passes/Tessera/*.h",
"Dialect/*.h",
"Dialect/Axis/*.h",
"Dialect/Distributed/*.h",
"Dialect/Tessera/*.h",
"Dialect/Perfify/*.h",
Expand All @@ -1232,9 +1391,16 @@ cc_library(
],
visibility = ["//visibility:public"],
deps = [
":AxisDialectIncGen",
":AxisOpsIncGen",
":AxisTypeInterfacesIncGen",
":AxisTypesIncGen",
":CheckedRewrite",
":DistributedDialectIncGen",
":DistributedInterfacesIncGen",
":DistributedOpsIncGen",
":DistributedPassesIncGen",
":DistributedTypeInterfacesIncGen",
":DistributedTypesIncGen",
":EnzymeHLOPatternsIncGen",
":EnzymeXLAAttrsIncGen",
Expand Down
22 changes: 22 additions & 0 deletions src/enzyme_ad/jax/Dialect/Axis/Dialect.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
#include "Dialect.h"

#include "llvm/ADT/TypeSwitch.h"

// Include the .cpp.inc files
#include "src/enzyme_ad/jax/Dialect/Axis/AxisDialect.cpp.inc"

#define GET_TYPEDEF_CLASSES
#include "src/enzyme_ad/jax/Dialect/Axis/AxisTypes.cpp.inc"

#include "src/enzyme_ad/jax/Dialect/Axis/AxisTypeInterfaces.cpp.inc"

void mlir::enzyme::axis::AxisDialect::initialize() {
addTypes<
#define GET_TYPEDEF_LIST
#include "src/enzyme_ad/jax/Dialect/Axis/AxisTypes.cpp.inc"
>();
addOperations<
#define GET_OP_LIST
#include "src/enzyme_ad/jax/Dialect/Axis/AxisOps.cpp.inc"
>();
}
27 changes: 27 additions & 0 deletions src/enzyme_ad/jax/Dialect/Axis/Dialect.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
#ifndef ENZYME_AD_JAX_DIALECT_AXIS_DIALECT_H
#define ENZYME_AD_JAX_DIALECT_AXIS_DIALECT_H

#include "mlir/Bytecode/BytecodeOpInterface.h"
#include "mlir/IR/Attributes.h"
#include "mlir/IR/Dialect.h"
#include "mlir/IR/DialectImplementation.h"
#include "mlir/IR/OpDefinition.h"
#include "mlir/IR/Types.h"
#include "mlir/Interfaces/InferTypeOpInterface.h"
#include "mlir/Interfaces/SideEffectInterfaces.h"
#include "mlir/Support/LLVM.h"

#include "Traits.h"

// Include the dialect
#include "src/enzyme_ad/jax/Dialect/Axis/AxisDialect.h.inc"
// Type interfaces
#include "src/enzyme_ad/jax/Dialect/Axis/AxisTypeInterfaces.h.inc"
// Types
#define GET_TYPEDEF_CLASSES
#include "src/enzyme_ad/jax/Dialect/Axis/AxisTypes.h.inc"
// Ops
#define GET_OP_CLASSES
#include "src/enzyme_ad/jax/Dialect/Axis/AxisOps.h.inc"

#endif // ENZYME_AD_JAX_DIALECT_AXIS_DIALECT_H
29 changes: 29 additions & 0 deletions src/enzyme_ad/jax/Dialect/Axis/Dialect.td
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
#ifndef ENZYME_AD_JAX_DIALECT_AXIS_DIALECT_TD
#define ENZYME_AD_JAX_DIALECT_AXIS_DIALECT_TD

include "mlir/IR/AttrTypeBase.td"
include "mlir/IR/DialectBase.td"
include "mlir/IR/OpBase.td"
include "mlir/IR/Traits.td"
include "mlir/Interfaces/SideEffectInterfaces.td"

def AxisDialect : Dialect {
let name = "axis";
let description = [{SSA algebra over canonical axes and axis factors.}];
let cppNamespace = "::mlir::enzyme::axis";
let useDefaultTypePrinterParser = 1;
}

class AxisDialectType<string name, string type_mnemonic, list<Trait> traits = [
]> : TypeDef<AxisDialect, name, traits> {
let mnemonic = type_mnemonic;
}

class AxisOp<string mnemonic, list<Trait> traits = []>
: Op<AxisDialect, mnemonic, traits>;

// Marker trait for static metadata computations. This implies Pure.
def MetadataTrait : NativeOpTrait<"enzyme::axis::MetadataTrait">;
def MetadataOpTrait : TraitList<[Pure, MetadataTrait]>;

#endif // ENZYME_AD_JAX_DIALECT_AXIS_DIALECT_TD
28 changes: 28 additions & 0 deletions src/enzyme_ad/jax/Dialect/Axis/Interfaces.td
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
#ifndef ENZYME_AD_JAX_DIALECT_AXIS_INTERFACES_TD
#define ENZYME_AD_JAX_DIALECT_AXIS_INTERFACES_TD

include "mlir/IR/OpBase.td"

def AxisTypeInterface : TypeInterface<"AxisTypeInterface"> {
let cppNamespace = "::mlir::enzyme::axis";
let description = [{
Interface for canonical axis types with statically decidable equivalence.
}];

let methods = [
InterfaceMethod<
[{Returns the compile-time extent of this canonical axis.}],
"unsigned", "extent", (ins), "",
[{
return $_type.getExtent();
}]>,
InterfaceMethod<
[{Returns true when two SSA values of this axis type are equivalent.
}],
"bool", "aliases",
(ins "::mlir::Value":$ax1, "::mlir::Value":$ax2)
>
];
}

#endif // ENZYME_AD_JAX_DIALECT_AXIS_INTERFACES_TD
Loading
Loading