Skip to content
Open
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
893 changes: 893 additions & 0 deletions docs/designs/ptoas-implicit-tmp-materialization-design.md

Large diffs are not rendered by default.

130 changes: 34 additions & 96 deletions include/PTO/IR/PTOOps.td
Original file line number Diff line number Diff line change
Expand Up @@ -861,16 +861,12 @@ def TTransOp : PTO_TOp<"ttrans", [
}];
let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);
let results = (outs);

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let hasVerifier = 1;

Expand Down Expand Up @@ -4070,7 +4066,7 @@ def TColArgMaxOp : PTO_TOp<"tcolargmax", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

Expand All @@ -4083,11 +4079,7 @@ def TColArgMaxOp : PTO_TOp<"tcolargmax", [
::mlir::MutableOperandRange getDpsInitsMutable() { return getDstMutable(); }
}];

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;
}

def TColMinOp : PTO_TOp<"tcolmin", [
Expand Down Expand Up @@ -4127,7 +4119,7 @@ def TColArgMinOp : PTO_TOp<"tcolargmin", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

Expand All @@ -4140,11 +4132,7 @@ def TColArgMinOp : PTO_TOp<"tcolargmin", [
::mlir::MutableOperandRange getDpsInitsMutable() { return getDstMutable(); }
}];

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;
}

def TColSumOp : PTO_TOp<"tcolsum", [
Expand Down Expand Up @@ -4214,6 +4202,7 @@ def TCvtOp : PTO_TOp<"tcvt", [

let arguments = (ins
PTODpsType:$src,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst,
DefaultValuedAttr<PTO_RoundModeAttr, "::mlir::pto::RoundMode::CAST_RINT">:$rmode,
DefaultValuedAttr<PTO_SaturationModeAttr, "::mlir::pto::SaturationMode::OFF">:$sat_mode
Expand Down Expand Up @@ -5149,6 +5138,7 @@ def TMrgSortOp: PTO_TOp<"tmrgsort", [
let extraClassDeclaration = [{
bool isFormat1() { return getSrcs().size() == 1u && getBlockLen() && getDsts().size() == 1u; }
bool isFormat2() { return getSrcs().size() >= 2u && getSrcs().size() <= 4u && getTmp() && getDsts().size() == 1u && getExcuted(); }
bool isFormat2WithoutTmp() { return getSrcs().size() >= 2u && getSrcs().size() <= 4u && !getTmp() && !getBlockLen() && getDsts().size() == 1u && getExcuted(); }
Value getSrc() { return getSrcs().front(); }
Value getDst() { return getDsts().front(); }
::mlir::MutableOperandRange getDpsInitsMutable() { return getDstsMutable(); }
Expand Down Expand Up @@ -5565,19 +5555,15 @@ def TPReluOp: PTO_TOp<"tprelu", [
let arguments = (ins
PTODpsType:$src0,
PTODpsType:$src1,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src0 `,` $src1 `,` $tmp `:` qualified(type($src0)) `,` qualified(type($src1)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand Down Expand Up @@ -5815,7 +5801,7 @@ def TRemOp: PTO_TOp<"trem", [
let arguments = (ins
PTODpsType:$src0,
PTODpsType:$src1,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst,
DefaultValuedAttr<PTO_RemPrecisionAttr, "::mlir::pto::RemPrecision::Default">:$precisionType
);
Expand All @@ -5824,11 +5810,7 @@ def TRemOp: PTO_TOp<"trem", [

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src0 `,` $src1 `,` $tmp `:` qualified(type($src0)) `,` qualified(type($src1)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -5849,19 +5831,15 @@ def TRemSOp: PTO_TOp<"trems", [
let arguments = (ins
PTODpsType:$src,
ScalarType:$scalar,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $scalar `,` $tmp `:` qualified(type($src)) `,` type($scalar) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand Down Expand Up @@ -6154,19 +6132,15 @@ def TRowMaxOp: PTO_TOp<"trowmax", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -6183,19 +6157,15 @@ def TRowArgMaxOp: PTO_TOp<"trowargmax", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -6215,19 +6185,15 @@ def TRowMinOp: PTO_TOp<"trowmin", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -6244,19 +6210,15 @@ def TRowArgMinOp: PTO_TOp<"trowargmin", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -6276,19 +6238,15 @@ def TRowSumOp: PTO_TOp<"trowsum", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -6305,19 +6263,15 @@ def TRowProdOp: PTO_TOp<"trowprod", [

let arguments = (ins
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $tmp `:` qualified(type($src)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand Down Expand Up @@ -6433,19 +6387,15 @@ def TSelOp: PTO_TOp<"tsel", [
PTODpsType:$mask,
PTODpsType:$src0,
PTODpsType:$src1,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $mask `,` $src0 `,` $src1 `,` $tmp `:` qualified(type($mask)) `,` qualified(type($src0)) `,` qualified(type($src1)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -6469,7 +6419,7 @@ def TSelSOp: PTO_TOp<"tsels", [
let arguments = (ins
PTODpsType:$mask,
PTODpsType:$src,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
ScalarType:$scalar,
PTODpsType:$dst
);
Expand All @@ -6478,11 +6428,7 @@ def TSelSOp: PTO_TOp<"tsels", [

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $mask `,` $src `,` $tmp `,` $scalar `:` qualified(type($mask)) `,` qualified(type($src)) `,` qualified(type($tmp)) `,` type($scalar) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand Down Expand Up @@ -6857,19 +6803,15 @@ def TXorSOp: PTO_TOp<"txors", [
let arguments = (ins
PTODpsType:$src,
AnySignlessInteger:$scalar,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src `,` $scalar `,` $tmp `:` qualified(type($src)) `,` type($scalar) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand All @@ -6896,19 +6838,15 @@ def TXorOp: PTO_TOp<"txor", [
let arguments = (ins
PTODpsType:$src0,
PTODpsType:$src1,
PTODpsType:$tmp,
Optional<PTODpsType>:$tmp,
PTODpsType:$dst
);

let results = (outs);

let hasVerifier = 1;

let assemblyFormat = [{
`ins` `(` $src0 `,` $src1 `,` $tmp `:` qualified(type($src0)) `,` qualified(type($src1)) `,` qualified(type($tmp)) `)`
`outs` `(` $dst `:` qualified(type($dst) ) `)`
attr-dict
}];
let hasCustomAssemblyFormat = 1;

let extraClassDeclaration = [{
::mlir::pto::PIPE getPipe() { return ::mlir::pto::PIPE::PIPE_V; }
Expand Down
2 changes: 2 additions & 0 deletions include/PTO/Transforms/Passes.h
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,8 @@ createPlanMemoryModernPass(const PlanMemoryOptions &options);
std::unique_ptr<Pass> createPTORemoveRedundantBarrierPass();
std::unique_ptr<Pass> createPTOValidateIntToPtrUsesPass();
std::unique_ptr<Pass> createPTORematerializeFixpipeVectorQuantPass();
std::unique_ptr<Pass>
createPTOMaterializeImplicitTmpPass(bool requireExplicitTmp = false);
std::unique_ptr<Pass> createPTOResolveBufferSelectPass();
std::unique_ptr<Pass> createInferPTOLayoutPass();
std::unique_ptr<Pass> createPTOA5NormalizeTMovPass();
Expand Down
16 changes: 16 additions & 0 deletions include/PTO/Transforms/Passes.td
Original file line number Diff line number Diff line change
Expand Up @@ -192,6 +192,22 @@ def PTORematerializeFixpipeVectorQuant
let dependentDialects = ["mlir::pto::PTODialect", "mlir::func::FuncDialect"];
}

def PTOMaterializeImplicitTmp
: Pass<"pto-materialize-implicit-tmp", "func::FuncOp"> {
let summary = "Materialize implicit tmp tiles for PTO ops before memplan";
let description = [{
Rewrites PTO ops with optional tmp operands into explicit tmp forms when
the backend tmp-aware overload is required. The synthesized tmp is emitted
as tile-native `pto.alloc_tile(no addr)` so PlanMemory can assign its local
address together with other tile buffers.
}];
let constructor = "mlir::pto::createPTOMaterializeImplicitTmpPass()";
let dependentDialects = [
"mlir::pto::PTODialect",
"mlir::func::FuncDialect"
];
}

def PlanMemory : Pass<"pto-plan-memory", "ModuleOp"> {
let summary = "Plan memory for PTO Ops";
let constructor = "mlir::pto::createPlanMemoryPass()";
Expand Down
Loading
Loading