diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp index dfda1157a720e..3dba45d0af32a 100644 --- a/llvm/lib/Target/X86/X86ISelLowering.cpp +++ b/llvm/lib/Target/X86/X86ISelLowering.cpp @@ -48564,21 +48564,44 @@ static SDValue commuteSelect(SDNode *N, SelectionDAG &DAG, const SDLoc &DL, ISD::CondCode CC; SDValue Cond, X, Y, LHS, RHS; - if (!sd_match(N, m_VSelect(m_AllOf(m_Value(Cond), - m_OneUse(m_SetCC(m_Value(X), m_Value(Y), - m_CondCode(CC)))), - m_Value(LHS), m_Value(RHS)))) + if (!sd_match( + N, m_VSelect(m_AllOf(m_Value(Cond), + m_SetCC(m_Value(X), m_Value(Y), m_CondCode(CC))), + m_Value(LHS), m_Value(RHS)))) return SDValue(); if (canCombineAsMaskOperation(LHS, Subtarget) || !canCombineAsMaskOperation(RHS, Subtarget)) return SDValue(); + // For multi-use setcc, check that all users are vselects that benefit. + bool CondHasOneUse = Cond.hasOneUse(); + if (!CondHasOneUse) { + if (!llvm::all_of(Cond->users(), [&](SDNode *User) { + SDValue UserLHS, UserRHS; + return sd_match(User, m_VSelect(m_Specific(Cond), m_Value(UserLHS), + m_Value(UserRHS))) && + !canCombineAsMaskOperation(UserLHS, Subtarget) && + canCombineAsMaskOperation(UserRHS, Subtarget); + })) + return SDValue(); + } + // Commute LHS and RHS to create opportunity to select mask instruction. // (vselect M, L, R) -> (vselect ~M, R, L) ISD::CondCode NewCC = ISD::getSetCCInverse(CC, X.getValueType()); - Cond = DAG.getSetCC(SDLoc(Cond), Cond.getValueType(), X, Y, NewCC); - return DAG.getSelect(DL, LHS.getValueType(), Cond, RHS, LHS); + SDValue NewCond = DAG.getSetCC(SDLoc(Cond), Cond.getValueType(), X, Y, NewCC); + if (CondHasOneUse) + return DAG.getSelect(DL, LHS.getValueType(), NewCond, RHS, LHS); + + // Invert the setcc for all users and commute all vselects. + DAG.ReplaceAllUsesOfValueWith(Cond, NewCond); + for (SDNode *User : NewCond->users()) { + SDValue UserLHS = User->getOperand(1); + SDValue UserRHS = User->getOperand(2); + DAG.UpdateNodeOperands(User, NewCond, UserRHS, UserLHS); + } + return SDValue(N, 0); } /// Do target-specific dag combines on SELECT and VSELECT nodes. diff --git a/llvm/test/CodeGen/X86/avx512-masked-op-fusion.ll b/llvm/test/CodeGen/X86/avx512-masked-op-fusion.ll index 9280366c78b9c..b11ad0df8bbe8 100644 --- a/llvm/test/CodeGen/X86/avx512-masked-op-fusion.ll +++ b/llvm/test/CodeGen/X86/avx512-masked-op-fusion.ll @@ -8,27 +8,23 @@ define void @masked_min_max(ptr %pSrc, ptr %pMsk, i64 %n, ptr %pMin, ptr %pMax) { ; CHECK-LABEL: masked_min_max: ; CHECK: # %bb.0: # %entry -; CHECK-NEXT: vbroadcastss {{.*#+}} zmm1 = [-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf] -; CHECK-NEXT: vbroadcastss {{.*#+}} zmm0 = [+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf] +; CHECK-NEXT: vbroadcastss {{.*#+}} zmm0 = [-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf,-Inf] +; CHECK-NEXT: vbroadcastss {{.*#+}} zmm1 = [+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf,+Inf] ; CHECK-NEXT: xorl %eax, %eax ; CHECK-NEXT: .p2align 4 ; CHECK-NEXT: .LBB0_1: # %loop ; CHECK-NEXT: # =>This Inner Loop Header: Depth=1 -; CHECK-NEXT: vmovaps %zmm1, %zmm2 -; CHECK-NEXT: vmovaps %zmm0, %zmm1 -; CHECK-NEXT: vmovdqu (%rsi,%rax), %xmm0 -; CHECK-NEXT: vptestnmb %xmm0, %xmm0, %k1 -; CHECK-NEXT: vmovups (%rdi,%rax,4), %zmm3 -; CHECK-NEXT: vminps %zmm3, %zmm1, %zmm0 -; CHECK-NEXT: vmovaps %zmm1, %zmm0 {%k1} -; CHECK-NEXT: vmaxps %zmm3, %zmm2, %zmm1 -; CHECK-NEXT: vmovaps %zmm2, %zmm1 {%k1} +; CHECK-NEXT: vmovdqu (%rsi,%rax), %xmm2 +; CHECK-NEXT: vptestmb %xmm2, %xmm2, %k1 +; CHECK-NEXT: vmovups (%rdi,%rax,4), %zmm2 +; CHECK-NEXT: vminps %zmm2, %zmm1, %zmm1 {%k1} +; CHECK-NEXT: vmaxps %zmm2, %zmm0, %zmm0 {%k1} ; CHECK-NEXT: addq $16, %rax ; CHECK-NEXT: cmpq %rdx, %rax ; CHECK-NEXT: jb .LBB0_1 ; CHECK-NEXT: # %bb.2: # %exit -; CHECK-NEXT: vmovaps %zmm0, (%rcx) -; CHECK-NEXT: vmovaps %zmm1, (%r8) +; CHECK-NEXT: vmovaps %zmm1, (%rcx) +; CHECK-NEXT: vmovaps %zmm0, (%r8) ; CHECK-NEXT: vzeroupper ; CHECK-NEXT: retq entry: