Skip to content

Commit 97b78d6

Browse files
authored
[MLIR][NVVM] Fix predicate operand index in BasicPtxBuilderInterface (#189552)
Predicate index computation was incorrect, it was not counting write/readwrite symbols. Wrong case ``` // CHECK: %{{.*}} = llvm.inline_asm has_side_effects asm_dialect = att "@$1 ex2.approx.ftz.f32 $0, $1;", "=f,f,b" %{{.*}}, %{{.*}} : (f32, i1) -> f32 %1 = nvvm.inline_ptx "ex2.approx.ftz.f32 {$w0}, {$r0};" ro (%input : f32), predicate = %pred -> f32 ``` PR fixes, predicate index became `@$2` ``` // CHECK: %{{.*}} = llvm.inline_asm has_side_effects asm_dialect = att "@$2 ex2.approx.ftz.f32 $0, $1;", "=f,f,b" %{{.*}}, %{{.*}} : (f32, i1) -> f32 %1 = nvvm.inline_ptx "ex2.approx.ftz.f32 {$w0}, {$r0};" ro (%input : f32), predicate = %pred -> f32 ```
1 parent 66c45a3 commit 97b78d6

2 files changed

Lines changed: 87 additions & 2 deletions

File tree

‎mlir/lib/Dialect/LLVMIR/IR/BasicPtxBuilderInterface.cpp‎

Lines changed: 27 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -423,6 +423,31 @@ static std::string rewriteAsmPlaceholders(llvm::StringRef ptxCode) {
423423
return out;
424424
}
425425

426+
/// Return the constraint index of the predicate operand. The predicate
427+
/// constraint ("b") is always the last non-tied token in the canonicalized
428+
/// constraint string. Tied constraints (digit-only tokens from read-write
429+
/// canonicalization) are appended at the end, so we walk backwards to skip
430+
/// them.
431+
static unsigned getPredicateConstraintIndex(StringRef constraints) {
432+
SmallVector<StringRef> tokens;
433+
constraints.split(tokens, ',');
434+
assert(!tokens.empty() && "expected at least a predicate constraint");
435+
436+
auto isTiedConstraint = [](StringRef tok) {
437+
unsigned idx;
438+
return !tok.trim().getAsInteger(10, idx);
439+
};
440+
441+
size_t numTied = 0;
442+
for (StringRef tok : llvm::reverse(tokens)) {
443+
if (!isTiedConstraint(tok))
444+
break;
445+
++numTied;
446+
}
447+
assert(numTied < tokens.size() && "all constraints are tied");
448+
return tokens.size() - numTied - 1;
449+
}
450+
426451
LLVM::InlineAsmOp PtxBuilder::build() {
427452
auto asmDialectAttr = LLVM::AsmDialectAttr::get(interfaceOp->getContext(),
428453
LLVM::AsmDialect::AD_ATT);
@@ -443,8 +468,9 @@ LLVM::InlineAsmOp PtxBuilder::build() {
443468
// Add the predicate to the asm string.
444469
if (interfaceOp.getPredicate().has_value() &&
445470
interfaceOp.getPredicate().value()) {
471+
unsigned predIdx = getPredicateConstraintIndex(registerConstraints);
446472
std::string predicateStr = "@%";
447-
predicateStr += std::to_string((ptxOperands.size() - 1));
473+
predicateStr += std::to_string(predIdx);
448474
ptxInstruction = predicateStr + " " + ptxInstruction;
449475
}
450476

‎mlir/test/Conversion/NVVMToLLVM/nvvm-to-llvm.mlir‎

Lines changed: 60 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -614,11 +614,25 @@ llvm.func @ex2(%input : f32, %pred : i1) {
614614
// CHECK: %{{.*}} = llvm.inline_asm has_side_effects asm_dialect = att "ex2.approx.ftz.f32 $0, $1;", "=f,f" %{{.*}} : (f32) -> f32
615615
%0 = nvvm.inline_ptx "ex2.approx.ftz.f32 {$w0}, {$r0};" ro (%input : f32) -> f32
616616

617-
// CHECK: %{{.*}} = llvm.inline_asm has_side_effects asm_dialect = att "@$1 ex2.approx.ftz.f32 $0, $1;", "=f,f,b" %{{.*}}, %{{.*}} : (f32, i1) -> f32
617+
// CHECK: %{{.*}} = llvm.inline_asm has_side_effects asm_dialect = att "@$2 ex2.approx.ftz.f32 $0, $1;", "=f,f,b" %{{.*}}, %{{.*}} : (f32, i1) -> f32
618618
%1 = nvvm.inline_ptx "ex2.approx.ftz.f32 {$w0}, {$r0};" ro (%input : f32), predicate = %pred -> f32
619619
llvm.return
620620
}
621621

622+
// CHECK-LABEL: @multi_return_pred(
623+
// CHECK-SAME: %[[arg0:[a-zA-Z0-9_]+]]: i32, %[[arg1:[a-zA-Z0-9_]+]]: i32, %[[pred:[a-zA-Z0-9_]+]]: i1)
624+
llvm.func @multi_return_pred(%a : i32, %b : i32, %pred : i1) -> i32 {
625+
// CHECK: %[[S1:.+]] = llvm.inline_asm has_side_effects asm_dialect = att "@$4 {.reg .pred p; setp.ge.s32 p, $2, $3; selp.s32 $0, $2,$3, p; selp.s32 $1, $2,$3, p;}", "=r,=r,r,r,b" %[[arg0]], %[[arg1]], %[[pred]] : (i32, i32, i1) -> !llvm.struct<(i32, i32)>
626+
// CHECK: %[[S2:.+]] = llvm.extractvalue %[[S1]][0] : !llvm.struct<(i32, i32)>
627+
// CHECK: %[[S3:.+]] = llvm.extractvalue %[[S1]][1] : !llvm.struct<(i32, i32)>
628+
// CHECK: %[[S4:.+]] = llvm.add %[[S2]], %[[S3]] : i32
629+
// CHECK: llvm.return %[[S4]] : i32
630+
%r1, %r2 = nvvm.inline_ptx "{.reg .pred p; setp.ge.s32 p, {$r0}, {$r1}; selp.s32 {$w0}, {$r0},{$r1}, p; selp.s32 {$w1}, {$r0},{$r1}, p;}"
631+
ro (%a, %b : i32,i32), predicate = %pred -> i32,i32
632+
%r3 = llvm.add %r1, %r2 : i32
633+
llvm.return %r3 : i32
634+
}
635+
622636
// CHECK-LABEL: @multi_return(
623637
// CHECK-SAME: %[[arg0:[a-zA-Z0-9_]+]]: i32, %[[arg1:[a-zA-Z0-9_]+]]: i32)
624638
llvm.func @multi_return(%a : i32, %b : i32) -> i32 {
@@ -651,6 +665,24 @@ llvm.func @inline_ptx_multi_rw(%a : i32, %b : i32, %rw_c : f32, %rw_d : f32) ->
651665
llvm.return %r4 : f32
652666
}
653667

668+
// CHECK-LABEL: @inline_ptx_multi_rw_pred(
669+
// CHECK-SAME: %[[arg0:[a-zA-Z0-9_]+]]: i32, %[[arg1:[a-zA-Z0-9_]+]]: i32, %[[arg2:[a-zA-Z0-9_]+]]: f32, %[[arg3:[a-zA-Z0-9_]+]]: f32, %[[pred:[a-zA-Z0-9_]+]]: i1)
670+
llvm.func @inline_ptx_multi_rw_pred(%a : i32, %b : i32, %rw_c : f32, %rw_d : f32, %pred : i1) -> f32 {
671+
// CHECK: %[[S0:.+]] = llvm.inline_asm has_side_effects asm_dialect = att "@$4 {.reg .pred p; setp.ge.s32 p, $2, $3; selp.s32 $0, $2,$3, p; selp.s32 $1, $2,$3, p;}",
672+
// CHECK-SAME: "=f,=f,r,r,b,0,1"
673+
// CHECK-SAME: %[[arg2]], %[[arg3]], %[[arg0]], %[[arg1]], %[[pred]]
674+
// CHECK-SAME: : (f32, f32, i32, i32, i1) -> !llvm.struct<(f32, f32)>
675+
// CHECK: %[[S1:.+]] = llvm.extractvalue %[[S0]][0] : !llvm.struct<(f32, f32)>
676+
// CHECK: %[[S2:.+]] = llvm.extractvalue %[[S0]][1] : !llvm.struct<(f32, f32)>
677+
// CHECK: %[[S3:.+]] = llvm.fadd %[[S1]], %[[S2]] : f32
678+
// CHECK: llvm.return %[[S3]] : f32
679+
nvvm.inline_ptx "{.reg .pred p; setp.ge.s32 p, {$r0}, {$r1}; selp.s32 {$rw0}, {$r0},{$r1}, p; selp.s32 {$rw1}, {$r0},{$r1}, p;}"
680+
ro (%a, %b : i32,i32)
681+
rw (%rw_c, %rw_d: f32,f32), predicate = %pred
682+
%r4 = llvm.fadd %rw_c, %rw_d : f32
683+
llvm.return %r4 : f32
684+
}
685+
654686
// CHECK-LABEL: @inline_ptx_multi_rw_r(
655687
// CHECK-SAME: %[[arg0:[a-zA-Z0-9_]+]]: i32, %[[arg1:[a-zA-Z0-9_]+]]: i32, %[[arg2:[a-zA-Z0-9_]+]]: f32, %[[arg3:[a-zA-Z0-9_]+]]: f32)
656688
llvm.func @inline_ptx_multi_rw_r(%a : i32, %b : i32, %rw_c : f32, %rw_d : f32) -> f32 {
@@ -678,6 +710,33 @@ llvm.func @inline_ptx_multi_rw_r(%a : i32, %b : i32, %rw_c : f32, %rw_d : f32)
678710
llvm.return %r5 : f32
679711
}
680712

713+
// CHECK-LABEL: @inline_ptx_multi_rw_r_pred(
714+
// CHECK-SAME: %[[arg0:[a-zA-Z0-9_]+]]: i32, %[[arg1:[a-zA-Z0-9_]+]]: i32, %[[arg2:[a-zA-Z0-9_]+]]: f32, %[[arg3:[a-zA-Z0-9_]+]]: f32, %[[pred:[a-zA-Z0-9_]+]]: i1)
715+
llvm.func @inline_ptx_multi_rw_r_pred(%a : i32, %b : i32, %rw_c : f32, %rw_d : f32, %pred : i1) -> f32 {
716+
// CHECK: %[[S0:.+]] = llvm.inline_asm has_side_effects asm_dialect = att "@$6 {.reg .pred p; setp.ge.s32 p, $4, $5; selp.s32 $0, $4,$5, p; selp.s32 $1, $4,$5, p; selp.s32 $2, $4,$5, p; selp.s32 $3, $4,$5, p;}",
717+
// CHECK-SAME: "=f,=f,=r,=r,r,r,b,0,1"
718+
// CHECK-SAME: %[[arg2]], %[[arg3]], %[[arg0]], %[[arg1]], %[[pred]] :
719+
// CHECK-SAME: (f32, f32, i32, i32, i1) -> !llvm.struct<(f32, f32, i32, i32)>
720+
// CHECK: %[[S1:.+]] = llvm.extractvalue %[[S0]][0] : !llvm.struct<(f32, f32, i32, i32)>
721+
// CHECK: %[[S2:.+]] = llvm.extractvalue %[[S0]][1] : !llvm.struct<(f32, f32, i32, i32)>
722+
// CHECK: %[[S3:.+]] = llvm.extractvalue %[[S0]][2] : !llvm.struct<(f32, f32, i32, i32)>
723+
// CHECK: %[[S4:.+]] = llvm.extractvalue %[[S0]][3] : !llvm.struct<(f32, f32, i32, i32)>
724+
// CHECK: %[[S5:.+]] = llvm.add %[[S3]], %[[S4]] : i32
725+
// CHECK: %[[S6:.+]] = llvm.sitofp %[[S5]] : i32 to f32
726+
// CHECK: %[[S7:.+]] = llvm.fadd %[[S1]], %[[S2]] : f32
727+
// CHECK: %[[S8:.+]] = llvm.fadd %[[S6]], %[[S2]] : f32
728+
// CHECK: llvm.return %[[S8]] : f32
729+
730+
%wo0, %wo1 = nvvm.inline_ptx "{.reg .pred p; setp.ge.s32 p, {$r0}, {$r1}; selp.s32 {$rw0}, {$r0},{$r1}, p; selp.s32 {$rw1}, {$r0},{$r1}, p; selp.s32 {$w0}, {$r0},{$r1}, p; selp.s32 {$w1}, {$r0},{$r1}, p;}"
731+
ro (%a, %b : i32,i32)
732+
rw (%rw_c, %rw_d: f32,f32), predicate = %pred -> i32,i32
733+
%r3 = llvm.add %wo0, %wo1 : i32
734+
%r3f = llvm.sitofp %r3 : i32 to f32
735+
%r4 = llvm.fadd %rw_c, %rw_d : f32
736+
%r5 = llvm.fadd %r3f, %rw_d : f32
737+
llvm.return %r5 : f32
738+
}
739+
681740
// -----
682741

683742
llvm.func @inline_ptx_pack_4i8(%src : vector<4xi8>, %mask : i32, %zero: i32) {

0 commit comments

Comments
 (0)