1//===- DAGCombiner.cpp - Implement a DAG node combiner --------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This pass combines dag nodes to form fewer, simpler DAG nodes. It can be run
10// both before and after the DAG is legalized.
11//
12// This pass is not a substitute for the LLVM IR instcombine pass. This pass is
13// primarily intended to handle simplification opportunities that are implicit
14// in the LLVM IR and exposed by the various codegen lowering phases.
15//
16//===----------------------------------------------------------------------===//
17
18#include "llvm/ADT/APFloat.h"
19#include "llvm/ADT/APInt.h"
20#include "llvm/ADT/APSInt.h"
21#include "llvm/ADT/ArrayRef.h"
22#include "llvm/ADT/DenseMap.h"
23#include "llvm/ADT/IntervalMap.h"
24#include "llvm/ADT/STLExtras.h"
25#include "llvm/ADT/SetVector.h"
26#include "llvm/ADT/SmallBitVector.h"
27#include "llvm/ADT/SmallPtrSet.h"
28#include "llvm/ADT/SmallSet.h"
29#include "llvm/ADT/SmallVector.h"
30#include "llvm/ADT/Statistic.h"
31#include "llvm/Analysis/AliasAnalysis.h"
32#include "llvm/Analysis/MemoryLocation.h"
33#include "llvm/Analysis/TargetLibraryInfo.h"
34#include "llvm/Analysis/ValueTracking.h"
35#include "llvm/Analysis/VectorUtils.h"
36#include "llvm/CodeGen/ByteProvider.h"
37#include "llvm/CodeGen/DAGCombine.h"
38#include "llvm/CodeGen/ISDOpcodes.h"
39#include "llvm/CodeGen/MachineFrameInfo.h"
40#include "llvm/CodeGen/MachineFunction.h"
41#include "llvm/CodeGen/MachineMemOperand.h"
42#include "llvm/CodeGen/SDPatternMatch.h"
43#include "llvm/CodeGen/SelectionDAG.h"
44#include "llvm/CodeGen/SelectionDAGAddressAnalysis.h"
45#include "llvm/CodeGen/SelectionDAGNodes.h"
46#include "llvm/CodeGen/SelectionDAGTargetInfo.h"
47#include "llvm/CodeGen/TargetLowering.h"
48#include "llvm/CodeGen/TargetRegisterInfo.h"
49#include "llvm/CodeGen/TargetSubtargetInfo.h"
50#include "llvm/CodeGen/ValueTypes.h"
51#include "llvm/CodeGenTypes/MachineValueType.h"
52#include "llvm/IR/Attributes.h"
53#include "llvm/IR/Constant.h"
54#include "llvm/IR/DataLayout.h"
55#include "llvm/IR/DebugInfoMetadata.h"
56#include "llvm/IR/DerivedTypes.h"
57#include "llvm/IR/Function.h"
58#include "llvm/IR/Metadata.h"
59#include "llvm/Support/Casting.h"
60#include "llvm/Support/CodeGen.h"
61#include "llvm/Support/CommandLine.h"
62#include "llvm/Support/Compiler.h"
63#include "llvm/Support/Debug.h"
64#include "llvm/Support/DebugCounter.h"
65#include "llvm/Support/ErrorHandling.h"
66#include "llvm/Support/KnownBits.h"
67#include "llvm/Support/MathExtras.h"
68#include "llvm/Support/raw_ostream.h"
69#include "llvm/Target/TargetMachine.h"
70#include "llvm/Target/TargetOptions.h"
71#include <algorithm>
72#include <cassert>
73#include <cstdint>
74#include <functional>
75#include <iterator>
76#include <optional>
77#include <string>
78#include <tuple>
79#include <utility>
80#include <variant>
81
82#include "SDNodeDbgValue.h"
83
84using namespace llvm;
85using namespace llvm::SDPatternMatch;
86
87#define DEBUG_TYPE "dagcombine"
88
89STATISTIC(NodesCombined , "Number of dag nodes combined");
90STATISTIC(PreIndexedNodes , "Number of pre-indexed nodes created");
91STATISTIC(PostIndexedNodes, "Number of post-indexed nodes created");
92STATISTIC(OpsNarrowed , "Number of load/op/store narrowed");
93STATISTIC(LdStFP2Int , "Number of fp load/store pairs transformed to int");
94STATISTIC(SlicedLoads, "Number of load sliced");
95STATISTIC(NumFPLogicOpsConv, "Number of logic ops converted to fp ops");
96
97DEBUG_COUNTER(DAGCombineCounter, "dagcombine",
98 "Controls whether a DAG combine is performed for a node");
99
100static cl::opt<bool>
101CombinerGlobalAA("combiner-global-alias-analysis", cl::Hidden,
102 cl::desc("Enable DAG combiner's use of IR alias analysis"));
103
104static cl::opt<bool>
105UseTBAA("combiner-use-tbaa", cl::Hidden, cl::init(Val: true),
106 cl::desc("Enable DAG combiner's use of TBAA"));
107
108#ifndef NDEBUG
109static cl::opt<std::string>
110CombinerAAOnlyFunc("combiner-aa-only-func", cl::Hidden,
111 cl::desc("Only use DAG-combiner alias analysis in this"
112 " function"));
113#endif
114
115/// Hidden option to stress test load slicing, i.e., when this option
116/// is enabled, load slicing bypasses most of its profitability guards.
117static cl::opt<bool>
118StressLoadSlicing("combiner-stress-load-slicing", cl::Hidden,
119 cl::desc("Bypass the profitability model of load slicing"),
120 cl::init(Val: false));
121
122static cl::opt<bool>
123 MaySplitLoadIndex("combiner-split-load-index", cl::Hidden, cl::init(Val: true),
124 cl::desc("DAG combiner may split indexing from loads"));
125
126static cl::opt<bool>
127 EnableStoreMerging("combiner-store-merging", cl::Hidden, cl::init(Val: true),
128 cl::desc("DAG combiner enable merging multiple stores "
129 "into a wider store"));
130
131static cl::opt<unsigned> TokenFactorInlineLimit(
132 "combiner-tokenfactor-inline-limit", cl::Hidden, cl::init(Val: 2048),
133 cl::desc("Limit the number of operands to inline for Token Factors"));
134
135static cl::opt<unsigned> StoreMergeDependenceLimit(
136 "combiner-store-merge-dependence-limit", cl::Hidden, cl::init(Val: 10),
137 cl::desc("Limit the number of times for the same StoreNode and RootNode "
138 "to bail out in store merging dependence check"));
139
140static cl::opt<bool> EnableReduceLoadOpStoreWidth(
141 "combiner-reduce-load-op-store-width", cl::Hidden, cl::init(Val: true),
142 cl::desc("DAG combiner enable reducing the width of load/op/store "
143 "sequence"));
144static cl::opt<bool> ReduceLoadOpStoreWidthForceNarrowingProfitable(
145 "combiner-reduce-load-op-store-width-force-narrowing-profitable",
146 cl::Hidden, cl::init(Val: false),
147 cl::desc("DAG combiner force override the narrowing profitable check when "
148 "reducing the width of load/op/store sequences"));
149
150static cl::opt<bool> EnableShrinkLoadReplaceStoreWithStore(
151 "combiner-shrink-load-replace-store-with-store", cl::Hidden, cl::init(Val: true),
152 cl::desc("DAG combiner enable load/<replace bytes>/store with "
153 "a narrower store"));
154
155static cl::opt<bool> EnableTopologicalSorting(
156 "combiner-topological-sorting", cl::Hidden, cl::init(Val: false),
157 cl::desc("DAG combiner nodes consistently processed in topological order"));
158
159static cl::opt<bool> DisableCombines("combiner-disabled", cl::Hidden,
160 cl::init(Val: false),
161 cl::desc("Disable the DAG combiner"));
162
163namespace {
164
165 class DAGCombiner {
166 SelectionDAG &DAG;
167 const TargetLowering &TLI;
168 const SelectionDAGTargetInfo *STI;
169 CombineLevel Level = BeforeLegalizeTypes;
170 CodeGenOptLevel OptLevel;
171 bool LegalDAG = false;
172 bool LegalOperations = false;
173 bool LegalTypes = false;
174 bool ForCodeSize;
175 bool DisableGenericCombines;
176
177 /// Worklist of all of the nodes that need to be simplified.
178 ///
179 /// This must behave as a stack -- new nodes to process are pushed onto the
180 /// back and when processing we pop off of the back.
181 ///
182 /// The worklist will not contain duplicates but may contain null entries
183 /// due to nodes being deleted from the underlying DAG. For fast lookup and
184 /// deduplication, the index of the node in this vector is stored in the
185 /// node in SDNode::CombinerWorklistIndex.
186 SmallVector<SDNode *, 64> Worklist;
187
188 /// This records all nodes attempted to be added to the worklist since we
189 /// considered a new worklist entry. As we keep do not add duplicate nodes
190 /// in the worklist, this is different from the tail of the worklist.
191 SmallSetVector<SDNode *, 32> PruningList;
192
193 /// Map from candidate StoreNode to the pair of RootNode and count.
194 /// The count is used to track how many times we have seen the StoreNode
195 /// with the same RootNode bail out in dependence check. If we have seen
196 /// the bail out for the same pair many times over a limit, we won't
197 /// consider the StoreNode with the same RootNode as store merging
198 /// candidate again.
199 DenseMap<SDNode *, std::pair<SDNode *, unsigned>> StoreRootCountMap;
200
201 // BatchAA - Used for DAG load/store alias analysis.
202 BatchAAResults *BatchAA;
203
204 /// This caches all chains that have already been processed in
205 /// DAGCombiner::getStoreMergeCandidates() and found to have no mergeable
206 /// stores candidates.
207 SmallPtrSet<SDNode *, 4> ChainsWithoutMergeableStores;
208
209 /// When an instruction is simplified, add all users of the instruction to
210 /// the work lists because they might get more simplified now.
211 void AddUsersToWorklist(SDNode *N) {
212 for (SDNode *Node : N->users())
213 AddToWorklist(N: Node);
214 }
215
216 /// Convenient shorthand to add a node and all of its user to the worklist.
217 void AddToWorklistWithUsers(SDNode *N) {
218 AddUsersToWorklist(N);
219 AddToWorklist(N);
220 }
221
222 // Prune potentially dangling nodes. This is called after
223 // any visit to a node, but should also be called during a visit after any
224 // failed combine which may have created a DAG node.
225 void clearAddedDanglingWorklistEntries() {
226 // Check any nodes added to the worklist to see if they are prunable.
227 while (!PruningList.empty()) {
228 auto *N = PruningList.pop_back_val();
229 if (N->use_empty())
230 recursivelyDeleteUnusedNodes(N);
231 }
232 }
233
234 SDNode *getNextWorklistEntry() {
235 // Before we do any work, remove nodes that are not in use.
236 clearAddedDanglingWorklistEntries();
237 SDNode *N = nullptr;
238 // The Worklist holds the SDNodes in order, but it may contain null
239 // entries.
240 while (!N && !Worklist.empty()) {
241 N = Worklist.pop_back_val();
242 }
243
244 if (N) {
245 assert(N->getCombinerWorklistIndex() >= 0 &&
246 "Found a worklist entry without a corresponding map entry!");
247 // Set to -2 to indicate that we combined the node.
248 N->setCombinerWorklistIndex(-2);
249 }
250 return N;
251 }
252
253 /// Call the node-specific routine that folds each particular type of node.
254 SDValue visit(SDNode *N);
255
256 public:
257 DAGCombiner(SelectionDAG &D, BatchAAResults *BatchAA, CodeGenOptLevel OL)
258 : DAG(D), TLI(D.getTargetLoweringInfo()),
259 STI(D.getSubtarget().getSelectionDAGInfo()), OptLevel(OL),
260 BatchAA(BatchAA) {
261 ForCodeSize = DAG.shouldOptForSize();
262 DisableGenericCombines =
263 DisableCombines || (STI && STI->disableGenericCombines(OptLevel));
264 }
265
266 void ConsiderForPruning(SDNode *N) {
267 // Mark this for potential pruning.
268 PruningList.insert(X: N);
269 }
270
271 /// Add to the worklist making sure its instance is at the back (next to be
272 /// processed.)
273 void AddToWorklist(SDNode *N, bool IsCandidateForPruning = true,
274 bool SkipIfCombinedBefore = false) {
275 assert(N->getOpcode() != ISD::DELETED_NODE &&
276 "Deleted Node added to Worklist");
277
278 // Skip handle nodes as they can't usefully be combined and confuse the
279 // zero-use deletion strategy.
280 if (N->getOpcode() == ISD::HANDLENODE)
281 return;
282
283 if (SkipIfCombinedBefore && N->getCombinerWorklistIndex() == -2)
284 return;
285
286 if (IsCandidateForPruning)
287 ConsiderForPruning(N);
288
289 if (N->getCombinerWorklistIndex() < 0) {
290 N->setCombinerWorklistIndex(Worklist.size());
291 Worklist.push_back(Elt: N);
292 }
293 }
294
295 /// Remove all instances of N from the worklist.
296 void removeFromWorklist(SDNode *N) {
297 PruningList.remove(X: N);
298 StoreRootCountMap.erase(Val: N);
299
300 int WorklistIndex = N->getCombinerWorklistIndex();
301 // If not in the worklist, the index might be -1 or -2 (was combined
302 // before). As the node gets deleted anyway, there's no need to update
303 // the index.
304 if (WorklistIndex < 0)
305 return; // Not in the worklist.
306
307 // Null out the entry rather than erasing it to avoid a linear operation.
308 Worklist[WorklistIndex] = nullptr;
309 N->setCombinerWorklistIndex(-1);
310 }
311
312 void deleteAndRecombine(SDNode *N);
313 bool recursivelyDeleteUnusedNodes(SDNode *N);
314
315 /// Replaces all uses of the results of one DAG node with new values.
316 SDValue CombineTo(SDNode *N, const SDValue *To, unsigned NumTo,
317 bool AddTo = true);
318
319 /// Replaces all uses of the results of one DAG node with new values.
320 SDValue CombineTo(SDNode *N, SDValue Res, bool AddTo = true) {
321 return CombineTo(N, To: &Res, NumTo: 1, AddTo);
322 }
323
324 /// Replaces all uses of the results of one DAG node with new values.
325 SDValue CombineTo(SDNode *N, SDValue Res0, SDValue Res1,
326 bool AddTo = true) {
327 SDValue To[] = { Res0, Res1 };
328 return CombineTo(N, To, NumTo: 2, AddTo);
329 }
330
331 SDValue CombineTo(SDNode *N, SmallVectorImpl<SDValue> *To,
332 bool AddTo = true) {
333 return CombineTo(N, To: To->data(), NumTo: To->size(), AddTo);
334 }
335
336 void CommitTargetLoweringOpt(const TargetLowering::TargetLoweringOpt &TLO);
337
338 private:
339 /// Check the specified integer node value to see if it can be simplified or
340 /// if things it uses can be simplified by bit propagation.
341 /// If so, return true.
342 bool SimplifyDemandedBits(SDValue Op) {
343 unsigned BitWidth = Op.getScalarValueSizeInBits();
344 APInt DemandedBits = APInt::getAllOnes(numBits: BitWidth);
345 return SimplifyDemandedBits(Op, DemandedBits);
346 }
347
348 bool SimplifyDemandedBits(SDValue Op, const APInt &DemandedBits) {
349 EVT VT = Op.getValueType();
350 APInt DemandedElts = VT.isFixedLengthVector()
351 ? APInt::getAllOnes(numBits: VT.getVectorNumElements())
352 : APInt(1, 1);
353 return SimplifyDemandedBits(Op, DemandedBits, DemandedElts, AssumeSingleUse: false);
354 }
355
356 /// Check the specified vector node value to see if it can be simplified or
357 /// if things it uses can be simplified as it only uses some of the
358 /// elements. If so, return true.
359 bool SimplifyDemandedVectorElts(SDValue Op) {
360 // TODO: For now just pretend it cannot be simplified.
361 if (Op.getValueType().isScalableVector())
362 return false;
363
364 unsigned NumElts = Op.getValueType().getVectorNumElements();
365 APInt DemandedElts = APInt::getAllOnes(numBits: NumElts);
366 return SimplifyDemandedVectorElts(Op, DemandedElts);
367 }
368
369 bool SimplifyDemandedBits(SDValue Op, const APInt &DemandedBits,
370 const APInt &DemandedElts,
371 bool AssumeSingleUse = false);
372 bool SimplifyDemandedVectorElts(SDValue Op, const APInt &DemandedElts,
373 bool AssumeSingleUse = false);
374
375 bool CombineToPreIndexedLoadStore(SDNode *N);
376 bool CombineToPostIndexedLoadStore(SDNode *N);
377 SDValue SplitIndexingFromLoad(LoadSDNode *LD);
378 bool SliceUpLoad(SDNode *N);
379
380 // Looks up the chain to find a unique (unaliased) store feeding the passed
381 // load. If no such store is found, returns a nullptr.
382 // Note: This will look past a CALLSEQ_START if the load is chained to it so
383 // so that it can find stack stores for byval params.
384 StoreSDNode *getUniqueStoreFeeding(LoadSDNode *LD, int64_t &Offset);
385 // Scalars have size 0 to distinguish from singleton vectors.
386 SDValue ForwardStoreValueToDirectLoad(LoadSDNode *LD);
387 bool getTruncatedStoreValue(StoreSDNode *ST, SDValue &Val);
388 bool extendLoadedValueToExtension(LoadSDNode *LD, SDValue &Val);
389
390 void ReplaceLoadWithPromotedLoad(SDNode *Load, SDNode *ExtLoad);
391 SDValue PromoteOperand(SDValue Op, EVT PVT, bool &Replace);
392 SDValue SExtPromoteOperand(SDValue Op, EVT PVT);
393 SDValue ZExtPromoteOperand(SDValue Op, EVT PVT);
394 SDValue PromoteIntBinOp(SDValue Op);
395 SDValue PromoteIntShiftOp(SDValue Op);
396 SDValue PromoteExtend(SDValue Op);
397 bool PromoteLoad(SDValue Op);
398
399 SDValue foldShiftToAvg(SDNode *N, const SDLoc &DL);
400 // Fold `a bitwiseop (~b +/- c)` -> `a bitwiseop ~(b -/+ c)`
401 SDValue foldBitwiseOpWithNeg(SDNode *N, const SDLoc &DL, EVT VT);
402
403 SDValue combineMinNumMaxNum(const SDLoc &DL, EVT VT, SDValue LHS,
404 SDValue RHS, SDValue True, SDValue False,
405 ISD::CondCode CC);
406
407 /// Call the node-specific routine that knows how to fold each
408 /// particular type of node. If that doesn't do anything, try the
409 /// target-specific DAG combines.
410 SDValue combine(SDNode *N);
411
412 // Visitation implementation - Implement dag node combining for different
413 // node types. The semantics are as follows:
414 // Return Value:
415 // SDValue.getNode() == 0 - No change was made
416 // SDValue.getNode() == N - N was replaced, is dead and has been handled.
417 // otherwise - N should be replaced by the returned Operand.
418 //
419 SDValue visitTokenFactor(SDNode *N);
420 SDValue visitMERGE_VALUES(SDNode *N);
421 SDValue visitADD(SDNode *N);
422 SDValue visitADDLike(SDNode *N);
423 SDValue visitADDLikeCommutative(SDValue N0, SDValue N1, const SDLoc &DL);
424 SDValue visitPTRADD(SDNode *N);
425 SDValue visitSUB(SDNode *N);
426 SDValue visitADDSAT(SDNode *N);
427 SDValue visitSUBSAT(SDNode *N);
428 SDValue visitADDC(SDNode *N);
429 SDValue visitADDO(SDNode *N);
430 SDValue visitUADDOLike(SDValue N0, SDValue N1, SDNode *N);
431 SDValue visitSUBC(SDNode *N);
432 SDValue visitSUBO(SDNode *N);
433 SDValue visitADDE(SDNode *N);
434 SDValue visitUADDO_CARRY(SDNode *N);
435 SDValue visitSADDO_CARRY(SDNode *N);
436 SDValue visitUADDO_CARRYLike(SDValue N0, SDValue N1, SDValue CarryIn,
437 SDNode *N);
438 SDValue visitSADDO_CARRYLike(SDValue N0, SDValue N1, SDValue CarryIn,
439 SDNode *N);
440 SDValue visitSUBE(SDNode *N);
441 SDValue visitUSUBO_CARRY(SDNode *N);
442 SDValue visitSSUBO_CARRY(SDNode *N);
443 SDValue visitMUL(SDNode *N);
444 SDValue visitMULFIX(SDNode *N);
445 SDValue useDivRem(SDNode *N);
446 SDValue visitSDIV(SDNode *N);
447 SDValue visitSDIVLike(SDValue N0, SDValue N1, SDNode *N);
448 SDValue visitUDIV(SDNode *N);
449 SDValue visitUDIVLike(SDValue N0, SDValue N1, SDNode *N);
450 SDValue visitREM(SDNode *N);
451 SDValue visitMULHU(SDNode *N);
452 SDValue visitMULHS(SDNode *N);
453 SDValue visitAVG(SDNode *N);
454 SDValue visitABD(SDNode *N);
455 SDValue visitSMUL_LOHI(SDNode *N);
456 SDValue visitUMUL_LOHI(SDNode *N);
457 SDValue visitMULO(SDNode *N);
458 SDValue visitIMINMAX(SDNode *N);
459 SDValue visitAND(SDNode *N);
460 SDValue visitANDLike(SDValue N0, SDValue N1, SDNode *N);
461 SDValue visitOR(SDNode *N);
462 SDValue visitORLike(SDValue N0, SDValue N1, const SDLoc &DL);
463 SDValue visitXOR(SDNode *N);
464 SDValue SimplifyVCastOp(SDNode *N, const SDLoc &DL);
465 SDValue SimplifyVBinOp(SDNode *N, const SDLoc &DL);
466 SDValue visitSHL(SDNode *N);
467 SDValue visitSRA(SDNode *N);
468 SDValue visitSRL(SDNode *N);
469 SDValue visitFunnelShift(SDNode *N);
470 SDValue visitSHLSAT(SDNode *N);
471 SDValue visitRotate(SDNode *N);
472 SDValue visitABS(SDNode *N);
473 SDValue visitABS_MIN_POISON(SDNode *N);
474 SDValue visitCLMUL(SDNode *N);
475 SDValue visitPEXT(SDNode *N);
476 SDValue visitPDEP(SDNode *N);
477 SDValue visitBSWAP(SDNode *N);
478 SDValue visitBITREVERSE(SDNode *N);
479 SDValue visitCTLZ(SDNode *N);
480 SDValue visitCTLZ_ZERO_POISON(SDNode *N);
481 SDValue visitCTTZ(SDNode *N);
482 SDValue visitCTTZ_ZERO_POISON(SDNode *N);
483 SDValue visitCTPOP(SDNode *N);
484 SDValue visitPARITY(SDNode *N);
485 SDValue visitSELECT(SDNode *N);
486 SDValue visitVSELECT(SDNode *N);
487 SDValue visitSELECT_CC(SDNode *N);
488 SDValue visitSETCC(SDNode *N);
489 SDValue visitSETCCCARRY(SDNode *N);
490 SDValue visitSIGN_EXTEND(SDNode *N);
491 SDValue visitZERO_EXTEND(SDNode *N);
492 SDValue visitANY_EXTEND(SDNode *N);
493 SDValue visitAssertExt(SDNode *N);
494 SDValue visitAssertAlign(SDNode *N);
495 SDValue visitIS_FPCLASS(SDNode *N);
496 SDValue visitSIGN_EXTEND_INREG(SDNode *N);
497 SDValue visitEXTEND_VECTOR_INREG(SDNode *N);
498 SDValue visitTRUNCATE(SDNode *N);
499 SDValue visitTRUNCATE_USAT_U(SDNode *N);
500 SDValue visitBITCAST(SDNode *N);
501 SDValue visitFREEZE(SDNode *N);
502 SDValue visitBUILD_PAIR(SDNode *N);
503 SDValue visitFADD(SDNode *N);
504 SDValue visitSTRICT_FADD(SDNode *N);
505 SDValue visitFSUB(SDNode *N);
506 SDValue visitFMUL(SDNode *N);
507 SDValue visitFMA(SDNode *N);
508 SDValue visitFMAD(SDNode *N);
509 SDValue visitFMULADD(SDNode *N);
510 SDValue visitFDIV(SDNode *N);
511 SDValue visitFREM(SDNode *N);
512 SDValue visitFSQRT(SDNode *N);
513 SDValue visitFCOPYSIGN(SDNode *N);
514 SDValue visitFPOW(SDNode *N);
515 SDValue visitFCANONICALIZE(SDNode *N);
516 SDValue visitSINT_TO_FP(SDNode *N);
517 SDValue visitUINT_TO_FP(SDNode *N);
518 SDValue visitFP_TO_SINT(SDNode *N);
519 SDValue visitFP_TO_UINT(SDNode *N);
520 SDValue visitXROUND(SDNode *N);
521 SDValue visitFP_ROUND(SDNode *N);
522 SDValue visitFP_EXTEND(SDNode *N);
523 SDValue visitFNEG(SDNode *N);
524 SDValue visitFABS(SDNode *N);
525 SDValue visitFCEIL(SDNode *N);
526 SDValue visitFTRUNC(SDNode *N);
527 SDValue visitFFREXP(SDNode *N);
528 SDValue visitFFLOOR(SDNode *N);
529 SDValue visitFMinMax(SDNode *N);
530 SDValue visitBRCOND(SDNode *N);
531 SDValue visitBR_CC(SDNode *N);
532 SDValue visitLOAD(SDNode *N);
533
534 SDValue replaceStoreChain(StoreSDNode *ST, SDValue BetterChain);
535 SDValue replaceStoreOfFPConstant(StoreSDNode *ST);
536 SDValue replaceStoreOfInsertLoad(StoreSDNode *ST);
537
538 bool refineExtractVectorEltIntoMultipleNarrowExtractVectorElts(SDNode *N);
539 SDValue combineStoreConcatTruncVector(StoreSDNode *N);
540 SDValue visitSTORE(SDNode *N);
541 SDValue visitATOMIC_STORE(SDNode *N);
542 SDValue visitLIFETIME_END(SDNode *N);
543 SDValue visitINSERT_VECTOR_ELT(SDNode *N);
544 SDValue visitEXTRACT_VECTOR_ELT(SDNode *N);
545 SDValue visitBUILD_VECTOR(SDNode *N);
546 SDValue visitCONCAT_VECTORS(SDNode *N);
547 SDValue visitVECTOR_INTERLEAVE(SDNode *N);
548 SDValue visitVECTOR_DEINTERLEAVE(SDNode *N);
549 SDValue visitEXTRACT_SUBVECTOR(SDNode *N);
550 SDValue visitVECTOR_SHUFFLE(SDNode *N);
551 SDValue visitSCALAR_TO_VECTOR(SDNode *N);
552 SDValue visitINSERT_SUBVECTOR(SDNode *N);
553 SDValue visitVECTOR_COMPRESS(SDNode *N);
554 SDValue visitMLOAD(SDNode *N);
555 SDValue visitMSTORE(SDNode *N);
556 SDValue visitMGATHER(SDNode *N);
557 SDValue visitMSCATTER(SDNode *N);
558 SDValue visitMHISTOGRAM(SDNode *N);
559 SDValue visitPARTIAL_REDUCE_MLA(SDNode *N);
560 SDValue visitLOOP_DEPENDENCE_MASK(SDNode *N);
561 SDValue visitVPGATHER(SDNode *N);
562 SDValue visitVPSCATTER(SDNode *N);
563 SDValue visitVP_STRIDED_LOAD(SDNode *N);
564 SDValue visitVP_STRIDED_STORE(SDNode *N);
565 SDValue visitFP_TO_FP16(SDNode *N);
566 SDValue visitFP16_TO_FP(SDNode *N);
567 SDValue visitFP_TO_BF16(SDNode *N);
568 SDValue visitBF16_TO_FP(SDNode *N);
569 SDValue visitVECREDUCE(SDNode *N);
570 SDValue visitVPOp(SDNode *N);
571 SDValue visitGET_FPENV_MEM(SDNode *N);
572 SDValue visitSET_FPENV_MEM(SDNode *N);
573
574 SDValue visitFADDForFMACombine(SDNode *N);
575 SDValue visitFSUBForFMACombine(SDNode *N);
576 SDValue visitFMULForFMADistributiveCombine(SDNode *N);
577
578 SDValue XformToShuffleWithZero(SDNode *N);
579 bool reassociationCanBreakAddressingModePattern(unsigned Opc,
580 const SDLoc &DL,
581 SDNode *N,
582 SDValue N0,
583 SDValue N1);
584 SDValue reassociateOpsCommutative(unsigned Opc, const SDLoc &DL, SDValue N0,
585 SDValue N1, SDNodeFlags Flags);
586 SDValue reassociateOps(unsigned Opc, const SDLoc &DL, SDValue N0,
587 SDValue N1, SDNodeFlags Flags);
588 SDValue reassociateReduction(unsigned RedOpc, unsigned Opc, const SDLoc &DL,
589 EVT VT, SDValue N0, SDValue N1,
590 SDNodeFlags Flags = SDNodeFlags());
591
592 SDValue visitShiftByConstant(SDNode *N);
593
594 SDValue foldSelectOfConstants(SDNode *N);
595 SDValue foldVSelectOfConstants(SDNode *N);
596 SDValue foldBinOpIntoSelect(SDNode *BO);
597 bool SimplifySelectOps(SDNode *SELECT, SDValue LHS, SDValue RHS);
598 SDValue hoistLogicOpWithSameOpcodeHands(SDNode *N);
599 SDValue SimplifySelect(const SDLoc &DL, SDValue N0, SDValue N1, SDValue N2);
600 SDValue SimplifySelectCC(const SDLoc &DL, SDValue N0, SDValue N1,
601 SDValue N2, SDValue N3, ISD::CondCode CC,
602 bool NotExtCompare = false);
603 SDValue convertSelectOfFPConstantsToLoadOffset(
604 const SDLoc &DL, SDValue N0, SDValue N1, SDValue N2, SDValue N3,
605 ISD::CondCode CC);
606 SDValue foldSignChangeInBitcast(SDNode *N);
607 SDValue foldSelectCCToShiftAnd(const SDLoc &DL, SDValue N0, SDValue N1,
608 SDValue N2, SDValue N3, ISD::CondCode CC);
609 SDValue foldSelectOfBinops(SDNode *N);
610 SDValue foldSextSetcc(SDNode *N);
611 SDValue foldLogicOfSetCCs(bool IsAnd, SDValue N0, SDValue N1,
612 const SDLoc &DL);
613 SDValue foldSubToUSubSat(EVT DstVT, SDNode *N, const SDLoc &DL);
614 SDValue foldABSToABD(SDNode *N, const SDLoc &DL);
615 SDValue foldSelectToABD(SDValue LHS, SDValue RHS, SDValue True,
616 SDValue False, ISD::CondCode CC, const SDLoc &DL);
617 SDValue foldSelectToUMin(SDValue LHS, SDValue RHS, SDValue True,
618 SDValue False, ISD::CondCode CC, const SDLoc &DL);
619 SDValue unfoldMaskedMerge(SDNode *N);
620 SDValue unfoldExtremeBitClearingToShifts(SDNode *N);
621 SDValue SimplifySetCC(EVT VT, SDValue N0, SDValue N1, ISD::CondCode Cond,
622 const SDLoc &DL, bool foldBooleans);
623 SDValue rebuildSetCC(SDValue N);
624
625 bool isSetCCEquivalent(SDValue N, SDValue &LHS, SDValue &RHS,
626 SDValue &CC, bool MatchStrict = false) const;
627 bool isOneUseSetCC(SDValue N) const;
628
629 SDValue foldAddToAvg(SDNode *N, const SDLoc &DL);
630 SDValue foldSubToAvg(SDNode *N, const SDLoc &DL);
631
632 SDValue foldCTLZToCTLS(SDValue Src, const SDLoc &DL);
633
634 SDValue SimplifyNodeWithTwoResults(SDNode *N, unsigned LoOp,
635 unsigned HiOp);
636 SDValue CombineConsecutiveLoads(SDNode *N, EVT VT);
637 SDValue foldBitcastedFPLogic(SDNode *N, SelectionDAG &DAG,
638 const TargetLowering &TLI);
639 SDValue foldPartialReduceMLAMulOp(SDNode *N);
640 SDValue foldPartialReduceAdd(SDNode *N);
641
642 SDValue CombineExtLoad(SDNode *N);
643 SDValue CombineZExtLogicopShiftLoad(SDNode *N);
644 SDValue combineRepeatedFPDivisors(SDNode *N);
645 SDValue combineFMulOrFDivWithIntPow2(SDNode *N);
646 SDValue replaceShuffleOfInsert(ShuffleVectorSDNode *Shuf);
647 SDValue mergeInsertEltWithShuffle(SDNode *N, unsigned InsIndex);
648 SDValue combineInsertEltToShuffle(SDNode *N, unsigned InsIndex);
649 SDValue combineInsertEltToLoad(SDNode *N, unsigned InsIndex);
650 SDValue foldExtractSubvectorFromConcatVectors(EVT VT, SDValue V,
651 uint64_t ExtIdx,
652 const SDLoc &DL);
653 SDValue BuildSDIV(SDNode *N);
654 SDValue BuildSDIVPow2(SDNode *N);
655 SDValue BuildUDIV(SDNode *N);
656 SDValue BuildSREMPow2(SDNode *N);
657 SDValue buildOptimizedSREM(SDValue N0, SDValue N1, SDNode *N);
658 SDValue BuildLogBase2(SDValue V, const SDLoc &DL,
659 bool KnownNeverZero = false,
660 bool InexpensiveOnly = false,
661 std::optional<EVT> OutVT = std::nullopt);
662 SDValue BuildDivEstimate(SDValue N, SDValue Op, SDNodeFlags Flags);
663 SDValue buildRsqrtEstimate(SDValue Op, SDNodeFlags Flags);
664 SDValue buildSqrtEstimate(SDValue Op, SDNodeFlags Flags);
665 SDValue buildSqrtEstimateImpl(SDValue Op, bool Recip, SDNodeFlags Flags);
666 SDValue buildSqrtNROneConst(SDValue Arg, SDValue Est, unsigned Iterations,
667 bool Reciprocal);
668 SDValue buildSqrtNRTwoConst(SDValue Arg, SDValue Est, unsigned Iterations,
669 bool Reciprocal);
670 SDValue MatchBSwapHWordLow(SDNode *N, SDValue N0, SDValue N1,
671 bool DemandHighBits = true);
672 SDValue MatchBSwapHWord(SDNode *N, SDValue N0, SDValue N1);
673 SDValue MatchRotatePosNeg(SDValue Shifted, SDValue Pos, SDValue Neg,
674 SDValue InnerPos, SDValue InnerNeg, bool FromAdd,
675 bool HasPos, unsigned PosOpcode,
676 unsigned NegOpcode, const SDLoc &DL);
677 SDValue MatchFunnelPosNeg(SDValue N0, SDValue N1, SDValue Pos, SDValue Neg,
678 SDValue InnerPos, SDValue InnerNeg, bool FromAdd,
679 bool HasPos, unsigned PosOpcode,
680 unsigned NegOpcode, const SDLoc &DL);
681 SDValue MatchRotate(SDValue LHS, SDValue RHS, const SDLoc &DL,
682 bool FromAdd);
683 SDValue MatchLoadCombine(SDNode *N);
684 SDValue mergeTruncStores(StoreSDNode *N);
685 SDValue reduceLoadWidth(SDNode *N);
686 SDValue ReduceLoadOpStoreWidth(SDNode *N);
687 SDValue splitMergedValStore(StoreSDNode *ST);
688 SDValue TransformFPLoadStorePair(SDNode *N);
689 SDValue convertBuildVecExtToExt(SDNode *N);
690 SDValue convertBuildVecZextToBuildVecWithZeros(SDNode *N);
691 SDValue reduceBuildVecExtToExtBuildVec(SDNode *N);
692 SDValue reduceBuildVecTruncToBitCast(SDNode *N);
693 SDValue reduceBuildVecToShuffle(SDNode *N);
694 SDValue createBuildVecShuffle(const SDLoc &DL, SDNode *N,
695 ArrayRef<int> VectorMask, SDValue VecIn1,
696 SDValue VecIn2, unsigned LeftIdx,
697 bool DidSplitVec);
698 SDValue matchVSelectOpSizesWithSetCC(SDNode *Cast);
699
700 /// Walk up chain skipping non-aliasing memory nodes,
701 /// looking for aliasing nodes and adding them to the Aliases vector.
702 void GatherAllAliases(SDNode *N, SDValue OriginalChain,
703 SmallVectorImpl<SDValue> &Aliases);
704
705 /// Return true if there is any possibility that the two addresses overlap.
706 bool mayAlias(SDNode *Op0, SDNode *Op1) const;
707
708 /// Walk up chain skipping non-aliasing memory nodes, looking for a better
709 /// chain (aliasing node.)
710 SDValue FindBetterChain(SDNode *N, SDValue Chain);
711
712 /// Try to replace a store and any possibly adjacent stores on
713 /// consecutive chains with better chains. Return true only if St is
714 /// replaced.
715 ///
716 /// Notice that other chains may still be replaced even if the function
717 /// returns false.
718 bool findBetterNeighborChains(StoreSDNode *St);
719
720 // Helper for findBetterNeighborChains. Walk up store chain add additional
721 // chained stores that do not overlap and can be parallelized.
722 bool parallelizeChainedStores(StoreSDNode *St);
723
724 /// Holds a pointer to an LSBaseSDNode as well as information on where it
725 /// is located in a sequence of memory operations connected by a chain.
726 struct MemOpLink {
727 // Ptr to the mem node.
728 LSBaseSDNode *MemNode;
729
730 // Offset from the base ptr.
731 int64_t OffsetFromBase;
732
733 MemOpLink(LSBaseSDNode *N, int64_t Offset)
734 : MemNode(N), OffsetFromBase(Offset) {}
735 };
736
737 // Classify the origin of a stored value.
738 enum class StoreSource { Unknown, Constant, Extract, Load };
739 StoreSource getStoreSource(SDValue StoreVal) {
740 switch (StoreVal.getOpcode()) {
741 case ISD::Constant:
742 case ISD::ConstantFP:
743 return StoreSource::Constant;
744 case ISD::BUILD_VECTOR:
745 if (ISD::isBuildVectorOfConstantSDNodes(N: StoreVal.getNode()) ||
746 ISD::isBuildVectorOfConstantFPSDNodes(N: StoreVal.getNode()))
747 return StoreSource::Constant;
748 return StoreSource::Unknown;
749 case ISD::EXTRACT_VECTOR_ELT:
750 case ISD::EXTRACT_SUBVECTOR:
751 return StoreSource::Extract;
752 case ISD::LOAD:
753 return StoreSource::Load;
754 default:
755 return StoreSource::Unknown;
756 }
757 }
758
759 /// This is a helper function for visitMUL to check the profitability
760 /// of folding (mul (add x, c1), c2) -> (add (mul x, c2), c1*c2).
761 /// MulNode is the original multiply, AddNode is (add x, c1),
762 /// and ConstNode is c2.
763 bool isMulAddWithConstProfitable(SDNode *MulNode, SDValue AddNode,
764 SDValue ConstNode);
765
766 /// This is a helper function for visitAND and visitZERO_EXTEND. Returns
767 /// true if the (and (load x) c) pattern matches an extload. ExtVT returns
768 /// the type of the loaded value to be extended.
769 bool isAndLoadExtLoad(ConstantSDNode *AndC, LoadSDNode *LoadN,
770 EVT LoadResultTy, EVT &ExtVT);
771
772 /// Helper function to calculate whether the given Load/Store can have its
773 /// width reduced to ExtVT.
774 bool isLegalNarrowLdSt(LSBaseSDNode *LDSTN, ISD::LoadExtType ExtType,
775 EVT &MemVT, unsigned ShAmt = 0);
776
777 /// Used by BackwardsPropagateMask to find suitable loads.
778 bool SearchForAndLoads(SDNode *N, SmallVectorImpl<LoadSDNode*> &Loads,
779 SmallPtrSetImpl<SDNode*> &NodesWithConsts,
780 ConstantSDNode *Mask, SDNode *&NodeToMask);
781 /// Attempt to propagate a given AND node back to load leaves so that they
782 /// can be combined into narrow loads.
783 bool BackwardsPropagateMask(SDNode *N);
784
785 /// Helper function for mergeConsecutiveStores which merges the component
786 /// store chains.
787 SDValue getMergeStoreChains(SmallVectorImpl<MemOpLink> &StoreNodes,
788 unsigned NumStores);
789
790 /// Helper function for mergeConsecutiveStores which checks if all the store
791 /// nodes have the same underlying object. We can still reuse the first
792 /// store's pointer info if all the stores are from the same object.
793 bool hasSameUnderlyingObj(ArrayRef<MemOpLink> StoreNodes);
794
795 /// This is a helper function for mergeConsecutiveStores. When the source
796 /// elements of the consecutive stores are all constants or all extracted
797 /// vector elements, try to merge them into one larger store introducing
798 /// bitcasts if necessary. \return True if a merged store was created.
799 bool mergeStoresOfConstantsOrVecElts(SmallVectorImpl<MemOpLink> &StoreNodes,
800 EVT MemVT, unsigned NumStores,
801 bool IsConstantSrc, bool UseVector,
802 bool UseTrunc);
803
804 /// This is a helper function for mergeConsecutiveStores. Stores that
805 /// potentially may be merged with St are placed in StoreNodes. On success,
806 /// returns a chain predecessor to all store candidates.
807 SDNode *getStoreMergeCandidates(StoreSDNode *St,
808 SmallVectorImpl<MemOpLink> &StoreNodes);
809
810 /// Helper function for mergeConsecutiveStores. Checks if candidate stores
811 /// have indirect dependency through their operands. RootNode is the
812 /// predecessor to all stores calculated by getStoreMergeCandidates and is
813 /// used to prune the dependency check. \return True if safe to merge.
814 bool checkMergeStoreCandidatesForDependencies(
815 SmallVectorImpl<MemOpLink> &StoreNodes, unsigned NumStores,
816 SDNode *RootNode);
817
818 /// Helper function for tryStoreMergeOfLoads. Checks if the load/store
819 /// chain has a call in it. \return True if a call is found.
820 bool hasCallInLdStChain(StoreSDNode *St, LoadSDNode *Ld);
821
822 /// This is a helper function for mergeConsecutiveStores. Given a list of
823 /// store candidates, find the first N that are consecutive in memory.
824 /// Returns 0 if there are not at least 2 consecutive stores to try merging.
825 unsigned getConsecutiveStores(SmallVectorImpl<MemOpLink> &StoreNodes,
826 int64_t ElementSizeBytes) const;
827
828 /// This is a helper function for mergeConsecutiveStores. It is used for
829 /// store chains that are composed entirely of constant values.
830 bool tryStoreMergeOfConstants(SmallVectorImpl<MemOpLink> &StoreNodes,
831 unsigned NumConsecutiveStores,
832 EVT MemVT, SDNode *Root, bool AllowVectors);
833
834 /// This is a helper function for mergeConsecutiveStores. It is used for
835 /// store chains that are composed entirely of extracted vector elements.
836 /// When extracting multiple vector elements, try to store them in one
837 /// vector store rather than a sequence of scalar stores.
838 bool tryStoreMergeOfExtracts(SmallVectorImpl<MemOpLink> &StoreNodes,
839 unsigned NumConsecutiveStores, EVT MemVT,
840 SDNode *Root);
841
842 /// This is a helper function for mergeConsecutiveStores. It is used for
843 /// store chains that are composed entirely of loaded values.
844 bool tryStoreMergeOfLoads(SmallVectorImpl<MemOpLink> &StoreNodes,
845 unsigned NumConsecutiveStores, EVT MemVT,
846 SDNode *Root, bool AllowVectors,
847 bool IsNonTemporalStore, bool IsNonTemporalLoad);
848
849 /// Merge consecutive store operations into a wide store.
850 /// This optimization uses wide integers or vectors when possible.
851 /// \return true if stores were merged.
852 bool mergeConsecutiveStores(StoreSDNode *St);
853
854 /// Try to transform a truncation where C is a constant:
855 /// (trunc (and X, C)) -> (and (trunc X), (trunc C))
856 ///
857 /// \p N needs to be a truncation and its first operand an AND. Other
858 /// requirements are checked by the function (e.g. that trunc is
859 /// single-use) and if missed an empty SDValue is returned.
860 SDValue distributeTruncateThroughAnd(SDNode *N);
861
862 /// Helper function to determine whether the target supports operation
863 /// given by \p Opcode for type \p VT, that is, whether the operation
864 /// is legal or custom before legalizing operations, and whether is
865 /// legal (but not custom) after legalization.
866 bool hasOperation(unsigned Opcode, EVT VT) {
867 return TLI.isOperationLegalOrCustom(Op: Opcode, VT, LegalOnly: LegalOperations);
868 }
869
870 bool hasUMin(EVT VT) const {
871 auto LK = TLI.getTypeConversion(Context&: *DAG.getContext(), VT);
872 return (LK.first == TargetLoweringBase::TypeLegal ||
873 LK.first == TargetLoweringBase::TypePromoteInteger) &&
874 TLI.isOperationLegalOrCustom(Op: ISD::UMIN, VT: LK.second);
875 }
876
877 public:
878 /// Runs the dag combiner on all nodes in the work list
879 void Run(CombineLevel AtLevel);
880
881 SelectionDAG &getDAG() const { return DAG; }
882
883 /// Convenience wrapper around TargetLowering::getShiftAmountTy.
884 EVT getShiftAmountTy(EVT LHSTy) {
885 return TLI.getShiftAmountTy(LHSTy, DL: DAG.getDataLayout());
886 }
887
888 /// This method returns true if we are running before type legalization or
889 /// if the specified VT is legal.
890 bool isTypeLegal(const EVT &VT) {
891 if (!LegalTypes) return true;
892 return TLI.isTypeLegal(VT);
893 }
894
895 /// Convenience wrapper around TargetLowering::getSetCCResultType
896 EVT getSetCCResultType(EVT VT) const {
897 return TLI.getSetCCResultType(DL: DAG.getDataLayout(), Context&: *DAG.getContext(), VT);
898 }
899
900 void ExtendSetCCUses(const SmallVectorImpl<SDNode *> &SetCCs,
901 SDValue OrigLoad, SDValue ExtLoad,
902 ISD::NodeType ExtType);
903 };
904
905/// This class is a DAGUpdateListener that removes any deleted
906/// nodes from the worklist.
907class WorklistRemover : public SelectionDAG::DAGUpdateListener {
908 DAGCombiner &DC;
909
910public:
911 explicit WorklistRemover(DAGCombiner &dc)
912 : SelectionDAG::DAGUpdateListener(dc.getDAG()), DC(dc) {}
913
914 void NodeDeleted(SDNode *N, SDNode *E) override {
915 DC.removeFromWorklist(N);
916 }
917};
918
919class WorklistInserter : public SelectionDAG::DAGUpdateListener {
920 DAGCombiner &DC;
921
922public:
923 explicit WorklistInserter(DAGCombiner &dc)
924 : SelectionDAG::DAGUpdateListener(dc.getDAG()), DC(dc) {}
925
926 // FIXME: Ideally we could add N to the worklist, but this causes exponential
927 // compile time costs in large DAGs, e.g. Halide.
928 void NodeInserted(SDNode *N) override { DC.ConsiderForPruning(N); }
929};
930
931} // end anonymous namespace
932
933//===----------------------------------------------------------------------===//
934// TargetLowering::DAGCombinerInfo implementation
935//===----------------------------------------------------------------------===//
936
937void TargetLowering::DAGCombinerInfo::AddToWorklist(SDNode *N) {
938 ((DAGCombiner*)DC)->AddToWorklist(N);
939}
940
941SDValue TargetLowering::DAGCombinerInfo::
942CombineTo(SDNode *N, ArrayRef<SDValue> To, bool AddTo) {
943 return ((DAGCombiner*)DC)->CombineTo(N, To: &To[0], NumTo: To.size(), AddTo);
944}
945
946SDValue TargetLowering::DAGCombinerInfo::
947CombineTo(SDNode *N, SDValue Res, bool AddTo) {
948 return ((DAGCombiner*)DC)->CombineTo(N, Res, AddTo);
949}
950
951SDValue TargetLowering::DAGCombinerInfo::
952CombineTo(SDNode *N, SDValue Res0, SDValue Res1, bool AddTo) {
953 return ((DAGCombiner*)DC)->CombineTo(N, Res0, Res1, AddTo);
954}
955
956bool TargetLowering::DAGCombinerInfo::
957recursivelyDeleteUnusedNodes(SDNode *N) {
958 return ((DAGCombiner*)DC)->recursivelyDeleteUnusedNodes(N);
959}
960
961void TargetLowering::DAGCombinerInfo::
962CommitTargetLoweringOpt(const TargetLowering::TargetLoweringOpt &TLO) {
963 return ((DAGCombiner*)DC)->CommitTargetLoweringOpt(TLO);
964}
965
966//===----------------------------------------------------------------------===//
967// Helper Functions
968//===----------------------------------------------------------------------===//
969
970void DAGCombiner::deleteAndRecombine(SDNode *N) {
971 removeFromWorklist(N);
972
973 // If the operands of this node are only used by the node, they will now be
974 // dead. Make sure to re-visit them and recursively delete dead nodes.
975 for (const SDValue &Op : N->ops())
976 // For an operand generating multiple values, one of the values may
977 // become dead allowing further simplification (e.g. split index
978 // arithmetic from an indexed load).
979 if (Op->hasOneUse() || Op->getNumValues() > 1)
980 AddToWorklist(N: Op.getNode());
981
982 DAG.DeleteNode(N);
983}
984
985// APInts must be the same size for most operations, this helper
986// function zero extends the shorter of the pair so that they match.
987// We provide an Offset so that we can create bitwidths that won't overflow.
988static void zeroExtendToMatch(APInt &LHS, APInt &RHS, unsigned Offset = 0) {
989 unsigned Bits = Offset + std::max(a: LHS.getBitWidth(), b: RHS.getBitWidth());
990 LHS = LHS.zext(width: Bits);
991 RHS = RHS.zext(width: Bits);
992}
993
994// Return true if this node is a setcc, or is a select_cc
995// that selects between the target values used for true and false, making it
996// equivalent to a setcc. Also, set the incoming LHS, RHS, and CC references to
997// the appropriate nodes based on the type of node we are checking. This
998// simplifies life a bit for the callers.
999bool DAGCombiner::isSetCCEquivalent(SDValue N, SDValue &LHS, SDValue &RHS,
1000 SDValue &CC, bool MatchStrict) const {
1001 if (N.getOpcode() == ISD::SETCC) {
1002 LHS = N.getOperand(i: 0);
1003 RHS = N.getOperand(i: 1);
1004 CC = N.getOperand(i: 2);
1005 return true;
1006 }
1007
1008 if (MatchStrict &&
1009 (N.getOpcode() == ISD::STRICT_FSETCC ||
1010 N.getOpcode() == ISD::STRICT_FSETCCS)) {
1011 LHS = N.getOperand(i: 1);
1012 RHS = N.getOperand(i: 2);
1013 CC = N.getOperand(i: 3);
1014 return true;
1015 }
1016
1017 if (N.getOpcode() != ISD::SELECT_CC || !TLI.isConstTrueVal(N: N.getOperand(i: 2)) ||
1018 !TLI.isConstFalseVal(N: N.getOperand(i: 3)))
1019 return false;
1020
1021 if (TLI.getBooleanContents(Type: N.getValueType()) ==
1022 TargetLowering::UndefinedBooleanContent)
1023 return false;
1024
1025 LHS = N.getOperand(i: 0);
1026 RHS = N.getOperand(i: 1);
1027 CC = N.getOperand(i: 4);
1028 return true;
1029}
1030
1031/// Return true if this is a SetCC-equivalent operation with only one use.
1032/// If this is true, it allows the users to invert the operation for free when
1033/// it is profitable to do so.
1034bool DAGCombiner::isOneUseSetCC(SDValue N) const {
1035 SDValue N0, N1, N2;
1036 if (isSetCCEquivalent(N, LHS&: N0, RHS&: N1, CC&: N2) && N->hasOneUse())
1037 return true;
1038 return false;
1039}
1040
1041static bool isConstantSplatVectorMaskForType(SDNode *N, EVT ScalarTy) {
1042 if (!ScalarTy.isSimple())
1043 return false;
1044
1045 uint64_t MaskForTy = 0ULL;
1046 switch (ScalarTy.getSimpleVT().SimpleTy) {
1047 case MVT::i8:
1048 MaskForTy = 0xFFULL;
1049 break;
1050 case MVT::i16:
1051 MaskForTy = 0xFFFFULL;
1052 break;
1053 case MVT::i32:
1054 MaskForTy = 0xFFFFFFFFULL;
1055 break;
1056 default:
1057 return false;
1058 break;
1059 }
1060
1061 APInt Val;
1062 if (ISD::isConstantSplatVector(N, SplatValue&: Val))
1063 return Val.getLimitedValue() == MaskForTy;
1064
1065 return false;
1066}
1067
1068// Determines if it is a constant integer or a splat/build vector of constant
1069// integers (and undefs).
1070// Do not permit build vector implicit truncation unless AllowTruncation is set.
1071static bool isConstantOrConstantVector(SDValue N, bool NoOpaques = false,
1072 bool AllowTruncation = false) {
1073 if (ConstantSDNode *Const = dyn_cast<ConstantSDNode>(Val&: N))
1074 return !(Const->isOpaque() && NoOpaques);
1075 if (N.getOpcode() != ISD::BUILD_VECTOR && N.getOpcode() != ISD::SPLAT_VECTOR)
1076 return false;
1077 unsigned BitWidth = N.getScalarValueSizeInBits();
1078 for (const SDValue &Op : N->op_values()) {
1079 if (Op.isUndef())
1080 continue;
1081 ConstantSDNode *Const = dyn_cast<ConstantSDNode>(Val: Op);
1082 if (!Const || (Const->isOpaque() && NoOpaques))
1083 return false;
1084 // When AllowTruncation is true, allow constants that have been promoted
1085 // during type legalization as long as the value fits in the target type.
1086 if ((AllowTruncation &&
1087 Const->getAPIntValue().getActiveBits() > BitWidth) ||
1088 (!AllowTruncation && Const->getAPIntValue().getBitWidth() != BitWidth))
1089 return false;
1090 }
1091 return true;
1092}
1093
1094// Determines if a BUILD_VECTOR is composed of all-constants possibly mixed with
1095// undef's.
1096static bool isAnyConstantBuildVector(SDValue V, bool NoOpaques = false) {
1097 if (V.getOpcode() != ISD::BUILD_VECTOR)
1098 return false;
1099 return isConstantOrConstantVector(N: V, NoOpaques) ||
1100 ISD::isBuildVectorOfConstantFPSDNodes(N: V.getNode());
1101}
1102
1103// Determine if this an indexed load with an opaque target constant index.
1104static bool canSplitIdx(LoadSDNode *LD) {
1105 return MaySplitLoadIndex &&
1106 (LD->getOperand(Num: 2).getOpcode() != ISD::TargetConstant ||
1107 !cast<ConstantSDNode>(Val: LD->getOperand(Num: 2))->isOpaque());
1108}
1109
1110bool DAGCombiner::reassociationCanBreakAddressingModePattern(unsigned Opc,
1111 const SDLoc &DL,
1112 SDNode *N,
1113 SDValue N0,
1114 SDValue N1) {
1115 // Currently this only tries to ensure we don't undo the GEP splits done by
1116 // CodeGenPrepare when shouldConsiderGEPOffsetSplit is true. To ensure this,
1117 // we check if the following transformation would be problematic:
1118 // (load/store (add, (add, x, offset1), offset2)) ->
1119 // (load/store (add, x, offset1+offset2)).
1120
1121 // (load/store (add, (add, x, y), offset2)) ->
1122 // (load/store (add, (add, x, offset2), y)).
1123
1124 if (!N0.isAnyAdd())
1125 return false;
1126
1127 // Check for vscale addressing modes.
1128 // (load/store (add/sub (add x, y), vscale))
1129 // (load/store (add/sub (add x, y), (lsl vscale, C)))
1130 // (load/store (add/sub (add x, y), (mul vscale, C)))
1131 if ((N1.getOpcode() == ISD::VSCALE ||
1132 ((N1.getOpcode() == ISD::SHL || N1.getOpcode() == ISD::MUL) &&
1133 N1.getOperand(i: 0).getOpcode() == ISD::VSCALE &&
1134 isa<ConstantSDNode>(Val: N1.getOperand(i: 1)))) &&
1135 N1.getValueType().getFixedSizeInBits() <= 64) {
1136 int64_t ScalableOffset = N1.getOpcode() == ISD::VSCALE
1137 ? N1.getConstantOperandVal(i: 0)
1138 : (N1.getOperand(i: 0).getConstantOperandVal(i: 0) *
1139 (N1.getOpcode() == ISD::SHL
1140 ? (1LL << N1.getConstantOperandVal(i: 1))
1141 : N1.getConstantOperandVal(i: 1)));
1142 if (Opc == ISD::SUB)
1143 ScalableOffset = -ScalableOffset;
1144 if (all_of(Range: N->users(), P: [&](SDNode *Node) {
1145 if (auto *LoadStore = dyn_cast<MemSDNode>(Val: Node);
1146 LoadStore && LoadStore->hasUniqueMemOperand() &&
1147 LoadStore->getBasePtr().getNode() == N) {
1148 TargetLoweringBase::AddrMode AM;
1149 AM.HasBaseReg = true;
1150 AM.ScalableOffset = ScalableOffset;
1151 EVT VT = LoadStore->getMemoryVT();
1152 unsigned AS = LoadStore->getAddressSpace();
1153 Type *AccessTy = VT.getTypeForEVT(Context&: *DAG.getContext());
1154 return TLI.isLegalAddressingMode(DL: DAG.getDataLayout(), AM, Ty: AccessTy,
1155 AddrSpace: AS);
1156 }
1157 return false;
1158 }))
1159 return true;
1160 }
1161
1162 if (Opc != ISD::ADD && Opc != ISD::PTRADD)
1163 return false;
1164
1165 auto *C2 = dyn_cast<ConstantSDNode>(Val&: N1);
1166 if (!C2)
1167 return false;
1168
1169 const APInt &C2APIntVal = C2->getAPIntValue();
1170 if (C2APIntVal.getSignificantBits() > 64)
1171 return false;
1172
1173 if (auto *C1 = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1))) {
1174 if (N0.hasOneUse())
1175 return false;
1176
1177 const APInt &C1APIntVal = C1->getAPIntValue();
1178 const APInt CombinedValueIntVal = C1APIntVal + C2APIntVal;
1179 if (CombinedValueIntVal.getSignificantBits() > 64)
1180 return false;
1181 const int64_t CombinedValue = CombinedValueIntVal.getSExtValue();
1182
1183 for (SDNode *Node : N->users()) {
1184 if (auto *LoadStore = dyn_cast<MemSDNode>(Val: Node)) {
1185 if (!LoadStore->hasUniqueMemOperand())
1186 continue;
1187 // Is x[offset2] already not a legal addressing mode? If so then
1188 // reassociating the constants breaks nothing (we test offset2 because
1189 // that's the one we hope to fold into the load or store).
1190 TargetLoweringBase::AddrMode AM;
1191 AM.HasBaseReg = true;
1192 AM.BaseOffs = C2APIntVal.getSExtValue();
1193 EVT VT = LoadStore->getMemoryVT();
1194 unsigned AS = LoadStore->getAddressSpace();
1195 Type *AccessTy = VT.getTypeForEVT(Context&: *DAG.getContext());
1196 if (!TLI.isLegalAddressingMode(DL: DAG.getDataLayout(), AM, Ty: AccessTy, AddrSpace: AS))
1197 continue;
1198
1199 // Would x[offset1+offset2] still be a legal addressing mode?
1200 AM.BaseOffs = CombinedValue;
1201 if (!TLI.isLegalAddressingMode(DL: DAG.getDataLayout(), AM, Ty: AccessTy, AddrSpace: AS))
1202 return true;
1203 }
1204 }
1205 } else {
1206 if (auto *GA = dyn_cast<GlobalAddressSDNode>(Val: N0.getOperand(i: 1)))
1207 if (GA->getOpcode() == ISD::GlobalAddress && TLI.isOffsetFoldingLegal(GA))
1208 return false;
1209
1210 for (SDNode *Node : N->users()) {
1211 auto *LoadStore = dyn_cast<MemSDNode>(Val: Node);
1212 if (!LoadStore || !LoadStore->hasUniqueMemOperand())
1213 return false;
1214
1215 // Is x[offset2] a legal addressing mode? If so then
1216 // reassociating the constants breaks address pattern
1217 TargetLoweringBase::AddrMode AM;
1218 AM.HasBaseReg = true;
1219 AM.BaseOffs = C2APIntVal.getSExtValue();
1220 EVT VT = LoadStore->getMemoryVT();
1221 unsigned AS = LoadStore->getAddressSpace();
1222 Type *AccessTy = VT.getTypeForEVT(Context&: *DAG.getContext());
1223 if (!TLI.isLegalAddressingMode(DL: DAG.getDataLayout(), AM, Ty: AccessTy, AddrSpace: AS))
1224 return false;
1225 }
1226 return true;
1227 }
1228
1229 return false;
1230}
1231
1232/// Helper for DAGCombiner::reassociateOps. Try to reassociate (Opc N0, N1) if
1233/// \p N0 is the same kind of operation as \p Opc.
1234SDValue DAGCombiner::reassociateOpsCommutative(unsigned Opc, const SDLoc &DL,
1235 SDValue N0, SDValue N1,
1236 SDNodeFlags Flags) {
1237 EVT VT = N0.getValueType();
1238
1239 if (N0.getOpcode() != Opc)
1240 return SDValue();
1241
1242 SDValue N00 = N0.getOperand(i: 0);
1243 SDValue N01 = N0.getOperand(i: 1);
1244
1245 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N01)) {
1246 SDNodeFlags NewFlags;
1247 if (N0.getOpcode() == ISD::ADD && N0->getFlags().hasNoUnsignedWrap() &&
1248 Flags.hasNoUnsignedWrap())
1249 NewFlags |= SDNodeFlags::NoUnsignedWrap;
1250
1251 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N1)) {
1252 // Reassociate: (op (op x, c1), c2) -> (op x, (op c1, c2))
1253 if (SDValue OpNode = DAG.FoldConstantArithmetic(Opcode: Opc, DL, VT, Ops: {N01, N1})) {
1254 NewFlags.setDisjoint(Flags.hasDisjoint() &&
1255 N0->getFlags().hasDisjoint());
1256 return DAG.getNode(Opcode: Opc, DL, VT, N1: N00, N2: OpNode, Flags: NewFlags);
1257 }
1258 return SDValue();
1259 }
1260 if (TLI.isReassocProfitable(DAG, N0, N1)) {
1261 // Reassociate: (op (op x, c1), y) -> (op (op x, y), c1)
1262 // iff (op x, c1) has one use
1263 SDValue OpNode = DAG.getNode(Opcode: Opc, DL: SDLoc(N0), VT, N1: N00, N2: N1, Flags: NewFlags);
1264 return DAG.getNode(Opcode: Opc, DL, VT, N1: OpNode, N2: N01, Flags: NewFlags);
1265 }
1266 }
1267
1268 // Check for repeated operand logic simplifications.
1269 if (Opc == ISD::AND || Opc == ISD::OR) {
1270 // (N00 & N01) & N00 --> N00 & N01
1271 // (N00 & N01) & N01 --> N00 & N01
1272 // (N00 | N01) | N00 --> N00 | N01
1273 // (N00 | N01) | N01 --> N00 | N01
1274 if (N1 == N00 || N1 == N01)
1275 return N0;
1276 }
1277 if (Opc == ISD::XOR) {
1278 // (N00 ^ N01) ^ N00 --> N01
1279 if (N1 == N00)
1280 return N01;
1281 // (N00 ^ N01) ^ N01 --> N00
1282 if (N1 == N01)
1283 return N00;
1284 }
1285
1286 if (TLI.isReassocProfitable(DAG, N0, N1)) {
1287 if (N1 != N01) {
1288 // Reassociate if (op N00, N1) already exist
1289 if (SDNode *NE = DAG.getNodeIfExists(Opcode: Opc, VTList: DAG.getVTList(VT), Ops: {N00, N1})) {
1290 // if Op (Op N00, N1), N01 already exist
1291 // we need to stop reassciate to avoid dead loop
1292 if (!DAG.doesNodeExist(Opcode: Opc, VTList: DAG.getVTList(VT), Ops: {SDValue(NE, 0), N01}))
1293 return DAG.getNode(Opcode: Opc, DL, VT, N1: SDValue(NE, 0), N2: N01);
1294 }
1295 }
1296
1297 if (N1 != N00) {
1298 // Reassociate if (op N01, N1) already exist
1299 if (SDNode *NE = DAG.getNodeIfExists(Opcode: Opc, VTList: DAG.getVTList(VT), Ops: {N01, N1})) {
1300 // if Op (Op N01, N1), N00 already exist
1301 // we need to stop reassciate to avoid dead loop
1302 if (!DAG.doesNodeExist(Opcode: Opc, VTList: DAG.getVTList(VT), Ops: {SDValue(NE, 0), N00}))
1303 return DAG.getNode(Opcode: Opc, DL, VT, N1: SDValue(NE, 0), N2: N00);
1304 }
1305 }
1306
1307 // Reassociate the operands from (OR/AND (OR/AND(N00, N001)), N1) to (OR/AND
1308 // (OR/AND(N00, N1)), N01) when N00 and N1 are comparisons with the same
1309 // predicate or to (OR/AND (OR/AND(N1, N01)), N00) when N01 and N1 are
1310 // comparisons with the same predicate. This enables optimizations as the
1311 // following one:
1312 // CMP(A,C)||CMP(B,C) => CMP(MIN/MAX(A,B), C)
1313 // CMP(A,C)&&CMP(B,C) => CMP(MIN/MAX(A,B), C)
1314 if (Opc == ISD::AND || Opc == ISD::OR) {
1315 if (N1->getOpcode() == ISD::SETCC && N00->getOpcode() == ISD::SETCC &&
1316 N01->getOpcode() == ISD::SETCC) {
1317 ISD::CondCode CC1 = cast<CondCodeSDNode>(Val: N1.getOperand(i: 2))->get();
1318 ISD::CondCode CC00 = cast<CondCodeSDNode>(Val: N00.getOperand(i: 2))->get();
1319 ISD::CondCode CC01 = cast<CondCodeSDNode>(Val: N01.getOperand(i: 2))->get();
1320 if (CC1 == CC00 && CC1 != CC01) {
1321 SDValue OpNode = DAG.getNode(Opcode: Opc, DL: SDLoc(N0), VT, N1: N00, N2: N1, Flags);
1322 return DAG.getNode(Opcode: Opc, DL, VT, N1: OpNode, N2: N01, Flags);
1323 }
1324 if (CC1 == CC01 && CC1 != CC00) {
1325 SDValue OpNode = DAG.getNode(Opcode: Opc, DL: SDLoc(N0), VT, N1: N01, N2: N1, Flags);
1326 return DAG.getNode(Opcode: Opc, DL, VT, N1: OpNode, N2: N00, Flags);
1327 }
1328 }
1329 }
1330 }
1331
1332 return SDValue();
1333}
1334
1335/// Try to reassociate commutative (Opc N0, N1) if either \p N0 or \p N1 is the
1336/// same kind of operation as \p Opc.
1337SDValue DAGCombiner::reassociateOps(unsigned Opc, const SDLoc &DL, SDValue N0,
1338 SDValue N1, SDNodeFlags Flags) {
1339 assert(TLI.isCommutativeBinOp(Opc) && "Operation not commutative.");
1340
1341 // Floating-point reassociation is not allowed without loose FP math.
1342 if (N0.getValueType().isFloatingPoint() ||
1343 N1.getValueType().isFloatingPoint())
1344 if (!Flags.hasAllowReassociation() || !Flags.hasNoSignedZeros())
1345 return SDValue();
1346
1347 if (SDValue Combined = reassociateOpsCommutative(Opc, DL, N0, N1, Flags))
1348 return Combined;
1349 if (SDValue Combined = reassociateOpsCommutative(Opc, DL, N0: N1, N1: N0, Flags))
1350 return Combined;
1351 return SDValue();
1352}
1353
1354// Try to fold Opc(vecreduce(x), vecreduce(y)) -> vecreduce(Opc(x, y))
1355// Note that we only expect Flags to be passed from FP operations. For integer
1356// operations they need to be dropped.
1357SDValue DAGCombiner::reassociateReduction(unsigned RedOpc, unsigned Opc,
1358 const SDLoc &DL, EVT VT, SDValue N0,
1359 SDValue N1, SDNodeFlags Flags) {
1360 if (N0.getOpcode() == RedOpc && N1.getOpcode() == RedOpc &&
1361 N0.getOperand(i: 0).getValueType() == N1.getOperand(i: 0).getValueType() &&
1362 N0->hasOneUse() && N1->hasOneUse() &&
1363 TLI.isOperationLegalOrCustom(Op: Opc, VT: N0.getOperand(i: 0).getValueType()) &&
1364 TLI.shouldReassociateReduction(RedOpc, VT: N0.getOperand(i: 0).getValueType())) {
1365 SelectionDAG::FlagInserter FlagsInserter(DAG, Flags);
1366 return DAG.getNode(Opcode: RedOpc, DL, VT,
1367 Operand: DAG.getNode(Opcode: Opc, DL, VT: N0.getOperand(i: 0).getValueType(),
1368 N1: N0.getOperand(i: 0), N2: N1.getOperand(i: 0)));
1369 }
1370
1371 // Reassociate op(op(vecreduce(a), b), op(vecreduce(c), d)) into
1372 // op(vecreduce(op(a, c)), op(b, d)), to combine the reductions into a
1373 // single node.
1374 SDValue A, B, C, D, RedA, RedB;
1375 if (sd_match(N: N0,
1376 P: m_OneUse(P: m_c_BinOp(
1377 Opc, L: m_Value(N&: RedA, P: m_OneUse(P: m_UnaryOp(Opc: RedOpc, Op: m_Value(N&: A)))),
1378 R: m_Value(N&: B, P: m_Unless(P: m_UnaryOp(Opc: RedOpc, Op: m_Value())))))) &&
1379 sd_match(N: N1,
1380 P: m_OneUse(P: m_c_BinOp(
1381 Opc, L: m_Value(N&: RedB, P: m_OneUse(P: m_UnaryOp(Opc: RedOpc, Op: m_Value(N&: C)))),
1382 R: m_Value(N&: D, P: m_Unless(P: m_UnaryOp(Opc: RedOpc, Op: m_Value())))))) &&
1383 A.getValueType() == C.getValueType() &&
1384 hasOperation(Opcode: Opc, VT: A.getValueType()) &&
1385 TLI.shouldReassociateReduction(RedOpc, VT)) {
1386 if ((Opc == ISD::FADD || Opc == ISD::FMUL) &&
1387 (!N0->getFlags().hasAllowReassociation() ||
1388 !N1->getFlags().hasAllowReassociation() ||
1389 !RedA->getFlags().hasAllowReassociation() ||
1390 !RedB->getFlags().hasAllowReassociation()))
1391 return SDValue();
1392 SelectionDAG::FlagInserter FlagsInserter(
1393 DAG, Flags & N0->getFlags() & N1->getFlags() & RedA->getFlags() &
1394 RedB->getFlags());
1395 SDValue Op = DAG.getNode(Opcode: Opc, DL, VT: A.getValueType(), N1: A, N2: C);
1396 SDValue Red = DAG.getNode(Opcode: RedOpc, DL, VT, Operand: Op);
1397 SDValue Op2 = DAG.getNode(Opcode: Opc, DL, VT, N1: B, N2: D);
1398 return DAG.getNode(Opcode: Opc, DL, VT, N1: Red, N2: Op2);
1399 }
1400
1401 // Reassociate a reduction chain so two reductions become adjacent and the
1402 // folds above can merge them:
1403 // op(vecreduce(X), op(vecreduce(Y), Z))
1404 // -> op(vecreduce(op(X, Y)), Z)
1405 // Applied to fixpoint by the combiner worklist, this collapses an
1406 // arbitrarily long chain of reductions (such as the left-leaning chain SLP
1407 // emits) into a single reduction.
1408 auto FoldReductionChain = [&](SDValue Red0, SDValue Chain) -> SDValue {
1409 SDValue X, Y, Z, RedY;
1410 if (!sd_match(N: Red0, P: m_OneUse(P: m_UnaryOp(Opc: RedOpc, Op: m_Value(N&: X)))) ||
1411 !sd_match(
1412 N: Chain,
1413 P: m_OneUse(P: m_c_BinOp(
1414 Opc, L: m_Value(N&: RedY, P: m_OneUse(P: m_UnaryOp(Opc: RedOpc, Op: m_Value(N&: Y)))),
1415 R: m_Value(N&: Z, P: m_Unless(P: m_UnaryOp(Opc: RedOpc, Op: m_Value())))))) ||
1416 X.getValueType() != Y.getValueType() ||
1417 !hasOperation(Opcode: Opc, VT: X.getValueType()) ||
1418 !TLI.shouldReassociateReduction(RedOpc, VT))
1419 return SDValue();
1420 if ((Opc == ISD::FADD || Opc == ISD::FMUL) &&
1421 (!Chain->getFlags().hasAllowReassociation() ||
1422 !Red0->getFlags().hasAllowReassociation() ||
1423 !RedY->getFlags().hasAllowReassociation()))
1424 return SDValue();
1425 SelectionDAG::FlagInserter FlagsInserter(
1426 DAG, Flags & Chain->getFlags() & Red0->getFlags() & RedY->getFlags());
1427 SDValue Op = DAG.getNode(Opcode: Opc, DL, VT: X.getValueType(), N1: X, N2: Y);
1428 SDValue Red = DAG.getNode(Opcode: RedOpc, DL, VT, Operand: Op);
1429 return DAG.getNode(Opcode: Opc, DL, VT, N1: Red, N2: Z);
1430 };
1431 if (SDValue V = FoldReductionChain(N0, N1))
1432 return V;
1433 if (SDValue V = FoldReductionChain(N1, N0))
1434 return V;
1435
1436 return SDValue();
1437}
1438
1439SDValue DAGCombiner::CombineTo(SDNode *N, const SDValue *To, unsigned NumTo,
1440 bool AddTo) {
1441 assert(N->getNumValues() == NumTo && "Broken CombineTo call!");
1442 ++NodesCombined;
1443 LLVM_DEBUG(dbgs() << "\nReplacing.1 "; N->dump(&DAG); dbgs() << "\nWith: ";
1444 To[0].dump(&DAG);
1445 dbgs() << " and " << NumTo - 1 << " other values\n");
1446 for (unsigned i = 0, e = NumTo; i != e; ++i)
1447 assert((!To[i].getNode() ||
1448 N->getValueType(i) == To[i].getValueType()) &&
1449 "Cannot combine value to value of different type!");
1450
1451 WorklistRemover DeadNodes(*this);
1452 DAG.ReplaceAllUsesWith(From: N, To);
1453 if (AddTo) {
1454 // Push the new nodes and any users onto the worklist
1455 for (unsigned i = 0, e = NumTo; i != e; ++i) {
1456 if (To[i].getNode())
1457 AddToWorklistWithUsers(N: To[i].getNode());
1458 }
1459 }
1460
1461 // Finally, if the node is now dead, remove it from the graph. The node
1462 // may not be dead if the replacement process recursively simplified to
1463 // something else needing this node.
1464 if (N->use_empty())
1465 deleteAndRecombine(N);
1466 return SDValue(N, 0);
1467}
1468
1469void DAGCombiner::
1470CommitTargetLoweringOpt(const TargetLowering::TargetLoweringOpt &TLO) {
1471 // Replace the old value with the new one.
1472 ++NodesCombined;
1473 LLVM_DEBUG(dbgs() << "\nReplacing.2 "; TLO.Old.dump(&DAG);
1474 dbgs() << "\nWith: "; TLO.New.dump(&DAG); dbgs() << '\n');
1475
1476 // Replace all uses.
1477 DAG.ReplaceAllUsesOfValueWith(From: TLO.Old, To: TLO.New);
1478
1479 // Push the new node and any (possibly new) users onto the worklist.
1480 AddToWorklistWithUsers(N: TLO.New.getNode());
1481
1482 // Finally, if the node is now dead, remove it from the graph.
1483 recursivelyDeleteUnusedNodes(N: TLO.Old.getNode());
1484}
1485
1486/// Check the specified integer node value to see if it can be simplified or if
1487/// things it uses can be simplified by bit propagation. If so, return true.
1488bool DAGCombiner::SimplifyDemandedBits(SDValue Op, const APInt &DemandedBits,
1489 const APInt &DemandedElts,
1490 bool AssumeSingleUse) {
1491 TargetLowering::TargetLoweringOpt TLO(DAG, LegalTypes, LegalOperations);
1492 KnownBits Known;
1493 if (!TLI.SimplifyDemandedBits(Op, DemandedBits, DemandedElts, Known, TLO, Depth: 0,
1494 AssumeSingleUse))
1495 return false;
1496
1497 // Revisit the node.
1498 AddToWorklist(N: Op.getNode());
1499
1500 CommitTargetLoweringOpt(TLO);
1501 return true;
1502}
1503
1504/// Check the specified vector node value to see if it can be simplified or
1505/// if things it uses can be simplified as it only uses some of the elements.
1506/// If so, return true.
1507bool DAGCombiner::SimplifyDemandedVectorElts(SDValue Op,
1508 const APInt &DemandedElts,
1509 bool AssumeSingleUse) {
1510 TargetLowering::TargetLoweringOpt TLO(DAG, LegalTypes, LegalOperations);
1511 APInt KnownUndef, KnownZero;
1512 if (!TLI.SimplifyDemandedVectorElts(Op, DemandedEltMask: DemandedElts, KnownUndef, KnownZero,
1513 TLO, Depth: 0, AssumeSingleUse))
1514 return false;
1515
1516 // Revisit the node.
1517 AddToWorklist(N: Op.getNode());
1518
1519 CommitTargetLoweringOpt(TLO);
1520 return true;
1521}
1522
1523void DAGCombiner::ReplaceLoadWithPromotedLoad(SDNode *Load, SDNode *ExtLoad) {
1524 SDLoc DL(Load);
1525 EVT VT = Load->getValueType(ResNo: 0);
1526 SDValue Trunc = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: SDValue(ExtLoad, 0));
1527
1528 LLVM_DEBUG(dbgs() << "\nReplacing.9 "; Load->dump(&DAG); dbgs() << "\nWith: ";
1529 Trunc.dump(&DAG); dbgs() << '\n');
1530
1531 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Load, 0), To: Trunc);
1532 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Load, 1), To: SDValue(ExtLoad, 1));
1533
1534 AddToWorklist(N: Trunc.getNode());
1535 recursivelyDeleteUnusedNodes(N: Load);
1536}
1537
1538SDValue DAGCombiner::PromoteOperand(SDValue Op, EVT PVT, bool &Replace) {
1539 Replace = false;
1540 SDLoc DL(Op);
1541 if (ISD::isUNINDEXEDLoad(N: Op.getNode())) {
1542 LoadSDNode *LD = cast<LoadSDNode>(Val&: Op);
1543 EVT MemVT = LD->getMemoryVT();
1544 ISD::LoadExtType ExtType = ISD::isNON_EXTLoad(N: LD) ? ISD::EXTLOAD
1545 : LD->getExtensionType();
1546 Replace = true;
1547 return DAG.getExtLoad(ExtType, dl: DL, VT: PVT,
1548 Chain: LD->getChain(), Ptr: LD->getBasePtr(),
1549 MemVT, MMO: LD->getMemOperand());
1550 }
1551
1552 unsigned Opc = Op.getOpcode();
1553 switch (Opc) {
1554 default: break;
1555 case ISD::AssertSext:
1556 if (SDValue Op0 = SExtPromoteOperand(Op: Op.getOperand(i: 0), PVT))
1557 return DAG.getNode(Opcode: ISD::AssertSext, DL, VT: PVT, N1: Op0, N2: Op.getOperand(i: 1));
1558 break;
1559 case ISD::AssertZext:
1560 if (SDValue Op0 = ZExtPromoteOperand(Op: Op.getOperand(i: 0), PVT))
1561 return DAG.getNode(Opcode: ISD::AssertZext, DL, VT: PVT, N1: Op0, N2: Op.getOperand(i: 1));
1562 break;
1563 case ISD::Constant: {
1564 unsigned ExtOpc =
1565 Op.getValueType().isByteSized() ? ISD::SIGN_EXTEND : ISD::ZERO_EXTEND;
1566 return DAG.getNode(Opcode: ExtOpc, DL, VT: PVT, Operand: Op);
1567 }
1568 }
1569
1570 if (!TLI.isOperationLegal(Op: ISD::ANY_EXTEND, VT: PVT))
1571 return SDValue();
1572 return DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT: PVT, Operand: Op);
1573}
1574
1575SDValue DAGCombiner::SExtPromoteOperand(SDValue Op, EVT PVT) {
1576 if (!TLI.isOperationLegal(Op: ISD::SIGN_EXTEND_INREG, VT: PVT))
1577 return SDValue();
1578 EVT OldVT = Op.getValueType();
1579 SDLoc DL(Op);
1580 bool Replace = false;
1581 SDValue NewOp = PromoteOperand(Op, PVT, Replace);
1582 if (!NewOp.getNode())
1583 return SDValue();
1584 AddToWorklist(N: NewOp.getNode());
1585
1586 if (Replace)
1587 ReplaceLoadWithPromotedLoad(Load: Op.getNode(), ExtLoad: NewOp.getNode());
1588 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT: NewOp.getValueType(), N1: NewOp,
1589 N2: DAG.getValueType(OldVT));
1590}
1591
1592SDValue DAGCombiner::ZExtPromoteOperand(SDValue Op, EVT PVT) {
1593 EVT OldVT = Op.getValueType();
1594 SDLoc DL(Op);
1595 bool Replace = false;
1596 SDValue NewOp = PromoteOperand(Op, PVT, Replace);
1597 if (!NewOp.getNode())
1598 return SDValue();
1599 AddToWorklist(N: NewOp.getNode());
1600
1601 if (Replace)
1602 ReplaceLoadWithPromotedLoad(Load: Op.getNode(), ExtLoad: NewOp.getNode());
1603 return DAG.getZeroExtendInReg(Op: NewOp, DL, VT: OldVT);
1604}
1605
1606/// Promote the specified integer binary operation if the target indicates it is
1607/// beneficial. e.g. On x86, it's usually better to promote i16 operations to
1608/// i32 since i16 instructions are longer.
1609SDValue DAGCombiner::PromoteIntBinOp(SDValue Op) {
1610 if (!LegalOperations)
1611 return SDValue();
1612
1613 EVT VT = Op.getValueType();
1614 if (VT.isVector() || !VT.isInteger())
1615 return SDValue();
1616
1617 // If operation type is 'undesirable', e.g. i16 on x86, consider
1618 // promoting it.
1619 unsigned Opc = Op.getOpcode();
1620 if (TLI.isTypeDesirableForOp(Opc, VT))
1621 return SDValue();
1622
1623 EVT PVT = VT;
1624 // Consult target whether it is a good idea to promote this operation and
1625 // what's the right type to promote it to.
1626 if (TLI.IsDesirableToPromoteOp(Op, PVT)) {
1627 assert(PVT != VT && "Don't know what type to promote to!");
1628
1629 LLVM_DEBUG(dbgs() << "\nPromoting "; Op.dump(&DAG));
1630
1631 bool Replace0 = false;
1632 SDValue N0 = Op.getOperand(i: 0);
1633 SDValue NN0 = PromoteOperand(Op: N0, PVT, Replace&: Replace0);
1634
1635 bool Replace1 = false;
1636 SDValue N1 = Op.getOperand(i: 1);
1637 SDValue NN1 = PromoteOperand(Op: N1, PVT, Replace&: Replace1);
1638 SDLoc DL(Op);
1639
1640 SDValue RV =
1641 DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: DAG.getNode(Opcode: Opc, DL, VT: PVT, N1: NN0, N2: NN1));
1642
1643 // We are always replacing N0/N1's use in N and only need additional
1644 // replacements if there are additional uses.
1645 // Note: We are checking uses of the *nodes* (SDNode) rather than values
1646 // (SDValue) here because the node may reference multiple values
1647 // (for example, the chain value of a load node).
1648 Replace0 &= !N0->hasOneUse();
1649 Replace1 &= (N0 != N1) && !N1->hasOneUse();
1650
1651 // Combine Op here so it is preserved past replacements.
1652 CombineTo(N: Op.getNode(), Res: RV);
1653
1654 // If operands have a use ordering, make sure we deal with
1655 // predecessor first.
1656 if (Replace0 && Replace1 && N0->isPredecessorOf(N: N1.getNode())) {
1657 std::swap(a&: N0, b&: N1);
1658 std::swap(a&: NN0, b&: NN1);
1659 }
1660
1661 if (Replace0) {
1662 AddToWorklist(N: NN0.getNode());
1663 ReplaceLoadWithPromotedLoad(Load: N0.getNode(), ExtLoad: NN0.getNode());
1664 }
1665 if (Replace1) {
1666 AddToWorklist(N: NN1.getNode());
1667 ReplaceLoadWithPromotedLoad(Load: N1.getNode(), ExtLoad: NN1.getNode());
1668 }
1669 return Op;
1670 }
1671 return SDValue();
1672}
1673
1674/// Promote the specified integer shift operation if the target indicates it is
1675/// beneficial. e.g. On x86, it's usually better to promote i16 operations to
1676/// i32 since i16 instructions are longer.
1677SDValue DAGCombiner::PromoteIntShiftOp(SDValue Op) {
1678 if (!LegalOperations)
1679 return SDValue();
1680
1681 EVT VT = Op.getValueType();
1682 if (VT.isVector() || !VT.isInteger())
1683 return SDValue();
1684
1685 // If operation type is 'undesirable', e.g. i16 on x86, consider
1686 // promoting it.
1687 unsigned Opc = Op.getOpcode();
1688 if (TLI.isTypeDesirableForOp(Opc, VT))
1689 return SDValue();
1690
1691 EVT PVT = VT;
1692 // Consult target whether it is a good idea to promote this operation and
1693 // what's the right type to promote it to.
1694 if (TLI.IsDesirableToPromoteOp(Op, PVT)) {
1695 assert(PVT != VT && "Don't know what type to promote to!");
1696
1697 LLVM_DEBUG(dbgs() << "\nPromoting "; Op.dump(&DAG));
1698
1699 SDNodeFlags TruncFlags;
1700 bool Replace = false;
1701 SDValue N0 = Op.getOperand(i: 0);
1702 if (Opc == ISD::SRA) {
1703 N0 = SExtPromoteOperand(Op: N0, PVT);
1704 } else if (Opc == ISD::SRL) {
1705 N0 = ZExtPromoteOperand(Op: N0, PVT);
1706 } else {
1707 if (Op->getFlags().hasNoUnsignedWrap()) {
1708 N0 = ZExtPromoteOperand(Op: N0, PVT);
1709 TruncFlags = SDNodeFlags::NoUnsignedWrap;
1710 } else if (Op->getFlags().hasNoSignedWrap()) {
1711 N0 = SExtPromoteOperand(Op: N0, PVT);
1712 TruncFlags = SDNodeFlags::NoSignedWrap;
1713 } else {
1714 N0 = PromoteOperand(Op: N0, PVT, Replace);
1715 }
1716 }
1717
1718 if (!N0.getNode())
1719 return SDValue();
1720
1721 SDLoc DL(Op);
1722 SDValue N1 = Op.getOperand(i: 1);
1723 SDValue RV = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT,
1724 Operand: DAG.getNode(Opcode: Opc, DL, VT: PVT, N1: N0, N2: N1), Flags: TruncFlags);
1725
1726 if (Replace)
1727 ReplaceLoadWithPromotedLoad(Load: Op.getOperand(i: 0).getNode(), ExtLoad: N0.getNode());
1728
1729 // Deal with Op being deleted.
1730 if (Op && Op.getOpcode() != ISD::DELETED_NODE)
1731 return RV;
1732 }
1733 return SDValue();
1734}
1735
1736SDValue DAGCombiner::PromoteExtend(SDValue Op) {
1737 if (!LegalOperations)
1738 return SDValue();
1739
1740 EVT VT = Op.getValueType();
1741 if (VT.isVector() || !VT.isInteger())
1742 return SDValue();
1743
1744 // If operation type is 'undesirable', e.g. i16 on x86, consider
1745 // promoting it.
1746 unsigned Opc = Op.getOpcode();
1747 if (TLI.isTypeDesirableForOp(Opc, VT))
1748 return SDValue();
1749
1750 EVT PVT = VT;
1751 // Consult target whether it is a good idea to promote this operation and
1752 // what's the right type to promote it to.
1753 if (TLI.IsDesirableToPromoteOp(Op, PVT)) {
1754 assert(PVT != VT && "Don't know what type to promote to!");
1755 // fold (aext (aext x)) -> (aext x)
1756 // fold (aext (zext x)) -> (zext x)
1757 // fold (aext (sext x)) -> (sext x)
1758 LLVM_DEBUG(dbgs() << "\nPromoting "; Op.dump(&DAG));
1759 return DAG.getNode(Opcode: Op.getOpcode(), DL: SDLoc(Op), VT, Operand: Op.getOperand(i: 0));
1760 }
1761 return SDValue();
1762}
1763
1764bool DAGCombiner::PromoteLoad(SDValue Op) {
1765 if (!LegalOperations)
1766 return false;
1767
1768 if (!ISD::isUNINDEXEDLoad(N: Op.getNode()))
1769 return false;
1770
1771 EVT VT = Op.getValueType();
1772 if (VT.isVector() || !VT.isInteger())
1773 return false;
1774
1775 // If operation type is 'undesirable', e.g. i16 on x86, consider
1776 // promoting it.
1777 unsigned Opc = Op.getOpcode();
1778 if (TLI.isTypeDesirableForOp(Opc, VT))
1779 return false;
1780
1781 EVT PVT = VT;
1782 // Consult target whether it is a good idea to promote this operation and
1783 // what's the right type to promote it to.
1784 if (TLI.IsDesirableToPromoteOp(Op, PVT)) {
1785 assert(PVT != VT && "Don't know what type to promote to!");
1786
1787 SDLoc DL(Op);
1788 SDNode *N = Op.getNode();
1789 LoadSDNode *LD = cast<LoadSDNode>(Val: N);
1790 EVT MemVT = LD->getMemoryVT();
1791 ISD::LoadExtType ExtType = ISD::isNON_EXTLoad(N: LD) ? ISD::EXTLOAD
1792 : LD->getExtensionType();
1793 SDValue NewLD = DAG.getExtLoad(ExtType, dl: DL, VT: PVT,
1794 Chain: LD->getChain(), Ptr: LD->getBasePtr(),
1795 MemVT, MMO: LD->getMemOperand());
1796 SDValue Result = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: NewLD);
1797
1798 LLVM_DEBUG(dbgs() << "\nPromoting "; N->dump(&DAG); dbgs() << "\nTo: ";
1799 Result.dump(&DAG); dbgs() << '\n');
1800
1801 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 0), To: Result);
1802 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 1), To: NewLD.getValue(R: 1));
1803
1804 AddToWorklist(N: Result.getNode());
1805 recursivelyDeleteUnusedNodes(N);
1806 return true;
1807 }
1808
1809 return false;
1810}
1811
1812/// Recursively delete a node which has no uses and any operands for
1813/// which it is the only use.
1814///
1815/// Note that this both deletes the nodes and removes them from the worklist.
1816/// It also adds any nodes who have had a user deleted to the worklist as they
1817/// may now have only one use and subject to other combines.
1818bool DAGCombiner::recursivelyDeleteUnusedNodes(SDNode *N) {
1819 if (!N->use_empty())
1820 return false;
1821
1822 SmallSetVector<SDNode *, 16> Nodes;
1823 Nodes.insert(X: N);
1824 do {
1825 N = Nodes.pop_back_val();
1826 if (!N)
1827 continue;
1828
1829 if (N->use_empty()) {
1830 for (const SDValue &ChildN : N->op_values())
1831 Nodes.insert(X: ChildN.getNode());
1832
1833 removeFromWorklist(N);
1834 DAG.DeleteNode(N);
1835 } else {
1836 AddToWorklist(N);
1837 }
1838 } while (!Nodes.empty());
1839 return true;
1840}
1841
1842//===----------------------------------------------------------------------===//
1843// Main DAG Combiner implementation
1844//===----------------------------------------------------------------------===//
1845
1846void DAGCombiner::Run(CombineLevel AtLevel) {
1847 // set the instance variables, so that the various visit routines may use it.
1848 Level = AtLevel;
1849 LegalDAG = Level >= AfterLegalizeDAG;
1850 LegalOperations = Level >= AfterLegalizeVectorOps;
1851 LegalTypes = Level >= AfterLegalizeTypes;
1852
1853 bool UseTopologicalSorting = EnableTopologicalSorting.getNumOccurrences() > 0
1854 ? EnableTopologicalSorting
1855 : TLI.useTopologicalSorting();
1856
1857 WorklistInserter AddNodes(*this);
1858
1859 if (UseTopologicalSorting)
1860 DAG.AssignTopologicalOrder();
1861
1862 // Add all the dag nodes to the worklist.
1863 //
1864 // Note: All nodes are not added to PruningList here, this is because the only
1865 // nodes which can be deleted are those which have no uses and all other nodes
1866 // which would otherwise be added to the worklist by the first call to
1867 // getNextWorklistEntry are already present in it.
1868 if (UseTopologicalSorting) {
1869 for (SDNode &Node : reverse(C: DAG.allnodes()))
1870 AddToWorklist(N: &Node, /* IsCandidateForPruning */ Node.use_empty());
1871 } else {
1872 for (SDNode &Node : DAG.allnodes())
1873 AddToWorklist(N: &Node, /* IsCandidateForPruning */ Node.use_empty());
1874 }
1875
1876 // Create a dummy node (which is not added to allnodes), that adds a reference
1877 // to the root node, preventing it from being deleted, and tracking any
1878 // changes of the root.
1879 HandleSDNode Dummy(DAG.getRoot());
1880
1881 // While we have a valid worklist entry node, try to combine it.
1882 while (SDNode *N = getNextWorklistEntry()) {
1883 // If N has no uses, it is dead. Make sure to revisit all N's operands once
1884 // N is deleted from the DAG, since they too may now be dead or may have a
1885 // reduced number of uses, allowing other xforms.
1886 if (recursivelyDeleteUnusedNodes(N))
1887 continue;
1888
1889 WorklistRemover DeadNodes(*this);
1890
1891 // If this combine is running after legalizing the DAG, re-legalize any
1892 // nodes pulled off the worklist.
1893 if (LegalDAG) {
1894 SmallSetVector<SDNode *, 16> UpdatedNodes;
1895 bool NIsValid = DAG.LegalizeOp(N, UpdatedNodes);
1896
1897 for (SDNode *LN : UpdatedNodes)
1898 AddToWorklistWithUsers(N: LN);
1899
1900 if (!NIsValid)
1901 continue;
1902 }
1903
1904 LLVM_DEBUG(dbgs() << "\nCombining: "; N->dump(&DAG));
1905
1906 // Add any operands of the new node which have not yet been combined to the
1907 // worklist as well. getNextWorklistEntry flags nodes that have been
1908 // combined before. Because the worklist uniques things already, this won't
1909 // repeatedly process the same operand.
1910 for (const SDValue &ChildN : N->op_values())
1911 AddToWorklist(N: ChildN.getNode(), /*IsCandidateForPruning=*/true,
1912 /*SkipIfCombinedBefore=*/true);
1913
1914 SDValue RV = combine(N);
1915
1916 if (!RV.getNode())
1917 continue;
1918
1919 ++NodesCombined;
1920
1921 // Invalidate cached info.
1922 ChainsWithoutMergeableStores.clear();
1923
1924 // If we get back the same node we passed in, rather than a new node or
1925 // zero, we know that the node must have defined multiple values and
1926 // CombineTo was used. Since CombineTo takes care of the worklist
1927 // mechanics for us, we have no work to do in this case.
1928 if (RV.getNode() == N)
1929 continue;
1930
1931 assert(N->getOpcode() != ISD::DELETED_NODE &&
1932 RV.getOpcode() != ISD::DELETED_NODE &&
1933 "Node was deleted but visit returned new node!");
1934
1935 LLVM_DEBUG(dbgs() << " ... into: "; RV.dump(&DAG));
1936
1937 if (N->getNumValues() == RV->getNumValues())
1938 DAG.ReplaceAllUsesWith(From: N, To: RV.getNode());
1939 else {
1940 assert(N->getValueType(0) == RV.getValueType() &&
1941 N->getNumValues() == 1 && "Type mismatch");
1942 DAG.ReplaceAllUsesWith(From: N, To: &RV);
1943 }
1944
1945 // Push the new node and any users onto the worklist. Omit this if the
1946 // new node is the EntryToken (e.g. if a store managed to get optimized
1947 // out), because re-visiting the EntryToken and its users will not uncover
1948 // any additional opportunities, but there may be a large number of such
1949 // users, potentially causing compile time explosion.
1950 if (RV.getOpcode() != ISD::EntryToken)
1951 AddToWorklistWithUsers(N: RV.getNode());
1952
1953 // Finally, if the node is now dead, remove it from the graph. The node
1954 // may not be dead if the replacement process recursively simplified to
1955 // something else needing this node. This will also take care of adding any
1956 // operands which have lost a user to the worklist.
1957 recursivelyDeleteUnusedNodes(N);
1958 }
1959
1960 // If the root changed (e.g. it was a dead load, update the root).
1961 DAG.setRoot(Dummy.getValue());
1962 DAG.RemoveDeadNodes();
1963}
1964
1965SDValue DAGCombiner::visit(SDNode *N) {
1966 // clang-format off
1967 switch (N->getOpcode()) {
1968 default: break;
1969 case ISD::TokenFactor: return visitTokenFactor(N);
1970 case ISD::MERGE_VALUES: return visitMERGE_VALUES(N);
1971 case ISD::ADD: return visitADD(N);
1972 case ISD::PTRADD: return visitPTRADD(N);
1973 case ISD::SUB: return visitSUB(N);
1974 case ISD::SADDSAT:
1975 case ISD::UADDSAT: return visitADDSAT(N);
1976 case ISD::SSUBSAT:
1977 case ISD::USUBSAT: return visitSUBSAT(N);
1978 case ISD::ADDC: return visitADDC(N);
1979 case ISD::SADDO:
1980 case ISD::UADDO: return visitADDO(N);
1981 case ISD::SUBC: return visitSUBC(N);
1982 case ISD::SSUBO:
1983 case ISD::USUBO: return visitSUBO(N);
1984 case ISD::ADDE: return visitADDE(N);
1985 case ISD::UADDO_CARRY: return visitUADDO_CARRY(N);
1986 case ISD::SADDO_CARRY: return visitSADDO_CARRY(N);
1987 case ISD::SUBE: return visitSUBE(N);
1988 case ISD::USUBO_CARRY: return visitUSUBO_CARRY(N);
1989 case ISD::SSUBO_CARRY: return visitSSUBO_CARRY(N);
1990 case ISD::SMULFIX:
1991 case ISD::SMULFIXSAT:
1992 case ISD::UMULFIX:
1993 case ISD::UMULFIXSAT: return visitMULFIX(N);
1994 case ISD::MUL: return visitMUL(N);
1995 case ISD::SDIV: return visitSDIV(N);
1996 case ISD::UDIV: return visitUDIV(N);
1997 case ISD::SREM:
1998 case ISD::UREM: return visitREM(N);
1999 case ISD::MULHU: return visitMULHU(N);
2000 case ISD::MULHS: return visitMULHS(N);
2001 case ISD::AVGFLOORS:
2002 case ISD::AVGFLOORU:
2003 case ISD::AVGCEILS:
2004 case ISD::AVGCEILU: return visitAVG(N);
2005 case ISD::ABDS:
2006 case ISD::ABDU: return visitABD(N);
2007 case ISD::SMUL_LOHI: return visitSMUL_LOHI(N);
2008 case ISD::UMUL_LOHI: return visitUMUL_LOHI(N);
2009 case ISD::SMULO:
2010 case ISD::UMULO: return visitMULO(N);
2011 case ISD::SMIN:
2012 case ISD::SMAX:
2013 case ISD::UMIN:
2014 case ISD::UMAX: return visitIMINMAX(N);
2015 case ISD::AND: return visitAND(N);
2016 case ISD::OR: return visitOR(N);
2017 case ISD::XOR: return visitXOR(N);
2018 case ISD::SHL: return visitSHL(N);
2019 case ISD::SRA: return visitSRA(N);
2020 case ISD::SRL: return visitSRL(N);
2021 case ISD::ROTR:
2022 case ISD::ROTL: return visitRotate(N);
2023 case ISD::FSHL:
2024 case ISD::FSHR: return visitFunnelShift(N);
2025 case ISD::SSHLSAT:
2026 case ISD::USHLSAT: return visitSHLSAT(N);
2027 case ISD::ABS: return visitABS(N);
2028 case ISD::ABS_MIN_POISON: return visitABS_MIN_POISON(N);
2029 case ISD::CLMUL:
2030 case ISD::CLMULR:
2031 case ISD::CLMULH: return visitCLMUL(N);
2032 case ISD::PEXT: return visitPEXT(N);
2033 case ISD::PDEP: return visitPDEP(N);
2034 case ISD::BSWAP: return visitBSWAP(N);
2035 case ISD::BITREVERSE: return visitBITREVERSE(N);
2036 case ISD::CTLZ: return visitCTLZ(N);
2037 case ISD::CTLZ_ZERO_POISON: return visitCTLZ_ZERO_POISON(N);
2038 case ISD::CTTZ: return visitCTTZ(N);
2039 case ISD::CTTZ_ZERO_POISON: return visitCTTZ_ZERO_POISON(N);
2040 case ISD::CTPOP: return visitCTPOP(N);
2041 case ISD::PARITY: return visitPARITY(N);
2042 case ISD::SELECT: return visitSELECT(N);
2043 case ISD::VSELECT: return visitVSELECT(N);
2044 case ISD::SELECT_CC: return visitSELECT_CC(N);
2045 case ISD::SETCC: return visitSETCC(N);
2046 case ISD::SETCCCARRY: return visitSETCCCARRY(N);
2047 case ISD::SIGN_EXTEND: return visitSIGN_EXTEND(N);
2048 case ISD::ZERO_EXTEND: return visitZERO_EXTEND(N);
2049 case ISD::ANY_EXTEND: return visitANY_EXTEND(N);
2050 case ISD::AssertSext:
2051 case ISD::AssertZext: return visitAssertExt(N);
2052 case ISD::AssertAlign: return visitAssertAlign(N);
2053 case ISD::IS_FPCLASS: return visitIS_FPCLASS(N);
2054 case ISD::SIGN_EXTEND_INREG: return visitSIGN_EXTEND_INREG(N);
2055 case ISD::SIGN_EXTEND_VECTOR_INREG:
2056 case ISD::ZERO_EXTEND_VECTOR_INREG:
2057 case ISD::ANY_EXTEND_VECTOR_INREG: return visitEXTEND_VECTOR_INREG(N);
2058 case ISD::TRUNCATE: return visitTRUNCATE(N);
2059 case ISD::TRUNCATE_USAT_U: return visitTRUNCATE_USAT_U(N);
2060 case ISD::BITCAST: return visitBITCAST(N);
2061 case ISD::BUILD_PAIR: return visitBUILD_PAIR(N);
2062 case ISD::FADD: return visitFADD(N);
2063 case ISD::STRICT_FADD: return visitSTRICT_FADD(N);
2064 case ISD::FSUB: return visitFSUB(N);
2065 case ISD::FMUL: return visitFMUL(N);
2066 case ISD::FMA: return visitFMA(N);
2067 case ISD::FMAD: return visitFMAD(N);
2068 case ISD::FMULADD: return visitFMULADD(N);
2069 case ISD::FDIV: return visitFDIV(N);
2070 case ISD::FREM: return visitFREM(N);
2071 case ISD::FSQRT: return visitFSQRT(N);
2072 case ISD::FCOPYSIGN: return visitFCOPYSIGN(N);
2073 case ISD::FPOW: return visitFPOW(N);
2074 case ISD::SINT_TO_FP: return visitSINT_TO_FP(N);
2075 case ISD::UINT_TO_FP: return visitUINT_TO_FP(N);
2076 case ISD::FP_TO_SINT: return visitFP_TO_SINT(N);
2077 case ISD::FP_TO_UINT: return visitFP_TO_UINT(N);
2078 case ISD::LROUND:
2079 case ISD::LLROUND:
2080 case ISD::LRINT:
2081 case ISD::LLRINT: return visitXROUND(N);
2082 case ISD::FP_ROUND: return visitFP_ROUND(N);
2083 case ISD::FP_EXTEND: return visitFP_EXTEND(N);
2084 case ISD::FNEG: return visitFNEG(N);
2085 case ISD::FABS: return visitFABS(N);
2086 case ISD::FFLOOR: return visitFFLOOR(N);
2087 case ISD::FMINNUM:
2088 case ISD::FMAXNUM:
2089 case ISD::FMINIMUM:
2090 case ISD::FMAXIMUM:
2091 case ISD::FMINIMUMNUM:
2092 case ISD::FMAXIMUMNUM: return visitFMinMax(N);
2093 case ISD::FCEIL: return visitFCEIL(N);
2094 case ISD::FTRUNC: return visitFTRUNC(N);
2095 case ISD::FFREXP: return visitFFREXP(N);
2096 case ISD::BRCOND: return visitBRCOND(N);
2097 case ISD::BR_CC: return visitBR_CC(N);
2098 case ISD::LOAD: return visitLOAD(N);
2099 case ISD::STORE: return visitSTORE(N);
2100 case ISD::ATOMIC_STORE: return visitATOMIC_STORE(N);
2101 case ISD::INSERT_VECTOR_ELT: return visitINSERT_VECTOR_ELT(N);
2102 case ISD::EXTRACT_VECTOR_ELT: return visitEXTRACT_VECTOR_ELT(N);
2103 case ISD::BUILD_VECTOR: return visitBUILD_VECTOR(N);
2104 case ISD::CONCAT_VECTORS: return visitCONCAT_VECTORS(N);
2105 case ISD::VECTOR_INTERLEAVE: return visitVECTOR_INTERLEAVE(N);
2106 case ISD::VECTOR_DEINTERLEAVE: return visitVECTOR_DEINTERLEAVE(N);
2107 case ISD::EXTRACT_SUBVECTOR: return visitEXTRACT_SUBVECTOR(N);
2108 case ISD::VECTOR_SHUFFLE: return visitVECTOR_SHUFFLE(N);
2109 case ISD::SCALAR_TO_VECTOR: return visitSCALAR_TO_VECTOR(N);
2110 case ISD::INSERT_SUBVECTOR: return visitINSERT_SUBVECTOR(N);
2111 case ISD::MGATHER: return visitMGATHER(N);
2112 case ISD::MLOAD: return visitMLOAD(N);
2113 case ISD::MSCATTER: return visitMSCATTER(N);
2114 case ISD::MSTORE: return visitMSTORE(N);
2115 case ISD::EXPERIMENTAL_VECTOR_HISTOGRAM: return visitMHISTOGRAM(N);
2116 case ISD::PARTIAL_REDUCE_SMLA:
2117 case ISD::PARTIAL_REDUCE_UMLA:
2118 case ISD::PARTIAL_REDUCE_SUMLA:
2119 case ISD::PARTIAL_REDUCE_FMLA:
2120 return visitPARTIAL_REDUCE_MLA(N);
2121 case ISD::LOOP_DEPENDENCE_RAW_MASK:
2122 case ISD::LOOP_DEPENDENCE_WAR_MASK:
2123 return visitLOOP_DEPENDENCE_MASK(N);
2124 case ISD::VECTOR_COMPRESS: return visitVECTOR_COMPRESS(N);
2125 case ISD::LIFETIME_END: return visitLIFETIME_END(N);
2126 case ISD::FP_TO_FP16: return visitFP_TO_FP16(N);
2127 case ISD::FP16_TO_FP: return visitFP16_TO_FP(N);
2128 case ISD::FP_TO_BF16: return visitFP_TO_BF16(N);
2129 case ISD::BF16_TO_FP: return visitBF16_TO_FP(N);
2130 case ISD::FREEZE: return visitFREEZE(N);
2131 case ISD::GET_FPENV_MEM: return visitGET_FPENV_MEM(N);
2132 case ISD::SET_FPENV_MEM: return visitSET_FPENV_MEM(N);
2133 case ISD::FCANONICALIZE: return visitFCANONICALIZE(N);
2134 case ISD::VECREDUCE_FADD:
2135 case ISD::VECREDUCE_FMUL:
2136 case ISD::VECREDUCE_ADD:
2137 case ISD::VECREDUCE_MUL:
2138 case ISD::VECREDUCE_AND:
2139 case ISD::VECREDUCE_OR:
2140 case ISD::VECREDUCE_XOR:
2141 case ISD::VECREDUCE_SMAX:
2142 case ISD::VECREDUCE_SMIN:
2143 case ISD::VECREDUCE_UMAX:
2144 case ISD::VECREDUCE_UMIN:
2145 case ISD::VECREDUCE_FMAX:
2146 case ISD::VECREDUCE_FMIN:
2147 case ISD::VECREDUCE_FMAXIMUM:
2148 case ISD::VECREDUCE_FMINIMUM:
2149 case ISD::VECREDUCE_FMAXIMUMNUM:
2150 case ISD::VECREDUCE_FMINIMUMNUM: return visitVECREDUCE(N);
2151#define BEGIN_REGISTER_VP_SDNODE(SDOPC, ...) case ISD::SDOPC:
2152#include "llvm/IR/VPIntrinsics.def"
2153 return visitVPOp(N);
2154 }
2155 // clang-format on
2156 return SDValue();
2157}
2158
2159SDValue DAGCombiner::combine(SDNode *N) {
2160 if (!DebugCounter::shouldExecute(Counter&: DAGCombineCounter))
2161 return SDValue();
2162
2163 SDValue RV;
2164 if (!DisableGenericCombines)
2165 RV = visit(N);
2166
2167 // If nothing happened, try a target-specific DAG combine.
2168 if (!RV.getNode()) {
2169 assert(N->getOpcode() != ISD::DELETED_NODE &&
2170 "Node was deleted but visit returned NULL!");
2171
2172 if (N->getOpcode() >= ISD::BUILTIN_OP_END ||
2173 TLI.hasTargetDAGCombine(NT: (ISD::NodeType)N->getOpcode())) {
2174
2175 // Expose the DAG combiner to the target combiner impls.
2176 TargetLowering::DAGCombinerInfo
2177 DagCombineInfo(DAG, Level, false, this);
2178
2179 RV = TLI.PerformDAGCombine(N, DCI&: DagCombineInfo);
2180 }
2181 }
2182
2183 // If nothing happened still, try promoting the operation.
2184 if (!RV.getNode()) {
2185 switch (N->getOpcode()) {
2186 default: break;
2187 case ISD::ADD:
2188 case ISD::SUB:
2189 case ISD::MUL:
2190 case ISD::AND:
2191 case ISD::OR:
2192 case ISD::XOR:
2193 RV = PromoteIntBinOp(Op: SDValue(N, 0));
2194 break;
2195 case ISD::SHL:
2196 case ISD::SRA:
2197 case ISD::SRL:
2198 RV = PromoteIntShiftOp(Op: SDValue(N, 0));
2199 break;
2200 case ISD::SIGN_EXTEND:
2201 case ISD::ZERO_EXTEND:
2202 case ISD::ANY_EXTEND:
2203 RV = PromoteExtend(Op: SDValue(N, 0));
2204 break;
2205 case ISD::LOAD:
2206 if (PromoteLoad(Op: SDValue(N, 0)))
2207 RV = SDValue(N, 0);
2208 break;
2209 }
2210 }
2211
2212 // If N is a commutative binary node, try to eliminate it if the commuted
2213 // version is already present in the DAG.
2214 if (!RV.getNode() && TLI.isCommutativeBinOp(Opcode: N->getOpcode())) {
2215 SDValue N0 = N->getOperand(Num: 0);
2216 SDValue N1 = N->getOperand(Num: 1);
2217
2218 // Constant operands are canonicalized to RHS.
2219 if (N0 != N1 && (isa<ConstantSDNode>(Val: N0) || !isa<ConstantSDNode>(Val: N1))) {
2220 SDValue Ops[] = {N1, N0};
2221 SDNode *CSENode = DAG.getNodeIfExists(Opcode: N->getOpcode(), VTList: N->getVTList(), Ops,
2222 Flags: N->getFlags());
2223 if (CSENode)
2224 return SDValue(CSENode, 0);
2225 }
2226 }
2227
2228 return RV;
2229}
2230
2231/// Given a node, return its input chain if it has one, otherwise return a null
2232/// sd operand.
2233static SDValue getInputChainForNode(SDNode *N) {
2234 if (unsigned NumOps = N->getNumOperands()) {
2235 if (N->getOperand(Num: 0).getValueType() == MVT::Other)
2236 return N->getOperand(Num: 0);
2237 if (N->getOperand(Num: NumOps-1).getValueType() == MVT::Other)
2238 return N->getOperand(Num: NumOps-1);
2239 for (unsigned i = 1; i < NumOps-1; ++i)
2240 if (N->getOperand(Num: i).getValueType() == MVT::Other)
2241 return N->getOperand(Num: i);
2242 }
2243 return SDValue();
2244}
2245
2246SDValue DAGCombiner::visitFCANONICALIZE(SDNode *N) {
2247 SDValue Operand = N->getOperand(Num: 0);
2248 EVT VT = Operand.getValueType();
2249 SDLoc dl(N);
2250
2251 // Canonicalize undef to quiet NaN.
2252 if (Operand.isUndef()) {
2253 APFloat CanonicalQNaN = APFloat::getQNaN(Sem: VT.getFltSemantics());
2254 return DAG.getConstantFP(Val: CanonicalQNaN, DL: dl, VT);
2255 }
2256 return SDValue();
2257}
2258
2259SDValue DAGCombiner::visitTokenFactor(SDNode *N) {
2260 // If N has two operands, where one has an input chain equal to the other,
2261 // the 'other' chain is redundant.
2262 if (N->getNumOperands() == 2) {
2263 if (getInputChainForNode(N: N->getOperand(Num: 0).getNode()) == N->getOperand(Num: 1))
2264 return N->getOperand(Num: 0);
2265 if (getInputChainForNode(N: N->getOperand(Num: 1).getNode()) == N->getOperand(Num: 0))
2266 return N->getOperand(Num: 1);
2267 }
2268
2269 // Don't simplify token factors if optnone.
2270 if (OptLevel == CodeGenOptLevel::None)
2271 return SDValue();
2272
2273 // Don't simplify the token factor if the node itself has too many operands.
2274 if (N->getNumOperands() > TokenFactorInlineLimit)
2275 return SDValue();
2276
2277 // If the sole user is a token factor, we should make sure we have a
2278 // chance to merge them together. This prevents TF chains from inhibiting
2279 // optimizations.
2280 if (N->hasOneUse() && N->user_begin()->getOpcode() == ISD::TokenFactor)
2281 AddToWorklist(N: *(N->user_begin()));
2282
2283 SmallVector<SDNode *, 8> TFs; // List of token factors to visit.
2284 SmallVector<SDValue, 8> Ops; // Ops for replacing token factor.
2285 SmallPtrSet<SDNode*, 16> SeenOps;
2286 bool Changed = false; // If we should replace this token factor.
2287
2288 // Start out with this token factor.
2289 TFs.push_back(Elt: N);
2290
2291 // Iterate through token factors. The TFs grows when new token factors are
2292 // encountered.
2293 for (unsigned i = 0; i < TFs.size(); ++i) {
2294 // Limit number of nodes to inline, to avoid quadratic compile times.
2295 // We have to add the outstanding Token Factors to Ops, otherwise we might
2296 // drop Ops from the resulting Token Factors.
2297 if (Ops.size() > TokenFactorInlineLimit) {
2298 for (unsigned j = i; j < TFs.size(); j++)
2299 Ops.emplace_back(Args&: TFs[j], Args: 0);
2300 // Drop unprocessed Token Factors from TFs, so we do not add them to the
2301 // combiner worklist later.
2302 TFs.resize(N: i);
2303 break;
2304 }
2305
2306 SDNode *TF = TFs[i];
2307 // Check each of the operands.
2308 for (const SDValue &Op : TF->op_values()) {
2309 switch (Op.getOpcode()) {
2310 case ISD::EntryToken:
2311 // Entry tokens don't need to be added to the list. They are
2312 // redundant.
2313 Changed = true;
2314 break;
2315
2316 case ISD::TokenFactor:
2317 if (Op.hasOneUse() && !is_contained(Range&: TFs, Element: Op.getNode())) {
2318 // Queue up for processing.
2319 TFs.push_back(Elt: Op.getNode());
2320 Changed = true;
2321 break;
2322 }
2323 [[fallthrough]];
2324
2325 default:
2326 // Only add if it isn't already in the list.
2327 if (SeenOps.insert(Ptr: Op.getNode()).second)
2328 Ops.push_back(Elt: Op);
2329 else
2330 Changed = true;
2331 break;
2332 }
2333 }
2334 }
2335
2336 // Re-visit inlined Token Factors, to clean them up in case they have been
2337 // removed. Skip the first Token Factor, as this is the current node.
2338 for (unsigned i = 1, e = TFs.size(); i < e; i++)
2339 AddToWorklist(N: TFs[i]);
2340
2341 // Remove Nodes that are chained to another node in the list. Do so
2342 // by walking up chains breath-first stopping when we've seen
2343 // another operand. In general we must climb to the EntryNode, but we can exit
2344 // early if we find all remaining work is associated with just one operand as
2345 // no further pruning is possible.
2346
2347 // List of nodes to search through and original Ops from which they originate.
2348 SmallVector<std::pair<SDNode *, unsigned>, 8> Worklist;
2349 SmallVector<unsigned, 8> OpWorkCount; // Count of work for each Op.
2350 SmallPtrSet<SDNode *, 16> SeenChains;
2351 bool DidPruneOps = false;
2352
2353 unsigned NumLeftToConsider = 0;
2354 for (const SDValue &Op : Ops) {
2355 Worklist.push_back(Elt: std::make_pair(x: Op.getNode(), y: NumLeftToConsider++));
2356 OpWorkCount.push_back(Elt: 1);
2357 }
2358
2359 auto AddToWorklist = [&](unsigned CurIdx, SDNode *Op, unsigned OpNumber) {
2360 // If this is an Op, we can remove the op from the list. Remark any
2361 // search associated with it as from the current OpNumber.
2362 if (SeenOps.contains(Ptr: Op)) {
2363 Changed = true;
2364 DidPruneOps = true;
2365 unsigned OrigOpNumber = 0;
2366 while (OrigOpNumber < Ops.size() && Ops[OrigOpNumber].getNode() != Op)
2367 OrigOpNumber++;
2368 assert((OrigOpNumber != Ops.size()) &&
2369 "expected to find TokenFactor Operand");
2370 // Re-mark worklist from OrigOpNumber to OpNumber
2371 for (unsigned i = CurIdx + 1; i < Worklist.size(); ++i) {
2372 if (Worklist[i].second == OrigOpNumber) {
2373 Worklist[i].second = OpNumber;
2374 }
2375 }
2376 OpWorkCount[OpNumber] += OpWorkCount[OrigOpNumber];
2377 OpWorkCount[OrigOpNumber] = 0;
2378 NumLeftToConsider--;
2379 }
2380 // Add if it's a new chain
2381 if (SeenChains.insert(Ptr: Op).second) {
2382 OpWorkCount[OpNumber]++;
2383 Worklist.push_back(Elt: std::make_pair(x&: Op, y&: OpNumber));
2384 }
2385 };
2386
2387 for (unsigned i = 0; i < Worklist.size() && i < 1024; ++i) {
2388 // We need at least be consider at least 2 Ops to prune.
2389 if (NumLeftToConsider <= 1)
2390 break;
2391 auto CurNode = Worklist[i].first;
2392 auto CurOpNumber = Worklist[i].second;
2393 assert((OpWorkCount[CurOpNumber] > 0) &&
2394 "Node should not appear in worklist");
2395 switch (CurNode->getOpcode()) {
2396 case ISD::EntryToken:
2397 // Hitting EntryToken is the only way for the search to terminate without
2398 // hitting
2399 // another operand's search. Prevent us from marking this operand
2400 // considered.
2401 NumLeftToConsider++;
2402 break;
2403 case ISD::TokenFactor:
2404 for (const SDValue &Op : CurNode->op_values())
2405 AddToWorklist(i, Op.getNode(), CurOpNumber);
2406 break;
2407 case ISD::LIFETIME_START:
2408 case ISD::LIFETIME_END:
2409 case ISD::CopyFromReg:
2410 case ISD::CopyToReg:
2411 AddToWorklist(i, CurNode->getOperand(Num: 0).getNode(), CurOpNumber);
2412 break;
2413 default:
2414 if (auto *MemNode = dyn_cast<MemSDNode>(Val: CurNode))
2415 AddToWorklist(i, MemNode->getChain().getNode(), CurOpNumber);
2416 break;
2417 }
2418 OpWorkCount[CurOpNumber]--;
2419 if (OpWorkCount[CurOpNumber] == 0)
2420 NumLeftToConsider--;
2421 }
2422
2423 // If we've changed things around then replace token factor.
2424 if (Changed) {
2425 SDValue Result;
2426 if (Ops.empty()) {
2427 // The entry token is the only possible outcome.
2428 Result = DAG.getEntryNode();
2429 } else {
2430 if (DidPruneOps) {
2431 SmallVector<SDValue, 8> PrunedOps;
2432 //
2433 for (const SDValue &Op : Ops) {
2434 if (SeenChains.count(Ptr: Op.getNode()) == 0)
2435 PrunedOps.push_back(Elt: Op);
2436 }
2437 Result = DAG.getTokenFactor(DL: SDLoc(N), Vals&: PrunedOps);
2438 } else {
2439 Result = DAG.getTokenFactor(DL: SDLoc(N), Vals&: Ops);
2440 }
2441 }
2442 return Result;
2443 }
2444 return SDValue();
2445}
2446
2447/// MERGE_VALUES can always be eliminated.
2448SDValue DAGCombiner::visitMERGE_VALUES(SDNode *N) {
2449 WorklistRemover DeadNodes(*this);
2450 // Replacing results may cause a different MERGE_VALUES to suddenly
2451 // be CSE'd with N, and carry its uses with it. Iterate until no
2452 // uses remain, to ensure that the node can be safely deleted.
2453 // First add the users of this node to the work list so that they
2454 // can be tried again once they have new operands.
2455 AddUsersToWorklist(N);
2456 do {
2457 // Do as a single replacement to avoid rewalking use lists.
2458 SmallVector<SDValue, 8> Ops(N->ops());
2459 DAG.ReplaceAllUsesWith(From: N, To: Ops.data());
2460 } while (!N->use_empty());
2461 deleteAndRecombine(N);
2462 return SDValue(N, 0); // Return N so it doesn't get rechecked!
2463}
2464
2465/// If \p N is a ConstantSDNode with isOpaque() == false return it casted to a
2466/// ConstantSDNode pointer else nullptr.
2467static ConstantSDNode *getAsNonOpaqueConstant(SDValue N) {
2468 ConstantSDNode *Const = dyn_cast<ConstantSDNode>(Val&: N);
2469 return Const != nullptr && !Const->isOpaque() ? Const : nullptr;
2470}
2471
2472// isTruncateOf - If N is a truncate of some other value, return true, record
2473// the value being truncated in Op and which of Op's bits are zero/one in Known.
2474// This function computes KnownBits to avoid a duplicated call to
2475// computeKnownBits in the caller.
2476static bool isTruncateOf(SelectionDAG &DAG, SDValue N, SDValue &Op,
2477 KnownBits &Known) {
2478 if (N->getOpcode() == ISD::TRUNCATE) {
2479 Op = N->getOperand(Num: 0);
2480 Known = DAG.computeKnownBits(Op);
2481 if (N->getFlags().hasNoUnsignedWrap())
2482 Known.Zero.setBitsFrom(N.getScalarValueSizeInBits());
2483 return true;
2484 }
2485
2486 if (N.getValueType().getScalarType() != MVT::i1 ||
2487 !sd_match(N, P: m_c_SpecificSetCC(CC: ISD::SETNE, LHS: m_Value(N&: Op), RHS: m_Zero())))
2488 return false;
2489
2490 Known = DAG.computeKnownBits(Op);
2491 return (Known.Zero | 1).isAllOnes();
2492}
2493
2494/// Return true if 'Use' is a load or a store that uses N as its base pointer
2495/// and that N may be folded in the load / store addressing mode.
2496static bool canFoldInAddressingMode(SDNode *N, SDNode *Use, SelectionDAG &DAG,
2497 const TargetLowering &TLI) {
2498 EVT VT;
2499 unsigned AS;
2500
2501 if (LoadSDNode *LD = dyn_cast<LoadSDNode>(Val: Use)) {
2502 if (LD->isIndexed() || LD->getBasePtr().getNode() != N)
2503 return false;
2504 VT = LD->getMemoryVT();
2505 AS = LD->getAddressSpace();
2506 } else if (StoreSDNode *ST = dyn_cast<StoreSDNode>(Val: Use)) {
2507 if (ST->isIndexed() || ST->getBasePtr().getNode() != N)
2508 return false;
2509 VT = ST->getMemoryVT();
2510 AS = ST->getAddressSpace();
2511 } else if (MaskedLoadSDNode *LD = dyn_cast<MaskedLoadSDNode>(Val: Use)) {
2512 if (LD->isIndexed() || LD->getBasePtr().getNode() != N)
2513 return false;
2514 VT = LD->getMemoryVT();
2515 AS = LD->getAddressSpace();
2516 } else if (MaskedStoreSDNode *ST = dyn_cast<MaskedStoreSDNode>(Val: Use)) {
2517 if (ST->isIndexed() || ST->getBasePtr().getNode() != N)
2518 return false;
2519 VT = ST->getMemoryVT();
2520 AS = ST->getAddressSpace();
2521 } else {
2522 return false;
2523 }
2524
2525 TargetLowering::AddrMode AM;
2526 if (N->isAnyAdd()) {
2527 AM.HasBaseReg = true;
2528 ConstantSDNode *Offset = dyn_cast<ConstantSDNode>(Val: N->getOperand(Num: 1));
2529 if (Offset)
2530 // [reg +/- imm]
2531 AM.BaseOffs = Offset->getSExtValue();
2532 else
2533 // [reg +/- reg]
2534 AM.Scale = 1;
2535 } else if (N->getOpcode() == ISD::SUB) {
2536 AM.HasBaseReg = true;
2537 ConstantSDNode *Offset = dyn_cast<ConstantSDNode>(Val: N->getOperand(Num: 1));
2538 if (Offset)
2539 // [reg +/- imm]
2540 AM.BaseOffs = -Offset->getSExtValue();
2541 else
2542 // [reg +/- reg]
2543 AM.Scale = 1;
2544 } else {
2545 return false;
2546 }
2547
2548 return TLI.isLegalAddressingMode(DL: DAG.getDataLayout(), AM,
2549 Ty: VT.getTypeForEVT(Context&: *DAG.getContext()), AddrSpace: AS);
2550}
2551
2552/// This inverts a canonicalization in IR that replaces a variable select arm
2553/// with an identity constant. Codegen improves if we re-use the variable
2554/// operand rather than load a constant. This can also be converted into a
2555/// masked vector operation if the target supports it.
2556static SDValue foldSelectWithIdentityConstant(SDNode *N, SelectionDAG &DAG,
2557 bool ShouldCommuteOperands) {
2558 SDValue N0 = N->getOperand(Num: 0);
2559 SDValue N1 = N->getOperand(Num: 1);
2560
2561 // Match a select as operand 1. The identity constant that we are looking for
2562 // is only valid as operand 1 of a non-commutative binop.
2563 if (ShouldCommuteOperands)
2564 std::swap(a&: N0, b&: N1);
2565
2566 SDValue Cond, TVal, FVal;
2567 if (!sd_match(N: N1, P: m_OneUse(P: m_SelectLike(Cond: m_Value(N&: Cond), T: m_Value(N&: TVal),
2568 F: m_Value(N&: FVal)))))
2569 return SDValue();
2570
2571 // We can't hoist all instructions because of immediate UB (not speculatable).
2572 // For example div/rem by zero.
2573 if (!DAG.isSafeToSpeculativelyExecuteNode(N))
2574 return SDValue();
2575
2576 unsigned SelOpcode = N1.getOpcode();
2577 unsigned Opcode = N->getOpcode();
2578 EVT VT = N->getValueType(ResNo: 0);
2579 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
2580
2581 // This transform increases uses of N0, so freeze it to be safe.
2582 // binop N0, (vselect Cond, IDC, FVal) --> vselect Cond, N0, (binop N0, FVal)
2583 unsigned OpNo = ShouldCommuteOperands ? 0 : 1;
2584 if (DAG.isIdentityElement(Opc: Opcode, Flags: N->getFlags(), V: TVal, OperandNo: OpNo) &&
2585 TLI.shouldFoldSelectWithIdentityConstant(BinOpcode: Opcode, VT, SelectOpcode: SelOpcode, X: N0,
2586 Y: FVal)) {
2587 SDValue F0 = DAG.getFreeze(V: N0);
2588 SDValue NewBO = DAG.getNode(Opcode, DL: SDLoc(N), VT, N1: F0, N2: FVal, Flags: N->getFlags());
2589 return DAG.getSelect(DL: SDLoc(N), VT, Cond, LHS: F0, RHS: NewBO);
2590 }
2591 // binop N0, (vselect Cond, TVal, IDC) --> vselect Cond, (binop N0, TVal), N0
2592 if (DAG.isIdentityElement(Opc: Opcode, Flags: N->getFlags(), V: FVal, OperandNo: OpNo) &&
2593 TLI.shouldFoldSelectWithIdentityConstant(BinOpcode: Opcode, VT, SelectOpcode: SelOpcode, X: N0,
2594 Y: TVal)) {
2595 SDValue F0 = DAG.getFreeze(V: N0);
2596 SDValue NewBO = DAG.getNode(Opcode, DL: SDLoc(N), VT, N1: F0, N2: TVal, Flags: N->getFlags());
2597 return DAG.getSelect(DL: SDLoc(N), VT, Cond, LHS: NewBO, RHS: F0);
2598 }
2599
2600 return SDValue();
2601}
2602
2603SDValue DAGCombiner::foldBinOpIntoSelect(SDNode *BO) {
2604 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
2605 assert(TLI.isBinOp(BO->getOpcode()) && BO->getNumValues() == 1 &&
2606 "Unexpected binary operator");
2607
2608 if (SDValue Sel = foldSelectWithIdentityConstant(N: BO, DAG, ShouldCommuteOperands: false))
2609 return Sel;
2610
2611 if (TLI.isCommutativeBinOp(Opcode: BO->getOpcode()))
2612 if (SDValue Sel = foldSelectWithIdentityConstant(N: BO, DAG, ShouldCommuteOperands: true))
2613 return Sel;
2614
2615 // Don't do this unless the old select is going away. We want to eliminate the
2616 // binary operator, not replace a binop with a select.
2617 // TODO: Handle ISD::SELECT_CC.
2618 unsigned SelOpNo = 0;
2619 SDValue Sel = BO->getOperand(Num: 0);
2620 auto BinOpcode = BO->getOpcode();
2621 if (Sel.getOpcode() != ISD::SELECT || !Sel.hasOneUse()) {
2622 SelOpNo = 1;
2623 Sel = BO->getOperand(Num: 1);
2624
2625 // Peek through trunc to shift amount type.
2626 if ((BinOpcode == ISD::SHL || BinOpcode == ISD::SRA ||
2627 BinOpcode == ISD::SRL) && Sel.hasOneUse()) {
2628 // This is valid when the truncated bits of x are already zero.
2629 SDValue Op;
2630 KnownBits Known;
2631 if (isTruncateOf(DAG, N: Sel, Op, Known) &&
2632 Known.countMaxActiveBits() < Sel.getScalarValueSizeInBits())
2633 Sel = Op;
2634 }
2635 }
2636
2637 if (Sel.getOpcode() != ISD::SELECT || !Sel.hasOneUse())
2638 return SDValue();
2639
2640 SDValue CT = Sel.getOperand(i: 1);
2641 if (!isConstantOrConstantVector(N: CT, NoOpaques: true) &&
2642 !DAG.isConstantFPBuildVectorOrConstantFP(N: CT))
2643 return SDValue();
2644
2645 SDValue CF = Sel.getOperand(i: 2);
2646 if (!isConstantOrConstantVector(N: CF, NoOpaques: true) &&
2647 !DAG.isConstantFPBuildVectorOrConstantFP(N: CF))
2648 return SDValue();
2649
2650 // Bail out if any constants are opaque because we can't constant fold those.
2651 // The exception is "and" and "or" with either 0 or -1 in which case we can
2652 // propagate non constant operands into select. I.e.:
2653 // and (select Cond, 0, -1), X --> select Cond, 0, X
2654 // or X, (select Cond, -1, 0) --> select Cond, -1, X
2655 bool CanFoldNonConst =
2656 (BinOpcode == ISD::AND || BinOpcode == ISD::OR) &&
2657 ((isNullOrNullSplat(V: CT) && isAllOnesOrAllOnesSplat(V: CF)) ||
2658 (isNullOrNullSplat(V: CF) && isAllOnesOrAllOnesSplat(V: CT)));
2659
2660 SDValue CBO = BO->getOperand(Num: SelOpNo ^ 1);
2661 if (!CanFoldNonConst &&
2662 !isConstantOrConstantVector(N: CBO, NoOpaques: true) &&
2663 !DAG.isConstantFPBuildVectorOrConstantFP(N: CBO))
2664 return SDValue();
2665
2666 SDLoc DL(Sel);
2667 SDValue NewCT, NewCF;
2668 EVT VT = BO->getValueType(ResNo: 0);
2669
2670 if (CanFoldNonConst) {
2671 // If CBO is an opaque constant, we can't rely on getNode to constant fold.
2672 if ((BinOpcode == ISD::AND && isNullOrNullSplat(V: CT)) ||
2673 (BinOpcode == ISD::OR && isAllOnesOrAllOnesSplat(V: CT)))
2674 NewCT = CT;
2675 else
2676 NewCT = CBO;
2677
2678 if ((BinOpcode == ISD::AND && isNullOrNullSplat(V: CF)) ||
2679 (BinOpcode == ISD::OR && isAllOnesOrAllOnesSplat(V: CF)))
2680 NewCF = CF;
2681 else
2682 NewCF = CBO;
2683 } else {
2684 // We have a select-of-constants followed by a binary operator with a
2685 // constant. Eliminate the binop by pulling the constant math into the
2686 // select. Example: add (select Cond, CT, CF), CBO --> select Cond, CT +
2687 // CBO, CF + CBO
2688 NewCT = SelOpNo ? DAG.FoldConstantArithmetic(Opcode: BinOpcode, DL, VT, Ops: {CBO, CT})
2689 : DAG.FoldConstantArithmetic(Opcode: BinOpcode, DL, VT, Ops: {CT, CBO});
2690 if (!NewCT)
2691 return SDValue();
2692
2693 NewCF = SelOpNo ? DAG.FoldConstantArithmetic(Opcode: BinOpcode, DL, VT, Ops: {CBO, CF})
2694 : DAG.FoldConstantArithmetic(Opcode: BinOpcode, DL, VT, Ops: {CF, CBO});
2695 if (!NewCF)
2696 return SDValue();
2697 }
2698
2699 return DAG.getSelect(DL, VT, Cond: Sel.getOperand(i: 0), LHS: NewCT, RHS: NewCF, Flags: BO->getFlags());
2700}
2701
2702static SDValue foldAddSubBoolOfMaskedVal(SDNode *N, const SDLoc &DL,
2703 SelectionDAG &DAG) {
2704 assert((N->getOpcode() == ISD::ADD || N->getOpcode() == ISD::SUB) &&
2705 "Expecting add or sub");
2706
2707 // Match a constant operand and a zext operand for the math instruction:
2708 // add Z, C
2709 // sub C, Z
2710 bool IsAdd = N->getOpcode() == ISD::ADD;
2711 SDValue C = IsAdd ? N->getOperand(Num: 1) : N->getOperand(Num: 0);
2712 SDValue Z = IsAdd ? N->getOperand(Num: 0) : N->getOperand(Num: 1);
2713 auto *CN = dyn_cast<ConstantSDNode>(Val&: C);
2714 if (!CN || Z.getOpcode() != ISD::ZERO_EXTEND)
2715 return SDValue();
2716
2717 // Match the zext operand as a setcc of a boolean.
2718 if (Z.getOperand(i: 0).getValueType() != MVT::i1)
2719 return SDValue();
2720
2721 // Match the compare as: setcc (X & 1), 0, eq.
2722 if (!sd_match(
2723 N: Z.getOperand(i: 0),
2724 P: m_SpecificSetCC(CC: ISD::SETEQ, LHS: m_And(L: m_Value(), R: m_One()), RHS: m_Zero())))
2725 return SDValue();
2726
2727 // We are adding/subtracting a constant and an inverted low bit. Turn that
2728 // into a subtract/add of the low bit with incremented/decremented constant:
2729 // add (zext i1 (seteq (X & 1), 0)), C --> sub C+1, (zext (X & 1))
2730 // sub C, (zext i1 (seteq (X & 1), 0)) --> add C-1, (zext (X & 1))
2731 EVT VT = C.getValueType();
2732 SDValue LowBit = DAG.getZExtOrTrunc(Op: Z.getOperand(i: 0).getOperand(i: 0), DL, VT);
2733 SDValue C1 = IsAdd ? DAG.getConstant(Val: CN->getAPIntValue() + 1, DL, VT)
2734 : DAG.getConstant(Val: CN->getAPIntValue() - 1, DL, VT);
2735 return DAG.getNode(Opcode: IsAdd ? ISD::SUB : ISD::ADD, DL, VT, N1: C1, N2: LowBit);
2736}
2737
2738// Attempt to form avgceil(A, B) from (A | B) - ((A ^ B) >> 1)
2739SDValue DAGCombiner::foldSubToAvg(SDNode *N, const SDLoc &DL) {
2740 SDValue N0 = N->getOperand(Num: 0);
2741 EVT VT = N0.getValueType();
2742 SDValue A, B;
2743
2744 if ((!LegalOperations || hasOperation(Opcode: ISD::AVGCEILU, VT)) &&
2745 sd_match(N, P: m_Sub(L: m_Or(L: m_Value(N&: A), R: m_Value(N&: B)),
2746 R: m_Srl(L: m_Xor(L: m_Deferred(V&: A), R: m_Deferred(V&: B)), R: m_One())))) {
2747 return DAG.getNode(Opcode: ISD::AVGCEILU, DL, VT, N1: A, N2: B);
2748 }
2749 if ((!LegalOperations || hasOperation(Opcode: ISD::AVGCEILS, VT)) &&
2750 sd_match(N, P: m_Sub(L: m_Or(L: m_Value(N&: A), R: m_Value(N&: B)),
2751 R: m_Sra(L: m_Xor(L: m_Deferred(V&: A), R: m_Deferred(V&: B)), R: m_One())))) {
2752 return DAG.getNode(Opcode: ISD::AVGCEILS, DL, VT, N1: A, N2: B);
2753 }
2754 return SDValue();
2755}
2756
2757/// Try to fold a pointer arithmetic node.
2758/// This needs to be done separately from normal addition, because pointer
2759/// addition is not commutative.
2760SDValue DAGCombiner::visitPTRADD(SDNode *N) {
2761 SDValue N0 = N->getOperand(Num: 0);
2762 SDValue N1 = N->getOperand(Num: 1);
2763 EVT PtrVT = N0.getValueType();
2764 EVT IntVT = N1.getValueType();
2765 SDLoc DL(N);
2766
2767 // This is already ensured by an assert in SelectionDAG::getNode(). Several
2768 // combines here depend on this assumption.
2769 assert(PtrVT == IntVT &&
2770 "PTRADD with different operand types is not supported");
2771
2772 // fold (ptradd x, 0) -> x
2773 if (isNullConstant(V: N1))
2774 return N0;
2775
2776 // fold (ptradd 0, x) -> x
2777 if (PtrVT == IntVT && isNullConstant(V: N0))
2778 return N1;
2779
2780 if (N0.getOpcode() == ISD::PTRADD &&
2781 !reassociationCanBreakAddressingModePattern(Opc: ISD::PTRADD, DL, N, N0, N1)) {
2782 SDValue X = N0.getOperand(i: 0);
2783 SDValue Y = N0.getOperand(i: 1);
2784 SDValue Z = N1;
2785 bool N0OneUse = N0.hasOneUse();
2786 bool YIsConstant = DAG.isConstantIntBuildVectorOrConstantInt(N: Y);
2787 bool ZIsConstant = DAG.isConstantIntBuildVectorOrConstantInt(N: Z);
2788
2789 // (ptradd (ptradd x, y), z) -> (ptradd x, (add y, z)) if:
2790 // * y is a constant and (ptradd x, y) has one use; or
2791 // * y and z are both constants.
2792 if ((YIsConstant && N0OneUse) || (YIsConstant && ZIsConstant)) {
2793 // If both additions in the original were NUW, the new ones are as well.
2794 SDNodeFlags Flags =
2795 (N->getFlags() & N0->getFlags()) & SDNodeFlags::NoUnsignedWrap;
2796 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT: IntVT, Ops: {Y, Z}, Flags);
2797 AddToWorklist(N: Add.getNode());
2798 // We can't set InBounds even if both original ptradds were InBounds and
2799 // NUW: SDAG usually represents pointers as integers, therefore, the
2800 // matched pattern behaves as if it had implicit casts:
2801 // (ptradd inbounds (inttoptr (ptrtoint (ptradd inbounds x, y))), z)
2802 // The outer inbounds ptradd might therefore rely on a provenance that x
2803 // does not have.
2804 return DAG.getMemBasePlusOffset(Base: X, Offset: Add, DL, Flags);
2805 }
2806 }
2807
2808 // The following combines can turn in-bounds pointer arithmetic out of bounds.
2809 // That is problematic for settings like AArch64's CPA, which checks that
2810 // intermediate results of pointer arithmetic remain in bounds. The target
2811 // therefore needs to opt-in to enable them.
2812 if (!TLI.canTransformPtrArithOutOfBounds(
2813 F: DAG.getMachineFunction().getFunction(), PtrVT))
2814 return SDValue();
2815
2816 if (N0.getOpcode() == ISD::PTRADD && isa<ConstantSDNode>(Val: N1)) {
2817 // Fold (ptradd (ptradd GA, v), c) -> (ptradd (ptradd GA, c) v) with
2818 // global address GA and constant c, such that c can be folded into GA.
2819 // TODO: Support constant vector splats.
2820 SDValue GAValue = N0.getOperand(i: 0);
2821 if (const GlobalAddressSDNode *GA =
2822 dyn_cast<GlobalAddressSDNode>(Val&: GAValue)) {
2823 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
2824 if (!LegalOperations && TLI.isOffsetFoldingLegal(GA)) {
2825 // If both additions in the original were NUW, reassociation preserves
2826 // that.
2827 SDNodeFlags Flags =
2828 (N->getFlags() & N0->getFlags()) & SDNodeFlags::NoUnsignedWrap;
2829 // We can't set InBounds even if both original ptradds were InBounds and
2830 // NUW: SDAG usually represents pointers as integers, therefore, the
2831 // matched pattern behaves as if it had implicit casts:
2832 // (ptradd inbounds (inttoptr (ptrtoint (ptradd inbounds GA, v))), c)
2833 // The outer inbounds ptradd might therefore rely on a provenance that
2834 // GA does not have.
2835 SDValue Inner = DAG.getMemBasePlusOffset(Base: GAValue, Offset: N1, DL, Flags);
2836 AddToWorklist(N: Inner.getNode());
2837 return DAG.getMemBasePlusOffset(Base: Inner, Offset: N0.getOperand(i: 1), DL, Flags);
2838 }
2839 }
2840 }
2841
2842 if (N1.getOpcode() == ISD::ADD && N1.hasOneUse()) {
2843 // (ptradd x, (add y, z)) -> (ptradd (ptradd x, y), z) if z is a constant,
2844 // y is not, and (add y, z) is used only once.
2845 // (ptradd x, (add y, z)) -> (ptradd (ptradd x, z), y) if y is a constant,
2846 // z is not, and (add y, z) is used only once.
2847 // The goal is to move constant offsets to the outermost ptradd, to create
2848 // more opportunities to fold offsets into memory instructions.
2849 // Together with the another combine above, this also implements
2850 // (ptradd (ptradd x, y), z) -> (ptradd (ptradd x, z), y)).
2851 SDValue X = N0;
2852 SDValue Y = N1.getOperand(i: 0);
2853 SDValue Z = N1.getOperand(i: 1);
2854 bool YIsConstant = DAG.isConstantIntBuildVectorOrConstantInt(N: Y);
2855 bool ZIsConstant = DAG.isConstantIntBuildVectorOrConstantInt(N: Z);
2856
2857 // If both additions in the original were NUW, reassociation preserves that.
2858 SDNodeFlags CommonFlags = N->getFlags() & N1->getFlags();
2859 SDNodeFlags ReassocFlags = CommonFlags & SDNodeFlags::NoUnsignedWrap;
2860 if (CommonFlags.hasNoUnsignedWrap()) {
2861 // If both operations are NUW and the PTRADD is inbounds, the offests are
2862 // both non-negative, so the reassociated PTRADDs are also inbounds.
2863 ReassocFlags |= N->getFlags() & SDNodeFlags::InBounds;
2864 }
2865
2866 if (ZIsConstant != YIsConstant) {
2867 if (YIsConstant)
2868 std::swap(a&: Y, b&: Z);
2869 SDValue Inner = DAG.getMemBasePlusOffset(Base: X, Offset: Y, DL, Flags: ReassocFlags);
2870 AddToWorklist(N: Inner.getNode());
2871 return DAG.getMemBasePlusOffset(Base: Inner, Offset: Z, DL, Flags: ReassocFlags);
2872 }
2873 }
2874
2875 // Transform (ptradd a, b) -> (or disjoint a, b) if it is equivalent and if
2876 // that transformation can't block an offset folding at any use of the ptradd.
2877 // This should be done late, after legalization, so that it doesn't block
2878 // other ptradd combines that could enable more offset folding.
2879 if (LegalOperations && DAG.haveNoCommonBitsSet(A: N0, B: N1)) {
2880 bool TransformCannotBreakAddrMode = none_of(Range: N->users(), P: [&](SDNode *User) {
2881 return canFoldInAddressingMode(N, Use: User, DAG, TLI);
2882 });
2883
2884 if (TransformCannotBreakAddrMode)
2885 return DAG.getNode(Opcode: ISD::OR, DL, VT: PtrVT, N1: N0, N2: N1, Flags: SDNodeFlags::Disjoint);
2886 }
2887
2888 return SDValue();
2889}
2890
2891/// Try to fold a 'not' shifted sign-bit with add/sub with constant operand into
2892/// a shift and add with a different constant.
2893static SDValue foldAddSubOfSignBit(SDNode *N, const SDLoc &DL,
2894 SelectionDAG &DAG) {
2895 assert((N->getOpcode() == ISD::ADD || N->getOpcode() == ISD::SUB) &&
2896 "Expecting add or sub");
2897
2898 // We need a constant operand for the add/sub, and the other operand is a
2899 // logical shift right: add (srl), C or sub C, (srl).
2900 bool IsAdd = N->getOpcode() == ISD::ADD;
2901 SDValue ConstantOp = IsAdd ? N->getOperand(Num: 1) : N->getOperand(Num: 0);
2902 SDValue ShiftOp = IsAdd ? N->getOperand(Num: 0) : N->getOperand(Num: 1);
2903 if (!DAG.isConstantIntBuildVectorOrConstantInt(N: ConstantOp) ||
2904 ShiftOp.getOpcode() != ISD::SRL)
2905 return SDValue();
2906
2907 // The shift must be of a 'not' value.
2908 SDValue Not = ShiftOp.getOperand(i: 0);
2909 if (!Not.hasOneUse() || !isBitwiseNot(V: Not))
2910 return SDValue();
2911
2912 // The shift must be moving the sign bit to the least-significant-bit.
2913 EVT VT = ShiftOp.getValueType();
2914 SDValue ShAmt = ShiftOp.getOperand(i: 1);
2915 ConstantSDNode *ShAmtC = isConstOrConstSplat(N: ShAmt);
2916 if (!ShAmtC || ShAmtC->getAPIntValue() != (VT.getScalarSizeInBits() - 1))
2917 return SDValue();
2918
2919 // Eliminate the 'not' by adjusting the shift and add/sub constant:
2920 // add (srl (not X), 31), C --> add (sra X, 31), (C + 1)
2921 // sub C, (srl (not X), 31) --> add (srl X, 31), (C - 1)
2922 if (SDValue NewC = DAG.FoldConstantArithmetic(
2923 Opcode: IsAdd ? ISD::ADD : ISD::SUB, DL, VT,
2924 Ops: {ConstantOp, DAG.getConstant(Val: 1, DL, VT)})) {
2925 SDValue NewShift = DAG.getNode(Opcode: IsAdd ? ISD::SRA : ISD::SRL, DL, VT,
2926 N1: Not.getOperand(i: 0), N2: ShAmt);
2927 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: NewShift, N2: NewC);
2928 }
2929
2930 return SDValue();
2931}
2932
2933static bool
2934areBitwiseNotOfEachother(SDValue Op0, SDValue Op1) {
2935 return (isBitwiseNot(V: Op0) && Op0.getOperand(i: 0) == Op1) ||
2936 (isBitwiseNot(V: Op1) && Op1.getOperand(i: 0) == Op0);
2937}
2938
2939/// Try to fold a node that behaves like an ADD (note that N isn't necessarily
2940/// an ISD::ADD here, it could for example be an ISD::OR if we know that there
2941/// are no common bits set in the operands).
2942SDValue DAGCombiner::visitADDLike(SDNode *N) {
2943 SDValue N0 = N->getOperand(Num: 0);
2944 SDValue N1 = N->getOperand(Num: 1);
2945 EVT VT = N0.getValueType();
2946 SDLoc DL(N);
2947
2948 // fold (add x, undef) -> undef
2949 if (N0.isUndef())
2950 return N0;
2951 if (N1.isUndef())
2952 return N1;
2953
2954 // fold (add c1, c2) -> c1+c2
2955 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::ADD, DL, VT, Ops: {N0, N1}))
2956 return C;
2957
2958 // canonicalize constant to RHS
2959 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
2960 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
2961 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1, N2: N0);
2962
2963 if (areBitwiseNotOfEachother(Op0: N0, Op1: N1))
2964 return DAG.getConstant(Val: APInt::getAllOnes(numBits: VT.getScalarSizeInBits()), DL, VT);
2965
2966 // fold vector ops
2967 if (VT.isVector()) {
2968 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
2969 return FoldedVOp;
2970
2971 // fold (add x, 0) -> x, vector edition
2972 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
2973 return N0;
2974 }
2975
2976 // fold (add x, 0) -> x
2977 if (isNullConstant(V: N1))
2978 return N0;
2979
2980 if (N0.getOpcode() == ISD::SUB) {
2981 SDValue N00 = N0.getOperand(i: 0);
2982 SDValue N01 = N0.getOperand(i: 1);
2983
2984 // fold ((A-c1)+c2) -> (A+(c2-c1))
2985 if (SDValue Sub = DAG.FoldConstantArithmetic(Opcode: ISD::SUB, DL, VT, Ops: {N1, N01}))
2986 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 0), N2: Sub);
2987
2988 // fold ((c1-A)+c2) -> (c1+c2)-A
2989 if (SDValue Add = DAG.FoldConstantArithmetic(Opcode: ISD::ADD, DL, VT, Ops: {N1, N00}))
2990 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Add, N2: N0.getOperand(i: 1));
2991 }
2992
2993 // add (sext i1 X), 1 -> zext (not i1 X)
2994 // We don't transform this pattern:
2995 // add (zext i1 X), -1 -> sext (not i1 X)
2996 // because most (?) targets generate better code for the zext form.
2997 if (N0.getOpcode() == ISD::SIGN_EXTEND && N0.hasOneUse() &&
2998 isOneOrOneSplat(V: N1)) {
2999 SDValue X = N0.getOperand(i: 0);
3000 if ((!LegalOperations ||
3001 (TLI.isOperationLegal(Op: ISD::XOR, VT: X.getValueType()) &&
3002 TLI.isOperationLegal(Op: ISD::ZERO_EXTEND, VT))) &&
3003 X.getScalarValueSizeInBits() == 1) {
3004 SDValue Not = DAG.getNOT(DL, Val: X, VT: X.getValueType());
3005 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: Not);
3006 }
3007 }
3008
3009 // Fold (add (or x, c0), c1) -> (add x, (c0 + c1))
3010 // iff (or x, c0) is equivalent to (add x, c0).
3011 // Fold (add (xor x, c0), c1) -> (add x, (c0 + c1))
3012 // iff (xor x, c0) is equivalent to (add x, c0).
3013 if (DAG.isADDLike(Op: N0)) {
3014 SDValue N01 = N0.getOperand(i: 1);
3015 if (SDValue Add = DAG.FoldConstantArithmetic(Opcode: ISD::ADD, DL, VT, Ops: {N1, N01}))
3016 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 0), N2: Add);
3017 }
3018
3019 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
3020 return NewSel;
3021
3022 // reassociate add
3023 if (!reassociationCanBreakAddressingModePattern(Opc: ISD::ADD, DL, N, N0, N1)) {
3024 if (SDValue RADD = reassociateOps(Opc: ISD::ADD, DL, N0, N1, Flags: N->getFlags()))
3025 return RADD;
3026
3027 // (X + Y) + X --> Y + (X + X)
3028 SDValue X, Y, InnerAdd;
3029 if (sd_match(
3030 N, P: m_Add(L: m_OneUse(P: m_Value(N&: InnerAdd, P: m_Add(L: m_Value(N&: X), R: m_Value(N&: Y)))),
3031 R: m_Deferred(V&: X)))) {
3032 if (X != Y) {
3033 // Redistribute shared NUW flag.
3034 // TODO: If NSW+NUW occurs on both adds, that can be redistributed too.
3035 SDNodeFlags NewFlags =
3036 N->getFlags() & InnerAdd->getFlags() & SDNodeFlags::NoUnsignedWrap;
3037 SDValue X2 = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: X, N2: X, Flags: NewFlags);
3038 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Y, N2: X2, Flags: NewFlags);
3039 }
3040 }
3041
3042 // Reassociate (add (or x, c), y) -> (add add(x, y), c)) if (or x, c) is
3043 // equivalent to (add x, c).
3044 // Reassociate (add (xor x, c), y) -> (add add(x, y), c)) if (xor x, c) is
3045 // equivalent to (add x, c).
3046 // Do this optimization only when adding c does not introduce instructions
3047 // for adding carries.
3048 auto ReassociateAddOr = [&](SDValue N0, SDValue N1) {
3049 if (DAG.isADDLike(Op: N0) && N0.hasOneUse() &&
3050 isConstantOrConstantVector(N: N0.getOperand(i: 1), /* NoOpaque */ NoOpaques: true)) {
3051 // If N0's type does not split or is a sign mask, it does not introduce
3052 // add carry.
3053 auto TyActn = TLI.getTypeAction(Context&: *DAG.getContext(), VT: N0.getValueType());
3054 bool NoAddCarry = TyActn == TargetLoweringBase::TypeLegal ||
3055 TyActn == TargetLoweringBase::TypePromoteInteger ||
3056 isMinSignedConstant(V: N0.getOperand(i: 1));
3057 if (NoAddCarry)
3058 return DAG.getNode(
3059 Opcode: ISD::ADD, DL, VT,
3060 N1: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1, N2: N0.getOperand(i: 0)),
3061 N2: N0.getOperand(i: 1));
3062 }
3063 return SDValue();
3064 };
3065 if (SDValue Add = ReassociateAddOr(N0, N1))
3066 return Add;
3067 if (SDValue Add = ReassociateAddOr(N1, N0))
3068 return Add;
3069
3070 // Fold add(vecreduce(x), vecreduce(y)) -> vecreduce(add(x, y))
3071 if (SDValue SD =
3072 reassociateReduction(RedOpc: ISD::VECREDUCE_ADD, Opc: ISD::ADD, DL, VT, N0, N1))
3073 return SD;
3074 }
3075
3076 SDValue A, B, C, D;
3077
3078 // fold ((0-A) + B) -> B-A
3079 if (sd_match(N: N0, P: m_Neg(V: m_Value(N&: A))))
3080 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1, N2: A);
3081
3082 // fold (A + (0-B)) -> A-B
3083 if (sd_match(N: N1, P: m_Neg(V: m_Value(N&: B))))
3084 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: B);
3085
3086 // fold (A+(B-A)) -> B
3087 if (sd_match(N: N1, P: m_Sub(L: m_Value(N&: B), R: m_Specific(N: N0))))
3088 return B;
3089
3090 // fold ((B-A)+A) -> B
3091 if (sd_match(N: N0, P: m_Sub(L: m_Value(N&: B), R: m_Specific(N: N1))))
3092 return B;
3093
3094 // fold ((A-B)+(C-A)) -> (C-B)
3095 if (sd_match(N: N0, P: m_Sub(L: m_Value(N&: A), R: m_Value(N&: B))) &&
3096 sd_match(N: N1, P: m_Sub(L: m_Value(N&: C), R: m_Specific(N: A))))
3097 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: C, N2: B);
3098
3099 // fold ((A-B)+(B-C)) -> (A-C)
3100 if (sd_match(N: N0, P: m_Sub(L: m_Value(N&: A), R: m_Value(N&: B))) &&
3101 sd_match(N: N1, P: m_Sub(L: m_Specific(N: B), R: m_Value(N&: C))))
3102 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: A, N2: C);
3103
3104 // fold (A+(B-(A+C))) to (B-C)
3105 // fold (A+(B-(C+A))) to (B-C)
3106 if (sd_match(N: N1, P: m_Sub(L: m_Value(N&: B), R: m_Add(L: m_Specific(N: N0), R: m_Value(N&: C)))))
3107 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: B, N2: C);
3108
3109 // fold (A+((B-A)+or-C)) to (B+or-C)
3110 if (sd_match(N: N1,
3111 P: m_AnyOf(preds: m_Add(L: m_Sub(L: m_Value(N&: B), R: m_Specific(N: N0)), R: m_Value(N&: C)),
3112 preds: m_Sub(L: m_Sub(L: m_Value(N&: B), R: m_Specific(N: N0)), R: m_Value(N&: C)))))
3113 return DAG.getNode(Opcode: N1.getOpcode(), DL, VT, N1: B, N2: C);
3114
3115 // fold (A-B)+(C-D) to (A+C)-(B+D) when A or C is constant
3116 if (sd_match(N: N0, P: m_OneUse(P: m_Sub(L: m_Value(N&: A), R: m_Value(N&: B)))) &&
3117 sd_match(N: N1, P: m_OneUse(P: m_Sub(L: m_Value(N&: C), R: m_Value(N&: D)))) &&
3118 (isConstantOrConstantVector(N: A) || isConstantOrConstantVector(N: C)))
3119 return DAG.getNode(Opcode: ISD::SUB, DL, VT,
3120 N1: DAG.getNode(Opcode: ISD::ADD, DL: SDLoc(N0), VT, N1: A, N2: C),
3121 N2: DAG.getNode(Opcode: ISD::ADD, DL: SDLoc(N1), VT, N1: B, N2: D));
3122
3123 // fold (add (umax X, C), -C) --> (usubsat X, C)
3124 if (N0.getOpcode() == ISD::UMAX && hasOperation(Opcode: ISD::USUBSAT, VT)) {
3125 auto MatchUSUBSAT = [](ConstantSDNode *Max, ConstantSDNode *Op) {
3126 return (!Max && !Op) ||
3127 (Max && Op && Max->getAPIntValue() == (-Op->getAPIntValue()));
3128 };
3129 if (ISD::matchBinaryPredicate(LHS: N0.getOperand(i: 1), RHS: N1, Match: MatchUSUBSAT,
3130 /*AllowUndefs*/ true))
3131 return DAG.getNode(Opcode: ISD::USUBSAT, DL, VT, N1: N0.getOperand(i: 0),
3132 N2: N0.getOperand(i: 1));
3133 }
3134
3135 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
3136 return SDValue(N, 0);
3137
3138 if (isOneOrOneSplat(V: N1)) {
3139 // fold (add (xor a, -1), 1) -> (sub 0, a)
3140 if (isBitwiseNot(V: N0))
3141 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: DAG.getConstant(Val: 0, DL, VT),
3142 N2: N0.getOperand(i: 0));
3143
3144 // fold (add (add (xor a, -1), b), 1) -> (sub b, a)
3145 if (N0.getOpcode() == ISD::ADD) {
3146 SDValue A, Xor;
3147
3148 if (isBitwiseNot(V: N0.getOperand(i: 0))) {
3149 A = N0.getOperand(i: 1);
3150 Xor = N0.getOperand(i: 0);
3151 } else if (isBitwiseNot(V: N0.getOperand(i: 1))) {
3152 A = N0.getOperand(i: 0);
3153 Xor = N0.getOperand(i: 1);
3154 }
3155
3156 if (Xor)
3157 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: A, N2: Xor.getOperand(i: 0));
3158 }
3159
3160 // Look for:
3161 // add (add x, y), 1
3162 // And if the target does not like this form then turn into:
3163 // sub x, (xor y, -1)
3164 if (!TLI.preferIncOfAddToSubOfNot(VT) && N0.getOpcode() == ISD::ADD &&
3165 N0.hasOneUse() &&
3166 // Limit this to after legalization if the add has wrap flags
3167 (Level >= AfterLegalizeDAG || (!N->getFlags().hasNoUnsignedWrap() &&
3168 !N->getFlags().hasNoSignedWrap()))) {
3169 SDValue Not = DAG.getNOT(DL, Val: N0.getOperand(i: 1), VT);
3170 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0.getOperand(i: 0), N2: Not);
3171 }
3172 }
3173
3174 // (x - y) + -1 -> add (xor y, -1), x
3175 if (N0.getOpcode() == ISD::SUB && N0.hasOneUse() &&
3176 isAllOnesOrAllOnesSplat(V: N1, /*AllowUndefs=*/true)) {
3177 SDValue Not = DAG.getNOT(DL, Val: N0.getOperand(i: 1), VT);
3178 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Not, N2: N0.getOperand(i: 0));
3179 }
3180
3181 // Fold add(mul(add(A, CA), CM), CB) -> add(mul(A, CM), CM*CA+CB).
3182 // This can help if the inner add has multiple uses.
3183 APInt CM, CA;
3184 if (ConstantSDNode *CB = dyn_cast<ConstantSDNode>(Val&: N1)) {
3185 if (VT.getScalarSizeInBits() <= 64) {
3186 if (sd_match(N: N0, P: m_OneUse(P: m_Mul(L: m_Add(L: m_Value(N&: A), R: m_ConstInt(V&: CA)),
3187 R: m_ConstInt(V&: CM)))) &&
3188 TLI.isLegalAddImmediate(
3189 (CA * CM + CB->getAPIntValue()).getSExtValue())) {
3190 SDNodeFlags Flags;
3191 // If all the inputs are nuw, the outputs can be nuw. If all the input
3192 // are _also_ nsw the outputs can be too.
3193 if (N->getFlags().hasNoUnsignedWrap() &&
3194 N0->getFlags().hasNoUnsignedWrap() &&
3195 N0.getOperand(i: 0)->getFlags().hasNoUnsignedWrap()) {
3196 Flags |= SDNodeFlags::NoUnsignedWrap;
3197 if (N->getFlags().hasNoSignedWrap() &&
3198 N0->getFlags().hasNoSignedWrap() &&
3199 N0.getOperand(i: 0)->getFlags().hasNoSignedWrap())
3200 Flags |= SDNodeFlags::NoSignedWrap;
3201 }
3202 SDValue Mul = DAG.getNode(Opcode: ISD::MUL, DL: SDLoc(N1), VT, N1: A,
3203 N2: DAG.getConstant(Val: CM, DL, VT), Flags);
3204 return DAG.getNode(
3205 Opcode: ISD::ADD, DL, VT, N1: Mul,
3206 N2: DAG.getConstant(Val: CA * CM + CB->getAPIntValue(), DL, VT), Flags);
3207 }
3208 // Also look in case there is an intermediate add.
3209 if (sd_match(N: N0, P: m_OneUse(P: m_Add(
3210 L: m_OneUse(P: m_Mul(L: m_Add(L: m_Value(N&: A), R: m_ConstInt(V&: CA)),
3211 R: m_ConstInt(V&: CM))),
3212 R: m_Value(N&: B)))) &&
3213 TLI.isLegalAddImmediate(
3214 (CA * CM + CB->getAPIntValue()).getSExtValue())) {
3215 SDNodeFlags Flags;
3216 // If all the inputs are nuw, the outputs can be nuw. If all the input
3217 // are _also_ nsw the outputs can be too.
3218 SDValue OMul =
3219 N0.getOperand(i: 0) == B ? N0.getOperand(i: 1) : N0.getOperand(i: 0);
3220 if (N->getFlags().hasNoUnsignedWrap() &&
3221 N0->getFlags().hasNoUnsignedWrap() &&
3222 OMul->getFlags().hasNoUnsignedWrap() &&
3223 OMul.getOperand(i: 0)->getFlags().hasNoUnsignedWrap()) {
3224 Flags |= SDNodeFlags::NoUnsignedWrap;
3225 if (N->getFlags().hasNoSignedWrap() &&
3226 N0->getFlags().hasNoSignedWrap() &&
3227 OMul->getFlags().hasNoSignedWrap() &&
3228 OMul.getOperand(i: 0)->getFlags().hasNoSignedWrap())
3229 Flags |= SDNodeFlags::NoSignedWrap;
3230 }
3231 SDValue Mul = DAG.getNode(Opcode: ISD::MUL, DL: SDLoc(N1), VT, N1: A,
3232 N2: DAG.getConstant(Val: CM, DL, VT), Flags);
3233 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL: SDLoc(N1), VT, N1: Mul, N2: B, Flags);
3234 return DAG.getNode(
3235 Opcode: ISD::ADD, DL, VT, N1: Add,
3236 N2: DAG.getConstant(Val: CA * CM + CB->getAPIntValue(), DL, VT), Flags);
3237 }
3238 }
3239 }
3240
3241 if (SDValue Combined = visitADDLikeCommutative(N0, N1, DL))
3242 return Combined;
3243
3244 if (SDValue Combined = visitADDLikeCommutative(N0: N1, N1: N0, DL))
3245 return Combined;
3246
3247 return SDValue();
3248}
3249
3250// Attempt to form avgfloor(A, B) from (A & B) + ((A ^ B) >> 1)
3251// Attempt to form avgfloor(A, B) from ((A >> 1) + (B >> 1)) + (A & B & 1)
3252// Attempt to form avgceil(A, B) from ((A >> 1) + (B >> 1)) + ((A | B) & 1)
3253SDValue DAGCombiner::foldAddToAvg(SDNode *N, const SDLoc &DL) {
3254 SDValue N0 = N->getOperand(Num: 0);
3255 EVT VT = N0.getValueType();
3256 SDValue A, B;
3257
3258 if ((!LegalOperations || hasOperation(Opcode: ISD::AVGFLOORU, VT)) &&
3259 (sd_match(N,
3260 P: m_Add(L: m_And(L: m_Value(N&: A), R: m_Value(N&: B)),
3261 R: m_Srl(L: m_Xor(L: m_Deferred(V&: A), R: m_Deferred(V&: B)), R: m_One()))) ||
3262 sd_match(N, P: m_ReassociatableAdd(
3263 Patterns: m_ReassociatableAnd(Patterns: m_Value(N&: A), Patterns: m_Value(N&: B), Patterns: m_One()),
3264 Patterns: m_Srl(L: m_Deferred(V&: A), R: m_One()),
3265 Patterns: m_Srl(L: m_Deferred(V&: B), R: m_One()))))) {
3266 return DAG.getNode(Opcode: ISD::AVGFLOORU, DL, VT, N1: A, N2: B);
3267 }
3268 if ((!LegalOperations || hasOperation(Opcode: ISD::AVGFLOORS, VT)) &&
3269 (sd_match(N,
3270 P: m_Add(L: m_And(L: m_Value(N&: A), R: m_Value(N&: B)),
3271 R: m_Sra(L: m_Xor(L: m_Deferred(V&: A), R: m_Deferred(V&: B)), R: m_One()))) ||
3272 sd_match(N, P: m_ReassociatableAdd(
3273 Patterns: m_ReassociatableAnd(Patterns: m_Value(N&: A), Patterns: m_Value(N&: B), Patterns: m_One()),
3274 Patterns: m_Sra(L: m_Deferred(V&: A), R: m_One()),
3275 Patterns: m_Sra(L: m_Deferred(V&: B), R: m_One()))))) {
3276 return DAG.getNode(Opcode: ISD::AVGFLOORS, DL, VT, N1: A, N2: B);
3277 }
3278
3279 if ((!LegalOperations || hasOperation(Opcode: ISD::AVGCEILU, VT)) &&
3280 sd_match(N,
3281 P: m_ReassociatableAdd(Patterns: m_And(L: m_Or(L: m_Value(N&: A), R: m_Value(N&: B)), R: m_One()),
3282 Patterns: m_Srl(L: m_Deferred(V&: A), R: m_One()),
3283 Patterns: m_Srl(L: m_Deferred(V&: B), R: m_One())))) {
3284 return DAG.getNode(Opcode: ISD::AVGCEILU, DL, VT, N1: A, N2: B);
3285 }
3286 if ((!LegalOperations || hasOperation(Opcode: ISD::AVGCEILS, VT)) &&
3287 sd_match(N,
3288 P: m_ReassociatableAdd(Patterns: m_And(L: m_Or(L: m_Value(N&: A), R: m_Value(N&: B)), R: m_One()),
3289 Patterns: m_Sra(L: m_Deferred(V&: A), R: m_One()),
3290 Patterns: m_Sra(L: m_Deferred(V&: B), R: m_One())))) {
3291 return DAG.getNode(Opcode: ISD::AVGCEILS, DL, VT, N1: A, N2: B);
3292 }
3293
3294 return SDValue();
3295}
3296
3297SDValue DAGCombiner::visitADD(SDNode *N) {
3298 SDValue N0 = N->getOperand(Num: 0);
3299 SDValue N1 = N->getOperand(Num: 1);
3300 EVT VT = N0.getValueType();
3301 SDLoc DL(N);
3302
3303 if (SDValue Combined = visitADDLike(N))
3304 return Combined;
3305
3306 if (SDValue V = foldAddSubBoolOfMaskedVal(N, DL, DAG))
3307 return V;
3308
3309 if (SDValue V = foldAddSubOfSignBit(N, DL, DAG))
3310 return V;
3311
3312 if (SDValue V = MatchRotate(LHS: N0, RHS: N1, DL: SDLoc(N), /*FromAdd=*/true))
3313 return V;
3314
3315 // Try to match AVGFLOOR fixedwidth pattern
3316 if (SDValue V = foldAddToAvg(N, DL))
3317 return V;
3318
3319 // fold (a+b) -> (a|b) iff a and b share no bits.
3320 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::OR, VT)) &&
3321 DAG.haveNoCommonBitsSet(A: N0, B: N1))
3322 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: N0, N2: N1, Flags: SDNodeFlags::Disjoint);
3323
3324 // Fold (add (vscale * C0), (vscale * C1)) to (vscale * (C0 + C1)).
3325 if (N0.getOpcode() == ISD::VSCALE && N1.getOpcode() == ISD::VSCALE) {
3326 const APInt &C0 = N0->getConstantOperandAPInt(Num: 0);
3327 const APInt &C1 = N1->getConstantOperandAPInt(Num: 0);
3328 return DAG.getVScale(DL, VT, MulImm: C0 + C1);
3329 }
3330
3331 // fold a+vscale(c1)+vscale(c2) -> a+vscale(c1+c2)
3332 if (N0.getOpcode() == ISD::ADD &&
3333 N0.getOperand(i: 1).getOpcode() == ISD::VSCALE &&
3334 N1.getOpcode() == ISD::VSCALE && TLI.isProfitableToFoldVScaleAdd(N: N0)) {
3335 const APInt &VS0 = N0.getOperand(i: 1)->getConstantOperandAPInt(Num: 0);
3336 const APInt &VS1 = N1->getConstantOperandAPInt(Num: 0);
3337 SDValue VS = DAG.getVScale(DL, VT, MulImm: VS0 + VS1);
3338 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 0), N2: VS);
3339 }
3340
3341 // Fold (add step_vector(c1), step_vector(c2) to step_vector(c1+c2))
3342 if (N0.getOpcode() == ISD::STEP_VECTOR &&
3343 N1.getOpcode() == ISD::STEP_VECTOR) {
3344 const APInt &C0 = N0->getConstantOperandAPInt(Num: 0);
3345 const APInt &C1 = N1->getConstantOperandAPInt(Num: 0);
3346 APInt NewStep = C0 + C1;
3347 return DAG.getStepVector(DL, ResVT: VT, StepVal: NewStep);
3348 }
3349
3350 // Fold a + step_vector(c1) + step_vector(c2) to a + step_vector(c1+c2)
3351 if (N0.getOpcode() == ISD::ADD &&
3352 N0.getOperand(i: 1).getOpcode() == ISD::STEP_VECTOR &&
3353 N1.getOpcode() == ISD::STEP_VECTOR) {
3354 const APInt &SV0 = N0.getOperand(i: 1)->getConstantOperandAPInt(Num: 0);
3355 const APInt &SV1 = N1->getConstantOperandAPInt(Num: 0);
3356 APInt NewStep = SV0 + SV1;
3357 SDValue SV = DAG.getStepVector(DL, ResVT: VT, StepVal: NewStep);
3358 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 0), N2: SV);
3359 }
3360
3361 return SDValue();
3362}
3363
3364SDValue DAGCombiner::visitADDSAT(SDNode *N) {
3365 unsigned Opcode = N->getOpcode();
3366 SDValue N0 = N->getOperand(Num: 0);
3367 SDValue N1 = N->getOperand(Num: 1);
3368 EVT VT = N0.getValueType();
3369 bool IsSigned = Opcode == ISD::SADDSAT;
3370 SDLoc DL(N);
3371
3372 // fold (add_sat x, undef) -> -1
3373 if (N0.isUndef() || N1.isUndef())
3374 return DAG.getAllOnesConstant(DL, VT);
3375
3376 // fold (add_sat c1, c2) -> c3
3377 if (SDValue C = DAG.FoldConstantArithmetic(Opcode, DL, VT, Ops: {N0, N1}))
3378 return C;
3379
3380 // canonicalize constant to RHS
3381 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
3382 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
3383 return DAG.getNode(Opcode, DL, VT, N1, N2: N0);
3384
3385 // fold vector ops
3386 if (VT.isVector()) {
3387 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
3388 return FoldedVOp;
3389
3390 // fold (add_sat x, 0) -> x, vector edition
3391 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
3392 return N0;
3393 }
3394
3395 // fold (add_sat x, 0) -> x
3396 if (isNullConstant(V: N1))
3397 return N0;
3398
3399 // If it cannot overflow, transform into an add.
3400 if (DAG.willNotOverflowAdd(IsSigned, N0, N1))
3401 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: N1);
3402
3403 return SDValue();
3404}
3405
3406static SDValue getAsCarry(const TargetLowering &TLI, SDValue V,
3407 bool ForceCarryReconstruction = false) {
3408 bool Masked = false;
3409
3410 // First, peel away TRUNCATE/ZERO_EXTEND/AND nodes due to legalization.
3411 while (true) {
3412 if (ForceCarryReconstruction && V.getValueType() == MVT::i1)
3413 return V;
3414
3415 if (V.getOpcode() == ISD::TRUNCATE || V.getOpcode() == ISD::ZERO_EXTEND) {
3416 V = V.getOperand(i: 0);
3417 continue;
3418 }
3419
3420 if (V.getOpcode() == ISD::AND && isOneConstant(V: V.getOperand(i: 1))) {
3421 if (ForceCarryReconstruction)
3422 return V;
3423
3424 Masked = true;
3425 V = V.getOperand(i: 0);
3426 continue;
3427 }
3428
3429 break;
3430 }
3431
3432 // If this is not a carry, return.
3433 if (V.getResNo() != 1)
3434 return SDValue();
3435
3436 if (V.getOpcode() != ISD::UADDO_CARRY && V.getOpcode() != ISD::USUBO_CARRY &&
3437 V.getOpcode() != ISD::UADDO && V.getOpcode() != ISD::USUBO)
3438 return SDValue();
3439
3440 EVT VT = V->getValueType(ResNo: 0);
3441 if (!TLI.isOperationLegalOrCustom(Op: V.getOpcode(), VT))
3442 return SDValue();
3443
3444 // If the result is masked, then no matter what kind of bool it is we can
3445 // return. If it isn't, then we need to make sure the bool type is either 0 or
3446 // 1 and not other values.
3447 if (Masked ||
3448 TLI.getBooleanContents(Type: V.getValueType()) ==
3449 TargetLoweringBase::ZeroOrOneBooleanContent)
3450 return V;
3451
3452 return SDValue();
3453}
3454
3455/// Given the operands of an add/sub operation, see if the 2nd operand is a
3456/// masked 0/1 whose source operand is actually known to be 0/-1. If so, invert
3457/// the opcode and bypass the mask operation.
3458static SDValue foldAddSubMasked1(bool IsAdd, SDValue N0, SDValue N1,
3459 SelectionDAG &DAG, const SDLoc &DL) {
3460 if (N1.getOpcode() == ISD::ZERO_EXTEND)
3461 N1 = N1.getOperand(i: 0);
3462
3463 if (N1.getOpcode() != ISD::AND || !isOneOrOneSplat(V: N1->getOperand(Num: 1)))
3464 return SDValue();
3465
3466 EVT VT = N0.getValueType();
3467 SDValue N10 = N1.getOperand(i: 0);
3468 if (N10.getValueType() != VT && N10.getOpcode() == ISD::TRUNCATE)
3469 N10 = N10.getOperand(i: 0);
3470
3471 if (N10.getValueType() != VT)
3472 return SDValue();
3473
3474 if (DAG.ComputeNumSignBits(Op: N10) != VT.getScalarSizeInBits())
3475 return SDValue();
3476
3477 // add N0, (and (AssertSext X, i1), 1) --> sub N0, X
3478 // sub N0, (and (AssertSext X, i1), 1) --> add N0, X
3479 return DAG.getNode(Opcode: IsAdd ? ISD::SUB : ISD::ADD, DL, VT, N1: N0, N2: N10);
3480}
3481
3482/// Helper for doing combines based on N0 and N1 being added to each other.
3483SDValue DAGCombiner::visitADDLikeCommutative(SDValue N0, SDValue N1,
3484 const SDLoc &DL) {
3485 EVT VT = N0.getValueType();
3486
3487 // fold (add x, shl(0 - y, n)) -> sub(x, shl(y, n))
3488 SDValue Y, N;
3489 if (sd_match(N: N1, P: m_Shl(L: m_Neg(V: m_Value(N&: Y)), R: m_Value(N))))
3490 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0,
3491 N2: DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Y, N2: N));
3492
3493 if (SDValue V = foldAddSubMasked1(IsAdd: true, N0, N1, DAG, DL))
3494 return V;
3495
3496 // Look for:
3497 // add (add x, 1), y
3498 // And if the target does not like this form then turn into:
3499 // sub x, (xor y, -1)
3500 if (!TLI.preferIncOfAddToSubOfNot(VT) && N0.getOpcode() == ISD::ADD &&
3501 N0.hasOneUse() && isOneOrOneSplat(V: N0.getOperand(i: 1)) &&
3502 // Limit this to after legalization if the add has wrap flags
3503 (Level >= AfterLegalizeDAG || (!N0->getFlags().hasNoUnsignedWrap() &&
3504 !N0->getFlags().hasNoSignedWrap()))) {
3505 SDValue Not = DAG.getNOT(DL, Val: N1, VT);
3506 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0.getOperand(i: 0), N2: Not);
3507 }
3508
3509 if (N0.getOpcode() == ISD::SUB && N0.hasOneUse()) {
3510 // Hoist one-use subtraction by non-opaque constant:
3511 // (x - C) + y -> (x + y) - C
3512 // This is necessary because SUB(X,C) -> ADD(X,-C) doesn't work for vectors.
3513 if (isConstantOrConstantVector(N: N0.getOperand(i: 1), /*NoOpaques=*/true)) {
3514 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 0), N2: N1);
3515 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Add, N2: N0.getOperand(i: 1));
3516 }
3517 // Hoist one-use subtraction from non-opaque constant:
3518 // (C - x) + y -> (y - x) + C
3519 if (isConstantOrConstantVector(N: N0.getOperand(i: 0), /*NoOpaques=*/true)) {
3520 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1, N2: N0.getOperand(i: 1));
3521 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Sub, N2: N0.getOperand(i: 0));
3522 }
3523 }
3524
3525 // add (mul x, C), x -> mul x, C+1
3526 if (N0.getOpcode() == ISD::MUL && N0.getOperand(i: 0) == N1 &&
3527 isConstantOrConstantVector(N: N0.getOperand(i: 1), /*NoOpaques=*/true) &&
3528 N0.hasOneUse()) {
3529 SDValue NewC = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 1),
3530 N2: DAG.getConstant(Val: 1, DL, VT));
3531 return DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: N0.getOperand(i: 0), N2: NewC);
3532 }
3533
3534 // If the target's bool is represented as 0/1, prefer to make this 'sub 0/1'
3535 // rather than 'add 0/-1' (the zext should get folded).
3536 // add (sext i1 Y), X --> sub X, (zext i1 Y)
3537 if (N0.getOpcode() == ISD::SIGN_EXTEND &&
3538 N0.getOperand(i: 0).getScalarValueSizeInBits() == 1 &&
3539 TLI.getBooleanContents(Type: VT) == TargetLowering::ZeroOrOneBooleanContent) {
3540 SDValue ZExt = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: N0.getOperand(i: 0));
3541 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1, N2: ZExt);
3542 }
3543
3544 // add X, (sextinreg Y i1) -> sub X, (and Y 1)
3545 if (N1.getOpcode() == ISD::SIGN_EXTEND_INREG) {
3546 VTSDNode *TN = cast<VTSDNode>(Val: N1.getOperand(i: 1));
3547 if (TN->getVT() == MVT::i1) {
3548 SDValue ZExt = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N1.getOperand(i: 0),
3549 N2: DAG.getConstant(Val: 1, DL, VT));
3550 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: ZExt);
3551 }
3552 }
3553
3554 // (add X, (uaddo_carry Y, 0, Carry)) -> (uaddo_carry X, Y, Carry)
3555 if (N1.getOpcode() == ISD::UADDO_CARRY && isNullConstant(V: N1.getOperand(i: 1)) &&
3556 N1.getResNo() == 0)
3557 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL, VTList: N1->getVTList(),
3558 N1: N0, N2: N1.getOperand(i: 0), N3: N1.getOperand(i: 2));
3559
3560 // (add X, Carry) -> (uaddo_carry X, 0, Carry)
3561 if (TLI.isOperationLegalOrCustom(Op: ISD::UADDO_CARRY, VT))
3562 if (SDValue Carry = getAsCarry(TLI, V: N1))
3563 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL,
3564 VTList: DAG.getVTList(VT1: VT, VT2: Carry.getValueType()), N1: N0,
3565 N2: DAG.getConstant(Val: 0, DL, VT), N3: Carry);
3566
3567 return SDValue();
3568}
3569
3570SDValue DAGCombiner::visitADDC(SDNode *N) {
3571 SDValue N0 = N->getOperand(Num: 0);
3572 SDValue N1 = N->getOperand(Num: 1);
3573 EVT VT = N0.getValueType();
3574 SDLoc DL(N);
3575
3576 // If the flag result is dead, turn this into an ADD.
3577 if (!N->hasAnyUseOfValue(Value: 1))
3578 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: N1),
3579 Res1: DAG.getNode(Opcode: ISD::CARRY_FALSE, DL, VT: MVT::Glue));
3580
3581 // canonicalize constant to RHS.
3582 ConstantSDNode *N0C = dyn_cast<ConstantSDNode>(Val&: N0);
3583 ConstantSDNode *N1C = dyn_cast<ConstantSDNode>(Val&: N1);
3584 if (N0C && !N1C)
3585 return DAG.getNode(Opcode: ISD::ADDC, DL, VTList: N->getVTList(), N1, N2: N0);
3586
3587 // fold (addc x, 0) -> x + no carry out
3588 if (isNullConstant(V: N1))
3589 return CombineTo(N, Res0: N0, Res1: DAG.getNode(Opcode: ISD::CARRY_FALSE,
3590 DL, VT: MVT::Glue));
3591
3592 // If it cannot overflow, transform into an add.
3593 if (DAG.computeOverflowForUnsignedAdd(N0, N1) == SelectionDAG::OFK_Never)
3594 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: N1),
3595 Res1: DAG.getNode(Opcode: ISD::CARRY_FALSE, DL, VT: MVT::Glue));
3596
3597 return SDValue();
3598}
3599
3600/**
3601 * Flips a boolean if it is cheaper to compute. If the Force parameters is set,
3602 * then the flip also occurs if computing the inverse is the same cost.
3603 * This function returns an empty SDValue in case it cannot flip the boolean
3604 * without increasing the cost of the computation. If you want to flip a boolean
3605 * no matter what, use DAG.getLogicalNOT.
3606 */
3607static SDValue extractBooleanFlip(SDValue V, SelectionDAG &DAG,
3608 const TargetLowering &TLI,
3609 bool Force) {
3610 if (Force && isa<ConstantSDNode>(Val: V))
3611 return DAG.getLogicalNOT(DL: SDLoc(V), Val: V, VT: V.getValueType());
3612
3613 if (V.getOpcode() != ISD::XOR)
3614 return SDValue();
3615
3616 if (DAG.isBoolConstant(N: V.getOperand(i: 1)) == true)
3617 return V.getOperand(i: 0);
3618 if (Force && isConstOrConstSplat(N: V.getOperand(i: 1), AllowUndefs: false))
3619 return DAG.getLogicalNOT(DL: SDLoc(V), Val: V, VT: V.getValueType());
3620 return SDValue();
3621}
3622
3623SDValue DAGCombiner::visitADDO(SDNode *N) {
3624 SDValue N0 = N->getOperand(Num: 0);
3625 SDValue N1 = N->getOperand(Num: 1);
3626 EVT VT = N0.getValueType();
3627 bool IsSigned = (ISD::SADDO == N->getOpcode());
3628
3629 EVT CarryVT = N->getValueType(ResNo: 1);
3630 SDLoc DL(N);
3631
3632 // If the flag result is dead, turn this into an ADD.
3633 if (!N->hasAnyUseOfValue(Value: 1))
3634 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: N1),
3635 Res1: DAG.getUNDEF(VT: CarryVT));
3636
3637 // canonicalize constant to RHS.
3638 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
3639 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
3640 return DAG.getNode(Opcode: N->getOpcode(), DL, VTList: N->getVTList(), N1, N2: N0);
3641
3642 // fold (addo x, 0) -> x + no carry out
3643 if (isNullOrNullSplat(V: N1))
3644 return CombineTo(N, Res0: N0, Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
3645
3646 // If it cannot overflow, transform into an add.
3647 if (DAG.willNotOverflowAdd(IsSigned, N0, N1))
3648 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: N1),
3649 Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
3650
3651 if (IsSigned) {
3652 // fold (saddo (xor a, -1), 1) -> (ssub 0, a).
3653 if (isBitwiseNot(V: N0) && isOneOrOneSplat(V: N1))
3654 return DAG.getNode(Opcode: ISD::SSUBO, DL, VTList: N->getVTList(),
3655 N1: DAG.getConstant(Val: 0, DL, VT), N2: N0.getOperand(i: 0));
3656 } else {
3657 // fold (uaddo (xor a, -1), 1) -> (usub 0, a) and flip carry.
3658 if (isBitwiseNot(V: N0) && isOneOrOneSplat(V: N1)) {
3659 SDValue Sub = DAG.getNode(Opcode: ISD::USUBO, DL, VTList: N->getVTList(),
3660 N1: DAG.getConstant(Val: 0, DL, VT), N2: N0.getOperand(i: 0));
3661 return CombineTo(
3662 N, Res0: Sub, Res1: DAG.getLogicalNOT(DL, Val: Sub.getValue(R: 1), VT: Sub->getValueType(ResNo: 1)));
3663 }
3664
3665 if (SDValue Combined = visitUADDOLike(N0, N1, N))
3666 return Combined;
3667
3668 if (SDValue Combined = visitUADDOLike(N0: N1, N1: N0, N))
3669 return Combined;
3670 }
3671
3672 return SDValue();
3673}
3674
3675SDValue DAGCombiner::visitUADDOLike(SDValue N0, SDValue N1, SDNode *N) {
3676 EVT VT = N0.getValueType();
3677 if (VT.isVector())
3678 return SDValue();
3679
3680 // (uaddo X, (uaddo_carry Y, 0, Carry)) -> (uaddo_carry X, Y, Carry)
3681 // If Y + 1 cannot overflow.
3682 if (N1.getOpcode() == ISD::UADDO_CARRY && isNullConstant(V: N1.getOperand(i: 1))) {
3683 SDValue Y = N1.getOperand(i: 0);
3684 SDValue One = DAG.getConstant(Val: 1, DL: SDLoc(N), VT: Y.getValueType());
3685 if (DAG.computeOverflowForUnsignedAdd(N0: Y, N1: One) == SelectionDAG::OFK_Never)
3686 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL: SDLoc(N), VTList: N->getVTList(), N1: N0, N2: Y,
3687 N3: N1.getOperand(i: 2));
3688 }
3689
3690 // (uaddo X, Carry) -> (uaddo_carry X, 0, Carry)
3691 if (TLI.isOperationLegalOrCustom(Op: ISD::UADDO_CARRY, VT))
3692 if (SDValue Carry = getAsCarry(TLI, V: N1))
3693 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL: SDLoc(N), VTList: N->getVTList(), N1: N0,
3694 N2: DAG.getConstant(Val: 0, DL: SDLoc(N), VT), N3: Carry);
3695
3696 return SDValue();
3697}
3698
3699SDValue DAGCombiner::visitADDE(SDNode *N) {
3700 SDValue N0 = N->getOperand(Num: 0);
3701 SDValue N1 = N->getOperand(Num: 1);
3702 SDValue CarryIn = N->getOperand(Num: 2);
3703
3704 // canonicalize constant to RHS
3705 ConstantSDNode *N0C = dyn_cast<ConstantSDNode>(Val&: N0);
3706 ConstantSDNode *N1C = dyn_cast<ConstantSDNode>(Val&: N1);
3707 if (N0C && !N1C)
3708 return DAG.getNode(Opcode: ISD::ADDE, DL: SDLoc(N), VTList: N->getVTList(),
3709 N1, N2: N0, N3: CarryIn);
3710
3711 // fold (adde x, y, false) -> (addc x, y)
3712 if (CarryIn.getOpcode() == ISD::CARRY_FALSE)
3713 return DAG.getNode(Opcode: ISD::ADDC, DL: SDLoc(N), VTList: N->getVTList(), N1: N0, N2: N1);
3714
3715 return SDValue();
3716}
3717
3718SDValue DAGCombiner::visitUADDO_CARRY(SDNode *N) {
3719 SDValue N0 = N->getOperand(Num: 0);
3720 SDValue N1 = N->getOperand(Num: 1);
3721 SDValue CarryIn = N->getOperand(Num: 2);
3722 SDLoc DL(N);
3723
3724 // canonicalize constant to RHS
3725 ConstantSDNode *N0C = dyn_cast<ConstantSDNode>(Val&: N0);
3726 ConstantSDNode *N1C = dyn_cast<ConstantSDNode>(Val&: N1);
3727 if (N0C && !N1C)
3728 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL, VTList: N->getVTList(), N1, N2: N0, N3: CarryIn);
3729
3730 // fold (uaddo_carry x, y, false) -> (uaddo x, y)
3731 if (isNullConstant(V: CarryIn)) {
3732 if (!LegalOperations ||
3733 TLI.isOperationLegalOrCustom(Op: ISD::UADDO, VT: N->getValueType(ResNo: 0)))
3734 return DAG.getNode(Opcode: ISD::UADDO, DL, VTList: N->getVTList(), N1: N0, N2: N1);
3735 }
3736
3737 // fold (uaddo_carry 0, 0, X) -> (and (ext/trunc X), 1) and no carry.
3738 if (isNullConstant(V: N0) && isNullConstant(V: N1)) {
3739 EVT VT = N0.getValueType();
3740 EVT CarryVT = CarryIn.getValueType();
3741 SDValue CarryExt = DAG.getBoolExtOrTrunc(Op: CarryIn, SL: DL, VT, OpVT: CarryVT);
3742 AddToWorklist(N: CarryExt.getNode());
3743 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::AND, DL, VT, N1: CarryExt,
3744 N2: DAG.getConstant(Val: 1, DL, VT)),
3745 Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
3746 }
3747
3748 if (SDValue Combined = visitUADDO_CARRYLike(N0, N1, CarryIn, N))
3749 return Combined;
3750
3751 if (SDValue Combined = visitUADDO_CARRYLike(N0: N1, N1: N0, CarryIn, N))
3752 return Combined;
3753
3754 // We want to avoid useless duplication.
3755 // TODO: This is done automatically for binary operations. As UADDO_CARRY is
3756 // not a binary operation, this is not really possible to leverage this
3757 // existing mechanism for it. However, if more operations require the same
3758 // deduplication logic, then it may be worth generalize.
3759 SDValue Ops[] = {N1, N0, CarryIn};
3760 SDNode *CSENode =
3761 DAG.getNodeIfExists(Opcode: ISD::UADDO_CARRY, VTList: N->getVTList(), Ops, Flags: N->getFlags());
3762 if (CSENode)
3763 return SDValue(CSENode, 0);
3764
3765 return SDValue();
3766}
3767
3768/**
3769 * If we are facing some sort of diamond carry propagation pattern try to
3770 * break it up to generate something like:
3771 * (uaddo_carry X, 0, (uaddo_carry A, B, Z):Carry)
3772 *
3773 * The end result is usually an increase in operation required, but because the
3774 * carry is now linearized, other transforms can kick in and optimize the DAG.
3775 *
3776 * Patterns typically look something like
3777 * (uaddo A, B)
3778 * / \
3779 * Carry Sum
3780 * | \
3781 * | (uaddo_carry *, 0, Z)
3782 * | /
3783 * \ Carry
3784 * | /
3785 * (uaddo_carry X, *, *)
3786 *
3787 * But numerous variation exist. Our goal is to identify A, B, X and Z and
3788 * produce a combine with a single path for carry propagation.
3789 */
3790static SDValue combineUADDO_CARRYDiamond(DAGCombiner &Combiner,
3791 SelectionDAG &DAG, SDValue X,
3792 SDValue Carry0, SDValue Carry1,
3793 SDNode *N) {
3794 if (Carry1.getResNo() != 1 || Carry0.getResNo() != 1)
3795 return SDValue();
3796 if (Carry1.getOpcode() != ISD::UADDO)
3797 return SDValue();
3798
3799 SDValue Z;
3800
3801 /**
3802 * First look for a suitable Z. It will present itself in the form of
3803 * (uaddo_carry Y, 0, Z) or its equivalent (uaddo Y, 1) for Z=true
3804 */
3805 if (Carry0.getOpcode() == ISD::UADDO_CARRY &&
3806 isNullConstant(V: Carry0.getOperand(i: 1))) {
3807 Z = Carry0.getOperand(i: 2);
3808 } else if (Carry0.getOpcode() == ISD::UADDO &&
3809 isOneConstant(V: Carry0.getOperand(i: 1))) {
3810 EVT VT = Carry0->getValueType(ResNo: 1);
3811 Z = DAG.getConstant(Val: 1, DL: SDLoc(Carry0.getOperand(i: 1)), VT);
3812 } else {
3813 // We couldn't find a suitable Z.
3814 return SDValue();
3815 }
3816
3817
3818 auto cancelDiamond = [&](SDValue A,SDValue B) {
3819 SDLoc DL(N);
3820 SDValue NewY =
3821 DAG.getNode(Opcode: ISD::UADDO_CARRY, DL, VTList: Carry0->getVTList(), N1: A, N2: B, N3: Z);
3822 Combiner.AddToWorklist(N: NewY.getNode());
3823 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL, VTList: N->getVTList(), N1: X,
3824 N2: DAG.getConstant(Val: 0, DL, VT: X.getValueType()),
3825 N3: NewY.getValue(R: 1));
3826 };
3827
3828 /**
3829 * (uaddo A, B)
3830 * |
3831 * Sum
3832 * |
3833 * (uaddo_carry *, 0, Z)
3834 */
3835 if (Carry0.getOperand(i: 0) == Carry1.getValue(R: 0)) {
3836 return cancelDiamond(Carry1.getOperand(i: 0), Carry1.getOperand(i: 1));
3837 }
3838
3839 /**
3840 * (uaddo_carry A, 0, Z)
3841 * |
3842 * Sum
3843 * |
3844 * (uaddo *, B)
3845 */
3846 if (Carry1.getOperand(i: 0) == Carry0.getValue(R: 0)) {
3847 return cancelDiamond(Carry0.getOperand(i: 0), Carry1.getOperand(i: 1));
3848 }
3849
3850 if (Carry1.getOperand(i: 1) == Carry0.getValue(R: 0)) {
3851 return cancelDiamond(Carry1.getOperand(i: 0), Carry0.getOperand(i: 0));
3852 }
3853
3854 return SDValue();
3855}
3856
3857// If we are facing some sort of diamond carry/borrow in/out pattern try to
3858// match patterns like:
3859//
3860// (uaddo A, B) CarryIn
3861// | \ |
3862// | \ |
3863// PartialSum PartialCarryOutX /
3864// | | /
3865// | ____|____________/
3866// | / |
3867// (uaddo *, *) \________
3868// | \ \
3869// | \ |
3870// | PartialCarryOutY |
3871// | \ |
3872// | \ /
3873// AddCarrySum | ______/
3874// | /
3875// CarryOut = (or *, *)
3876//
3877// And generate UADDO_CARRY (or USUBO_CARRY) with two result values:
3878//
3879// {AddCarrySum, CarryOut} = (uaddo_carry A, B, CarryIn)
3880//
3881// Our goal is to identify A, B, and CarryIn and produce UADDO_CARRY/USUBO_CARRY
3882// with a single path for carry/borrow out propagation.
3883static SDValue combineCarryDiamond(SelectionDAG &DAG, const TargetLowering &TLI,
3884 SDValue N0, SDValue N1, SDNode *N) {
3885 SDValue Carry0 = getAsCarry(TLI, V: N0);
3886 if (!Carry0)
3887 return SDValue();
3888 SDValue Carry1 = getAsCarry(TLI, V: N1);
3889 if (!Carry1)
3890 return SDValue();
3891
3892 unsigned Opcode = Carry0.getOpcode();
3893 if (Opcode != Carry1.getOpcode())
3894 return SDValue();
3895 if (Opcode != ISD::UADDO && Opcode != ISD::USUBO)
3896 return SDValue();
3897 // Guarantee identical type of CarryOut
3898 EVT CarryOutType = N->getValueType(ResNo: 0);
3899 if (CarryOutType != Carry0.getValue(R: 1).getValueType() ||
3900 CarryOutType != Carry1.getValue(R: 1).getValueType())
3901 return SDValue();
3902
3903 // Canonicalize the add/sub of A and B (the top node in the above ASCII art)
3904 // as Carry0 and the add/sub of the carry in as Carry1 (the middle node).
3905 if (Carry1.getNode()->isOperandOf(N: Carry0.getNode()))
3906 std::swap(a&: Carry0, b&: Carry1);
3907
3908 // Check if nodes are connected in expected way.
3909 if (Carry1.getOperand(i: 0) != Carry0.getValue(R: 0) &&
3910 Carry1.getOperand(i: 1) != Carry0.getValue(R: 0))
3911 return SDValue();
3912
3913 // The carry in value must be on the righthand side for subtraction.
3914 unsigned CarryInOperandNum =
3915 Carry1.getOperand(i: 0) == Carry0.getValue(R: 0) ? 1 : 0;
3916 if (Opcode == ISD::USUBO && CarryInOperandNum != 1)
3917 return SDValue();
3918 SDValue CarryIn = Carry1.getOperand(i: CarryInOperandNum);
3919
3920 unsigned NewOp = Opcode == ISD::UADDO ? ISD::UADDO_CARRY : ISD::USUBO_CARRY;
3921 if (!TLI.isOperationLegalOrCustom(Op: NewOp, VT: Carry0.getValue(R: 0).getValueType()))
3922 return SDValue();
3923
3924 // Verify that the carry/borrow in is plausibly a carry/borrow bit.
3925 CarryIn = getAsCarry(TLI, V: CarryIn, ForceCarryReconstruction: true);
3926 if (!CarryIn)
3927 return SDValue();
3928
3929 SDLoc DL(N);
3930 CarryIn = DAG.getBoolExtOrTrunc(Op: CarryIn, SL: DL, VT: Carry1->getValueType(ResNo: 1),
3931 OpVT: Carry1->getValueType(ResNo: 0));
3932 SDValue Merged =
3933 DAG.getNode(Opcode: NewOp, DL, VTList: Carry1->getVTList(), N1: Carry0.getOperand(i: 0),
3934 N2: Carry0.getOperand(i: 1), N3: CarryIn);
3935
3936 // Please note that because we have proven that the result of the UADDO/USUBO
3937 // of A and B feeds into the UADDO/USUBO that does the carry/borrow in, we can
3938 // therefore prove that if the first UADDO/USUBO overflows, the second
3939 // UADDO/USUBO cannot. For example consider 8-bit numbers where 0xFF is the
3940 // maximum value.
3941 //
3942 // 0xFF + 0xFF == 0xFE with carry but 0xFE + 1 does not carry
3943 // 0x00 - 0xFF == 1 with a carry/borrow but 1 - 1 == 0 (no carry/borrow)
3944 //
3945 // This is important because it means that OR and XOR can be used to merge
3946 // carry flags; and that AND can return a constant zero.
3947 //
3948 // TODO: match other operations that can merge flags (ADD, etc)
3949 DAG.ReplaceAllUsesOfValueWith(From: Carry1.getValue(R: 0), To: Merged.getValue(R: 0));
3950 if (N->getOpcode() == ISD::AND)
3951 return DAG.getConstant(Val: 0, DL, VT: CarryOutType);
3952 return Merged.getValue(R: 1);
3953}
3954
3955// Reconstruct a subtract-with-borrow chain from its canonicalized icmp form:
3956// carry_out = or(icmp ult A, B, and(icmp eq A, B, carry_in))
3957// InstCombine folds usub.with.overflow chains into this, losing the
3958// USUBO_CARRY that lowers to sbb/sbcs.
3959static SDValue combineOrOfSetCCToUSUBOCarry(SDNode *N, SelectionDAG &DAG,
3960 const TargetLowering &TLI) {
3961 SDValue A, B, CarryIn;
3962 if (!sd_match(N, P: m_Or(L: m_SpecificSetCC(CC: ISD::SETULT, LHS: m_Value(N&: A), RHS: m_Value(N&: B)),
3963 R: m_And(L: m_c_SpecificSetCC(CC: ISD::SETEQ, LHS: m_Deferred(V&: A),
3964 RHS: m_Deferred(V&: B)),
3965 R: m_Value(N&: CarryIn)))))
3966 return SDValue();
3967
3968 EVT IntVT = A.getValueType();
3969 // Skip vectors: USUBO_CARRY on a vector type has no legalization path and
3970 // would crash.
3971 if (IntVT.isVector() || !TLI.isOperationLegalOrCustom(
3972 Op: ISD::USUBO_CARRY, VT: TLI.getLegalTypeToTransformTo(
3973 Context&: *DAG.getContext(), VT: IntVT)))
3974 return SDValue();
3975
3976 // USUBO_CARRY's carry-in must match boolean contents, which the matched
3977 // pattern does not guarantee.
3978 // TODO: Extend to other boolean contents.
3979 if (TLI.getBooleanContents(Type: IntVT) !=
3980 TargetLowering::ZeroOrOneBooleanContent ||
3981 !DAG.MaskedValueIsZero(
3982 Op: CarryIn,
3983 Mask: APInt::getBitsSetFrom(numBits: CarryIn.getScalarValueSizeInBits(), loBit: 1)))
3984 return SDValue();
3985
3986 SDLoc DL(N);
3987 SDVTList VTs = DAG.getVTList(VT1: IntVT, VT2: N->getValueType(ResNo: 0));
3988 return DAG.getNode(Opcode: ISD::USUBO_CARRY, DL, VTList: VTs, N1: A, N2: B, N3: CarryIn).getValue(R: 1);
3989}
3990
3991SDValue DAGCombiner::visitUADDO_CARRYLike(SDValue N0, SDValue N1,
3992 SDValue CarryIn, SDNode *N) {
3993 // fold (uaddo_carry (xor a, -1), b, c) -> (usubo_carry b, a, !c) and flip
3994 // carry.
3995 if (isBitwiseNot(V: N0))
3996 if (SDValue NotC = extractBooleanFlip(V: CarryIn, DAG, TLI, Force: true)) {
3997 SDLoc DL(N);
3998 SDValue Sub = DAG.getNode(Opcode: ISD::USUBO_CARRY, DL, VTList: N->getVTList(), N1,
3999 N2: N0.getOperand(i: 0), N3: NotC);
4000 return CombineTo(
4001 N, Res0: Sub, Res1: DAG.getLogicalNOT(DL, Val: Sub.getValue(R: 1), VT: Sub->getValueType(ResNo: 1)));
4002 }
4003
4004 // Iff the flag result is dead:
4005 // (uaddo_carry (add|uaddo X, Y), 0, Carry) -> (uaddo_carry X, Y, Carry)
4006 // Don't do this if the Carry comes from the uaddo. It won't remove the uaddo
4007 // or the dependency between the instructions.
4008 if ((N0.getOpcode() == ISD::ADD ||
4009 (N0.getOpcode() == ISD::UADDO && N0.getResNo() == 0 &&
4010 N0.getValue(R: 1) != CarryIn)) &&
4011 isNullConstant(V: N1) && !N->hasAnyUseOfValue(Value: 1))
4012 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL: SDLoc(N), VTList: N->getVTList(),
4013 N1: N0.getOperand(i: 0), N2: N0.getOperand(i: 1), N3: CarryIn);
4014
4015 /**
4016 * When one of the uaddo_carry argument is itself a carry, we may be facing
4017 * a diamond carry propagation. In which case we try to transform the DAG
4018 * to ensure linear carry propagation if that is possible.
4019 */
4020 if (auto Y = getAsCarry(TLI, V: N1)) {
4021 // Because both are carries, Y and Z can be swapped.
4022 if (auto R = combineUADDO_CARRYDiamond(Combiner&: *this, DAG, X: N0, Carry0: Y, Carry1: CarryIn, N))
4023 return R;
4024 if (auto R = combineUADDO_CARRYDiamond(Combiner&: *this, DAG, X: N0, Carry0: CarryIn, Carry1: Y, N))
4025 return R;
4026 }
4027
4028 return SDValue();
4029}
4030
4031SDValue DAGCombiner::visitSADDO_CARRYLike(SDValue N0, SDValue N1,
4032 SDValue CarryIn, SDNode *N) {
4033 // fold (saddo_carry (xor a, -1), b, c) -> (ssubo_carry b, a, !c)
4034 if (isBitwiseNot(V: N0)) {
4035 if (SDValue NotC = extractBooleanFlip(V: CarryIn, DAG, TLI, Force: true))
4036 return DAG.getNode(Opcode: ISD::SSUBO_CARRY, DL: SDLoc(N), VTList: N->getVTList(), N1,
4037 N2: N0.getOperand(i: 0), N3: NotC);
4038 }
4039
4040 return SDValue();
4041}
4042
4043SDValue DAGCombiner::visitSADDO_CARRY(SDNode *N) {
4044 SDValue N0 = N->getOperand(Num: 0);
4045 SDValue N1 = N->getOperand(Num: 1);
4046 SDValue CarryIn = N->getOperand(Num: 2);
4047 SDLoc DL(N);
4048
4049 // canonicalize constant to RHS
4050 ConstantSDNode *N0C = dyn_cast<ConstantSDNode>(Val&: N0);
4051 ConstantSDNode *N1C = dyn_cast<ConstantSDNode>(Val&: N1);
4052 if (N0C && !N1C)
4053 return DAG.getNode(Opcode: ISD::SADDO_CARRY, DL, VTList: N->getVTList(), N1, N2: N0, N3: CarryIn);
4054
4055 // fold (saddo_carry x, y, false) -> (saddo x, y)
4056 if (isNullConstant(V: CarryIn)) {
4057 if (!LegalOperations ||
4058 TLI.isOperationLegalOrCustom(Op: ISD::SADDO, VT: N->getValueType(ResNo: 0)))
4059 return DAG.getNode(Opcode: ISD::SADDO, DL, VTList: N->getVTList(), N1: N0, N2: N1);
4060 }
4061
4062 if (SDValue Combined = visitSADDO_CARRYLike(N0, N1, CarryIn, N))
4063 return Combined;
4064
4065 if (SDValue Combined = visitSADDO_CARRYLike(N0: N1, N1: N0, CarryIn, N))
4066 return Combined;
4067
4068 return SDValue();
4069}
4070
4071// Attempt to create a USUBSAT(LHS, RHS) node with DstVT, performing a
4072// clamp/truncation if necessary.
4073static SDValue getTruncatedUSUBSAT(EVT DstVT, EVT SrcVT, SDValue LHS,
4074 SDValue RHS, SelectionDAG &DAG,
4075 const SDLoc &DL) {
4076 assert(DstVT.getScalarSizeInBits() <= SrcVT.getScalarSizeInBits() &&
4077 "Illegal truncation");
4078
4079 if (DstVT == SrcVT)
4080 return DAG.getNode(Opcode: ISD::USUBSAT, DL, VT: DstVT, N1: LHS, N2: RHS);
4081
4082 // If the LHS is zero-extended then we can perform the USUBSAT as DstVT by
4083 // clamping RHS.
4084 APInt UpperBits = APInt::getBitsSetFrom(numBits: SrcVT.getScalarSizeInBits(),
4085 loBit: DstVT.getScalarSizeInBits());
4086 if (!DAG.MaskedValueIsZero(Op: LHS, Mask: UpperBits))
4087 return SDValue();
4088
4089 SDValue SatLimit =
4090 DAG.getConstant(Val: APInt::getLowBitsSet(numBits: SrcVT.getScalarSizeInBits(),
4091 loBitsSet: DstVT.getScalarSizeInBits()),
4092 DL, VT: SrcVT);
4093 RHS = DAG.getNode(Opcode: ISD::UMIN, DL, VT: SrcVT, N1: RHS, N2: SatLimit);
4094 RHS = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: DstVT, Operand: RHS);
4095 LHS = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: DstVT, Operand: LHS);
4096 return DAG.getNode(Opcode: ISD::USUBSAT, DL, VT: DstVT, N1: LHS, N2: RHS);
4097}
4098
4099// Try to find umax(a,b) - b or a - umin(a,b) patterns that may be converted to
4100// usubsat(a,b), optionally as a truncated type.
4101SDValue DAGCombiner::foldSubToUSubSat(EVT DstVT, SDNode *N, const SDLoc &DL) {
4102 if (N->getOpcode() != ISD::SUB ||
4103 !(!LegalOperations || hasOperation(Opcode: ISD::USUBSAT, VT: DstVT)))
4104 return SDValue();
4105
4106 EVT SubVT = N->getValueType(ResNo: 0);
4107 SDValue Op0 = N->getOperand(Num: 0);
4108 SDValue Op1 = N->getOperand(Num: 1);
4109
4110 // Try to find umax(a,b) - b or a - umin(a,b) patterns
4111 // they may be converted to usubsat(a,b).
4112 if (Op0.getOpcode() == ISD::UMAX && Op0.hasOneUse()) {
4113 SDValue MaxLHS = Op0.getOperand(i: 0);
4114 SDValue MaxRHS = Op0.getOperand(i: 1);
4115 if (MaxLHS == Op1)
4116 return getTruncatedUSUBSAT(DstVT, SrcVT: SubVT, LHS: MaxRHS, RHS: Op1, DAG, DL);
4117 if (MaxRHS == Op1)
4118 return getTruncatedUSUBSAT(DstVT, SrcVT: SubVT, LHS: MaxLHS, RHS: Op1, DAG, DL);
4119 }
4120
4121 if (Op1.getOpcode() == ISD::UMIN && Op1.hasOneUse()) {
4122 SDValue MinLHS = Op1.getOperand(i: 0);
4123 SDValue MinRHS = Op1.getOperand(i: 1);
4124 if (MinLHS == Op0)
4125 return getTruncatedUSUBSAT(DstVT, SrcVT: SubVT, LHS: Op0, RHS: MinRHS, DAG, DL);
4126 if (MinRHS == Op0)
4127 return getTruncatedUSUBSAT(DstVT, SrcVT: SubVT, LHS: Op0, RHS: MinLHS, DAG, DL);
4128 }
4129
4130 // sub(a,trunc(umin(zext(a),b))) -> usubsat(a,trunc(umin(b,SatLimit)))
4131 if (Op1.getOpcode() == ISD::TRUNCATE &&
4132 Op1.getOperand(i: 0).getOpcode() == ISD::UMIN &&
4133 Op1.getOperand(i: 0).hasOneUse()) {
4134 SDValue MinLHS = Op1.getOperand(i: 0).getOperand(i: 0);
4135 SDValue MinRHS = Op1.getOperand(i: 0).getOperand(i: 1);
4136 if (MinLHS.getOpcode() == ISD::ZERO_EXTEND && MinLHS.getOperand(i: 0) == Op0)
4137 return getTruncatedUSUBSAT(DstVT, SrcVT: MinLHS.getValueType(), LHS: MinLHS, RHS: MinRHS,
4138 DAG, DL);
4139 if (MinRHS.getOpcode() == ISD::ZERO_EXTEND && MinRHS.getOperand(i: 0) == Op0)
4140 return getTruncatedUSUBSAT(DstVT, SrcVT: MinLHS.getValueType(), LHS: MinRHS, RHS: MinLHS,
4141 DAG, DL);
4142 }
4143
4144 return SDValue();
4145}
4146
4147// Refinement of DAG/Type Legalisation (promotion) when CTLZ is used for
4148// counting leading ones. Broadly, it replaces the substraction with a left
4149// shift.
4150//
4151// * DAG Legalisation Pattern:
4152//
4153// (sub (ctlz (zeroextend (not Src)))
4154// BitWidthDiff)
4155//
4156// if BitWidthDiff == BitWidth(Node) - BitWidth(Src)
4157// -->
4158//
4159// (ctlz_zero_poison (not (shl (anyextend Src)
4160// BitWidthDiff)))
4161//
4162// * Type Legalisation Pattern:
4163//
4164// (sub (ctlz (and (xor Src XorMask)
4165// AndMask))
4166// BitWidthDiff)
4167//
4168// if AndMask has only trailing ones
4169// and MaskBitWidth(AndMask) == BitWidth(Node) - BitWidthDiff
4170// and XorMask has more trailing ones than AndMask
4171// -->
4172//
4173// (ctlz_zero_poison (not (shl Src BitWidthDiff)))
4174static SDValue foldSubCtlzNot(SDNode *N, SelectionDAG &DAG) {
4175 const SDLoc DL(N);
4176 SDValue N0 = N->getOperand(Num: 0);
4177 EVT VT = N0.getValueType();
4178 unsigned BitWidth = VT.getScalarSizeInBits();
4179
4180 APInt AndMask;
4181 APInt XorMask;
4182 uint64_t BitWidthDiff;
4183
4184 SDValue CtlzOp;
4185 SDValue Src;
4186
4187 if (!sd_match(N, P: m_Sub(L: m_Ctlz(Op: m_Value(N&: CtlzOp)), R: m_ConstInt(V&: BitWidthDiff))))
4188 return SDValue();
4189
4190 if (sd_match(N: CtlzOp, P: m_ZExt(Op: m_Not(V: m_Value(N&: Src))))) {
4191 // DAG Legalisation Pattern:
4192 // (sub (ctlz (zero_extend (not Op)) BitWidthDiff))
4193 if ((BitWidth - Src.getValueType().getScalarSizeInBits()) != BitWidthDiff)
4194 return SDValue();
4195
4196 Src = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT, Operand: Src);
4197 } else if (sd_match(N: CtlzOp, P: m_And(L: m_Xor(L: m_Value(N&: Src), R: m_ConstInt(V&: XorMask)),
4198 R: m_ConstInt(V&: AndMask)))) {
4199 // Type Legalisation Pattern:
4200 // (sub (ctlz (and (xor Op XorMask) AndMask)) BitWidthDiff)
4201 if (BitWidthDiff >= BitWidth)
4202 return SDValue();
4203 unsigned AndMaskWidth = BitWidth - BitWidthDiff;
4204 if (!(AndMask.isMask(numBits: AndMaskWidth) && XorMask.countr_one() >= AndMaskWidth))
4205 return SDValue();
4206 } else
4207 return SDValue();
4208
4209 SDValue ShiftConst = DAG.getShiftAmountConstant(Val: BitWidthDiff, VT, DL);
4210 SDValue LShift = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Src, N2: ShiftConst);
4211 SDValue Not =
4212 DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: LShift, N2: DAG.getAllOnesConstant(DL, VT));
4213
4214 return DAG.getNode(Opcode: ISD::CTLZ_ZERO_POISON, DL, VT, Operand: Not);
4215}
4216
4217// Fold sub(x, mul(divrem(x,y)[0], y)) to divrem(x, y)[1]
4218static SDValue foldRemainderIdiom(SDNode *N, SelectionDAG &DAG,
4219 const SDLoc &DL) {
4220 assert(N->getOpcode() == ISD::SUB && "Node must be a SUB");
4221 SDValue Sub0 = N->getOperand(Num: 0);
4222 SDValue Sub1 = N->getOperand(Num: 1);
4223
4224 auto CheckAndFoldMulCase = [&](SDValue DivRem, SDValue MaybeY) -> SDValue {
4225 if ((DivRem.getOpcode() == ISD::SDIVREM ||
4226 DivRem.getOpcode() == ISD::UDIVREM) &&
4227 DivRem.getResNo() == 0 && DivRem.getOperand(i: 0) == Sub0 &&
4228 DivRem.getOperand(i: 1) == MaybeY) {
4229 return SDValue(DivRem.getNode(), 1);
4230 }
4231 return SDValue();
4232 };
4233
4234 if (Sub1.getOpcode() == ISD::MUL) {
4235 // (sub x, (mul divrem(x,y)[0], y))
4236 SDValue Mul0 = Sub1.getOperand(i: 0);
4237 SDValue Mul1 = Sub1.getOperand(i: 1);
4238
4239 if (SDValue Res = CheckAndFoldMulCase(Mul0, Mul1))
4240 return Res;
4241
4242 if (SDValue Res = CheckAndFoldMulCase(Mul1, Mul0))
4243 return Res;
4244
4245 } else if (Sub1.getOpcode() == ISD::SHL) {
4246 // Handle (sub x, (shl divrem(x,y)[0], C)) where y = 1 << C
4247 SDValue Shl0 = Sub1.getOperand(i: 0);
4248 SDValue Shl1 = Sub1.getOperand(i: 1);
4249 // Check if Shl0 is divrem(x, Y)[0]
4250 if ((Shl0.getOpcode() == ISD::SDIVREM ||
4251 Shl0.getOpcode() == ISD::UDIVREM) &&
4252 Shl0.getResNo() == 0 && Shl0.getOperand(i: 0) == Sub0) {
4253
4254 SDValue Divisor = Shl0.getOperand(i: 1);
4255
4256 ConstantSDNode *DivC = isConstOrConstSplat(N: Divisor);
4257 ConstantSDNode *ShC = isConstOrConstSplat(N: Shl1);
4258 if (!DivC || !ShC)
4259 return SDValue();
4260
4261 if (DivC->getAPIntValue().isPowerOf2() &&
4262 DivC->getAPIntValue().logBase2() == ShC->getAPIntValue())
4263 return SDValue(Shl0.getNode(), 1);
4264 }
4265 }
4266 return SDValue();
4267}
4268
4269// Since it may not be valid to emit a fold to zero for vector initializers
4270// check if we can before folding.
4271static SDValue tryFoldToZero(const SDLoc &DL, const TargetLowering &TLI, EVT VT,
4272 SelectionDAG &DAG, bool LegalOperations) {
4273 if (!VT.isVector())
4274 return DAG.getConstant(Val: 0, DL, VT);
4275 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::BUILD_VECTOR, VT))
4276 return DAG.getConstant(Val: 0, DL, VT);
4277 return SDValue();
4278}
4279
4280SDValue DAGCombiner::visitSUB(SDNode *N) {
4281 SDValue N0 = N->getOperand(Num: 0);
4282 SDValue N1 = N->getOperand(Num: 1);
4283 EVT VT = N0.getValueType();
4284 unsigned BitWidth = VT.getScalarSizeInBits();
4285 SDLoc DL(N);
4286
4287 if (SDValue V = foldSubCtlzNot(N, DAG))
4288 return V;
4289
4290 // fold (sub x, x) -> 0
4291 if (N0 == N1)
4292 return tryFoldToZero(DL, TLI, VT, DAG, LegalOperations);
4293
4294 // fold (sub c1, c2) -> c3
4295 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::SUB, DL, VT, Ops: {N0, N1}))
4296 return C;
4297
4298 // fold vector ops
4299 if (VT.isVector()) {
4300 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
4301 return FoldedVOp;
4302
4303 // fold (sub x, 0) -> x, vector edition
4304 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
4305 return N0;
4306 }
4307
4308 // (sub x, ([v]select (ult x, y), 0, y)) -> (umin x, (sub x, y))
4309 // (sub x, ([v]select (uge x, y), y, 0)) -> (umin x, (sub x, y))
4310 if (N1.hasOneUse() && hasUMin(VT)) {
4311 SDValue Y;
4312 auto MS0 = m_Specific(N: N0);
4313 auto MVY = m_Value(N&: Y);
4314 auto MZ = m_Zero();
4315
4316 if (sd_match(N: N1, P: m_SpecificSelectCCLike(CC: ISD::SETULT, L: MS0, R: MVY, T: MZ,
4317 F: m_Deferred(V&: Y))) ||
4318 sd_match(N: N1, P: m_SpecificSelectCCLike(CC: ISD::SETUGE, L: MS0, R: MVY,
4319 T: m_Deferred(V&: Y), F: MZ)) ||
4320 sd_match(N: N1, P: m_VSelect(Cond: m_SpecificSetCC(CC: ISD::SETULT, LHS: MS0, RHS: MVY), T: MZ,
4321 F: m_Deferred(V&: Y))) ||
4322 sd_match(N: N1, P: m_VSelect(Cond: m_SpecificSetCC(CC: ISD::SETUGE, LHS: MS0, RHS: MVY),
4323 T: m_Deferred(V&: Y), F: MZ)))
4324
4325 return DAG.getNode(Opcode: ISD::UMIN, DL, VT, N1: N0,
4326 N2: DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: Y));
4327 }
4328
4329 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
4330 return NewSel;
4331
4332 // fold (sub x, c) -> (add x, -c)
4333 if (ConstantSDNode *N1C = getAsNonOpaqueConstant(N: N1))
4334 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0,
4335 N2: DAG.getConstant(Val: -N1C->getAPIntValue(), DL, VT));
4336
4337 if (isNullOrNullSplat(V: N0)) {
4338 // Right-shifting everything out but the sign bit followed by negation is
4339 // the same as flipping arithmetic/logical shift type without the negation:
4340 // -(X >>u 31) -> (X >>s 31)
4341 // -(X >>s 31) -> (X >>u 31)
4342 if (N1->getOpcode() == ISD::SRA || N1->getOpcode() == ISD::SRL) {
4343 ConstantSDNode *ShiftAmt = isConstOrConstSplat(N: N1.getOperand(i: 1));
4344 if (ShiftAmt && ShiftAmt->getAPIntValue() == (BitWidth - 1)) {
4345 auto NewSh = N1->getOpcode() == ISD::SRA ? ISD::SRL : ISD::SRA;
4346 if (!LegalOperations || TLI.isOperationLegal(Op: NewSh, VT))
4347 return DAG.getNode(Opcode: NewSh, DL, VT, N1: N1.getOperand(i: 0), N2: N1.getOperand(i: 1));
4348 }
4349 }
4350
4351 // 0 - X --> 0 if the sub is NUW.
4352 if (N->getFlags().hasNoUnsignedWrap())
4353 return N0;
4354
4355 if (DAG.MaskedValueIsZero(Op: N1, Mask: ~APInt::getSignMask(BitWidth))) {
4356 // N1 is either 0 or the minimum signed value. If the sub is NSW, then
4357 // N1 must be 0 because negating the minimum signed value is undefined.
4358 if (N->getFlags().hasNoSignedWrap())
4359 return N0;
4360
4361 // 0 - X --> X if X is 0 or the minimum signed value.
4362 return N1;
4363 }
4364
4365 // Convert 0 - abs(x).
4366 if (ISD::isAbsOpcode(Opcode: N1.getOpcode()) && N1.hasOneUse() &&
4367 !TLI.isOperationLegalOrCustom(Op: N1.getOpcode(), VT))
4368 if (SDValue Result = TLI.expandABS(N: N1.getNode(), DAG, IsNegative: true))
4369 return Result;
4370
4371 // Similar to the previous rule, but this time targeting an expanded abs.
4372 // (sub 0, (max X, (sub 0, X))) --> (min X, (sub 0, X))
4373 // as well as
4374 // (sub 0, (min X, (sub 0, X))) --> (max X, (sub 0, X))
4375 // Note that these two are applicable to both signed and unsigned min/max.
4376 SDValue X;
4377 SDValue S0;
4378 auto NegPat = m_Value(N&: S0, P: m_Neg(V: m_Deferred(V&: X)));
4379 if (sd_match(N: N1, P: m_OneUse(P: m_AnyOf(preds: m_SMax(L: m_Value(N&: X), R: NegPat),
4380 preds: m_UMax(L: m_Value(N&: X), R: NegPat),
4381 preds: m_SMin(L: m_Value(N&: X), R: NegPat),
4382 preds: m_UMin(L: m_Value(N&: X), R: NegPat))))) {
4383 unsigned NewOpc = ISD::getInverseMinMaxOpcode(MinMaxOpc: N1->getOpcode());
4384 if (hasOperation(Opcode: NewOpc, VT))
4385 return DAG.getNode(Opcode: NewOpc, DL, VT, N1: X, N2: S0);
4386 }
4387
4388 // Fold neg(splat(neg(x)) -> splat(x)
4389 if (VT.isVector()) {
4390 SDValue N1S = DAG.getSplatValue(V: N1, LegalTypes: true);
4391 if (N1S && N1S.getOpcode() == ISD::SUB &&
4392 isNullConstant(V: N1S.getOperand(i: 0)))
4393 return DAG.getSplat(VT, DL, Op: N1S.getOperand(i: 1));
4394 }
4395
4396 // sub 0, (and x, 1) --> SIGN_EXTEND_INREG x, i1
4397 if (N1.getOpcode() == ISD::AND && N1.hasOneUse() &&
4398 isOneOrOneSplat(V: N1->getOperand(Num: 1))) {
4399 EVT ExtVT = VT.changeElementType(Context&: *DAG.getContext(), EltVT: MVT::i1);
4400 if (TLI.getOperationAction(Op: ISD::SIGN_EXTEND_INREG, VT: ExtVT) ==
4401 TargetLowering::Legal) {
4402 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT, N1: N1->getOperand(Num: 0),
4403 N2: DAG.getValueType(ExtVT));
4404 }
4405 }
4406 }
4407
4408 // Canonicalize (sub -1, x) -> ~x, i.e. (xor x, -1)
4409 if (isAllOnesOrAllOnesSplat(V: N0))
4410 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1, N2: N0);
4411
4412 // fold (A - (0-B)) -> A+B
4413 if (N1.getOpcode() == ISD::SUB && isNullOrNullSplat(V: N1.getOperand(i: 0)))
4414 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: N1.getOperand(i: 1));
4415
4416 // fold A-(A-B) -> B
4417 if (N1.getOpcode() == ISD::SUB && N0 == N1.getOperand(i: 0))
4418 return N1.getOperand(i: 1);
4419
4420 // fold (A+B)-A -> B
4421 if (N0.getOpcode() == ISD::ADD && N0.getOperand(i: 0) == N1)
4422 return N0.getOperand(i: 1);
4423
4424 // fold (A+B)-B -> A
4425 if (N0.getOpcode() == ISD::ADD && N0.getOperand(i: 1) == N1)
4426 return N0.getOperand(i: 0);
4427
4428 // fold (A+C1)-C2 -> A+(C1-C2)
4429 if (N0.getOpcode() == ISD::ADD) {
4430 SDValue N01 = N0.getOperand(i: 1);
4431 if (SDValue NewC = DAG.FoldConstantArithmetic(Opcode: ISD::SUB, DL, VT, Ops: {N01, N1}))
4432 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 0), N2: NewC);
4433 }
4434
4435 // fold C2-(A+C1) -> (C2-C1)-A
4436 if (N1.getOpcode() == ISD::ADD) {
4437 SDValue N11 = N1.getOperand(i: 1);
4438 if (SDValue NewC = DAG.FoldConstantArithmetic(Opcode: ISD::SUB, DL, VT, Ops: {N0, N11}))
4439 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: NewC, N2: N1.getOperand(i: 0));
4440 }
4441
4442 // fold (A-C1)-C2 -> A-(C1+C2)
4443 if (N0.getOpcode() == ISD::SUB) {
4444 SDValue N01 = N0.getOperand(i: 1);
4445 if (SDValue NewC = DAG.FoldConstantArithmetic(Opcode: ISD::ADD, DL, VT, Ops: {N01, N1}))
4446 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0.getOperand(i: 0), N2: NewC);
4447 }
4448
4449 // fold (c1-A)-c2 -> (c1-c2)-A
4450 if (N0.getOpcode() == ISD::SUB) {
4451 SDValue N00 = N0.getOperand(i: 0);
4452 if (SDValue NewC = DAG.FoldConstantArithmetic(Opcode: ISD::SUB, DL, VT, Ops: {N00, N1}))
4453 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: NewC, N2: N0.getOperand(i: 1));
4454 }
4455
4456 SDValue A, B, C;
4457
4458 // fold ((A+(B+C))-B) -> A+C
4459 if (sd_match(N: N0, P: m_Add(L: m_Value(N&: A), R: m_Add(L: m_Specific(N: N1), R: m_Value(N&: C)))))
4460 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: A, N2: C);
4461
4462 // fold ((A+(B-C))-B) -> A-C
4463 if (sd_match(N: N0, P: m_Add(L: m_Value(N&: A), R: m_Sub(L: m_Specific(N: N1), R: m_Value(N&: C)))))
4464 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: A, N2: C);
4465
4466 // fold ((A-(B-C))-C) -> A-B
4467 if (sd_match(N: N0, P: m_Sub(L: m_Value(N&: A), R: m_Sub(L: m_Value(N&: B), R: m_Specific(N: N1)))))
4468 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: A, N2: B);
4469
4470 // fold (A-(B-C)) -> A+(C-B)
4471 if (sd_match(N: N1, P: m_OneUse(P: m_Sub(L: m_Value(N&: B), R: m_Value(N&: C)))))
4472 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0,
4473 N2: DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: C, N2: B));
4474
4475 // A - (A & B) -> A & (~B)
4476 if (sd_match(N: N1, P: m_And(L: m_Specific(N: N0), R: m_Value(N&: B))) &&
4477 (N1.hasOneUse() || isConstantOrConstantVector(N: B, /*NoOpaques=*/true)))
4478 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N0, N2: DAG.getNOT(DL, Val: B, VT));
4479
4480 // fold (A - (-B * C)) -> (A + (B * C))
4481 if (sd_match(N: N1, P: m_OneUse(P: m_Mul(L: m_Neg(V: m_Value(N&: B)), R: m_Value(N&: C)))))
4482 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0,
4483 N2: DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: B, N2: C));
4484
4485 // If either operand of a sub is undef, the result is undef
4486 if (N0.isUndef())
4487 return N0;
4488 if (N1.isUndef())
4489 return N1;
4490
4491 if (SDValue V = foldAddSubBoolOfMaskedVal(N, DL, DAG))
4492 return V;
4493
4494 if (SDValue V = foldAddSubOfSignBit(N, DL, DAG))
4495 return V;
4496
4497 // Try to match AVGCEIL fixedwidth pattern
4498 if (SDValue V = foldSubToAvg(N, DL))
4499 return V;
4500
4501 if (SDValue V = foldAddSubMasked1(IsAdd: false, N0, N1, DAG, DL))
4502 return V;
4503
4504 if (SDValue V = foldSubToUSubSat(DstVT: VT, N, DL))
4505 return V;
4506
4507 if (SDValue V = foldRemainderIdiom(N, DAG, DL))
4508 return V;
4509
4510 // (A - B) - 1 -> add (xor B, -1), A
4511 if (sd_match(N, P: m_Sub(L: m_OneUse(P: m_Sub(L: m_Value(N&: A), R: m_Value(N&: B))),
4512 R: m_One(/*AllowUndefs=*/true))))
4513 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: A, N2: DAG.getNOT(DL, Val: B, VT));
4514
4515 // Look for:
4516 // sub y, (xor x, -1)
4517 // And if the target does not like this form then turn into:
4518 // add (add x, y), 1
4519 if (TLI.preferIncOfAddToSubOfNot(VT) && N1.hasOneUse() && isBitwiseNot(V: N1)) {
4520 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: N1.getOperand(i: 0));
4521 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Add, N2: DAG.getConstant(Val: 1, DL, VT));
4522 }
4523
4524 // Hoist one-use addition by non-opaque constant:
4525 // (x + C) - y -> (x - y) + C
4526 if (!reassociationCanBreakAddressingModePattern(Opc: ISD::SUB, DL, N, N0, N1) &&
4527 N0.getOpcode() == ISD::ADD && N0.hasOneUse() &&
4528 isConstantOrConstantVector(N: N0.getOperand(i: 1), /*NoOpaques=*/true)) {
4529 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0.getOperand(i: 0), N2: N1);
4530 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Sub, N2: N0.getOperand(i: 1));
4531 }
4532 // y - (x + C) -> (y - x) - C
4533 if (N1.getOpcode() == ISD::ADD && N1.hasOneUse() &&
4534 isConstantOrConstantVector(N: N1.getOperand(i: 1), /*NoOpaques=*/true)) {
4535 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: N1.getOperand(i: 0));
4536 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Sub, N2: N1.getOperand(i: 1));
4537 }
4538 // (x - C) - y -> (x - y) - C
4539 // This is necessary because SUB(X,C) -> ADD(X,-C) doesn't work for vectors.
4540 if (N0.getOpcode() == ISD::SUB && N0.hasOneUse() &&
4541 isConstantOrConstantVector(N: N0.getOperand(i: 1), /*NoOpaques=*/true)) {
4542 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0.getOperand(i: 0), N2: N1);
4543 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Sub, N2: N0.getOperand(i: 1));
4544 }
4545 // (C - x) - y -> C - (x + y)
4546 if (N0.getOpcode() == ISD::SUB && N0.hasOneUse() &&
4547 isConstantOrConstantVector(N: N0.getOperand(i: 0), /*NoOpaques=*/true)) {
4548 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0.getOperand(i: 1), N2: N1);
4549 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0.getOperand(i: 0), N2: Add);
4550 }
4551
4552 // If the target's bool is represented as 0/-1, prefer to make this 'add 0/-1'
4553 // rather than 'sub 0/1' (the sext should get folded).
4554 // sub X, (zext i1 Y) --> add X, (sext i1 Y)
4555 if (N1.getOpcode() == ISD::ZERO_EXTEND &&
4556 N1.getOperand(i: 0).getScalarValueSizeInBits() == 1 &&
4557 TLI.getBooleanContents(Type: VT) ==
4558 TargetLowering::ZeroOrNegativeOneBooleanContent) {
4559 SDValue SExt = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: N1.getOperand(i: 0));
4560 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: SExt);
4561 }
4562
4563 // fold B = sra (A, size(A)-1); sub (xor (A, B), B) -> (abs A)
4564 if ((!LegalOperations || hasOperation(Opcode: ISD::ABS, VT)) &&
4565 sd_match(N: N1, P: m_Sra(L: m_Value(N&: A), R: m_SpecificInt(V: BitWidth - 1))) &&
4566 sd_match(N: N0, P: m_Xor(L: m_Specific(N: A), R: m_Specific(N: N1))))
4567 return DAG.getNode(Opcode: ISD::ABS, DL, VT, Operand: A);
4568
4569 // If the relocation model supports it, consider symbol offsets.
4570 if (GlobalAddressSDNode *GA = dyn_cast<GlobalAddressSDNode>(Val&: N0))
4571 if (!LegalOperations && TLI.isOffsetFoldingLegal(GA)) {
4572 // fold (sub Sym+c1, Sym+c2) -> c1-c2
4573 if (GlobalAddressSDNode *GB = dyn_cast<GlobalAddressSDNode>(Val&: N1))
4574 if (GA->getGlobal() == GB->getGlobal())
4575 return DAG.getConstant(
4576 Val: APInt(VT.getScalarSizeInBits(), GA->getOffset() - GB->getOffset(),
4577 /*isSigned=*/false, /*implicitTrunc=*/true),
4578 DL, VT);
4579 }
4580
4581 // sub X, (sextinreg Y i1) -> add X, (and Y 1)
4582 if (N1.getOpcode() == ISD::SIGN_EXTEND_INREG) {
4583 VTSDNode *TN = cast<VTSDNode>(Val: N1.getOperand(i: 1));
4584 if (TN->getVT() == MVT::i1) {
4585 SDValue ZExt = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N1.getOperand(i: 0),
4586 N2: DAG.getConstant(Val: 1, DL, VT));
4587 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: ZExt);
4588 }
4589 }
4590
4591 // canonicalize (sub X, (vscale * C)) to (add X, (vscale * -C)) if this is the
4592 // only use of the vscale value or if (vscale * -C) is a valid add immediate.
4593 // avoid if ISD::MUL handling is poor and ISD::SHL isn't an option.
4594 if (N1.getOpcode() == ISD::VSCALE) {
4595 const APInt &IntVal = N1.getConstantOperandAPInt(i: 0);
4596 if ((N1.hasOneUse() ||
4597 TLI.isLegalAddScalableImmediate(-IntVal.getSExtValue())) &&
4598 (!IntVal.isPowerOf2() ||
4599 hasOperation(Opcode: ISD::MUL, VT: N1.getOperand(i: 0).getValueType())))
4600 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: DAG.getVScale(DL, VT, MulImm: -IntVal));
4601 }
4602
4603 // canonicalize (sub X, step_vector(C)) to (add X, step_vector(-C))
4604 if (N1.getOpcode() == ISD::STEP_VECTOR && N1.hasOneUse()) {
4605 APInt NewStep = -N1.getConstantOperandAPInt(i: 0);
4606 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0,
4607 N2: DAG.getStepVector(DL, ResVT: VT, StepVal: NewStep));
4608 }
4609
4610 // Prefer an add for more folding potential and possibly better codegen:
4611 // sub N0, (lshr N10, width-1) --> add N0, (ashr N10, width-1)
4612 if (!LegalOperations && N1.getOpcode() == ISD::SRL && N1.hasOneUse()) {
4613 SDValue ShAmt = N1.getOperand(i: 1);
4614 ConstantSDNode *ShAmtC = isConstOrConstSplat(N: ShAmt);
4615 if (ShAmtC && ShAmtC->getAPIntValue() == (BitWidth - 1)) {
4616 SDValue SRA = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N1.getOperand(i: 0), N2: ShAmt);
4617 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: SRA);
4618 }
4619 }
4620
4621 // As with the previous fold, prefer add for more folding potential.
4622 // Subtracting SMIN/0 is the same as adding SMIN/0:
4623 // N0 - (X << BW-1) --> N0 + (X << BW-1)
4624 if (N1.getOpcode() == ISD::SHL) {
4625 ConstantSDNode *ShlC = isConstOrConstSplat(N: N1.getOperand(i: 1));
4626 if (ShlC && ShlC->getAPIntValue() == (BitWidth - 1))
4627 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1, N2: N0);
4628 }
4629
4630 // (sub (usubo_carry X, 0, Carry), Y) -> (usubo_carry X, Y, Carry)
4631 if (N0.getOpcode() == ISD::USUBO_CARRY && isNullConstant(V: N0.getOperand(i: 1)) &&
4632 N0.getResNo() == 0 && N0.hasOneUse())
4633 return DAG.getNode(Opcode: ISD::USUBO_CARRY, DL, VTList: N0->getVTList(),
4634 N1: N0.getOperand(i: 0), N2: N1, N3: N0.getOperand(i: 2));
4635
4636 if (TLI.isOperationLegalOrCustom(Op: ISD::UADDO_CARRY, VT)) {
4637 // (sub Carry, X) -> (uaddo_carry (sub 0, X), 0, Carry)
4638 if (SDValue Carry = getAsCarry(TLI, V: N0)) {
4639 SDValue X = N1;
4640 SDValue Zero = DAG.getConstant(Val: 0, DL, VT);
4641 SDValue NegX = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Zero, N2: X);
4642 return DAG.getNode(Opcode: ISD::UADDO_CARRY, DL,
4643 VTList: DAG.getVTList(VT1: VT, VT2: Carry.getValueType()), N1: NegX, N2: Zero,
4644 N3: Carry);
4645 }
4646 }
4647
4648 if (ConstantSDNode *C0 = isConstOrConstSplat(N: N0)) {
4649 const APInt &C0Val = C0->getAPIntValue();
4650
4651 // sub nuw C, x --> xor x, C when C is a mask (2^k - 1)
4652 if (N->getFlags().hasNoUnsignedWrap() && C0Val.isMask())
4653 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1, N2: N0);
4654
4655 // If there's no chance of borrowing from adjacent bits, then sub is xor:
4656 // sub C0, X --> xor X, C0
4657 if (!C0->isOpaque()) {
4658 const APInt &MaybeOnes = ~DAG.computeKnownBits(Op: N1).Zero;
4659 if ((C0Val - MaybeOnes) == (C0Val ^ MaybeOnes))
4660 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1, N2: N0);
4661 }
4662 }
4663
4664 // smax(a,b) - smin(a,b) --> abds(a,b)
4665 if ((!LegalOperations || hasOperation(Opcode: ISD::ABDS, VT)) &&
4666 sd_match(N: N0, P: m_SMaxLike(L: m_Value(N&: A), R: m_Value(N&: B))) &&
4667 sd_match(N: N1, P: m_SMinLike(L: m_Specific(N: A), R: m_Specific(N: B))))
4668 return DAG.getNode(Opcode: ISD::ABDS, DL, VT, N1: A, N2: B);
4669
4670 // smin(a,b) - smax(a,b) --> neg(abds(a,b))
4671 if ((!LegalOperations || hasOperation(Opcode: ISD::ABDS, VT)) &&
4672 sd_match(N: N0, P: m_SMinLike(L: m_Value(N&: A), R: m_Value(N&: B))) &&
4673 sd_match(N: N1, P: m_SMaxLike(L: m_Specific(N: A), R: m_Specific(N: B))))
4674 return DAG.getNegative(Val: DAG.getNode(Opcode: ISD::ABDS, DL, VT, N1: A, N2: B), DL, VT);
4675
4676 // umax(a,b) - umin(a,b) --> abdu(a,b)
4677 if ((!LegalOperations || hasOperation(Opcode: ISD::ABDU, VT)) &&
4678 sd_match(N: N0, P: m_UMaxLike(L: m_Value(N&: A), R: m_Value(N&: B))) &&
4679 sd_match(N: N1, P: m_UMinLike(L: m_Specific(N: A), R: m_Specific(N: B))))
4680 return DAG.getNode(Opcode: ISD::ABDU, DL, VT, N1: A, N2: B);
4681
4682 // umin(a,b) - umax(a,b) --> neg(abdu(a,b))
4683 if ((!LegalOperations || hasOperation(Opcode: ISD::ABDU, VT)) &&
4684 sd_match(N: N0, P: m_UMinLike(L: m_Value(N&: A), R: m_Value(N&: B))) &&
4685 sd_match(N: N1, P: m_UMaxLike(L: m_Specific(N: A), R: m_Specific(N: B))))
4686 return DAG.getNegative(Val: DAG.getNode(Opcode: ISD::ABDU, DL, VT, N1: A, N2: B), DL, VT);
4687
4688 return SDValue();
4689}
4690
4691SDValue DAGCombiner::visitSUBSAT(SDNode *N) {
4692 unsigned Opcode = N->getOpcode();
4693 SDValue N0 = N->getOperand(Num: 0);
4694 SDValue N1 = N->getOperand(Num: 1);
4695 EVT VT = N0.getValueType();
4696 bool IsSigned = Opcode == ISD::SSUBSAT;
4697 SDLoc DL(N);
4698
4699 // fold (sub_sat x, undef) -> 0
4700 if (N0.isUndef() || N1.isUndef())
4701 return DAG.getConstant(Val: 0, DL, VT);
4702
4703 // fold (sub_sat x, x) -> 0
4704 if (N0 == N1)
4705 return DAG.getConstant(Val: 0, DL, VT);
4706
4707 // fold (sub_sat c1, c2) -> c3
4708 if (SDValue C = DAG.FoldConstantArithmetic(Opcode, DL, VT, Ops: {N0, N1}))
4709 return C;
4710
4711 // fold vector ops
4712 if (VT.isVector()) {
4713 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
4714 return FoldedVOp;
4715
4716 // fold (sub_sat x, 0) -> x, vector edition
4717 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
4718 return N0;
4719 }
4720
4721 // fold (sub_sat x, 0) -> x
4722 if (isNullConstant(V: N1))
4723 return N0;
4724
4725 // If it cannot overflow, transform into an sub.
4726 if (DAG.willNotOverflowSub(IsSigned, N0, N1))
4727 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: N1);
4728
4729 // Narrow a vXiN USUBSAT to a smaller type when both operands are known
4730 // to fit in fewer bits. This allows targets with native narrow USUBSAT
4731 // (e.g. vpsubusb/vpsubusw) to avoid emulation with vpmaxu + vsub.
4732 if (!IsSigned && VT.isVector() && VT.isSimple()) {
4733 unsigned ScalarBits = VT.getScalarSizeInBits();
4734 if (ScalarBits > 8 && isPowerOf2_32(Value: ScalarBits) &&
4735 !TLI.isOperationLegal(Op: ISD::USUBSAT, VT)) {
4736 KnownBits Known0 = DAG.computeKnownBits(Op: N0);
4737 unsigned ActiveBits = Known0.countMaxActiveBits();
4738 for (unsigned NarrowBits = PowerOf2Ceil(A: ActiveBits);
4739 NarrowBits != 0 && NarrowBits < ScalarBits; NarrowBits *= 2) {
4740 unsigned Scale = ScalarBits / NarrowBits;
4741 ElementCount ScaledEC = VT.getVectorElementCount() * Scale;
4742 MVT NarrowSVT = MVT::getIntegerVT(BitWidth: NarrowBits);
4743 EVT NarrowVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: NarrowSVT, EC: ScaledEC);
4744
4745 if (!TLI.isOperationLegalOrCustom(Op: ISD::USUBSAT, VT: NarrowVT))
4746 continue;
4747 KnownBits Known1 = DAG.computeKnownBits(Op: N1);
4748 if (Known1.countMaxActiveBits() <= NarrowBits) {
4749 SDValue NarrowN0 = DAG.getBitcast(VT: NarrowVT, V: N0);
4750 SDValue NarrowN1 = DAG.getBitcast(VT: NarrowVT, V: N1);
4751 SDValue NarrowSub =
4752 DAG.getNode(Opcode: ISD::USUBSAT, DL, VT: NarrowVT, N1: NarrowN0, N2: NarrowN1);
4753 return DAG.getBitcast(VT, V: NarrowSub);
4754 }
4755 // TODO: If N1 doesn't fit in NarrowBits, we could OR the upper bits
4756 // of N1 with 1s to force saturation in those lanes, allowing the
4757 // narrow USUBSAT to still be used. This requires a TLI hook to check
4758 // whether the constant can be folded as a broadcast memory operand
4759 // (profitable on AVX512, not on SSE/AVX), to avoid introducing an
4760 // extra register and instruction on non-AVX512 targets.
4761 break;
4762 }
4763 }
4764 }
4765 return SDValue();
4766}
4767
4768SDValue DAGCombiner::visitSUBC(SDNode *N) {
4769 SDValue N0 = N->getOperand(Num: 0);
4770 SDValue N1 = N->getOperand(Num: 1);
4771 EVT VT = N0.getValueType();
4772 SDLoc DL(N);
4773
4774 // If the flag result is dead, turn this into an SUB.
4775 if (!N->hasAnyUseOfValue(Value: 1))
4776 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: N1),
4777 Res1: DAG.getNode(Opcode: ISD::CARRY_FALSE, DL, VT: MVT::Glue));
4778
4779 // fold (subc x, x) -> 0 + no borrow
4780 if (N0 == N1)
4781 return CombineTo(N, Res0: DAG.getConstant(Val: 0, DL, VT),
4782 Res1: DAG.getNode(Opcode: ISD::CARRY_FALSE, DL, VT: MVT::Glue));
4783
4784 // fold (subc x, 0) -> x + no borrow
4785 if (isNullConstant(V: N1))
4786 return CombineTo(N, Res0: N0, Res1: DAG.getNode(Opcode: ISD::CARRY_FALSE, DL, VT: MVT::Glue));
4787
4788 // Canonicalize (sub -1, x) -> ~x, i.e. (xor x, -1) + no borrow
4789 if (isAllOnesConstant(V: N0))
4790 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::XOR, DL, VT, N1, N2: N0),
4791 Res1: DAG.getNode(Opcode: ISD::CARRY_FALSE, DL, VT: MVT::Glue));
4792
4793 return SDValue();
4794}
4795
4796SDValue DAGCombiner::visitSUBO(SDNode *N) {
4797 SDValue N0 = N->getOperand(Num: 0);
4798 SDValue N1 = N->getOperand(Num: 1);
4799 EVT VT = N0.getValueType();
4800 bool IsSigned = (ISD::SSUBO == N->getOpcode());
4801
4802 EVT CarryVT = N->getValueType(ResNo: 1);
4803 SDLoc DL(N);
4804
4805 // If the flag result is dead, turn this into an SUB.
4806 if (!N->hasAnyUseOfValue(Value: 1))
4807 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: N1),
4808 Res1: DAG.getUNDEF(VT: CarryVT));
4809
4810 // fold (subo x, x) -> 0 + no borrow
4811 if (N0 == N1)
4812 return CombineTo(N, Res0: DAG.getConstant(Val: 0, DL, VT),
4813 Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
4814
4815 // fold (subox, c) -> (addo x, -c)
4816 if (ConstantSDNode *N1C = getAsNonOpaqueConstant(N: N1))
4817 if (IsSigned && !N1C->isMinSignedValue())
4818 return DAG.getNode(Opcode: ISD::SADDO, DL, VTList: N->getVTList(), N1: N0,
4819 N2: DAG.getConstant(Val: -N1C->getAPIntValue(), DL, VT));
4820
4821 // fold (subo x, 0) -> x + no borrow
4822 if (isNullOrNullSplat(V: N1))
4823 return CombineTo(N, Res0: N0, Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
4824
4825 // If it cannot overflow, transform into an sub.
4826 if (DAG.willNotOverflowSub(IsSigned, N0, N1))
4827 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: N1),
4828 Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
4829
4830 // Canonicalize (usubo -1, x) -> ~x, i.e. (xor x, -1) + no borrow
4831 if (!IsSigned && isAllOnesOrAllOnesSplat(V: N0))
4832 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::XOR, DL, VT, N1, N2: N0),
4833 Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
4834
4835 return SDValue();
4836}
4837
4838SDValue DAGCombiner::visitSUBE(SDNode *N) {
4839 SDValue N0 = N->getOperand(Num: 0);
4840 SDValue N1 = N->getOperand(Num: 1);
4841 SDValue CarryIn = N->getOperand(Num: 2);
4842
4843 // fold (sube x, y, false) -> (subc x, y)
4844 if (CarryIn.getOpcode() == ISD::CARRY_FALSE)
4845 return DAG.getNode(Opcode: ISD::SUBC, DL: SDLoc(N), VTList: N->getVTList(), N1: N0, N2: N1);
4846
4847 return SDValue();
4848}
4849
4850SDValue DAGCombiner::visitUSUBO_CARRY(SDNode *N) {
4851 SDValue N0 = N->getOperand(Num: 0);
4852 SDValue N1 = N->getOperand(Num: 1);
4853 SDValue CarryIn = N->getOperand(Num: 2);
4854
4855 // fold (usubo_carry x, y, false) -> (usubo x, y)
4856 if (isNullConstant(V: CarryIn)) {
4857 if (!LegalOperations ||
4858 TLI.isOperationLegalOrCustom(Op: ISD::USUBO, VT: N->getValueType(ResNo: 0)))
4859 return DAG.getNode(Opcode: ISD::USUBO, DL: SDLoc(N), VTList: N->getVTList(), N1: N0, N2: N1);
4860 }
4861
4862 // Iff the flag result is dead:
4863 // (usubo_carry (sub X, Y), 0, Carry) -> (usubo_carry X, Y, Carry)
4864 if (N0.getOpcode() == ISD::SUB && isNullConstant(V: N1) &&
4865 !N->hasAnyUseOfValue(Value: 1))
4866 return DAG.getNode(Opcode: ISD::USUBO_CARRY, DL: SDLoc(N), VTList: N->getVTList(),
4867 N1: N0.getOperand(i: 0), N2: N0.getOperand(i: 1), N3: CarryIn);
4868
4869 return SDValue();
4870}
4871
4872SDValue DAGCombiner::visitSSUBO_CARRY(SDNode *N) {
4873 SDValue N0 = N->getOperand(Num: 0);
4874 SDValue N1 = N->getOperand(Num: 1);
4875 SDValue CarryIn = N->getOperand(Num: 2);
4876
4877 // fold (ssubo_carry x, y, false) -> (ssubo x, y)
4878 if (isNullConstant(V: CarryIn)) {
4879 if (!LegalOperations ||
4880 TLI.isOperationLegalOrCustom(Op: ISD::SSUBO, VT: N->getValueType(ResNo: 0)))
4881 return DAG.getNode(Opcode: ISD::SSUBO, DL: SDLoc(N), VTList: N->getVTList(), N1: N0, N2: N1);
4882 }
4883
4884 return SDValue();
4885}
4886
4887// Notice that "mulfix" can be any of SMULFIX, SMULFIXSAT, UMULFIX and
4888// UMULFIXSAT here.
4889SDValue DAGCombiner::visitMULFIX(SDNode *N) {
4890 SDValue N0 = N->getOperand(Num: 0);
4891 SDValue N1 = N->getOperand(Num: 1);
4892 SDValue Scale = N->getOperand(Num: 2);
4893 EVT VT = N0.getValueType();
4894
4895 // fold (mulfix x, undef, scale) -> 0
4896 if (N0.isUndef() || N1.isUndef())
4897 return DAG.getConstant(Val: 0, DL: SDLoc(N), VT);
4898
4899 // Canonicalize constant to RHS (vector doesn't have to splat)
4900 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
4901 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
4902 return DAG.getNode(Opcode: N->getOpcode(), DL: SDLoc(N), VT, N1, N2: N0, N3: Scale);
4903
4904 // fold (mulfix x, 0, scale) -> 0
4905 if (isNullConstant(V: N1))
4906 return DAG.getConstant(Val: 0, DL: SDLoc(N), VT);
4907
4908 return SDValue();
4909}
4910
4911SDValue DAGCombiner::visitMUL(SDNode *N) {
4912 SDValue N0 = N->getOperand(Num: 0);
4913 SDValue N1 = N->getOperand(Num: 1);
4914 EVT VT = N0.getValueType();
4915 unsigned BitWidth = VT.getScalarSizeInBits();
4916 SDLoc DL(N);
4917
4918 // fold (mul x, undef) -> 0
4919 if (N0.isUndef() || N1.isUndef())
4920 return DAG.getConstant(Val: 0, DL, VT);
4921
4922 // fold (mul c1, c2) -> c1*c2
4923 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::MUL, DL, VT, Ops: {N0, N1}))
4924 return C;
4925
4926 // canonicalize constant to RHS (vector doesn't have to splat). An opaque
4927 // constant on the RHS is treated as non-constant so that a foldable constant
4928 // still ends up on the RHS.
4929 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0, /*AllowOpaques=*/false) &&
4930 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1, /*AllowOpaques=*/false))
4931 return DAG.getNode(Opcode: ISD::MUL, DL, VT, N1, N2: N0);
4932
4933 bool N1IsConst = false;
4934 bool N1IsOpaqueConst = false;
4935 APInt ConstValue1;
4936
4937 // fold vector ops
4938 if (VT.isVector()) {
4939 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
4940 return FoldedVOp;
4941
4942 N1IsConst = ISD::isConstantSplatVector(N: N1.getNode(), SplatValue&: ConstValue1);
4943 assert((!N1IsConst || ConstValue1.getBitWidth() == BitWidth) &&
4944 "Splat APInt should be element width");
4945 } else {
4946 N1IsConst = isa<ConstantSDNode>(Val: N1);
4947 if (N1IsConst) {
4948 ConstValue1 = N1->getAsAPIntVal();
4949 N1IsOpaqueConst = cast<ConstantSDNode>(Val&: N1)->isOpaque();
4950 }
4951 }
4952
4953 // fold (mul x, 0) -> 0
4954 if (N1IsConst && ConstValue1.isZero())
4955 return N1;
4956
4957 // fold (mul x, 1) -> x
4958 if (N1IsConst && ConstValue1.isOne())
4959 return N0;
4960
4961 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
4962 return NewSel;
4963
4964 // fold (mul x, -1) -> 0-x
4965 if (N1IsConst && ConstValue1.isAllOnes())
4966 return DAG.getNegative(Val: N0, DL, VT);
4967
4968 // fold (mul x, (1 << c)) -> x << c
4969 if (isConstantOrConstantVector(N: N1, /*NoOpaques*/ true) &&
4970 (!VT.isVector() || Level <= AfterLegalizeVectorOps)) {
4971 if (SDValue LogBase2 = BuildLogBase2(V: N1, DL)) {
4972 EVT ShiftVT = getShiftAmountTy(LHSTy: N0.getValueType());
4973 SDValue Trunc = DAG.getZExtOrTrunc(Op: LogBase2, DL, VT: ShiftVT);
4974 SDNodeFlags Flags;
4975 Flags.setNoUnsignedWrap(N->getFlags().hasNoUnsignedWrap());
4976 // Preserve nsw when the shift amount is strictly less than BitWidth - 1,
4977 // i.e. the multiplier is not the signed minimum value.
4978 if (N->getFlags().hasNoSignedWrap() && N1IsConst &&
4979 ConstValue1.logBase2() < BitWidth - 1)
4980 Flags.setNoSignedWrap(true);
4981 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0, N2: Trunc, Flags);
4982 }
4983 }
4984
4985 // fold (mul x, -(1 << c)) -> -(x << c) or (-x) << c
4986 if (N1IsConst && !N1IsOpaqueConst && ConstValue1.isNegatedPowerOf2()) {
4987 unsigned Log2Val = (-ConstValue1).logBase2();
4988
4989 // FIXME: If the input is something that is easily negated (e.g. a
4990 // single-use add), we should put the negate there.
4991 return DAG.getNode(
4992 Opcode: ISD::SUB, DL, VT, N1: DAG.getConstant(Val: 0, DL, VT),
4993 N2: DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0,
4994 N2: DAG.getShiftAmountConstant(Val: Log2Val, VT, DL)));
4995 }
4996
4997 // Attempt to reuse an existing umul_lohi/smul_lohi node, but only if the
4998 // hi result is in use in case we hit this mid-legalization.
4999 for (unsigned LoHiOpc : {ISD::UMUL_LOHI, ISD::SMUL_LOHI}) {
5000 if (!LegalOperations || TLI.isOperationLegalOrCustom(Op: LoHiOpc, VT)) {
5001 SDVTList LoHiVT = DAG.getVTList(VT1: VT, VT2: VT);
5002 // TODO: Can we match commutable operands with getNodeIfExists?
5003 if (SDNode *LoHi = DAG.getNodeIfExists(Opcode: LoHiOpc, VTList: LoHiVT, Ops: {N0, N1}))
5004 if (LoHi->hasAnyUseOfValue(Value: 1))
5005 return SDValue(LoHi, 0);
5006 if (SDNode *LoHi = DAG.getNodeIfExists(Opcode: LoHiOpc, VTList: LoHiVT, Ops: {N1, N0}))
5007 if (LoHi->hasAnyUseOfValue(Value: 1))
5008 return SDValue(LoHi, 0);
5009 }
5010 }
5011
5012 // Try to transform:
5013 // (1) multiply-by-(power-of-2 +/- 1) into shift and add/sub.
5014 // mul x, (2^N + 1) --> add (shl x, N), x
5015 // mul x, (2^N - 1) --> sub (shl x, N), x
5016 // Examples: x * 33 --> (x << 5) + x
5017 // x * 15 --> (x << 4) - x
5018 // x * -33 --> -((x << 5) + x)
5019 // x * -15 --> -((x << 4) - x) ; this reduces --> x - (x << 4)
5020 // (2) multiply-by-(power-of-2 +/- power-of-2) into shifts and add/sub.
5021 // mul x, (2^N + 2^M) --> (add (shl x, N), (shl x, M))
5022 // mul x, (2^N - 2^M) --> (sub (shl x, N), (shl x, M))
5023 // Examples: x * 0x8800 --> (x << 15) + (x << 11)
5024 // x * 0xf800 --> (x << 16) - (x << 11)
5025 // x * -0x8800 --> -((x << 15) + (x << 11))
5026 // x * -0xf800 --> -((x << 16) - (x << 11)) ; (x << 11) - (x << 16)
5027 if (N1IsConst && TLI.decomposeMulByConstant(Context&: *DAG.getContext(), VT, C: N1)) {
5028 // TODO: We could handle more general decomposition of any constant by
5029 // having the target set a limit on number of ops and making a
5030 // callback to determine that sequence (similar to sqrt expansion).
5031 unsigned MathOp = ISD::DELETED_NODE;
5032 APInt MulC = ConstValue1.abs();
5033 // The constant `2` should be treated as (2^0 + 1).
5034 unsigned TZeros = MulC == 2 ? 0 : MulC.countr_zero();
5035 MulC.lshrInPlace(ShiftAmt: TZeros);
5036 if ((MulC - 1).isPowerOf2())
5037 MathOp = ISD::ADD;
5038 else if ((MulC + 1).isPowerOf2())
5039 MathOp = ISD::SUB;
5040
5041 if (MathOp != ISD::DELETED_NODE) {
5042 unsigned ShAmt =
5043 MathOp == ISD::ADD ? (MulC - 1).logBase2() : (MulC + 1).logBase2();
5044 ShAmt += TZeros;
5045 assert(ShAmt < BitWidth &&
5046 "multiply-by-constant generated out of bounds shift");
5047 SDValue Shl = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0,
5048 N2: DAG.getShiftAmountConstant(Val: ShAmt, VT, DL));
5049 SDValue R = N0;
5050 if (TZeros)
5051 R = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0,
5052 N2: DAG.getShiftAmountConstant(Val: TZeros, VT, DL));
5053 R = DAG.getNode(Opcode: MathOp, DL, VT, N1: Shl, N2: R);
5054 if (ConstValue1.isNegative())
5055 R = DAG.getNegative(Val: R, DL, VT);
5056 return R;
5057 }
5058 }
5059
5060 // (mul (shl X, c1), c2) -> (mul X, c2 << c1)
5061 {
5062 SDValue X, C1;
5063 if (sd_match(N: N0, P: m_Shl(L: m_Value(N&: X), R: m_Value(N&: C1))))
5064 if (SDValue C3 = DAG.FoldConstantArithmetic(Opcode: ISD::SHL, DL, VT, Ops: {N1, C1}))
5065 return DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: X, N2: C3);
5066 }
5067
5068 // Change (mul (shl X, C), Y) -> (shl (mul X, Y), C) when the shift has one
5069 // use.
5070 {
5071 SDValue X, C, Y;
5072 if (sd_match(N,
5073 P: m_Mul(L: m_OneUse(P: m_Shl(L: m_Value(N&: X), R: m_Value(N&: C))), R: m_Value(N&: Y))) &&
5074 isConstantOrConstantVector(N: C)) {
5075 SDValue Mul = DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: X, N2: Y);
5076 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Mul, N2: C);
5077 }
5078 }
5079
5080 // fold (mul (add x, c1), c2) -> (add (mul x, c2), c1*c2)
5081 if (N0.getOpcode() == ISD::ADD && isConstantOrConstantVector(N: N1) &&
5082 isConstantOrConstantVector(N: N0.getOperand(i: 1)) &&
5083 isMulAddWithConstProfitable(MulNode: N, AddNode: N0, ConstNode: N1))
5084 return DAG.getNode(
5085 Opcode: ISD::ADD, DL, VT,
5086 N1: DAG.getNode(Opcode: ISD::MUL, DL: SDLoc(N0), VT, N1: N0.getOperand(i: 0), N2: N1),
5087 N2: DAG.getNode(Opcode: ISD::MUL, DL: SDLoc(N1), VT, N1: N0.getOperand(i: 1), N2: N1));
5088
5089 // Fold (mul (vscale * C0), C1) to (vscale * (C0 * C1)).
5090 // avoid if ISD::MUL handling is poor and ISD::SHL isn't an option.
5091 ConstantSDNode *NC1 = isConstOrConstSplat(N: N1);
5092 if (N0.getOpcode() == ISD::VSCALE && NC1) {
5093 const APInt &C0 = N0.getConstantOperandAPInt(i: 0);
5094 const APInt &C1 = NC1->getAPIntValue();
5095 if (!C0.isPowerOf2() || C1.isPowerOf2() ||
5096 hasOperation(Opcode: ISD::MUL, VT: NC1->getValueType(ResNo: 0)))
5097 return DAG.getVScale(DL, VT, MulImm: C0 * C1);
5098 }
5099
5100 // Fold (mul step_vector(C0), C1) to (step_vector(C0 * C1)).
5101 APInt MulVal;
5102 if (N0.getOpcode() == ISD::STEP_VECTOR &&
5103 ISD::isConstantSplatVector(N: N1.getNode(), SplatValue&: MulVal)) {
5104 const APInt &C0 = N0.getConstantOperandAPInt(i: 0);
5105 APInt NewStep = C0 * MulVal;
5106 return DAG.getStepVector(DL, ResVT: VT, StepVal: NewStep);
5107 }
5108
5109 // Fold Y = sra (X, size(X)-1); mul (or (Y, 1), X) -> (abs X)
5110 SDValue X;
5111 if ((!LegalOperations || hasOperation(Opcode: ISD::ABS, VT)) &&
5112 sd_match(N, P: m_Mul(L: m_Or(L: m_Sra(L: m_Value(N&: X), R: m_SpecificInt(V: BitWidth - 1)),
5113 R: m_One()),
5114 R: m_Deferred(V&: X)))) {
5115 return DAG.getNode(Opcode: ISD::ABS, DL, VT, Operand: X);
5116 }
5117
5118 // Fold ((mul x, 0/undef) -> 0,
5119 // (mul x, 1) -> x) -> x)
5120 // -> and(x, mask)
5121 // We can replace vectors with '0' and '1' factors with a clearing mask.
5122 if (VT.isFixedLengthVector()) {
5123 unsigned NumElts = VT.getVectorNumElements();
5124 SmallBitVector ClearMask;
5125 ClearMask.reserve(N: NumElts);
5126 auto IsClearMask = [&ClearMask](ConstantSDNode *V) {
5127 if (!V || V->isZero()) {
5128 ClearMask.push_back(Val: true);
5129 return true;
5130 }
5131 ClearMask.push_back(Val: false);
5132 return V->isOne();
5133 };
5134 if ((!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::AND, VT)) &&
5135 ISD::matchUnaryPredicate(Op: N1, Match: IsClearMask, /*AllowUndefs*/ true)) {
5136 assert(N1.getOpcode() == ISD::BUILD_VECTOR && "Unknown constant vector");
5137 EVT LegalSVT = N1.getOperand(i: 0).getValueType();
5138 SDValue Zero = DAG.getConstant(Val: 0, DL, VT: LegalSVT);
5139 SDValue AllOnes = DAG.getAllOnesConstant(DL, VT: LegalSVT);
5140 SmallVector<SDValue, 16> Mask(NumElts, AllOnes);
5141 for (unsigned I = 0; I != NumElts; ++I)
5142 if (ClearMask[I])
5143 Mask[I] = Zero;
5144 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N0, N2: DAG.getBuildVector(VT, DL, Ops: Mask));
5145 }
5146 }
5147
5148 // reassociate mul
5149 if (SDValue RMUL = reassociateOps(Opc: ISD::MUL, DL, N0, N1, Flags: N->getFlags()))
5150 return RMUL;
5151
5152 // Fold mul(vecreduce(x), vecreduce(y)) -> vecreduce(mul(x, y))
5153 if (SDValue SD =
5154 reassociateReduction(RedOpc: ISD::VECREDUCE_MUL, Opc: ISD::MUL, DL, VT, N0, N1))
5155 return SD;
5156
5157 // Simplify the operands using demanded-bits information.
5158 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
5159 return SDValue(N, 0);
5160
5161 return SDValue();
5162}
5163
5164/// Return true if divmod libcall is available.
5165static bool isDivRemLibcallAvailable(SDNode *Node, bool isSigned,
5166 const SelectionDAG &DAG) {
5167 RTLIB::Libcall LC;
5168 EVT NodeType = Node->getValueType(ResNo: 0);
5169 if (!NodeType.isSimple())
5170 return false;
5171 switch (NodeType.getSimpleVT().SimpleTy) {
5172 default: return false; // No libcall for vector types.
5173 case MVT::i8: LC= isSigned ? RTLIB::SDIVREM_I8 : RTLIB::UDIVREM_I8; break;
5174 case MVT::i16: LC= isSigned ? RTLIB::SDIVREM_I16 : RTLIB::UDIVREM_I16; break;
5175 case MVT::i32: LC= isSigned ? RTLIB::SDIVREM_I32 : RTLIB::UDIVREM_I32; break;
5176 case MVT::i64: LC= isSigned ? RTLIB::SDIVREM_I64 : RTLIB::UDIVREM_I64; break;
5177 case MVT::i128: LC= isSigned ? RTLIB::SDIVREM_I128:RTLIB::UDIVREM_I128; break;
5178 }
5179
5180 return DAG.getLibcalls().getLibcallImpl(Call: LC) != RTLIB::Unsupported;
5181}
5182
5183/// Issue divrem if both quotient and remainder are needed.
5184SDValue DAGCombiner::useDivRem(SDNode *Node) {
5185 if (Node->use_empty())
5186 return SDValue(); // This is a dead node, leave it alone.
5187
5188 unsigned Opcode = Node->getOpcode();
5189 bool isSigned = (Opcode == ISD::SDIV) || (Opcode == ISD::SREM);
5190 unsigned DivRemOpc = isSigned ? ISD::SDIVREM : ISD::UDIVREM;
5191
5192 // DivMod lib calls can still work on non-legal types if using lib-calls.
5193 EVT VT = Node->getValueType(ResNo: 0);
5194 if (VT.isVector() || !VT.isInteger())
5195 return SDValue();
5196
5197 if (!TLI.isTypeLegal(VT) && !TLI.isOperationCustom(Op: DivRemOpc, VT))
5198 return SDValue();
5199
5200 // If DIVREM is going to get expanded into a libcall,
5201 // but there is no libcall available, then don't combine.
5202 if (!TLI.isOperationLegalOrCustom(Op: DivRemOpc, VT) &&
5203 !isDivRemLibcallAvailable(Node, isSigned, DAG))
5204 return SDValue();
5205
5206 // If div is legal, it's better to do the normal expansion
5207 unsigned OtherOpcode = 0;
5208 if ((Opcode == ISD::SDIV) || (Opcode == ISD::UDIV)) {
5209 OtherOpcode = isSigned ? ISD::SREM : ISD::UREM;
5210 if (TLI.isOperationLegalOrCustom(Op: Opcode, VT))
5211 return SDValue();
5212 } else {
5213 OtherOpcode = isSigned ? ISD::SDIV : ISD::UDIV;
5214 if (TLI.isOperationLegalOrCustom(Op: OtherOpcode, VT))
5215 return SDValue();
5216 }
5217
5218 SDValue Op0 = Node->getOperand(Num: 0);
5219 SDValue Op1 = Node->getOperand(Num: 1);
5220 SDValue combined;
5221 for (SDNode *User : Op0->users()) {
5222 if (User == Node || User->getOpcode() == ISD::DELETED_NODE ||
5223 User->use_empty())
5224 continue;
5225 // Convert the other matching node(s), too;
5226 // otherwise, the DIVREM may get target-legalized into something
5227 // target-specific that we won't be able to recognize.
5228 unsigned UserOpc = User->getOpcode();
5229 if ((UserOpc == Opcode || UserOpc == OtherOpcode || UserOpc == DivRemOpc) &&
5230 User->getOperand(Num: 0) == Op0 &&
5231 User->getOperand(Num: 1) == Op1) {
5232 if (!combined) {
5233 if (UserOpc == OtherOpcode) {
5234 SDVTList VTs = DAG.getVTList(VT1: VT, VT2: VT);
5235 combined = DAG.getNode(Opcode: DivRemOpc, DL: SDLoc(Node), VTList: VTs, N1: Op0, N2: Op1);
5236 } else if (UserOpc == DivRemOpc) {
5237 combined = SDValue(User, 0);
5238 } else {
5239 assert(UserOpc == Opcode);
5240 continue;
5241 }
5242 }
5243 if (UserOpc == ISD::SDIV || UserOpc == ISD::UDIV)
5244 CombineTo(N: User, Res: combined);
5245 else if (UserOpc == ISD::SREM || UserOpc == ISD::UREM)
5246 CombineTo(N: User, Res: combined.getValue(R: 1));
5247 }
5248 }
5249 return combined;
5250}
5251
5252static SDValue simplifyDivRem(SDNode *N, SelectionDAG &DAG) {
5253 SDValue N0 = N->getOperand(Num: 0);
5254 SDValue N1 = N->getOperand(Num: 1);
5255 EVT VT = N->getValueType(ResNo: 0);
5256 SDLoc DL(N);
5257
5258 unsigned Opc = N->getOpcode();
5259 bool IsDiv = (ISD::SDIV == Opc) || (ISD::UDIV == Opc);
5260
5261 // X / undef -> undef
5262 // X % undef -> undef
5263 // X / 0 -> undef
5264 // X % 0 -> undef
5265 // NOTE: This includes vectors where any divisor element is zero/undef.
5266 if (DAG.isUndef(Opcode: Opc, Ops: {N0, N1}))
5267 return DAG.getUNDEF(VT);
5268
5269 // undef / X -> 0
5270 // undef % X -> 0
5271 if (N0.isUndef())
5272 return DAG.getConstant(Val: 0, DL, VT);
5273
5274 // 0 / X -> 0
5275 // 0 % X -> 0
5276 ConstantSDNode *N0C = isConstOrConstSplat(N: N0);
5277 if (N0C && N0C->isZero())
5278 return N0;
5279
5280 // X / X -> 1
5281 // X % X -> 0
5282 if (N0 == N1)
5283 return DAG.getConstant(Val: IsDiv ? 1 : 0, DL, VT);
5284
5285 // X / 1 -> X
5286 // X % 1 -> 0
5287 // If this is a boolean op (single-bit element type), we can't have
5288 // division-by-zero or remainder-by-zero, so assume the divisor is 1.
5289 // TODO: Similarly, if we're zero-extending a boolean divisor, then assume
5290 // it's a 1.
5291 if (isOneOrOneSplat(V: N1) || (VT.getScalarType() == MVT::i1))
5292 return IsDiv ? N0 : DAG.getConstant(Val: 0, DL, VT);
5293
5294 return SDValue();
5295}
5296
5297SDValue DAGCombiner::visitSDIV(SDNode *N) {
5298 SDValue N0 = N->getOperand(Num: 0);
5299 SDValue N1 = N->getOperand(Num: 1);
5300 EVT VT = N->getValueType(ResNo: 0);
5301 EVT CCVT = getSetCCResultType(VT);
5302 SDLoc DL(N);
5303
5304 // fold (sdiv c1, c2) -> c1/c2
5305 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::SDIV, DL, VT, Ops: {N0, N1}))
5306 return C;
5307
5308 // fold vector ops
5309 if (VT.isVector())
5310 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
5311 return FoldedVOp;
5312
5313 // fold (sdiv X, -1) -> 0-X
5314 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
5315 if (N1C && N1C->isAllOnes())
5316 return DAG.getNegative(Val: N0, DL, VT);
5317
5318 // fold (sdiv X, MIN_SIGNED) -> select(X == MIN_SIGNED, 1, 0)
5319 if (N1C && N1C->isMinSignedValue())
5320 return DAG.getSelect(DL, VT, Cond: DAG.getSetCC(DL, VT: CCVT, LHS: N0, RHS: N1, Cond: ISD::SETEQ),
5321 LHS: DAG.getConstant(Val: 1, DL, VT),
5322 RHS: DAG.getConstant(Val: 0, DL, VT));
5323
5324 if (SDValue V = simplifyDivRem(N, DAG))
5325 return V;
5326
5327 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
5328 return NewSel;
5329
5330 // If we know the sign bits of both operands are zero, strength reduce to a
5331 // udiv instead. Handles (X&15) /s 4 -> X&15 >> 2
5332 if (DAG.SignBitIsZero(Op: N1) && DAG.SignBitIsZero(Op: N0))
5333 return DAG.getNode(Opcode: ISD::UDIV, DL, VT: N1.getValueType(), N1: N0, N2: N1);
5334
5335 if (SDValue V = visitSDIVLike(N0, N1, N)) {
5336 // If the corresponding remainder node exists, update its users with
5337 // (Dividend - (Quotient * Divisor).
5338 if (SDNode *RemNode = DAG.getNodeIfExists(Opcode: ISD::SREM, VTList: N->getVTList(),
5339 Ops: { N0, N1 })) {
5340 // If the sdiv has the exact flag we shouldn't propagate it to the
5341 // remainder node.
5342 if (!N->getFlags().hasExact()) {
5343 SDValue Mul = DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: V, N2: N1);
5344 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: Mul);
5345 AddToWorklist(N: Mul.getNode());
5346 AddToWorklist(N: Sub.getNode());
5347 CombineTo(N: RemNode, Res: Sub);
5348 }
5349 }
5350 return V;
5351 }
5352
5353 // sdiv, srem -> sdivrem
5354 // If the divisor is constant, then return DIVREM only if isIntDivCheap() is
5355 // true. Otherwise, we break the simplification logic in visitREM().
5356 AttributeList Attr = DAG.getMachineFunction().getFunction().getAttributes();
5357 if (!N1C || TLI.isIntDivCheap(VT: N->getValueType(ResNo: 0), Attr))
5358 if (SDValue DivRem = useDivRem(Node: N))
5359 return DivRem;
5360
5361 return SDValue();
5362}
5363
5364static bool isDivisorPowerOfTwo(SDValue Divisor) {
5365 // Helper for determining whether a value is a power-2 constant scalar or a
5366 // vector of such elements.
5367 auto IsPowerOfTwo = [](ConstantSDNode *C) {
5368 if (C->isZero() || C->isOpaque())
5369 return false;
5370 if (C->getAPIntValue().isPowerOf2())
5371 return true;
5372 if (C->getAPIntValue().isNegatedPowerOf2())
5373 return true;
5374 return false;
5375 };
5376
5377 return ISD::matchUnaryPredicate(Op: Divisor, Match: IsPowerOfTwo, /*AllowUndefs=*/false,
5378 /*AllowTruncation=*/true);
5379}
5380
5381SDValue DAGCombiner::visitSDIVLike(SDValue N0, SDValue N1, SDNode *N) {
5382 SDLoc DL(N);
5383 EVT VT = N->getValueType(ResNo: 0);
5384 EVT CCVT = getSetCCResultType(VT);
5385 unsigned BitWidth = VT.getScalarSizeInBits();
5386 unsigned MaxLegalDivRemBitWidth = TLI.getMaxDivRemBitWidthSupported();
5387
5388 // fold (sdiv X, pow2) -> simple ops after legalize
5389 // FIXME: We check for the exact bit here because the generic lowering gives
5390 // better results in that case. The target-specific lowering should learn how
5391 // to handle exact sdivs efficiently. An exception is made for large bitwidths
5392 // exceeding what the target can natively support, as division expansion was
5393 // skipped in favor of this optimization.
5394 if ((!N->getFlags().hasExact() || BitWidth > MaxLegalDivRemBitWidth) &&
5395 isDivisorPowerOfTwo(Divisor: N1)) {
5396 // Target-specific implementation of sdiv x, pow2.
5397 if (SDValue Res = BuildSDIVPow2(N))
5398 return Res;
5399
5400 // Create constants that are functions of the shift amount value.
5401 EVT ShiftAmtTy = getShiftAmountTy(LHSTy: N0.getValueType());
5402 SDValue Bits = DAG.getConstant(Val: BitWidth, DL, VT: ShiftAmtTy);
5403 SDValue C1 = DAG.getNode(Opcode: ISD::CTTZ, DL, VT, Operand: N1);
5404 C1 = DAG.getZExtOrTrunc(Op: C1, DL, VT: ShiftAmtTy);
5405 SDValue Inexact = DAG.getNode(Opcode: ISD::SUB, DL, VT: ShiftAmtTy, N1: Bits, N2: C1);
5406 if (!isConstantOrConstantVector(N: Inexact))
5407 return SDValue();
5408
5409 // Splat the sign bit into the register
5410 SDValue Sign = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N0,
5411 N2: DAG.getConstant(Val: BitWidth - 1, DL, VT: ShiftAmtTy));
5412 AddToWorklist(N: Sign.getNode());
5413
5414 // Add (N0 < 0) ? abs2 - 1 : 0;
5415 SDValue Srl = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: Sign, N2: Inexact);
5416 AddToWorklist(N: Srl.getNode());
5417 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: Srl);
5418 AddToWorklist(N: Add.getNode());
5419 SDValue Sra = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: Add, N2: C1);
5420 AddToWorklist(N: Sra.getNode());
5421
5422 // Special case: (sdiv X, 1) -> X
5423 // Special Case: (sdiv X, -1) -> 0-X
5424 SDValue One = DAG.getConstant(Val: 1, DL, VT);
5425 SDValue AllOnes = DAG.getAllOnesConstant(DL, VT);
5426 SDValue IsOne = DAG.getSetCC(DL, VT: CCVT, LHS: N1, RHS: One, Cond: ISD::SETEQ);
5427 SDValue IsAllOnes = DAG.getSetCC(DL, VT: CCVT, LHS: N1, RHS: AllOnes, Cond: ISD::SETEQ);
5428 SDValue IsOneOrAllOnes = DAG.getNode(Opcode: ISD::OR, DL, VT: CCVT, N1: IsOne, N2: IsAllOnes);
5429 Sra = DAG.getSelect(DL, VT, Cond: IsOneOrAllOnes, LHS: N0, RHS: Sra);
5430
5431 // If dividing by a positive value, we're done. Otherwise, the result must
5432 // be negated.
5433 SDValue Zero = DAG.getConstant(Val: 0, DL, VT);
5434 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Zero, N2: Sra);
5435
5436 // FIXME: Use SELECT_CC once we improve SELECT_CC constant-folding.
5437 SDValue IsNeg = DAG.getSetCC(DL, VT: CCVT, LHS: N1, RHS: Zero, Cond: ISD::SETLT);
5438 SDValue Res = DAG.getSelect(DL, VT, Cond: IsNeg, LHS: Sub, RHS: Sra);
5439 return Res;
5440 }
5441
5442 // If integer divide is expensive and we satisfy the requirements, emit an
5443 // alternate sequence. Targets may check function attributes for size/speed
5444 // trade-offs.
5445 AttributeList Attr = DAG.getMachineFunction().getFunction().getAttributes();
5446 if (isConstantOrConstantVector(N: N1, /*NoOpaques=*/false,
5447 /*AllowTruncation=*/true) &&
5448 !TLI.isIntDivCheap(VT: N->getValueType(ResNo: 0), Attr))
5449 if (SDValue Op = BuildSDIV(N))
5450 return Op;
5451
5452 return SDValue();
5453}
5454
5455SDValue DAGCombiner::visitUDIV(SDNode *N) {
5456 SDValue N0 = N->getOperand(Num: 0);
5457 SDValue N1 = N->getOperand(Num: 1);
5458 EVT VT = N->getValueType(ResNo: 0);
5459 EVT CCVT = getSetCCResultType(VT);
5460 SDLoc DL(N);
5461
5462 // fold (udiv c1, c2) -> c1/c2
5463 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::UDIV, DL, VT, Ops: {N0, N1}))
5464 return C;
5465
5466 // fold vector ops
5467 if (VT.isVector())
5468 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
5469 return FoldedVOp;
5470
5471 // fold (udiv X, -1) -> select(X == -1, 1, 0)
5472 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
5473 if (N1C && N1C->isAllOnes() && CCVT.isVector() == VT.isVector()) {
5474 return DAG.getSelect(DL, VT, Cond: DAG.getSetCC(DL, VT: CCVT, LHS: N0, RHS: N1, Cond: ISD::SETEQ),
5475 LHS: DAG.getConstant(Val: 1, DL, VT),
5476 RHS: DAG.getConstant(Val: 0, DL, VT));
5477 }
5478
5479 if (SDValue V = simplifyDivRem(N, DAG))
5480 return V;
5481
5482 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
5483 return NewSel;
5484
5485 if (SDValue V = visitUDIVLike(N0, N1, N)) {
5486 // If the corresponding remainder node exists, update its users with
5487 // (Dividend - (Quotient * Divisor).
5488 if (SDNode *RemNode = DAG.getNodeIfExists(Opcode: ISD::UREM, VTList: N->getVTList(),
5489 Ops: { N0, N1 })) {
5490 // If the udiv has the exact flag we shouldn't propagate it to the
5491 // remainder node.
5492 if (!N->getFlags().hasExact()) {
5493 SDValue Mul = DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: V, N2: N1);
5494 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: Mul);
5495 AddToWorklist(N: Mul.getNode());
5496 AddToWorklist(N: Sub.getNode());
5497 CombineTo(N: RemNode, Res: Sub);
5498 }
5499 }
5500 return V;
5501 }
5502
5503 // sdiv, srem -> sdivrem
5504 // If the divisor is constant, then return DIVREM only if isIntDivCheap() is
5505 // true. Otherwise, we break the simplification logic in visitREM().
5506 AttributeList Attr = DAG.getMachineFunction().getFunction().getAttributes();
5507 if (!N1C || TLI.isIntDivCheap(VT: N->getValueType(ResNo: 0), Attr))
5508 if (SDValue DivRem = useDivRem(Node: N))
5509 return DivRem;
5510
5511 // Simplify the operands using demanded-bits information.
5512 // We don't have demanded bits support for UDIV so this just enables constant
5513 // folding based on known bits.
5514 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
5515 return SDValue(N, 0);
5516
5517 return SDValue();
5518}
5519
5520SDValue DAGCombiner::visitUDIVLike(SDValue N0, SDValue N1, SDNode *N) {
5521 SDLoc DL(N);
5522 EVT VT = N->getValueType(ResNo: 0);
5523
5524 // fold (udiv x, (1 << c)) -> x >>u c
5525 if (isConstantOrConstantVector(N: N1, /*NoOpaques=*/true,
5526 /*AllowTruncation=*/true)) {
5527 if (SDValue LogBase2 = BuildLogBase2(V: N1, DL)) {
5528 AddToWorklist(N: LogBase2.getNode());
5529
5530 EVT ShiftVT = getShiftAmountTy(LHSTy: N0.getValueType());
5531 SDValue Trunc = DAG.getZExtOrTrunc(Op: LogBase2, DL, VT: ShiftVT);
5532 AddToWorklist(N: Trunc.getNode());
5533 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0, N2: Trunc);
5534 }
5535 }
5536
5537 // fold (udiv x, (shl c, y)) -> x >>u (log2(c)+y) iff c is power of 2
5538 if (N1.getOpcode() == ISD::SHL) {
5539 SDValue N10 = N1.getOperand(i: 0);
5540 if (isConstantOrConstantVector(N: N10, /*NoOpaques=*/true,
5541 /*AllowTruncation=*/true)) {
5542 if (SDValue LogBase2 = BuildLogBase2(V: N10, DL)) {
5543 AddToWorklist(N: LogBase2.getNode());
5544
5545 EVT ADDVT = N1.getOperand(i: 1).getValueType();
5546 SDValue Trunc = DAG.getZExtOrTrunc(Op: LogBase2, DL, VT: ADDVT);
5547 AddToWorklist(N: Trunc.getNode());
5548 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT: ADDVT, N1: N1.getOperand(i: 1), N2: Trunc);
5549 AddToWorklist(N: Add.getNode());
5550 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0, N2: Add);
5551 }
5552 }
5553 }
5554
5555 // fold (udiv x, c) -> alternate
5556 AttributeList Attr = DAG.getMachineFunction().getFunction().getAttributes();
5557 if (isConstantOrConstantVector(N: N1, /*NoOpaques=*/false,
5558 /*AllowTruncation=*/true) &&
5559 !TLI.isIntDivCheap(VT: N->getValueType(ResNo: 0), Attr))
5560 if (SDValue Op = BuildUDIV(N))
5561 return Op;
5562
5563 return SDValue();
5564}
5565
5566SDValue DAGCombiner::buildOptimizedSREM(SDValue N0, SDValue N1, SDNode *N) {
5567 if (!N->getFlags().hasExact() && isDivisorPowerOfTwo(Divisor: N1) &&
5568 !DAG.doesNodeExist(Opcode: ISD::SDIV, VTList: N->getVTList(), Ops: {N0, N1})) {
5569 // Target-specific implementation of srem x, pow2.
5570 if (SDValue Res = BuildSREMPow2(N))
5571 return Res;
5572 }
5573 return SDValue();
5574}
5575
5576// handles ISD::SREM and ISD::UREM
5577SDValue DAGCombiner::visitREM(SDNode *N) {
5578 unsigned Opcode = N->getOpcode();
5579 SDValue N0 = N->getOperand(Num: 0);
5580 SDValue N1 = N->getOperand(Num: 1);
5581 EVT VT = N->getValueType(ResNo: 0);
5582 EVT CCVT = getSetCCResultType(VT);
5583
5584 bool isSigned = (Opcode == ISD::SREM);
5585 unsigned DivOpcode = isSigned ? ISD::SDIV : ISD::UDIV;
5586 SDLoc DL(N);
5587
5588 // fold (rem c1, c2) -> c1%c2
5589 if (SDValue C = DAG.FoldConstantArithmetic(Opcode, DL, VT, Ops: {N0, N1}))
5590 return C;
5591
5592 // fold (urem X, -1) -> select(FX == -1, 0, FX)
5593 // Freeze the numerator to avoid a miscompile with an undefined value.
5594 if (!isSigned && llvm::isAllOnesOrAllOnesSplat(V: N1, /*AllowUndefs*/ false) &&
5595 CCVT.isVector() == VT.isVector()) {
5596 SDValue F0 = DAG.getFreeze(V: N0);
5597 SDValue EqualsNeg1 = DAG.getSetCC(DL, VT: CCVT, LHS: F0, RHS: N1, Cond: ISD::SETEQ);
5598 return DAG.getSelect(DL, VT, Cond: EqualsNeg1, LHS: DAG.getConstant(Val: 0, DL, VT), RHS: F0);
5599 }
5600
5601 if (SDValue V = simplifyDivRem(N, DAG))
5602 return V;
5603
5604 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
5605 return NewSel;
5606
5607 if (isSigned) {
5608 // If we know the sign bits of both operands are zero, strength reduce to a
5609 // urem instead. Handles (X & 0x0FFFFFFF) %s 16 -> X&15
5610 if (DAG.SignBitIsZero(Op: N1) && DAG.SignBitIsZero(Op: N0))
5611 return DAG.getNode(Opcode: ISD::UREM, DL, VT, N1: N0, N2: N1);
5612 } else {
5613 if (DAG.isKnownToBeAPowerOfTwo(Val: N1, /*OrZero=*/true)) {
5614 // fold (urem x, pow2) -> (and x, pow2-1)
5615 SDValue NegOne = DAG.getAllOnesConstant(DL, VT);
5616 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1, N2: NegOne);
5617 AddToWorklist(N: Add.getNode());
5618 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N0, N2: Add);
5619 }
5620 }
5621
5622 AttributeList Attr = DAG.getMachineFunction().getFunction().getAttributes();
5623
5624 // If X/C can be simplified by the division-by-constant logic, lower
5625 // X%C to the equivalent of X-X/C*C.
5626 // Reuse the SDIVLike/UDIVLike combines - to avoid mangling nodes, the
5627 // speculative DIV must not cause a DIVREM conversion. We guard against this
5628 // by skipping the simplification if isIntDivCheap(). When div is not cheap,
5629 // combine will not return a DIVREM. Regardless, checking cheapness here
5630 // makes sense since the simplification results in fatter code.
5631 if (DAG.isKnownNeverZero(Op: N1) && !TLI.isIntDivCheap(VT, Attr)) {
5632 if (isSigned) {
5633 // check if we can build faster implementation for srem
5634 if (SDValue OptimizedRem = buildOptimizedSREM(N0, N1, N))
5635 return OptimizedRem;
5636 }
5637
5638 SDValue OptimizedDiv =
5639 isSigned ? visitSDIVLike(N0, N1, N) : visitUDIVLike(N0, N1, N);
5640 if (OptimizedDiv.getNode() && OptimizedDiv.getNode() != N) {
5641 // If the equivalent Div node also exists, update its users.
5642 if (SDNode *DivNode = DAG.getNodeIfExists(Opcode: DivOpcode, VTList: N->getVTList(),
5643 Ops: { N0, N1 }))
5644 CombineTo(N: DivNode, Res: OptimizedDiv);
5645 SDValue Mul = DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: OptimizedDiv, N2: N1);
5646 SDValue Sub = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: Mul);
5647 AddToWorklist(N: OptimizedDiv.getNode());
5648 AddToWorklist(N: Mul.getNode());
5649 return Sub;
5650 }
5651 }
5652
5653 // Fold Num % Den -> Num - (Num / Den) * Den, if (Num / Den) is already
5654 // computed. Defer for types that will be promoted and do not fold if DIVREM
5655 // is available
5656 unsigned DivRemOpc = isSigned ? ISD::SDIVREM : ISD::UDIVREM;
5657 if (!ForCodeSize &&
5658 !TLI.isOperationLegalOrCustom(Op: DivRemOpc, VT: VT.getScalarType()) &&
5659 !isDivRemLibcallAvailable(Node: N, isSigned, DAG) &&
5660 TLI.getTypeAction(Context&: *DAG.getContext(), VT) !=
5661 TargetLowering::TypePromoteInteger) {
5662 if (SDNode *Div =
5663 DAG.getNodeIfExists(Opcode: DivOpcode, VTList: N->getVTList(), Ops: {N0, N1})) {
5664 SDValue Mul = DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: SDValue(Div, 0), N2: N1);
5665 return DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: N0, N2: Mul);
5666 }
5667 }
5668
5669 // sdiv, srem -> sdivrem
5670 if (SDValue DivRem = useDivRem(Node: N))
5671 return DivRem.getValue(R: 1);
5672
5673 // fold urem(urem(A, BCst), Op1Cst) -> urem(A, Op1Cst)
5674 // iff urem(BCst, Op1Cst) == 0
5675 SDValue A;
5676 APInt Op1Cst, BCst;
5677 if (sd_match(N, P: m_URem(L: m_URem(L: m_Value(N&: A), R: m_ConstInt(V&: BCst)),
5678 R: m_ConstInt(V&: Op1Cst))) &&
5679 BCst.urem(RHS: Op1Cst).isZero()) {
5680 return DAG.getNode(Opcode: ISD::UREM, DL, VT, N1: A, N2: DAG.getConstant(Val: Op1Cst, DL, VT));
5681 }
5682
5683 // fold srem(srem(A, BCst), Op1Cst) -> srem(A, Op1Cst)
5684 // iff srem(BCst, Op1Cst) == 0 && Op1Cst != 1
5685 if (sd_match(N, P: m_SRem(L: m_SRem(L: m_Value(N&: A), R: m_ConstInt(V&: BCst)),
5686 R: m_ConstInt(V&: Op1Cst))) &&
5687 BCst.srem(RHS: Op1Cst).isZero() && !Op1Cst.isAllOnes()) {
5688 return DAG.getNode(Opcode: ISD::SREM, DL, VT, N1: A, N2: DAG.getConstant(Val: Op1Cst, DL, VT));
5689 }
5690
5691 return SDValue();
5692}
5693
5694SDValue DAGCombiner::visitMULHS(SDNode *N) {
5695 SDValue N0 = N->getOperand(Num: 0);
5696 SDValue N1 = N->getOperand(Num: 1);
5697 EVT VT = N->getValueType(ResNo: 0);
5698 SDLoc DL(N);
5699
5700 // fold (mulhs c1, c2)
5701 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::MULHS, DL, VT, Ops: {N0, N1}))
5702 return C;
5703
5704 // canonicalize constant to RHS.
5705 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
5706 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
5707 return DAG.getNode(Opcode: ISD::MULHS, DL, VTList: N->getVTList(), N1, N2: N0);
5708
5709 if (VT.isVector()) {
5710 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
5711 return FoldedVOp;
5712
5713 // fold (mulhs x, 0) -> 0
5714 // do not return N1, because undef node may exist.
5715 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
5716 return DAG.getConstant(Val: 0, DL, VT);
5717 }
5718
5719 // fold (mulhs x, 0) -> 0
5720 if (isNullConstant(V: N1))
5721 return N1;
5722
5723 // fold (mulhs x, 1) -> (sra x, size(x)-1)
5724 if (isOneConstant(V: N1))
5725 return DAG.getNode(
5726 Opcode: ISD::SRA, DL, VT, N1: N0,
5727 N2: DAG.getShiftAmountConstant(Val: N0.getScalarValueSizeInBits() - 1, VT, DL));
5728
5729 // fold (mulhs x, undef) -> 0
5730 if (N0.isUndef() || N1.isUndef())
5731 return DAG.getConstant(Val: 0, DL, VT);
5732
5733 // If the type twice as wide is legal, transform the mulhs to a wider multiply
5734 // plus a shift.
5735 if (!TLI.isOperationLegalOrCustom(Op: ISD::MULHS, VT) && VT.isSimple() &&
5736 !VT.isVector()) {
5737 MVT Simple = VT.getSimpleVT();
5738 unsigned SimpleSize = Simple.getSizeInBits();
5739 EVT NewVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: SimpleSize*2);
5740 if (TLI.isOperationLegal(Op: ISD::MUL, VT: NewVT)) {
5741 N0 = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT: NewVT, Operand: N0);
5742 N1 = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT: NewVT, Operand: N1);
5743 N1 = DAG.getNode(Opcode: ISD::MUL, DL, VT: NewVT, N1: N0, N2: N1);
5744 N1 = DAG.getNode(Opcode: ISD::SRL, DL, VT: NewVT, N1,
5745 N2: DAG.getShiftAmountConstant(Val: SimpleSize, VT: NewVT, DL));
5746 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N1);
5747 }
5748 }
5749
5750 return SDValue();
5751}
5752
5753SDValue DAGCombiner::visitMULHU(SDNode *N) {
5754 SDValue N0 = N->getOperand(Num: 0);
5755 SDValue N1 = N->getOperand(Num: 1);
5756 EVT VT = N->getValueType(ResNo: 0);
5757 SDLoc DL(N);
5758
5759 // fold (mulhu c1, c2)
5760 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::MULHU, DL, VT, Ops: {N0, N1}))
5761 return C;
5762
5763 // canonicalize constant to RHS.
5764 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
5765 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
5766 return DAG.getNode(Opcode: ISD::MULHU, DL, VTList: N->getVTList(), N1, N2: N0);
5767
5768 if (VT.isVector()) {
5769 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
5770 return FoldedVOp;
5771
5772 // fold (mulhu x, 0) -> 0
5773 // do not return N1, because undef node may exist.
5774 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
5775 return DAG.getConstant(Val: 0, DL, VT);
5776 }
5777
5778 // fold (mulhu x, 0) -> 0
5779 if (isNullConstant(V: N1))
5780 return N1;
5781
5782 // fold (mulhu x, 1) -> 0
5783 if (isOneConstant(V: N1))
5784 return DAG.getConstant(Val: 0, DL, VT);
5785
5786 // fold (mulhu x, undef) -> 0
5787 if (N0.isUndef() || N1.isUndef())
5788 return DAG.getConstant(Val: 0, DL, VT);
5789
5790 // fold (mulhu x, (1 << c)) -> x >> (bitwidth - c)
5791 if (isConstantOrConstantVector(N: N1, /*NoOpaques=*/true,
5792 /*AllowTruncation=*/true) &&
5793 (!LegalOperations || hasOperation(Opcode: ISD::SRL, VT))) {
5794 if (SDValue LogBase2 = BuildLogBase2(V: N1, DL)) {
5795 unsigned NumEltBits = VT.getScalarSizeInBits();
5796 SDValue SRLAmt = DAG.getNode(
5797 Opcode: ISD::SUB, DL, VT, N1: DAG.getConstant(Val: NumEltBits, DL, VT), N2: LogBase2);
5798 EVT ShiftVT = getShiftAmountTy(LHSTy: N0.getValueType());
5799 SDValue Trunc = DAG.getZExtOrTrunc(Op: SRLAmt, DL, VT: ShiftVT);
5800 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0, N2: Trunc);
5801 }
5802 }
5803
5804 // If the type twice as wide is legal, transform the mulhu to a wider multiply
5805 // plus a shift.
5806 if (!TLI.isOperationLegalOrCustom(Op: ISD::MULHU, VT) && VT.isSimple() &&
5807 !VT.isVector()) {
5808 MVT Simple = VT.getSimpleVT();
5809 unsigned SimpleSize = Simple.getSizeInBits();
5810 EVT NewVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: SimpleSize*2);
5811 if (TLI.isOperationLegal(Op: ISD::MUL, VT: NewVT)) {
5812 N0 = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: NewVT, Operand: N0);
5813 N1 = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: NewVT, Operand: N1);
5814 N1 = DAG.getNode(Opcode: ISD::MUL, DL, VT: NewVT, N1: N0, N2: N1);
5815 N1 = DAG.getNode(Opcode: ISD::SRL, DL, VT: NewVT, N1,
5816 N2: DAG.getShiftAmountConstant(Val: SimpleSize, VT: NewVT, DL));
5817 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N1);
5818 }
5819 }
5820
5821 // Simplify the operands using demanded-bits information.
5822 // We don't have demanded bits support for MULHU so this just enables constant
5823 // folding based on known bits.
5824 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
5825 return SDValue(N, 0);
5826
5827 return SDValue();
5828}
5829
5830SDValue DAGCombiner::visitAVG(SDNode *N) {
5831 unsigned Opcode = N->getOpcode();
5832 SDValue N0 = N->getOperand(Num: 0);
5833 SDValue N1 = N->getOperand(Num: 1);
5834 EVT VT = N->getValueType(ResNo: 0);
5835 SDLoc DL(N);
5836 bool IsSigned = Opcode == ISD::AVGCEILS || Opcode == ISD::AVGFLOORS;
5837
5838 // fold (avg c1, c2)
5839 if (SDValue C = DAG.FoldConstantArithmetic(Opcode, DL, VT, Ops: {N0, N1}))
5840 return C;
5841
5842 // canonicalize constant to RHS.
5843 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
5844 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
5845 return DAG.getNode(Opcode, DL, VTList: N->getVTList(), N1, N2: N0);
5846
5847 if (VT.isVector())
5848 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
5849 return FoldedVOp;
5850
5851 // fold (avg x, undef) -> x
5852 if (N0.isUndef())
5853 return N1;
5854 if (N1.isUndef())
5855 return N0;
5856
5857 // fold (avg x, x) --> x
5858 if (N0 == N1 && Level >= AfterLegalizeTypes)
5859 return N0;
5860
5861 // fold (avgfloor x, 0) -> x >> 1
5862 SDValue X, Y;
5863 if (sd_match(N, P: m_c_BinOp(Opc: ISD::AVGFLOORS, L: m_Value(N&: X), R: m_Zero())))
5864 return DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: X,
5865 N2: DAG.getShiftAmountConstant(Val: 1, VT, DL));
5866 if (sd_match(N, P: m_c_BinOp(Opc: ISD::AVGFLOORU, L: m_Value(N&: X), R: m_Zero())))
5867 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: X,
5868 N2: DAG.getShiftAmountConstant(Val: 1, VT, DL));
5869
5870 // fold avgu(zext(x), zext(y)) -> zext(avgu(x, y))
5871 // fold avgs(sext(x), sext(y)) -> sext(avgs(x, y))
5872 if (!IsSigned &&
5873 sd_match(N, P: m_BinOp(Opc: Opcode, L: m_ZExt(Op: m_Value(N&: X)), R: m_ZExt(Op: m_Value(N&: Y)))) &&
5874 X.getValueType() == Y.getValueType() &&
5875 hasOperation(Opcode, VT: X.getValueType())) {
5876 SDValue AvgU = DAG.getNode(Opcode, DL, VT: X.getValueType(), N1: X, N2: Y);
5877 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: AvgU);
5878 }
5879 if (IsSigned &&
5880 sd_match(N, P: m_BinOp(Opc: Opcode, L: m_SExt(Op: m_Value(N&: X)), R: m_SExt(Op: m_Value(N&: Y)))) &&
5881 X.getValueType() == Y.getValueType() &&
5882 hasOperation(Opcode, VT: X.getValueType())) {
5883 SDValue AvgS = DAG.getNode(Opcode, DL, VT: X.getValueType(), N1: X, N2: Y);
5884 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: AvgS);
5885 }
5886
5887 // Fold avgflooru(x,y) -> avgceilu(x,y-1) iff y != 0
5888 // Fold avgflooru(x,y) -> avgceilu(x-1,y) iff x != 0
5889 // Check if avgflooru isn't legal/custom but avgceilu is.
5890 if (Opcode == ISD::AVGFLOORU && !hasOperation(Opcode: ISD::AVGFLOORU, VT) &&
5891 (!LegalOperations || hasOperation(Opcode: ISD::AVGCEILU, VT))) {
5892 if (DAG.isKnownNeverZero(Op: N1))
5893 return DAG.getNode(
5894 Opcode: ISD::AVGCEILU, DL, VT, N1: N0,
5895 N2: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1, N2: DAG.getAllOnesConstant(DL, VT)));
5896 if (DAG.isKnownNeverZero(Op: N0))
5897 return DAG.getNode(
5898 Opcode: ISD::AVGCEILU, DL, VT, N1,
5899 N2: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: N0, N2: DAG.getAllOnesConstant(DL, VT)));
5900 }
5901
5902 // Fold avgfloor((add nw x,y), 1) -> avgceil(x,y)
5903 // Fold avgfloor((add nw x,1), y) -> avgceil(x,y)
5904 if ((Opcode == ISD::AVGFLOORU && hasOperation(Opcode: ISD::AVGCEILU, VT)) ||
5905 (Opcode == ISD::AVGFLOORS && hasOperation(Opcode: ISD::AVGCEILS, VT))) {
5906 SDValue Add;
5907 if (sd_match(N,
5908 P: m_c_BinOp(Opc: Opcode, L: m_Value(N&: Add, P: m_Add(L: m_Value(N&: X), R: m_Value(N&: Y))),
5909 R: m_One())) ||
5910 sd_match(N, P: m_c_BinOp(Opc: Opcode, L: m_Value(N&: Add, P: m_Add(L: m_Value(N&: X), R: m_One())),
5911 R: m_Value(N&: Y)))) {
5912
5913 if (IsSigned && Add->getFlags().hasNoSignedWrap())
5914 return DAG.getNode(Opcode: ISD::AVGCEILS, DL, VT, N1: X, N2: Y);
5915
5916 if (!IsSigned && Add->getFlags().hasNoUnsignedWrap())
5917 return DAG.getNode(Opcode: ISD::AVGCEILU, DL, VT, N1: X, N2: Y);
5918 }
5919 }
5920
5921 // Fold avgfloors(x,y) -> avgflooru(x,y) if both x and y are non-negative
5922 if (Opcode == ISD::AVGFLOORS && hasOperation(Opcode: ISD::AVGFLOORU, VT)) {
5923 if (DAG.SignBitIsZero(Op: N0) && DAG.SignBitIsZero(Op: N1))
5924 return DAG.getNode(Opcode: ISD::AVGFLOORU, DL, VT, N1: N0, N2: N1);
5925 }
5926
5927 return SDValue();
5928}
5929
5930SDValue DAGCombiner::visitABD(SDNode *N) {
5931 unsigned Opcode = N->getOpcode();
5932 SDValue N0 = N->getOperand(Num: 0);
5933 SDValue N1 = N->getOperand(Num: 1);
5934 EVT VT = N->getValueType(ResNo: 0);
5935 SDLoc DL(N);
5936
5937 // fold (abd c1, c2)
5938 if (SDValue C = DAG.FoldConstantArithmetic(Opcode, DL, VT, Ops: {N0, N1}))
5939 return C;
5940
5941 // canonicalize constant to RHS.
5942 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
5943 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
5944 return DAG.getNode(Opcode, DL, VTList: N->getVTList(), N1, N2: N0);
5945
5946 if (VT.isVector())
5947 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
5948 return FoldedVOp;
5949
5950 // fold (abd x, undef) -> 0
5951 if (N0.isUndef() || N1.isUndef())
5952 return DAG.getConstant(Val: 0, DL, VT);
5953
5954 // fold (abd x, x) -> 0
5955 if (N0 == N1)
5956 return DAG.getConstant(Val: 0, DL, VT);
5957
5958 SDValue X, Y;
5959
5960 // fold (abds x, 0) -> abs x
5961 if (sd_match(N, P: m_c_BinOp(Opc: ISD::ABDS, L: m_Value(N&: X), R: m_Zero())) &&
5962 (!LegalOperations || hasOperation(Opcode: ISD::ABS, VT)))
5963 return DAG.getNode(Opcode: ISD::ABS, DL, VT, Operand: X);
5964
5965 // fold (abdu x, 0) -> x
5966 if (sd_match(N, P: m_c_BinOp(Opc: ISD::ABDU, L: m_Value(N&: X), R: m_Zero())))
5967 return X;
5968
5969 // fold (abds x, y) -> (abdu x, y) iff both args are known positive
5970 if (Opcode == ISD::ABDS && hasOperation(Opcode: ISD::ABDU, VT) &&
5971 DAG.SignBitIsZero(Op: N0) && DAG.SignBitIsZero(Op: N1))
5972 return DAG.getNode(Opcode: ISD::ABDU, DL, VT, N1, N2: N0);
5973
5974 // fold (abd? (?ext x), (?ext y)) -> (zext (abd? x, y))
5975 if (sd_match(N, P: m_BinOp(Opc: ISD::ABDU, L: m_ZExt(Op: m_Value(N&: X)), R: m_ZExt(Op: m_Value(N&: Y)))) ||
5976 sd_match(N, P: m_BinOp(Opc: ISD::ABDS, L: m_SExt(Op: m_Value(N&: X)), R: m_SExt(Op: m_Value(N&: Y))))) {
5977 EVT SmallVT = X.getScalarValueSizeInBits() > Y.getScalarValueSizeInBits()
5978 ? X.getValueType()
5979 : Y.getValueType();
5980 if (!LegalOperations || hasOperation(Opcode, VT: SmallVT)) {
5981 SDValue ExtedX = DAG.getExtOrTrunc(Op: X, DL: SDLoc(X), VT: SmallVT, Opcode: N0->getOpcode());
5982 SDValue ExtedY = DAG.getExtOrTrunc(Op: Y, DL: SDLoc(Y), VT: SmallVT, Opcode: N0->getOpcode());
5983 SDValue SmallABD = DAG.getNode(Opcode, DL, VT: SmallVT, Ops: {ExtedX, ExtedY});
5984 SDValue ZExted = DAG.getZExtOrTrunc(Op: SmallABD, DL, VT);
5985 return ZExted;
5986 }
5987 }
5988
5989 // fold (abd? (?ext ty:x), small_const:c) -> (zext (abd? x, c))
5990 if (sd_match(N, P: m_c_BinOp(Opc: ISD::ABDU, L: m_ZExt(Op: m_Value(N&: X)), R: m_Value(N&: Y))) ||
5991 sd_match(N, P: m_c_BinOp(Opc: ISD::ABDS, L: m_SExt(Op: m_Value(N&: X)), R: m_Value(N&: Y)))) {
5992 EVT SmallVT = X.getValueType();
5993 if (!LegalOperations || hasOperation(Opcode, VT: SmallVT)) {
5994 uint64_t Bits = SmallVT.getScalarSizeInBits();
5995 unsigned RelevantBits =
5996 (Opcode == ISD::ABDS) ? DAG.ComputeMaxSignificantBits(Op: Y)
5997 : DAG.computeKnownBits(Op: Y).countMaxActiveBits();
5998 bool TruncatingYIsCheap = TLI.isTruncateFree(Val: Y, VT2: SmallVT) ||
5999 ISD::matchUnaryPredicate(
6000 Op: Y,
6001 Match: [&](auto *C) {
6002 if (!C)
6003 return true;
6004 const APInt &YConst = C->getAsAPIntVal();
6005 return (Opcode == ISD::ABDS)
6006 ? YConst.isSignedIntN(N: Bits)
6007 : YConst.isIntN(N: Bits);
6008 },
6009 /*AllowUndefs=*/true);
6010
6011 if (RelevantBits <= Bits && TruncatingYIsCheap) {
6012 SDValue NewY = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(Y), VT: SmallVT, Operand: Y);
6013 SDValue SmallABD = DAG.getNode(Opcode, DL, VT: SmallVT, Ops: {X, NewY});
6014 return DAG.getZExtOrTrunc(Op: SmallABD, DL, VT);
6015 }
6016 }
6017 }
6018
6019 return SDValue();
6020}
6021
6022/// Perform optimizations common to nodes that compute two values. LoOp and HiOp
6023/// give the opcodes for the two computations that are being performed. Return
6024/// true if a simplification was made.
6025SDValue DAGCombiner::SimplifyNodeWithTwoResults(SDNode *N, unsigned LoOp,
6026 unsigned HiOp) {
6027 // If the high half is not needed, just compute the low half.
6028 bool HiExists = N->hasAnyUseOfValue(Value: 1);
6029 if (!HiExists && (!LegalOperations ||
6030 TLI.isOperationLegalOrCustom(Op: LoOp, VT: N->getValueType(ResNo: 0)))) {
6031 SDValue Res = DAG.getNode(Opcode: LoOp, DL: SDLoc(N), VT: N->getValueType(ResNo: 0), Ops: N->ops());
6032 return CombineTo(N, Res0: Res, Res1: Res);
6033 }
6034
6035 // If the low half is not needed, just compute the high half.
6036 bool LoExists = N->hasAnyUseOfValue(Value: 0);
6037 if (!LoExists && (!LegalOperations ||
6038 TLI.isOperationLegalOrCustom(Op: HiOp, VT: N->getValueType(ResNo: 1)))) {
6039 SDValue Res = DAG.getNode(Opcode: HiOp, DL: SDLoc(N), VT: N->getValueType(ResNo: 1), Ops: N->ops());
6040 return CombineTo(N, Res0: Res, Res1: Res);
6041 }
6042
6043 // If both halves are used, return as it is.
6044 if (LoExists && HiExists)
6045 return SDValue();
6046
6047 // If the two computed results can be simplified separately, separate them.
6048 if (LoExists) {
6049 SDValue Lo = DAG.getNode(Opcode: LoOp, DL: SDLoc(N), VT: N->getValueType(ResNo: 0), Ops: N->ops());
6050 AddToWorklist(N: Lo.getNode());
6051 SDValue LoOpt = combine(N: Lo.getNode());
6052 if (LoOpt.getNode() && LoOpt.getNode() != Lo.getNode() &&
6053 (!LegalOperations ||
6054 TLI.isOperationLegalOrCustom(Op: LoOpt.getOpcode(), VT: LoOpt.getValueType())))
6055 return CombineTo(N, Res0: LoOpt, Res1: LoOpt);
6056 }
6057
6058 if (HiExists) {
6059 SDValue Hi = DAG.getNode(Opcode: HiOp, DL: SDLoc(N), VT: N->getValueType(ResNo: 1), Ops: N->ops());
6060 AddToWorklist(N: Hi.getNode());
6061 SDValue HiOpt = combine(N: Hi.getNode());
6062 if (HiOpt.getNode() && HiOpt != Hi &&
6063 (!LegalOperations ||
6064 TLI.isOperationLegalOrCustom(Op: HiOpt.getOpcode(), VT: HiOpt.getValueType())))
6065 return CombineTo(N, Res0: HiOpt, Res1: HiOpt);
6066 }
6067
6068 return SDValue();
6069}
6070
6071SDValue DAGCombiner::visitSMUL_LOHI(SDNode *N) {
6072 if (SDValue Res = SimplifyNodeWithTwoResults(N, LoOp: ISD::MUL, HiOp: ISD::MULHS))
6073 return Res;
6074
6075 SDValue N0 = N->getOperand(Num: 0);
6076 SDValue N1 = N->getOperand(Num: 1);
6077 EVT VT = N->getValueType(ResNo: 0);
6078 SDLoc DL(N);
6079
6080 // Constant fold.
6081 if (isa<ConstantSDNode>(Val: N0) && isa<ConstantSDNode>(Val: N1))
6082 return DAG.getNode(Opcode: ISD::SMUL_LOHI, DL, VTList: N->getVTList(), N1: N0, N2: N1);
6083
6084 // canonicalize constant to RHS (vector doesn't have to splat)
6085 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
6086 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
6087 return DAG.getNode(Opcode: ISD::SMUL_LOHI, DL, VTList: N->getVTList(), N1, N2: N0);
6088
6089 // If the type is twice as wide is legal, transform the mulhu to a wider
6090 // multiply plus a shift.
6091 if (VT.isSimple() && !VT.isVector()) {
6092 MVT Simple = VT.getSimpleVT();
6093 unsigned SimpleSize = Simple.getSizeInBits();
6094 EVT NewVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: SimpleSize*2);
6095 if (TLI.isOperationLegal(Op: ISD::MUL, VT: NewVT)) {
6096 SDValue Lo = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT: NewVT, Operand: N0);
6097 SDValue Hi = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT: NewVT, Operand: N1);
6098 Lo = DAG.getNode(Opcode: ISD::MUL, DL, VT: NewVT, N1: Lo, N2: Hi);
6099 // Compute the high part as N1.
6100 Hi = DAG.getNode(Opcode: ISD::SRL, DL, VT: NewVT, N1: Lo,
6101 N2: DAG.getShiftAmountConstant(Val: SimpleSize, VT: NewVT, DL));
6102 Hi = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: Hi);
6103 // Compute the low part as N0.
6104 Lo = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: Lo);
6105 return CombineTo(N, Res0: Lo, Res1: Hi);
6106 }
6107 }
6108
6109 return SDValue();
6110}
6111
6112SDValue DAGCombiner::visitUMUL_LOHI(SDNode *N) {
6113 if (SDValue Res = SimplifyNodeWithTwoResults(N, LoOp: ISD::MUL, HiOp: ISD::MULHU))
6114 return Res;
6115
6116 SDValue N0 = N->getOperand(Num: 0);
6117 SDValue N1 = N->getOperand(Num: 1);
6118 EVT VT = N->getValueType(ResNo: 0);
6119 SDLoc DL(N);
6120
6121 // Constant fold.
6122 if (isa<ConstantSDNode>(Val: N0) && isa<ConstantSDNode>(Val: N1))
6123 return DAG.getNode(Opcode: ISD::UMUL_LOHI, DL, VTList: N->getVTList(), N1: N0, N2: N1);
6124
6125 // canonicalize constant to RHS (vector doesn't have to splat)
6126 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
6127 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
6128 return DAG.getNode(Opcode: ISD::UMUL_LOHI, DL, VTList: N->getVTList(), N1, N2: N0);
6129
6130 // (umul_lohi N0, 0) -> (0, 0)
6131 if (isNullConstant(V: N1)) {
6132 SDValue Zero = DAG.getConstant(Val: 0, DL, VT);
6133 return CombineTo(N, Res0: Zero, Res1: Zero);
6134 }
6135
6136 // (umul_lohi N0, 1) -> (N0, 0)
6137 if (isOneConstant(V: N1)) {
6138 SDValue Zero = DAG.getConstant(Val: 0, DL, VT);
6139 return CombineTo(N, Res0: N0, Res1: Zero);
6140 }
6141
6142 // If the type is twice as wide is legal, transform the mulhu to a wider
6143 // multiply plus a shift.
6144 if (VT.isSimple() && !VT.isVector()) {
6145 MVT Simple = VT.getSimpleVT();
6146 unsigned SimpleSize = Simple.getSizeInBits();
6147 EVT NewVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: SimpleSize*2);
6148 if (TLI.isOperationLegal(Op: ISD::MUL, VT: NewVT)) {
6149 SDValue Lo = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: NewVT, Operand: N0);
6150 SDValue Hi = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: NewVT, Operand: N1);
6151 Lo = DAG.getNode(Opcode: ISD::MUL, DL, VT: NewVT, N1: Lo, N2: Hi);
6152 // Compute the high part as N1.
6153 Hi = DAG.getNode(Opcode: ISD::SRL, DL, VT: NewVT, N1: Lo,
6154 N2: DAG.getShiftAmountConstant(Val: SimpleSize, VT: NewVT, DL));
6155 Hi = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: Hi);
6156 // Compute the low part as N0.
6157 Lo = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: Lo);
6158 return CombineTo(N, Res0: Lo, Res1: Hi);
6159 }
6160 }
6161
6162 return SDValue();
6163}
6164
6165SDValue DAGCombiner::visitMULO(SDNode *N) {
6166 SDValue N0 = N->getOperand(Num: 0);
6167 SDValue N1 = N->getOperand(Num: 1);
6168 EVT VT = N0.getValueType();
6169 bool IsSigned = (ISD::SMULO == N->getOpcode());
6170
6171 EVT CarryVT = N->getValueType(ResNo: 1);
6172 SDLoc DL(N);
6173
6174 ConstantSDNode *N0C = isConstOrConstSplat(N: N0);
6175 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
6176
6177 // fold operation with constant operands.
6178 // TODO: Move this to FoldConstantArithmetic when it supports nodes with
6179 // multiple results.
6180 if (N0C && N1C && !N0C->isOpaque() && !N1C->isOpaque()) {
6181 bool Overflow;
6182 APInt Result =
6183 IsSigned ? N0C->getAPIntValue().smul_ov(RHS: N1C->getAPIntValue(), Overflow)
6184 : N0C->getAPIntValue().umul_ov(RHS: N1C->getAPIntValue(), Overflow);
6185 return CombineTo(N, Res0: DAG.getConstant(Val: Result, DL, VT),
6186 Res1: DAG.getBoolConstant(V: Overflow, DL, VT: CarryVT, OpVT: CarryVT));
6187 }
6188
6189 // canonicalize constant to RHS.
6190 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
6191 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
6192 return DAG.getNode(Opcode: N->getOpcode(), DL, VTList: N->getVTList(), N1, N2: N0);
6193
6194 // fold (mulo x, 0) -> 0 + no carry out
6195 if (isNullOrNullSplat(V: N1))
6196 return CombineTo(N, Res0: DAG.getConstant(Val: 0, DL, VT),
6197 Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
6198
6199 // (mulo x, 2) -> (addo x, x)
6200 // FIXME: This needs a freeze.
6201 if (N1C && N1C->getAPIntValue() == 2 &&
6202 (!IsSigned || VT.getScalarSizeInBits() > 2))
6203 return DAG.getNode(Opcode: IsSigned ? ISD::SADDO : ISD::UADDO, DL,
6204 VTList: N->getVTList(), N1: N0, N2: N0);
6205
6206 // A 1 bit SMULO overflows if both inputs are 1.
6207 if (IsSigned && VT.getScalarSizeInBits() == 1) {
6208 SDValue And = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N0, N2: N1);
6209 SDValue Cmp = DAG.getSetCC(DL, VT: CarryVT, LHS: And,
6210 RHS: DAG.getConstant(Val: 0, DL, VT), Cond: ISD::SETNE);
6211 return CombineTo(N, Res0: And, Res1: Cmp);
6212 }
6213
6214 // If it cannot overflow, transform into a mul.
6215 if (DAG.willNotOverflowMul(IsSigned, N0, N1))
6216 return CombineTo(N, Res0: DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: N0, N2: N1),
6217 Res1: DAG.getConstant(Val: 0, DL, VT: CarryVT));
6218 return SDValue();
6219}
6220
6221// Function to calculate whether the Min/Max pair of SDNodes (potentially
6222// swapped around) make a signed saturate pattern, clamping to between a signed
6223// saturate of -2^(BW-1) and 2^(BW-1)-1, or an unsigned saturate of 0 and 2^BW.
6224// Returns the node being clamped and the bitwidth of the clamp in BW. Should
6225// work with both SMIN/SMAX nodes and setcc/select combo. The operands are the
6226// same as SimplifySelectCC. N0<N1 ? N2 : N3.
6227static SDValue isSaturatingMinMax(SDValue N0, SDValue N1, SDValue N2,
6228 SDValue N3, ISD::CondCode CC, unsigned &BW,
6229 bool &Unsigned, SelectionDAG &DAG) {
6230 auto isSignedMinMax = [&](SDValue N0, SDValue N1, SDValue N2, SDValue N3,
6231 ISD::CondCode CC) {
6232 // The compare and select operand should be the same or the select operands
6233 // should be truncated versions of the comparison.
6234 if (N0 != N2 && (N2.getOpcode() != ISD::TRUNCATE || N0 != N2.getOperand(i: 0)))
6235 return 0;
6236 // The constants need to be the same or a truncated version of each other.
6237 ConstantSDNode *N1C = isConstOrConstSplat(N: peekThroughTruncates(V: N1));
6238 ConstantSDNode *N3C = isConstOrConstSplat(N: peekThroughTruncates(V: N3));
6239 if (!N1C || !N3C)
6240 return 0;
6241 const APInt &C1 = N1C->getAPIntValue().trunc(width: N1.getScalarValueSizeInBits());
6242 const APInt &C2 = N3C->getAPIntValue().trunc(width: N3.getScalarValueSizeInBits());
6243 if (C1.getBitWidth() < C2.getBitWidth() || C1 != C2.sext(width: C1.getBitWidth()))
6244 return 0;
6245 return CC == ISD::SETLT ? ISD::SMIN : (CC == ISD::SETGT ? ISD::SMAX : 0);
6246 };
6247
6248 // Check the initial value is a SMIN/SMAX equivalent.
6249 unsigned Opcode0 = isSignedMinMax(N0, N1, N2, N3, CC);
6250 if (!Opcode0)
6251 return SDValue();
6252
6253 // We could only need one range check, if the fptosi could never produce
6254 // the upper value.
6255 if (N0.getOpcode() == ISD::FP_TO_SINT && Opcode0 == ISD::SMAX) {
6256 if (isNullOrNullSplat(V: N3)) {
6257 EVT IntVT = N0.getValueType().getScalarType();
6258 EVT FPVT = N0.getOperand(i: 0).getValueType().getScalarType();
6259 if (FPVT.isSimple()) {
6260 Type *InputTy = FPVT.getTypeForEVT(Context&: *DAG.getContext());
6261 const fltSemantics &Semantics = InputTy->getFltSemantics();
6262 uint32_t MinBitWidth =
6263 APFloatBase::semanticsIntSizeInBits(Semantics, /*isSigned*/ true);
6264 if (IntVT.getSizeInBits() >= MinBitWidth) {
6265 Unsigned = true;
6266 BW = PowerOf2Ceil(A: MinBitWidth);
6267 return N0;
6268 }
6269 }
6270 }
6271 }
6272
6273 SDValue N00, N01, N02, N03;
6274 ISD::CondCode N0CC;
6275 switch (N0.getOpcode()) {
6276 case ISD::SMIN:
6277 case ISD::SMAX:
6278 N00 = N02 = N0.getOperand(i: 0);
6279 N01 = N03 = N0.getOperand(i: 1);
6280 N0CC = N0.getOpcode() == ISD::SMIN ? ISD::SETLT : ISD::SETGT;
6281 break;
6282 case ISD::SELECT_CC:
6283 N00 = N0.getOperand(i: 0);
6284 N01 = N0.getOperand(i: 1);
6285 N02 = N0.getOperand(i: 2);
6286 N03 = N0.getOperand(i: 3);
6287 N0CC = cast<CondCodeSDNode>(Val: N0.getOperand(i: 4))->get();
6288 break;
6289 case ISD::SELECT:
6290 case ISD::VSELECT:
6291 if (N0.getOperand(i: 0).getOpcode() != ISD::SETCC)
6292 return SDValue();
6293 N00 = N0.getOperand(i: 0).getOperand(i: 0);
6294 N01 = N0.getOperand(i: 0).getOperand(i: 1);
6295 N02 = N0.getOperand(i: 1);
6296 N03 = N0.getOperand(i: 2);
6297 N0CC = cast<CondCodeSDNode>(Val: N0.getOperand(i: 0).getOperand(i: 2))->get();
6298 break;
6299 default:
6300 return SDValue();
6301 }
6302
6303 unsigned Opcode1 = isSignedMinMax(N00, N01, N02, N03, N0CC);
6304 if (!Opcode1 || Opcode0 == Opcode1)
6305 return SDValue();
6306
6307 ConstantSDNode *MinCOp = isConstOrConstSplat(N: Opcode0 == ISD::SMIN ? N1 : N01);
6308 ConstantSDNode *MaxCOp = isConstOrConstSplat(N: Opcode0 == ISD::SMIN ? N01 : N1);
6309 if (!MinCOp || !MaxCOp || MinCOp->getValueType(ResNo: 0) != MaxCOp->getValueType(ResNo: 0))
6310 return SDValue();
6311
6312 const APInt &MinC = MinCOp->getAPIntValue();
6313 const APInt &MaxC = MaxCOp->getAPIntValue();
6314 APInt MinCPlus1 = MinC + 1;
6315 if (-MaxC == MinCPlus1 && MinCPlus1.isPowerOf2()) {
6316 BW = MinCPlus1.exactLogBase2() + 1;
6317 Unsigned = false;
6318 return N02;
6319 }
6320
6321 if (MaxC == 0 && MinC != 0 && MinCPlus1.isPowerOf2()) {
6322 BW = MinCPlus1.exactLogBase2();
6323 Unsigned = true;
6324 return N02;
6325 }
6326
6327 return SDValue();
6328}
6329
6330static SDValue PerformMinMaxFpToSatCombine(SDValue N0, SDValue N1, SDValue N2,
6331 SDValue N3, ISD::CondCode CC,
6332 SelectionDAG &DAG) {
6333 unsigned BW;
6334 bool Unsigned;
6335 SDValue Fp = isSaturatingMinMax(N0, N1, N2, N3, CC, BW, Unsigned, DAG);
6336 if (!Fp || Fp.getOpcode() != ISD::FP_TO_SINT)
6337 return SDValue();
6338 EVT FPVT = Fp.getOperand(i: 0).getValueType();
6339 EVT NewVT = FPVT.changeElementType(Context&: *DAG.getContext(),
6340 EltVT: EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: BW));
6341 unsigned NewOpc = Unsigned ? ISD::FP_TO_UINT_SAT : ISD::FP_TO_SINT_SAT;
6342 if (!DAG.getTargetLoweringInfo().shouldConvertFpToSat(Op: NewOpc, FPVT, VT: NewVT))
6343 return SDValue();
6344 SDLoc DL(Fp);
6345 SDValue Sat = DAG.getNode(Opcode: NewOpc, DL, VT: NewVT, N1: Fp.getOperand(i: 0),
6346 N2: DAG.getValueType(NewVT.getScalarType()));
6347 return DAG.getExtOrTrunc(IsSigned: !Unsigned, Op: Sat, DL, VT: N2->getValueType(ResNo: 0));
6348}
6349
6350static SDValue PerformUMinFpToSatCombine(SDValue N0, SDValue N1, SDValue N2,
6351 SDValue N3, ISD::CondCode CC,
6352 SelectionDAG &DAG) {
6353 // We are looking for UMIN(FPTOUI(X), (2^n)-1), which may have come via a
6354 // select/vselect/select_cc. The two operands pairs for the select (N2/N3) may
6355 // be truncated versions of the setcc (N0/N1).
6356 if ((N0 != N2 &&
6357 (N2.getOpcode() != ISD::TRUNCATE || N0 != N2.getOperand(i: 0))) ||
6358 N0.getOpcode() != ISD::FP_TO_UINT || CC != ISD::SETULT)
6359 return SDValue();
6360 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
6361 ConstantSDNode *N3C = isConstOrConstSplat(N: N3);
6362 if (!N1C || !N3C)
6363 return SDValue();
6364 const APInt &C1 = N1C->getAPIntValue();
6365 const APInt &C3 = N3C->getAPIntValue();
6366 if (!(C1 + 1).isPowerOf2() || C1.getBitWidth() < C3.getBitWidth() ||
6367 C1 != C3.zext(width: C1.getBitWidth()))
6368 return SDValue();
6369
6370 unsigned BW = (C1 + 1).exactLogBase2();
6371 EVT FPVT = N0.getOperand(i: 0).getValueType();
6372 EVT NewVT = FPVT.changeElementType(Context&: *DAG.getContext(),
6373 EltVT: EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: BW));
6374 if (!DAG.getTargetLoweringInfo().shouldConvertFpToSat(Op: ISD::FP_TO_UINT_SAT,
6375 FPVT, VT: NewVT))
6376 return SDValue();
6377
6378 SDValue Sat =
6379 DAG.getNode(Opcode: ISD::FP_TO_UINT_SAT, DL: SDLoc(N0), VT: NewVT, N1: N0.getOperand(i: 0),
6380 N2: DAG.getValueType(NewVT.getScalarType()));
6381 return DAG.getZExtOrTrunc(Op: Sat, DL: SDLoc(N0), VT: N3.getValueType());
6382}
6383
6384// Fold a NaN-guard select of fp_to_sint/fp_to_uint into the saturating
6385// variant, which returns 0 for NaN.
6386static SDValue performNanGuardFpToSatCombine(SDNode *N, SelectionDAG &DAG) {
6387 EVT VT = N->getValueType(ResNo: 0);
6388 SDLoc DL(N);
6389
6390 // Match an isnan-guarded select, requiring the compare to be single-use.
6391 // The guarded value is fp_to_sint/fp_to_uint of X, optionally masked by an
6392 // AND:
6393 // select (setcc X, 0.0, uno), 0, (fp_to_sint/uint X)
6394 // select (setcc X, 0.0, ord), (fp_to_sint/uint X), 0
6395 // select (setcc X, 0.0, uno), 0, (and (fp_to_sint/uint X), M)
6396 // select (setcc X, 0.0, ord), (and (fp_to_sint/uint X), M), 0
6397 SDValue X, GuardedVal;
6398 if (!sd_match(N, P: m_SelectLike(Cond: m_OneUse(P: m_SpecificSetCC(CC: ISD::SETUO, LHS: m_Value(N&: X),
6399 RHS: m_AnyZeroFP())),
6400 T: m_Zero(), F: m_Value(N&: GuardedVal))) &&
6401 !sd_match(N, P: m_SelectLike(Cond: m_OneUse(P: m_SpecificSetCC(CC: ISD::SETO, LHS: m_Value(N&: X),
6402 RHS: m_AnyZeroFP())),
6403 T: m_Value(N&: GuardedVal), F: m_Zero())))
6404 return SDValue();
6405
6406 // The guarded value must be fp_to_sint/fp_to_uint of the same X, optionally
6407 // masked by a (commutative) AND.
6408 SDValue Mask;
6409 unsigned NewOpc;
6410 if (sd_match(N: GuardedVal, P: m_FPToSI(Op: m_Specific(N: X))) ||
6411 sd_match(N: GuardedVal, P: m_And(L: m_FPToSI(Op: m_Specific(N: X)), R: m_Value(N&: Mask))))
6412 NewOpc = ISD::FP_TO_SINT_SAT;
6413 else if (sd_match(N: GuardedVal, P: m_FPToUI(Op: m_Specific(N: X))) ||
6414 sd_match(N: GuardedVal, P: m_And(L: m_FPToUI(Op: m_Specific(N: X)), R: m_Value(N&: Mask))))
6415 NewOpc = ISD::FP_TO_UINT_SAT;
6416 else
6417 return SDValue();
6418
6419 if (!DAG.getTargetLoweringInfo().shouldConvertFpToSat(Op: NewOpc,
6420 FPVT: X.getValueType(), VT))
6421 return SDValue();
6422
6423 SDValue Sat =
6424 DAG.getNode(Opcode: NewOpc, DL, VT, N1: X, N2: DAG.getValueType(VT.getScalarType()));
6425 if (Mask) {
6426 // For NaN inputs the saturating conversion yields 0, so (and 0, Mask) must
6427 // stay 0 to match the original select. A poison Mask would make it poison,
6428 // so freeze Mask to guarantee a defined value.
6429 Sat = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Sat, N2: DAG.getFreeze(V: Mask));
6430 }
6431 return Sat;
6432}
6433
6434SDValue DAGCombiner::visitIMINMAX(SDNode *N) {
6435 SDValue N0 = N->getOperand(Num: 0);
6436 SDValue N1 = N->getOperand(Num: 1);
6437 EVT VT = N0.getValueType();
6438 unsigned Opcode = N->getOpcode();
6439 SDLoc DL(N);
6440
6441 // fold operation with constant operands.
6442 if (SDValue C = DAG.FoldConstantArithmetic(Opcode, DL, VT, Ops: {N0, N1}))
6443 return C;
6444
6445 // If the operands are the same, this is a no-op.
6446 if (N0 == N1)
6447 return N0;
6448
6449 // canonicalize constant to RHS
6450 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
6451 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
6452 return DAG.getNode(Opcode, DL, VT, N1, N2: N0);
6453
6454 // fold vector ops
6455 if (VT.isVector())
6456 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
6457 return FoldedVOp;
6458
6459 // reassociate minmax
6460 if (SDValue RMINMAX = reassociateOps(Opc: Opcode, DL, N0, N1, Flags: N->getFlags()))
6461 return RMINMAX;
6462
6463 // Fold sign-extension masks using arithmetic shift:
6464 // smax(X, -1) -> or(X, ashr(X, BW-1))
6465 // smin(X, 0) -> and(X, ashr(X, BW-1))
6466 // ashr(X, BW-1) sign-extends the sign bit: 0 for X>=0, -1 for X<0.
6467 // OR with X yields X (non-negative) or -1 (negative) = smax(X,-1).
6468 // AND with X yields 0 (non-negative) or X (negative) = smin(X, 0).
6469 // Both reduce to two instructions vs. a compare+cmov on x86-64.
6470 // Only fold when the target has no native SMAX/SMIN instruction for this
6471 // type (isOperationExpand), the type is legal (not needing splitting),
6472 // the operand is not a min/max chain (preserving target combine patterns
6473 // that fold smax(smin(x,C),D) into a single saturation instruction), and
6474 // for smax(X,-1) the operand is not a sign extension (doubling its use
6475 // count can cause the target to lower the extension less efficiently).
6476 APInt C;
6477 if (TLI.isTypeLegal(VT) &&
6478 !TLI.shouldAvoidTransformToShift(VT, Amount: VT.getScalarSizeInBits() - 1) &&
6479 sd_match(N: N1, P: m_ConstInt(V&: C))) {
6480 if (Opcode == ISD::SMAX && TLI.isOperationExpand(Op: ISD::SMAX, VT) &&
6481 N0.getOpcode() != ISD::SMIN && N0.getOpcode() != ISD::SIGN_EXTEND &&
6482 C.isAllOnes()) {
6483 SDValue ShiftAmt =
6484 DAG.getShiftAmountConstant(Val: VT.getScalarSizeInBits() - 1, VT, DL);
6485 SDValue Shift = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N0, N2: ShiftAmt);
6486 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: N0, N2: Shift);
6487 }
6488 if (Opcode == ISD::SMIN && TLI.isOperationExpand(Op: ISD::SMIN, VT) &&
6489 N0.getOpcode() != ISD::SMAX && C.isZero()) {
6490 SDValue ShiftAmt =
6491 DAG.getShiftAmountConstant(Val: VT.getScalarSizeInBits() - 1, VT, DL);
6492 SDValue Shift = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N0, N2: ShiftAmt);
6493 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N0, N2: Shift);
6494 }
6495 }
6496
6497 // If both operands are known to have the same sign (both non-negative or both
6498 // negative), flip between UMIN/UMAX and SMIN/SMAX.
6499 // Only do this if:
6500 // 1. The current op isn't legal and the flipped is.
6501 // 2. The saturation pattern is broken by canonicalization in InstCombine.
6502 bool IsOpIllegal = !TLI.isOperationLegal(Op: Opcode, VT);
6503 bool IsSatBroken = Opcode == ISD::UMIN && N0.getOpcode() == ISD::SMAX;
6504
6505 if (IsSatBroken || IsOpIllegal) {
6506 auto HasKnownSameSign = [&](SDValue A, SDValue B) {
6507 if (A.isUndef() || B.isUndef())
6508 return true;
6509
6510 KnownBits KA = DAG.computeKnownBits(Op: A);
6511 if (!KA.isNonNegative() && !KA.isNegative())
6512 return false;
6513
6514 KnownBits KB = DAG.computeKnownBits(Op: B);
6515 if (KA.isNonNegative())
6516 return KB.isNonNegative();
6517 return KB.isNegative();
6518 };
6519
6520 if (HasKnownSameSign(N0, N1)) {
6521 unsigned AltOpcode = ISD::getOppositeSignednessMinMaxOpcode(MinMaxOpc: Opcode);
6522 if ((IsSatBroken && IsOpIllegal) || TLI.isOperationLegal(Op: AltOpcode, VT))
6523 return DAG.getNode(Opcode: AltOpcode, DL, VT, N1: N0, N2: N1);
6524 }
6525 }
6526
6527 if (Opcode == ISD::SMIN || Opcode == ISD::SMAX)
6528 if (SDValue S = PerformMinMaxFpToSatCombine(
6529 N0, N1, N2: N0, N3: N1, CC: Opcode == ISD::SMIN ? ISD::SETLT : ISD::SETGT, DAG))
6530 return S;
6531 if (Opcode == ISD::UMIN)
6532 if (SDValue S = PerformUMinFpToSatCombine(N0, N1, N2: N0, N3: N1, CC: ISD::SETULT, DAG))
6533 return S;
6534
6535 // Fold min/max(vecreduce(x), vecreduce(y)) -> vecreduce(min/max(x, y))
6536 auto ReductionOpcode = [](unsigned Opcode) {
6537 switch (Opcode) {
6538 case ISD::SMIN:
6539 return ISD::VECREDUCE_SMIN;
6540 case ISD::SMAX:
6541 return ISD::VECREDUCE_SMAX;
6542 case ISD::UMIN:
6543 return ISD::VECREDUCE_UMIN;
6544 case ISD::UMAX:
6545 return ISD::VECREDUCE_UMAX;
6546 default:
6547 llvm_unreachable("Unexpected opcode");
6548 }
6549 };
6550 if (SDValue SD = reassociateReduction(RedOpc: ReductionOpcode(Opcode), Opc: Opcode,
6551 DL: SDLoc(N), VT, N0, N1))
6552 return SD;
6553
6554 // Fold operation with vscale operands.
6555 if (N0.getOpcode() == ISD::VSCALE && N1.getOpcode() == ISD::VSCALE) {
6556 uint64_t C0 = N0->getConstantOperandVal(Num: 0);
6557 uint64_t C1 = N1->getConstantOperandVal(Num: 0);
6558 if (Opcode == ISD::UMAX)
6559 return C0 > C1 ? N0 : N1;
6560 else if (Opcode == ISD::UMIN)
6561 return C0 > C1 ? N1 : N0;
6562 }
6563
6564 // If we know the range of vscale, see if we can fold it given a constant.
6565 if (N0.getOpcode() == ISD::VSCALE) {
6566 if (auto *C1 = dyn_cast<ConstantSDNode>(Val&: N1)) {
6567 bool ForSigned = (Opcode == ISD::SMAX || Opcode == ISD::SMIN);
6568 ConstantRange Range = DAG.computeConstantRange(Op: N0, ForSigned);
6569
6570 const APInt &C1V = C1->getAPIntValue();
6571 if ((Opcode == ISD::UMAX && Range.getUnsignedMax().ule(RHS: C1V)) ||
6572 (Opcode == ISD::UMIN && Range.getUnsignedMin().uge(RHS: C1V)) ||
6573 (Opcode == ISD::SMAX && Range.getSignedMax().sle(RHS: C1V)) ||
6574 (Opcode == ISD::SMIN && Range.getSignedMin().sge(RHS: C1V))) {
6575 return N1;
6576 }
6577 }
6578 }
6579
6580 // Simplify the operands using demanded-bits information.
6581 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
6582 return SDValue(N, 0);
6583
6584 return SDValue();
6585}
6586
6587/// If this is a bitwise logic instruction and both operands have the same
6588/// opcode, try to sink the other opcode after the logic instruction.
6589SDValue DAGCombiner::hoistLogicOpWithSameOpcodeHands(SDNode *N) {
6590 SDValue N0 = N->getOperand(Num: 0), N1 = N->getOperand(Num: 1);
6591 EVT VT = N0.getValueType();
6592 unsigned LogicOpcode = N->getOpcode();
6593 unsigned HandOpcode = N0.getOpcode();
6594 assert(ISD::isBitwiseLogicOp(LogicOpcode) && "Expected logic opcode");
6595 assert(HandOpcode == N1.getOpcode() && "Bad input!");
6596
6597 // Bail early if none of these transforms apply.
6598 if (N0.getNumOperands() == 0)
6599 return SDValue();
6600
6601 // FIXME: We should check number of uses of the operands to not increase
6602 // the instruction count for all transforms.
6603
6604 // Handle size-changing casts (or sign_extend_inreg).
6605 SDValue X = N0.getOperand(i: 0);
6606 SDValue Y = N1.getOperand(i: 0);
6607 EVT XVT = X.getValueType();
6608 SDLoc DL(N);
6609 if (ISD::isExtOpcode(Opcode: HandOpcode) || ISD::isExtVecInRegOpcode(Opcode: HandOpcode) ||
6610 (HandOpcode == ISD::SIGN_EXTEND_INREG &&
6611 N0.getOperand(i: 1) == N1.getOperand(i: 1))) {
6612 // If both operands have other uses, this transform would create extra
6613 // instructions without eliminating anything.
6614 if (!N0.hasOneUse() && !N1.hasOneUse())
6615 return SDValue();
6616 // We need matching integer source types.
6617 if (XVT != Y.getValueType())
6618 return SDValue();
6619 // Don't create an illegal op during or after legalization. Don't ever
6620 // create an unsupported vector op.
6621 if ((VT.isVector() || LegalOperations) &&
6622 !TLI.isOperationLegalOrCustom(Op: LogicOpcode, VT: XVT))
6623 return SDValue();
6624 // Avoid infinite looping with PromoteIntBinOp.
6625 // TODO: Should we apply desirable/legal constraints to all opcodes?
6626 if ((HandOpcode == ISD::ANY_EXTEND ||
6627 HandOpcode == ISD::ANY_EXTEND_VECTOR_INREG) &&
6628 LegalTypes && !TLI.isTypeDesirableForOp(LogicOpcode, VT: XVT))
6629 return SDValue();
6630 // logic_op (hand_op X), (hand_op Y) --> hand_op (logic_op X, Y)
6631 SDNodeFlags LogicFlags;
6632 LogicFlags.setDisjoint(N->getFlags().hasDisjoint() &&
6633 ISD::isExtOpcode(Opcode: HandOpcode));
6634 SDValue Logic = DAG.getNode(Opcode: LogicOpcode, DL, VT: XVT, N1: X, N2: Y, Flags: LogicFlags);
6635 if (HandOpcode == ISD::SIGN_EXTEND_INREG)
6636 return DAG.getNode(Opcode: HandOpcode, DL, VT, N1: Logic, N2: N0.getOperand(i: 1));
6637 return DAG.getNode(Opcode: HandOpcode, DL, VT, Operand: Logic);
6638 }
6639
6640 // logic_op (truncate x), (truncate y) --> truncate (logic_op x, y)
6641 if (HandOpcode == ISD::TRUNCATE) {
6642 // If both operands have other uses, this transform would create extra
6643 // instructions without eliminating anything.
6644 if (!N0.hasOneUse() && !N1.hasOneUse())
6645 return SDValue();
6646 // We need matching source types.
6647 if (XVT != Y.getValueType())
6648 return SDValue();
6649 // Don't create an illegal op during or after legalization.
6650 if (LegalOperations && !TLI.isOperationLegal(Op: LogicOpcode, VT: XVT))
6651 return SDValue();
6652 // Be extra careful sinking truncate. If it's free, there's no benefit in
6653 // widening a binop. Also, don't create a logic op on an illegal type.
6654 if (TLI.isZExtFree(FromTy: VT, ToTy: XVT) && TLI.isTruncateFree(FromVT: XVT, ToVT: VT))
6655 return SDValue();
6656 if (!TLI.isTypeLegal(VT: XVT))
6657 return SDValue();
6658 SDValue Logic = DAG.getNode(Opcode: LogicOpcode, DL, VT: XVT, N1: X, N2: Y);
6659 return DAG.getNode(Opcode: HandOpcode, DL, VT, Operand: Logic);
6660 }
6661
6662 // For binops SHL/SRL/SRA/AND:
6663 // logic_op (OP x, z), (OP y, z) --> OP (logic_op x, y), z
6664 if ((HandOpcode == ISD::SHL || HandOpcode == ISD::SRL ||
6665 HandOpcode == ISD::SRA || HandOpcode == ISD::AND) &&
6666 N0.getOperand(i: 1) == N1.getOperand(i: 1)) {
6667 // If either operand has other uses, this transform is not an improvement.
6668 if (!N0.hasOneUse() || !N1.hasOneUse())
6669 return SDValue();
6670 SDValue Logic = DAG.getNode(Opcode: LogicOpcode, DL, VT: XVT, N1: X, N2: Y);
6671 return DAG.getNode(Opcode: HandOpcode, DL, VT, N1: Logic, N2: N0.getOperand(i: 1));
6672 }
6673
6674 // Unary ops: logic_op (bswap x), (bswap y) --> bswap (logic_op x, y)
6675 if (HandOpcode == ISD::BSWAP) {
6676 // If either operand has other uses, this transform is not an improvement.
6677 if (!N0.hasOneUse() || !N1.hasOneUse())
6678 return SDValue();
6679 SDValue Logic = DAG.getNode(Opcode: LogicOpcode, DL, VT: XVT, N1: X, N2: Y);
6680 return DAG.getNode(Opcode: HandOpcode, DL, VT, Operand: Logic);
6681 }
6682
6683 // For funnel shifts FSHL/FSHR:
6684 // logic_op (OP x, x1, s), (OP y, y1, s) -->
6685 // --> OP (logic_op x, y), (logic_op, x1, y1), s
6686 if ((HandOpcode == ISD::FSHL || HandOpcode == ISD::FSHR) &&
6687 N0.getOperand(i: 2) == N1.getOperand(i: 2)) {
6688 if (!N0.hasOneUse() || !N1.hasOneUse())
6689 return SDValue();
6690 SDValue X1 = N0.getOperand(i: 1);
6691 SDValue Y1 = N1.getOperand(i: 1);
6692 SDValue S = N0.getOperand(i: 2);
6693 SDValue Logic0 = DAG.getNode(Opcode: LogicOpcode, DL, VT, N1: X, N2: Y);
6694 SDValue Logic1 = DAG.getNode(Opcode: LogicOpcode, DL, VT, N1: X1, N2: Y1);
6695 return DAG.getNode(Opcode: HandOpcode, DL, VT, N1: Logic0, N2: Logic1, N3: S);
6696 }
6697
6698 // Simplify xor/and/or (bitcast(A), bitcast(B)) -> bitcast(op (A,B))
6699 // Only perform this optimization up until type legalization, before
6700 // LegalizeVectorOprs. LegalizeVectorOprs promotes vector operations by
6701 // adding bitcasts. For example (xor v4i32) is promoted to (v2i64), and
6702 // we don't want to undo this promotion.
6703 // We also handle SCALAR_TO_VECTOR because xor/or/and operations are cheaper
6704 // on scalars.
6705 if ((HandOpcode == ISD::BITCAST || HandOpcode == ISD::SCALAR_TO_VECTOR) &&
6706 Level <= AfterLegalizeTypes) {
6707 // Input types must be integer and the same.
6708 if (XVT.isInteger() && XVT == Y.getValueType() &&
6709 !(VT.isVector() && TLI.isTypeLegal(VT) &&
6710 !XVT.isVector() && !TLI.isTypeLegal(VT: XVT))) {
6711 SDValue Logic = DAG.getNode(Opcode: LogicOpcode, DL, VT: XVT, N1: X, N2: Y);
6712 return DAG.getNode(Opcode: HandOpcode, DL, VT, Operand: Logic);
6713 }
6714 }
6715
6716 // Xor/and/or are indifferent to the swizzle operation (shuffle of one value).
6717 // Simplify xor/and/or (shuff(A), shuff(B)) -> shuff(op (A,B))
6718 // If both shuffles use the same mask, and both shuffle within a single
6719 // vector, then it is worthwhile to move the swizzle after the operation.
6720 // The type-legalizer generates this pattern when loading illegal
6721 // vector types from memory. In many cases this allows additional shuffle
6722 // optimizations.
6723 // There are other cases where moving the shuffle after the xor/and/or
6724 // is profitable even if shuffles don't perform a swizzle.
6725 // If both shuffles use the same mask, and both shuffles have the same first
6726 // or second operand, then it might still be profitable to move the shuffle
6727 // after the xor/and/or operation.
6728 if (HandOpcode == ISD::VECTOR_SHUFFLE && Level < AfterLegalizeDAG) {
6729 auto *SVN0 = cast<ShuffleVectorSDNode>(Val&: N0);
6730 auto *SVN1 = cast<ShuffleVectorSDNode>(Val&: N1);
6731 assert(X.getValueType() == Y.getValueType() &&
6732 "Inputs to shuffles are not the same type");
6733
6734 // Check that both shuffles use the same mask. The masks are known to be of
6735 // the same length because the result vector type is the same.
6736 // Check also that shuffles have only one use to avoid introducing extra
6737 // instructions.
6738 if (!SVN0->hasOneUse() || !SVN1->hasOneUse() ||
6739 !SVN0->getMask().equals(RHS: SVN1->getMask()))
6740 return SDValue();
6741
6742 // Don't try to fold this node if it requires introducing a
6743 // build vector of all zeros that might be illegal at this stage.
6744 SDValue ShOp = N0.getOperand(i: 1);
6745 if (LogicOpcode == ISD::XOR && !ShOp.isUndef())
6746 ShOp = tryFoldToZero(DL, TLI, VT, DAG, LegalOperations);
6747
6748 // (logic_op (shuf (A, C), shuf (B, C))) --> shuf (logic_op (A, B), C)
6749 if (N0.getOperand(i: 1) == N1.getOperand(i: 1) && ShOp.getNode()) {
6750 SDValue Logic = DAG.getNode(Opcode: LogicOpcode, DL, VT,
6751 N1: N0.getOperand(i: 0), N2: N1.getOperand(i: 0));
6752 return DAG.getVectorShuffle(VT, dl: DL, N1: Logic, N2: ShOp, Mask: SVN0->getMask());
6753 }
6754
6755 // Don't try to fold this node if it requires introducing a
6756 // build vector of all zeros that might be illegal at this stage.
6757 ShOp = N0.getOperand(i: 0);
6758 if (LogicOpcode == ISD::XOR && !ShOp.isUndef())
6759 ShOp = tryFoldToZero(DL, TLI, VT, DAG, LegalOperations);
6760
6761 // (logic_op (shuf (C, A), shuf (C, B))) --> shuf (C, logic_op (A, B))
6762 if (N0.getOperand(i: 0) == N1.getOperand(i: 0) && ShOp.getNode()) {
6763 SDValue Logic = DAG.getNode(Opcode: LogicOpcode, DL, VT, N1: N0.getOperand(i: 1),
6764 N2: N1.getOperand(i: 1));
6765 return DAG.getVectorShuffle(VT, dl: DL, N1: ShOp, N2: Logic, Mask: SVN0->getMask());
6766 }
6767 }
6768
6769 return SDValue();
6770}
6771
6772/// Try to make (and/or setcc (LL, LR), setcc (RL, RR)) more efficient.
6773SDValue DAGCombiner::foldLogicOfSetCCs(bool IsAnd, SDValue N0, SDValue N1,
6774 const SDLoc &DL) {
6775 SDValue LL, LR, RL, RR, N0CC, N1CC;
6776 if (!isSetCCEquivalent(N: N0, LHS&: LL, RHS&: LR, CC&: N0CC) ||
6777 !isSetCCEquivalent(N: N1, LHS&: RL, RHS&: RR, CC&: N1CC))
6778 return SDValue();
6779
6780 assert(N0.getValueType() == N1.getValueType() &&
6781 "Unexpected operand types for bitwise logic op");
6782 assert(LL.getValueType() == LR.getValueType() &&
6783 RL.getValueType() == RR.getValueType() &&
6784 "Unexpected operand types for setcc");
6785
6786 // If we're here post-legalization or the logic op type is not i1, the logic
6787 // op type must match a setcc result type. Also, all folds require new
6788 // operations on the left and right operands, so those types must match.
6789 EVT VT = N0.getValueType();
6790 EVT OpVT = LL.getValueType();
6791 if (LegalOperations || VT.getScalarType() != MVT::i1)
6792 if (VT != getSetCCResultType(VT: OpVT))
6793 return SDValue();
6794 if (OpVT != RL.getValueType())
6795 return SDValue();
6796
6797 ISD::CondCode CC0 = cast<CondCodeSDNode>(Val&: N0CC)->get();
6798 ISD::CondCode CC1 = cast<CondCodeSDNode>(Val&: N1CC)->get();
6799 bool IsInteger = OpVT.isInteger();
6800 if (LR == RR && CC0 == CC1 && IsInteger) {
6801 bool IsZero = isNullOrNullSplat(V: LR);
6802 bool IsNeg1 = isAllOnesOrAllOnesSplat(V: LR);
6803
6804 // All bits clear?
6805 bool AndEqZero = IsAnd && CC1 == ISD::SETEQ && IsZero;
6806 // All sign bits clear?
6807 bool AndGtNeg1 = IsAnd && CC1 == ISD::SETGT && IsNeg1;
6808 // Any bits set?
6809 bool OrNeZero = !IsAnd && CC1 == ISD::SETNE && IsZero;
6810 // Any sign bits set?
6811 bool OrLtZero = !IsAnd && CC1 == ISD::SETLT && IsZero;
6812
6813 // (and (seteq X, 0), (seteq Y, 0)) --> (seteq (or X, Y), 0)
6814 // (and (setgt X, -1), (setgt Y, -1)) --> (setgt (or X, Y), -1)
6815 // (or (setne X, 0), (setne Y, 0)) --> (setne (or X, Y), 0)
6816 // (or (setlt X, 0), (setlt Y, 0)) --> (setlt (or X, Y), 0)
6817 if (AndEqZero || AndGtNeg1 || OrNeZero || OrLtZero) {
6818 SDValue Or = DAG.getNode(Opcode: ISD::OR, DL: SDLoc(N0), VT: OpVT, N1: LL, N2: RL);
6819 AddToWorklist(N: Or.getNode());
6820 return DAG.getSetCC(DL, VT, LHS: Or, RHS: LR, Cond: CC1);
6821 }
6822
6823 // All bits set?
6824 bool AndEqNeg1 = IsAnd && CC1 == ISD::SETEQ && IsNeg1;
6825 // All sign bits set?
6826 bool AndLtZero = IsAnd && CC1 == ISD::SETLT && IsZero;
6827 // Any bits clear?
6828 bool OrNeNeg1 = !IsAnd && CC1 == ISD::SETNE && IsNeg1;
6829 // Any sign bits clear?
6830 bool OrGtNeg1 = !IsAnd && CC1 == ISD::SETGT && IsNeg1;
6831
6832 // (and (seteq X, -1), (seteq Y, -1)) --> (seteq (and X, Y), -1)
6833 // (and (setlt X, 0), (setlt Y, 0)) --> (setlt (and X, Y), 0)
6834 // (or (setne X, -1), (setne Y, -1)) --> (setne (and X, Y), -1)
6835 // (or (setgt X, -1), (setgt Y -1)) --> (setgt (and X, Y), -1)
6836 if (AndEqNeg1 || AndLtZero || OrNeNeg1 || OrGtNeg1) {
6837 SDValue And = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(N0), VT: OpVT, N1: LL, N2: RL);
6838 AddToWorklist(N: And.getNode());
6839 return DAG.getSetCC(DL, VT, LHS: And, RHS: LR, Cond: CC1);
6840 }
6841 }
6842
6843 // (and (setne (and X, LL1), 0), (setne (and X, RL1), 0))
6844 // --> (seteq (and X, (LL1|RL1)), (LL1|RL1))
6845 // (or (seteq (and X, LL1), 0), (seteq (and X, RL1), 0))
6846 // --> (setne (and X, (LL1|RL1)), (LL1|RL1))
6847 if (LL.getOpcode() == ISD::AND && RL.getOpcode() == ISD::AND &&
6848 isNullConstant(V: LR) && isNullConstant(V: RR) && CC0 == CC1 &&
6849 (CC0 == ISD::SETNE || CC0 == ISD::SETEQ)) {
6850 SDValue LL0, LL1, RL0, RL1;
6851 LL0 = LL.getOperand(i: 0);
6852 RL0 = RL.getOperand(i: 0);
6853 LL1 = LL.getOperand(i: 1);
6854 RL1 = RL.getOperand(i: 1);
6855 if (LL0 == RL0 && DAG.isKnownToBeAPowerOfTwo(Val: LL1) &&
6856 DAG.isKnownToBeAPowerOfTwo(Val: RL1)) {
6857 SDValue Or = DAG.getNode(Opcode: ISD::OR, DL: SDLoc(N0), VT: OpVT, N1: LL1, N2: RL1);
6858 SDValue And = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(N0), VT: OpVT, N1: LL0, N2: Or);
6859 return DAG.getSetCC(DL, VT, LHS: And, RHS: Or, Cond: IsAnd ? ISD::SETEQ : ISD::SETNE);
6860 }
6861 }
6862
6863 // (and (setne X, 0), (setne X, -1)) --> (setuge (add X, 1), 2)
6864 // (or (seteq X, 0), (seteq X, -1)) --> (setult (add X, 1), 2)
6865 if (LL == RL && CC0 == CC1 && OpVT.getScalarSizeInBits() > 1 && IsInteger &&
6866 ((IsAnd && CC0 == ISD::SETNE) || (!IsAnd && CC0 == ISD::SETEQ)) &&
6867 ((isNullConstant(V: LR) && isAllOnesConstant(V: RR)) ||
6868 (isAllOnesConstant(V: LR) && isNullConstant(V: RR)))) {
6869 SDValue One = DAG.getConstant(Val: 1, DL, VT: OpVT);
6870 SDValue Two = DAG.getConstant(Val: 2, DL, VT: OpVT);
6871 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL: SDLoc(N0), VT: OpVT, N1: LL, N2: One);
6872 AddToWorklist(N: Add.getNode());
6873 return DAG.getSetCC(DL, VT, LHS: Add, RHS: Two, Cond: IsAnd ? ISD::SETUGE : ISD::SETULT);
6874 }
6875
6876 // Try more general transforms if the predicates match and the only user of
6877 // the compares is the 'and' or 'or'.
6878 if (IsInteger && TLI.convertSetCCLogicToBitwiseLogic(VT: OpVT) && CC0 == CC1 &&
6879 N0.hasOneUse() && N1.hasOneUse()) {
6880 // and (seteq A, B), (seteq C, D) --> seteq (or (xor A, B), (xor C, D)), 0
6881 // or (setne A, B), (setne C, D) --> setne (or (xor A, B), (xor C, D)), 0
6882 if ((IsAnd && CC1 == ISD::SETEQ) || (!IsAnd && CC1 == ISD::SETNE)) {
6883 SDValue XorL = DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N0), VT: OpVT, N1: LL, N2: LR);
6884 SDValue XorR = DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N1), VT: OpVT, N1: RL, N2: RR);
6885 SDValue Or = DAG.getNode(Opcode: ISD::OR, DL, VT: OpVT, N1: XorL, N2: XorR);
6886 SDValue Zero = DAG.getConstant(Val: 0, DL, VT: OpVT);
6887 return DAG.getSetCC(DL, VT, LHS: Or, RHS: Zero, Cond: CC1);
6888 }
6889
6890 // Turn compare of constants whose difference is 1 bit into add+and+setcc.
6891 if ((IsAnd && CC1 == ISD::SETNE) || (!IsAnd && CC1 == ISD::SETEQ)) {
6892 // Match a shared variable operand and 2 non-opaque constant operands.
6893 auto MatchDiffPow2 = [&](ConstantSDNode *C0, ConstantSDNode *C1) {
6894 // The difference of the constants must be a single bit.
6895 const APInt &CMax =
6896 APIntOps::umax(A: C0->getAPIntValue(), B: C1->getAPIntValue());
6897 const APInt &CMin =
6898 APIntOps::umin(A: C0->getAPIntValue(), B: C1->getAPIntValue());
6899 return !C0->isOpaque() && !C1->isOpaque() && (CMax - CMin).isPowerOf2();
6900 };
6901 if (LL == RL && ISD::matchBinaryPredicate(LHS: LR, RHS: RR, Match: MatchDiffPow2)) {
6902 // and/or (setcc X, CMax, ne), (setcc X, CMin, ne/eq) -->
6903 // setcc ((sub X, CMin), ~(CMax - CMin)), 0, ne/eq
6904 SDValue Max = DAG.getNode(Opcode: ISD::UMAX, DL, VT: OpVT, N1: LR, N2: RR);
6905 SDValue Min = DAG.getNode(Opcode: ISD::UMIN, DL, VT: OpVT, N1: LR, N2: RR);
6906 SDValue Offset = DAG.getNode(Opcode: ISD::SUB, DL, VT: OpVT, N1: LL, N2: Min);
6907 SDValue Diff = DAG.getNode(Opcode: ISD::SUB, DL, VT: OpVT, N1: Max, N2: Min);
6908 SDValue Mask = DAG.getNOT(DL, Val: Diff, VT: OpVT);
6909 SDValue And = DAG.getNode(Opcode: ISD::AND, DL, VT: OpVT, N1: Offset, N2: Mask);
6910 SDValue Zero = DAG.getConstant(Val: 0, DL, VT: OpVT);
6911 return DAG.getSetCC(DL, VT, LHS: And, RHS: Zero, Cond: CC0);
6912 }
6913 }
6914 }
6915
6916 // Canonicalize equivalent operands to LL == RL.
6917 if (LL == RR && LR == RL) {
6918 CC1 = ISD::getSetCCSwappedOperands(Operation: CC1);
6919 std::swap(a&: RL, b&: RR);
6920 }
6921
6922 // (and (setcc X, Y, CC0), (setcc X, Y, CC1)) --> (setcc X, Y, NewCC)
6923 // (or (setcc X, Y, CC0), (setcc X, Y, CC1)) --> (setcc X, Y, NewCC)
6924 if (LL == RL && LR == RR) {
6925 ISD::CondCode NewCC = IsAnd ? ISD::getSetCCAndOperation(Op1: CC0, Op2: CC1, Type: OpVT)
6926 : ISD::getSetCCOrOperation(Op1: CC0, Op2: CC1, Type: OpVT);
6927 if (NewCC != ISD::SETCC_INVALID &&
6928 (!LegalOperations ||
6929 (TLI.isCondCodeLegal(CC: NewCC, VT: LL.getSimpleValueType()) &&
6930 TLI.isOperationLegal(Op: ISD::SETCC, VT: OpVT))))
6931 return DAG.getSetCC(DL, VT, LHS: LL, RHS: LR, Cond: NewCC);
6932 }
6933
6934 return SDValue();
6935}
6936
6937static bool arebothOperandsNotSNan(SDValue Operand1, SDValue Operand2,
6938 SelectionDAG &DAG) {
6939 return DAG.isKnownNeverSNaN(Op: Operand2) && DAG.isKnownNeverSNaN(Op: Operand1);
6940}
6941
6942static bool arebothOperandsNotNan(SDValue Operand1, SDValue Operand2,
6943 SelectionDAG &DAG) {
6944 return DAG.isKnownNeverNaN(Op: Operand2) && DAG.isKnownNeverNaN(Op: Operand1);
6945}
6946
6947/// Returns an appropriate FP min/max opcode for clamping operations.
6948static unsigned getMinMaxOpcodeForClamp(bool IsMin, SDValue Operand1,
6949 SDValue Operand2, SelectionDAG &DAG,
6950 const TargetLowering &TLI) {
6951 EVT VT = Operand1.getValueType();
6952 unsigned IEEEOp = IsMin ? ISD::FMINNUM_IEEE : ISD::FMAXNUM_IEEE;
6953 if (TLI.isOperationLegalOrCustom(Op: IEEEOp, VT) &&
6954 arebothOperandsNotNan(Operand1, Operand2, DAG))
6955 return IEEEOp;
6956 unsigned PreferredOp = IsMin ? ISD::FMINNUM : ISD::FMAXNUM;
6957 if (TLI.isOperationLegalOrCustom(Op: PreferredOp, VT))
6958 return PreferredOp;
6959 return ISD::DELETED_NODE;
6960}
6961
6962// FIXME: use FMINIMUMNUM if possible, such as for RISC-V.
6963static unsigned getMinMaxOpcodeForCompareFold(
6964 SDValue Operand1, SDValue Operand2, bool SetCCNoNaNs, ISD::CondCode CC,
6965 unsigned OrAndOpcode, SelectionDAG &DAG, bool isFMAXNUMFMINNUM_IEEE,
6966 bool isFMAXNUMFMINNUM) {
6967 // The optimization cannot be applied for all the predicates because
6968 // of the way FMINNUM/FMAXNUM and FMINNUM_IEEE/FMAXNUM_IEEE handle
6969 // NaNs. For FMINNUM_IEEE/FMAXNUM_IEEE, the optimization cannot be
6970 // applied at all if one of the operands is a signaling NaN.
6971
6972 // It is safe to use FMINNUM_IEEE/FMAXNUM_IEEE if all the operands
6973 // are non NaN values.
6974 if (((CC == ISD::SETLT || CC == ISD::SETLE) && (OrAndOpcode == ISD::OR)) ||
6975 ((CC == ISD::SETGT || CC == ISD::SETGE) && (OrAndOpcode == ISD::AND))) {
6976 return (SetCCNoNaNs || arebothOperandsNotNan(Operand1, Operand2, DAG)) &&
6977 isFMAXNUMFMINNUM_IEEE
6978 ? ISD::FMINNUM_IEEE
6979 : ISD::DELETED_NODE;
6980 }
6981
6982 if (((CC == ISD::SETGT || CC == ISD::SETGE) && (OrAndOpcode == ISD::OR)) ||
6983 ((CC == ISD::SETLT || CC == ISD::SETLE) && (OrAndOpcode == ISD::AND))) {
6984 return (SetCCNoNaNs || arebothOperandsNotNan(Operand1, Operand2, DAG)) &&
6985 isFMAXNUMFMINNUM_IEEE
6986 ? ISD::FMAXNUM_IEEE
6987 : ISD::DELETED_NODE;
6988 }
6989
6990 // Both FMINNUM/FMAXNUM and FMINNUM_IEEE/FMAXNUM_IEEE handle quiet
6991 // NaNs in the same way. But, FMINNUM/FMAXNUM and FMINNUM_IEEE/
6992 // FMAXNUM_IEEE handle signaling NaNs differently. If we cannot prove
6993 // that there are not any sNaNs, then the optimization is not valid
6994 // for FMINNUM_IEEE/FMAXNUM_IEEE. In the presence of sNaNs, we apply
6995 // the optimization using FMINNUM/FMAXNUM for the following cases. If
6996 // we can prove that we do not have any sNaNs, then we can do the
6997 // optimization using FMINNUM_IEEE/FMAXNUM_IEEE for the following
6998 // cases.
6999 if (((CC == ISD::SETOLT || CC == ISD::SETOLE) && (OrAndOpcode == ISD::OR)) ||
7000 ((CC == ISD::SETUGT || CC == ISD::SETUGE) && (OrAndOpcode == ISD::AND))) {
7001 return isFMAXNUMFMINNUM ? ISD::FMINNUM
7002 : arebothOperandsNotSNan(Operand1, Operand2, DAG) &&
7003 isFMAXNUMFMINNUM_IEEE
7004 ? ISD::FMINNUM_IEEE
7005 : ISD::DELETED_NODE;
7006 }
7007
7008 if (((CC == ISD::SETOGT || CC == ISD::SETOGE) && (OrAndOpcode == ISD::OR)) ||
7009 ((CC == ISD::SETULT || CC == ISD::SETULE) && (OrAndOpcode == ISD::AND))) {
7010 return isFMAXNUMFMINNUM ? ISD::FMAXNUM
7011 : arebothOperandsNotSNan(Operand1, Operand2, DAG) &&
7012 isFMAXNUMFMINNUM_IEEE
7013 ? ISD::FMAXNUM_IEEE
7014 : ISD::DELETED_NODE;
7015 }
7016
7017 return ISD::DELETED_NODE;
7018}
7019
7020static SDValue foldAndOrOfSETCC(SDNode *LogicOp, SelectionDAG &DAG) {
7021 using AndOrSETCCFoldKind = TargetLowering::AndOrSETCCFoldKind;
7022 assert(
7023 (LogicOp->getOpcode() == ISD::AND || LogicOp->getOpcode() == ISD::OR) &&
7024 "Invalid Op to combine SETCC with");
7025
7026 // TODO: Search past casts/truncates.
7027 SDValue LHS = LogicOp->getOperand(Num: 0);
7028 SDValue RHS = LogicOp->getOperand(Num: 1);
7029 if (LHS->getOpcode() != ISD::SETCC || RHS->getOpcode() != ISD::SETCC ||
7030 !LHS->hasOneUse() || !RHS->hasOneUse())
7031 return SDValue();
7032
7033 SDNodeFlags LHSSetCCFlags = LHS->getFlags();
7034 SDNodeFlags RHSSetCCFlags = RHS->getFlags();
7035 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
7036 AndOrSETCCFoldKind TargetPreference = TLI.isDesirableToCombineLogicOpOfSETCC(
7037 LogicOp, SETCC0: LHS.getNode(), SETCC1: RHS.getNode());
7038
7039 SDValue LHS0 = LHS->getOperand(Num: 0);
7040 SDValue RHS0 = RHS->getOperand(Num: 0);
7041 SDValue LHS1 = LHS->getOperand(Num: 1);
7042 SDValue RHS1 = RHS->getOperand(Num: 1);
7043 // TODO: We don't actually need a splat here, for vectors we just need the
7044 // invariants to hold for each element.
7045 auto *LHS1C = isConstOrConstSplat(N: LHS1);
7046 auto *RHS1C = isConstOrConstSplat(N: RHS1);
7047 ISD::CondCode CCL = cast<CondCodeSDNode>(Val: LHS.getOperand(i: 2))->get();
7048 ISD::CondCode CCR = cast<CondCodeSDNode>(Val: RHS.getOperand(i: 2))->get();
7049 EVT VT = LogicOp->getValueType(ResNo: 0);
7050 EVT OpVT = LHS0.getValueType();
7051 SDLoc DL(LogicOp);
7052
7053 // Check if the operands of an and/or operation are comparisons and if they
7054 // compare against the same value. Replace the and/or-cmp-cmp sequence with
7055 // min/max cmp sequence. If LHS1 is equal to RHS1, then the or-cmp-cmp
7056 // sequence will be replaced with min-cmp sequence:
7057 // (LHS0 < LHS1) | (RHS0 < RHS1) -> min(LHS0, RHS0) < LHS1
7058 // and and-cmp-cmp will be replaced with max-cmp sequence:
7059 // (LHS0 < LHS1) & (RHS0 < RHS1) -> max(LHS0, RHS0) < LHS1
7060 // The optimization does not work for `==` or `!=` .
7061 // The two comparisons should have either the same predicate or the
7062 // predicate of one of the comparisons is the opposite of the other one.
7063 bool isFMAXNUMFMINNUM_IEEE = TLI.isOperationLegal(Op: ISD::FMAXNUM_IEEE, VT: OpVT) &&
7064 TLI.isOperationLegal(Op: ISD::FMINNUM_IEEE, VT: OpVT);
7065 bool isFMAXNUMFMINNUM = TLI.isOperationLegalOrCustom(Op: ISD::FMAXNUM, VT: OpVT) &&
7066 TLI.isOperationLegalOrCustom(Op: ISD::FMINNUM, VT: OpVT);
7067 if (((OpVT.isInteger() && TLI.isOperationLegal(Op: ISD::UMAX, VT: OpVT) &&
7068 TLI.isOperationLegal(Op: ISD::SMAX, VT: OpVT) &&
7069 TLI.isOperationLegal(Op: ISD::UMIN, VT: OpVT) &&
7070 TLI.isOperationLegal(Op: ISD::SMIN, VT: OpVT)) ||
7071 (OpVT.isFloatingPoint() &&
7072 (isFMAXNUMFMINNUM_IEEE || isFMAXNUMFMINNUM))) &&
7073 !ISD::isIntEqualitySetCC(Code: CCL) && !ISD::isFPEqualitySetCC(Code: CCL) &&
7074 CCL != ISD::SETFALSE && CCL != ISD::SETO && CCL != ISD::SETUO &&
7075 CCL != ISD::SETTRUE &&
7076 (CCL == CCR || CCL == ISD::getSetCCSwappedOperands(Operation: CCR))) {
7077
7078 SDValue CommonValue, Operand1, Operand2;
7079 ISD::CondCode CC = ISD::SETCC_INVALID;
7080 if (CCL == CCR) {
7081 if (LHS0 == RHS0) {
7082 CommonValue = LHS0;
7083 Operand1 = LHS1;
7084 Operand2 = RHS1;
7085 CC = ISD::getSetCCSwappedOperands(Operation: CCL);
7086 } else if (LHS1 == RHS1) {
7087 CommonValue = LHS1;
7088 Operand1 = LHS0;
7089 Operand2 = RHS0;
7090 CC = CCL;
7091 }
7092 } else {
7093 assert(CCL == ISD::getSetCCSwappedOperands(CCR) && "Unexpected CC");
7094 if (LHS0 == RHS1) {
7095 CommonValue = LHS0;
7096 Operand1 = LHS1;
7097 Operand2 = RHS0;
7098 CC = CCR;
7099 } else if (RHS0 == LHS1) {
7100 CommonValue = LHS1;
7101 Operand1 = LHS0;
7102 Operand2 = RHS1;
7103 CC = CCL;
7104 }
7105 }
7106
7107 // Don't do this transform for sign bit tests. Let foldLogicOfSetCCs
7108 // handle it using OR/AND.
7109 if (CC == ISD::SETLT && isNullOrNullSplat(V: CommonValue))
7110 CC = ISD::SETCC_INVALID;
7111 else if (CC == ISD::SETGT && isAllOnesOrAllOnesSplat(V: CommonValue))
7112 CC = ISD::SETCC_INVALID;
7113
7114 if (CC != ISD::SETCC_INVALID) {
7115 unsigned NewOpcode = ISD::DELETED_NODE;
7116 bool IsSigned = isSignedIntSetCC(Code: CC);
7117 if (OpVT.isInteger()) {
7118 bool IsLess = (CC == ISD::SETLE || CC == ISD::SETULE ||
7119 CC == ISD::SETLT || CC == ISD::SETULT);
7120 bool IsOr = (LogicOp->getOpcode() == ISD::OR);
7121 if (IsLess == IsOr)
7122 NewOpcode = IsSigned ? ISD::SMIN : ISD::UMIN;
7123 else
7124 NewOpcode = IsSigned ? ISD::SMAX : ISD::UMAX;
7125 } else if (OpVT.isFloatingPoint())
7126 NewOpcode = getMinMaxOpcodeForCompareFold(
7127 Operand1, Operand2,
7128 SetCCNoNaNs: LHSSetCCFlags.hasNoNaNs() && RHSSetCCFlags.hasNoNaNs(), CC,
7129 OrAndOpcode: LogicOp->getOpcode(), DAG, isFMAXNUMFMINNUM_IEEE, isFMAXNUMFMINNUM);
7130
7131 if (NewOpcode != ISD::DELETED_NODE) {
7132 // Propagate fast-math flags from setcc.
7133 SDNodeFlags Flags = LHS->getFlags() & RHS->getFlags();
7134 SDValue MinMaxValue =
7135 DAG.getNode(Opcode: NewOpcode, DL, VT: OpVT, N1: Operand1, N2: Operand2, Flags);
7136 return DAG.getSetCC(DL, VT, LHS: MinMaxValue, RHS: CommonValue, Cond: CC, /*Chain=*/{},
7137 /*IsSignaling=*/false, Flags);
7138 }
7139 }
7140 }
7141
7142 if (LHS0 == LHS1 && RHS0 == RHS1 && CCL == CCR &&
7143 LHS0.getValueType() == RHS0.getValueType() &&
7144 ((LogicOp->getOpcode() == ISD::AND && CCL == ISD::SETO) ||
7145 (LogicOp->getOpcode() == ISD::OR && CCL == ISD::SETUO)))
7146 return DAG.getSetCC(DL, VT, LHS: LHS0, RHS: RHS0, Cond: CCL);
7147
7148 if (TargetPreference == AndOrSETCCFoldKind::None)
7149 return SDValue();
7150
7151 if (CCL == CCR &&
7152 CCL == (LogicOp->getOpcode() == ISD::AND ? ISD::SETNE : ISD::SETEQ) &&
7153 LHS0 == RHS0 && LHS1C && RHS1C && OpVT.isInteger()) {
7154 const APInt &APLhs = LHS1C->getAPIntValue();
7155 const APInt &APRhs = RHS1C->getAPIntValue();
7156
7157 // Preference is to use ISD::ABS or we already have an ISD::ABS (in which
7158 // case this is just a compare).
7159 if (APLhs == (-APRhs) &&
7160 ((TargetPreference & AndOrSETCCFoldKind::ABS) ||
7161 DAG.doesNodeExist(Opcode: ISD::ABS, VTList: DAG.getVTList(VT: OpVT), Ops: {LHS0}))) {
7162 const APInt &C = APLhs.isNegative() ? APRhs : APLhs;
7163 // (icmp eq A, C) | (icmp eq A, -C)
7164 // -> (icmp eq Abs(A), C)
7165 // (icmp ne A, C) & (icmp ne A, -C)
7166 // -> (icmp ne Abs(A), C)
7167 SDValue AbsOp = DAG.getNode(Opcode: ISD::ABS, DL, VT: OpVT, Operand: LHS0);
7168 return DAG.getNode(Opcode: ISD::SETCC, DL, VT, N1: AbsOp,
7169 N2: DAG.getConstant(Val: C, DL, VT: OpVT), N3: LHS.getOperand(i: 2));
7170 } else if (TargetPreference &
7171 (AndOrSETCCFoldKind::AddAnd | AndOrSETCCFoldKind::NotAnd)) {
7172
7173 // AndOrSETCCFoldKind::AddAnd:
7174 // A == C0 | A == C1
7175 // IF IsPow2(smax(C0, C1)-smin(C0, C1))
7176 // -> ((A - smin(C0, C1)) & ~(smax(C0, C1)-smin(C0, C1))) == 0
7177 // A != C0 & A != C1
7178 // IF IsPow2(smax(C0, C1)-smin(C0, C1))
7179 // -> ((A - smin(C0, C1)) & ~(smax(C0, C1)-smin(C0, C1))) != 0
7180
7181 // AndOrSETCCFoldKind::NotAnd:
7182 // A == C0 | A == C1
7183 // IF smax(C0, C1) == -1 AND IsPow2(smax(C0, C1) - smin(C0, C1))
7184 // -> ~A & smin(C0, C1) == 0
7185 // A != C0 & A != C1
7186 // IF smax(C0, C1) == -1 AND IsPow2(smax(C0, C1) - smin(C0, C1))
7187 // -> ~A & smin(C0, C1) != 0
7188
7189 const APInt &MaxC = APIntOps::smax(A: APRhs, B: APLhs);
7190 const APInt &MinC = APIntOps::smin(A: APRhs, B: APLhs);
7191 APInt Dif = MaxC - MinC;
7192 if (!Dif.isZero() && Dif.isPowerOf2()) {
7193 if (MaxC.isAllOnes() &&
7194 (TargetPreference & AndOrSETCCFoldKind::NotAnd)) {
7195 SDValue NotOp = DAG.getNOT(DL, Val: LHS0, VT: OpVT);
7196 SDValue AndOp = DAG.getNode(Opcode: ISD::AND, DL, VT: OpVT, N1: NotOp,
7197 N2: DAG.getConstant(Val: MinC, DL, VT: OpVT));
7198 return DAG.getNode(Opcode: ISD::SETCC, DL, VT, N1: AndOp,
7199 N2: DAG.getConstant(Val: 0, DL, VT: OpVT), N3: LHS.getOperand(i: 2));
7200 } else if (TargetPreference & AndOrSETCCFoldKind::AddAnd) {
7201
7202 SDValue AddOp = DAG.getNode(Opcode: ISD::ADD, DL, VT: OpVT, N1: LHS0,
7203 N2: DAG.getConstant(Val: -MinC, DL, VT: OpVT));
7204 SDValue AndOp = DAG.getNode(Opcode: ISD::AND, DL, VT: OpVT, N1: AddOp,
7205 N2: DAG.getConstant(Val: ~Dif, DL, VT: OpVT));
7206 return DAG.getNode(Opcode: ISD::SETCC, DL, VT, N1: AndOp,
7207 N2: DAG.getConstant(Val: 0, DL, VT: OpVT), N3: LHS.getOperand(i: 2));
7208 }
7209 }
7210 }
7211 }
7212
7213 return SDValue();
7214}
7215
7216// Combine `(select c, (X & 1), 0)` -> `(and (zext c), X)`.
7217// We canonicalize to the `select` form in the middle end, but the `and` form
7218// gets better codegen and all tested targets (arm, x86, riscv)
7219static SDValue combineSelectAsExtAnd(SDValue Cond, SDValue T, SDValue F,
7220 const SDLoc &DL, SelectionDAG &DAG) {
7221 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
7222 if (!isNullConstant(V: F))
7223 return SDValue();
7224
7225 EVT CondVT = Cond.getValueType();
7226 if (TLI.getBooleanContents(Type: CondVT) !=
7227 TargetLoweringBase::ZeroOrOneBooleanContent)
7228 return SDValue();
7229
7230 if (T.getOpcode() != ISD::AND)
7231 return SDValue();
7232
7233 if (!isOneConstant(V: T.getOperand(i: 1)))
7234 return SDValue();
7235
7236 EVT OpVT = T.getValueType();
7237
7238 SDValue CondMask =
7239 OpVT == CondVT ? Cond : DAG.getBoolExtOrTrunc(Op: Cond, SL: DL, VT: OpVT, OpVT: CondVT);
7240 return DAG.getNode(Opcode: ISD::AND, DL, VT: OpVT, N1: CondMask, N2: T.getOperand(i: 0));
7241}
7242
7243/// This contains all DAGCombine rules which reduce two values combined by
7244/// an And operation to a single value. This makes them reusable in the context
7245/// of visitSELECT(). Rules involving constants are not included as
7246/// visitSELECT() already handles those cases.
7247SDValue DAGCombiner::visitANDLike(SDValue N0, SDValue N1, SDNode *N) {
7248 EVT VT = N1.getValueType();
7249 SDLoc DL(N);
7250
7251 // fold (and x, undef) -> 0
7252 if (N0.isUndef() || N1.isUndef())
7253 return DAG.getConstant(Val: 0, DL, VT);
7254
7255 if (SDValue V = foldLogicOfSetCCs(IsAnd: true, N0, N1, DL))
7256 return V;
7257
7258 // Canonicalize:
7259 // and(x, add) -> and(add, x)
7260 if (N1.getOpcode() == ISD::ADD)
7261 std::swap(a&: N0, b&: N1);
7262
7263 // TODO: Rewrite this to return a new 'AND' instead of using CombineTo.
7264 if (N0.getOpcode() == ISD::ADD && N1.getOpcode() == ISD::SRL &&
7265 VT.isScalarInteger() && VT.getSizeInBits() <= 64 && N0->hasOneUse()) {
7266 if (ConstantSDNode *ADDI = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1))) {
7267 if (ConstantSDNode *SRLI = dyn_cast<ConstantSDNode>(Val: N1.getOperand(i: 1))) {
7268 // Look for (and (add x, c1), (lshr y, c2)). If C1 wasn't a legal
7269 // immediate for an add, but it is legal if its top c2 bits are set,
7270 // transform the ADD so the immediate doesn't need to be materialized
7271 // in a register.
7272 APInt ADDC = ADDI->getAPIntValue();
7273 APInt SRLC = SRLI->getAPIntValue();
7274 if (ADDC.getSignificantBits() <= 64 && SRLC.ult(RHS: VT.getSizeInBits()) &&
7275 !TLI.isLegalAddImmediate(ADDC.getSExtValue())) {
7276 APInt Mask = APInt::getHighBitsSet(numBits: VT.getSizeInBits(),
7277 hiBitsSet: SRLC.getZExtValue());
7278 if (DAG.MaskedValueIsZero(Op: N0.getOperand(i: 1), Mask)) {
7279 ADDC |= Mask;
7280 if (TLI.isLegalAddImmediate(ADDC.getSExtValue())) {
7281 SDLoc DL0(N0);
7282 SDValue NewAdd =
7283 DAG.getNode(Opcode: ISD::ADD, DL: DL0, VT,
7284 N1: N0.getOperand(i: 0), N2: DAG.getConstant(Val: ADDC, DL, VT));
7285 CombineTo(N: N0.getNode(), Res: NewAdd);
7286 // Return N so it doesn't get rechecked!
7287 return SDValue(N, 0);
7288 }
7289 }
7290 }
7291 }
7292 }
7293 }
7294
7295 return SDValue();
7296}
7297
7298bool DAGCombiner::isAndLoadExtLoad(ConstantSDNode *AndC, LoadSDNode *LoadN,
7299 EVT LoadResultTy, EVT &ExtVT) {
7300 if (!AndC->getAPIntValue().isMask())
7301 return false;
7302
7303 unsigned ActiveBits = AndC->getAPIntValue().countr_one();
7304
7305 ExtVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: ActiveBits);
7306 EVT LoadedVT = LoadN->getMemoryVT();
7307
7308 if (ExtVT == LoadedVT &&
7309 (!LegalOperations ||
7310 TLI.isLoadLegal(ValVT: LoadResultTy, MemVT: ExtVT, Alignment: LoadN->getAlign(),
7311 AddrSpace: LoadN->getAddressSpace(), ExtType: ISD::ZEXTLOAD, Atomic: false))) {
7312 // ZEXTLOAD will match without needing to change the size of the value being
7313 // loaded.
7314 return true;
7315 }
7316
7317 // Do not change the width of a volatile or atomic loads.
7318 if (!LoadN->isSimple())
7319 return false;
7320
7321 // Do not generate loads of non-round integer types since these can
7322 // be expensive (and would be wrong if the type is not byte sized).
7323 if (!LoadedVT.bitsGT(VT: ExtVT) || !ExtVT.isRound())
7324 return false;
7325
7326 if (LegalOperations &&
7327 !TLI.isLoadLegal(ValVT: LoadResultTy, MemVT: ExtVT, Alignment: LoadN->getAlign(),
7328 AddrSpace: LoadN->getAddressSpace(), ExtType: ISD::ZEXTLOAD, Atomic: false))
7329 return false;
7330
7331 if (!TLI.shouldReduceLoadWidth(Load: LoadN, ExtTy: ISD::ZEXTLOAD, NewVT: ExtVT, /*ByteOffset=*/0))
7332 return false;
7333
7334 return true;
7335}
7336
7337bool DAGCombiner::isLegalNarrowLdSt(LSBaseSDNode *LDST,
7338 ISD::LoadExtType ExtType, EVT &MemVT,
7339 unsigned ShAmt) {
7340 if (!LDST)
7341 return false;
7342
7343 // Only allow byte offsets.
7344 if (ShAmt % 8)
7345 return false;
7346 const unsigned ByteShAmt = ShAmt / 8;
7347
7348 // Do not generate loads of non-round integer types since these can
7349 // be expensive (and would be wrong if the type is not byte sized).
7350 if (!MemVT.isRound())
7351 return false;
7352
7353 // Don't change the width of a volatile or atomic loads.
7354 if (!LDST->isSimple())
7355 return false;
7356
7357 EVT LdStMemVT = LDST->getMemoryVT();
7358
7359 // Bail out when changing the scalable property, since we can't be sure that
7360 // we're actually narrowing here.
7361 if (LdStMemVT.isScalableVector() != MemVT.isScalableVector())
7362 return false;
7363
7364 // Verify that we are actually reducing a load width here.
7365 if (LdStMemVT.bitsLT(VT: MemVT))
7366 return false;
7367
7368 // Ensure that this isn't going to produce an unsupported memory access.
7369 if (ShAmt) {
7370 const Align LDSTAlign = LDST->getAlign();
7371 const Align NarrowAlign = commonAlignment(A: LDSTAlign, Offset: ByteShAmt);
7372 if (!TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT: MemVT,
7373 AddrSpace: LDST->getAddressSpace(), Alignment: NarrowAlign,
7374 Flags: LDST->getMemOperand()->getFlags()))
7375 return false;
7376 }
7377
7378 // It's not possible to generate a constant of extended or untyped type.
7379 EVT PtrType = LDST->getBasePtr().getValueType();
7380 if (PtrType == MVT::Untyped || PtrType.isExtended())
7381 return false;
7382
7383 if (isa<LoadSDNode>(Val: LDST)) {
7384 LoadSDNode *Load = cast<LoadSDNode>(Val: LDST);
7385 // Don't transform one with multiple uses, this would require adding a new
7386 // load.
7387 if (!SDValue(Load, 0).hasOneUse())
7388 return false;
7389
7390 if (LegalOperations &&
7391 !TLI.isLoadLegal(ValVT: Load->getValueType(ResNo: 0), MemVT, Alignment: Load->getAlign(),
7392 AddrSpace: Load->getAddressSpace(), ExtType, Atomic: false))
7393 return false;
7394
7395 // For the transform to be legal, the load must produce only two values
7396 // (the value loaded and the chain). Don't transform a pre-increment
7397 // load, for example, which produces an extra value. Otherwise the
7398 // transformation is not equivalent, and the downstream logic to replace
7399 // uses gets things wrong.
7400 if (Load->getNumValues() > 2)
7401 return false;
7402
7403 // If the load that we're shrinking is an extload and we're not just
7404 // discarding the extension we can't simply shrink the load. Bail.
7405 // TODO: It would be possible to merge the extensions in some cases.
7406 if (Load->getExtensionType() != ISD::NON_EXTLOAD &&
7407 Load->getMemoryVT().getSizeInBits() < MemVT.getSizeInBits() + ShAmt)
7408 return false;
7409
7410 if (!TLI.shouldReduceLoadWidth(Load, ExtTy: ExtType, NewVT: MemVT, ByteOffset: ByteShAmt))
7411 return false;
7412 } else {
7413 assert(isa<StoreSDNode>(LDST) && "It is not a Load nor a Store SDNode");
7414 StoreSDNode *Store = cast<StoreSDNode>(Val: LDST);
7415 // Can't write outside the original store
7416 if (Store->getMemoryVT().getSizeInBits() < MemVT.getSizeInBits() + ShAmt)
7417 return false;
7418
7419 if (LegalOperations &&
7420 !TLI.isTruncStoreLegal(ValVT: Store->getValue().getValueType(), MemVT,
7421 Alignment: Store->getAlign(), AddrSpace: Store->getAddressSpace()))
7422 return false;
7423 }
7424 return true;
7425}
7426
7427bool DAGCombiner::SearchForAndLoads(SDNode *N,
7428 SmallVectorImpl<LoadSDNode*> &Loads,
7429 SmallPtrSetImpl<SDNode*> &NodesWithConsts,
7430 ConstantSDNode *Mask,
7431 SDNode *&NodeToMask) {
7432 // Recursively search for the operands, looking for loads which can be
7433 // narrowed.
7434 for (SDValue Op : N->op_values()) {
7435 if (Op.getValueType().isVector())
7436 return false;
7437
7438 // Some constants may need fixing up later if they are too large.
7439 if (auto *C = dyn_cast<ConstantSDNode>(Val&: Op)) {
7440 assert(ISD::isBitwiseLogicOp(N->getOpcode()) &&
7441 "Expected bitwise logic operation");
7442 if (!C->getAPIntValue().isSubsetOf(RHS: Mask->getAPIntValue()))
7443 NodesWithConsts.insert(Ptr: N);
7444 continue;
7445 }
7446
7447 if (!Op.hasOneUse())
7448 return false;
7449
7450 switch(Op.getOpcode()) {
7451 case ISD::LOAD: {
7452 auto *Load = cast<LoadSDNode>(Val&: Op);
7453 EVT ExtVT;
7454 if (isAndLoadExtLoad(AndC: Mask, LoadN: Load, LoadResultTy: Load->getValueType(ResNo: 0), ExtVT) &&
7455 isLegalNarrowLdSt(LDST: Load, ExtType: ISD::ZEXTLOAD, MemVT&: ExtVT)) {
7456
7457 // ZEXTLOAD is already small enough.
7458 if (Load->getExtensionType() == ISD::ZEXTLOAD &&
7459 ExtVT.bitsGE(VT: Load->getMemoryVT()))
7460 continue;
7461
7462 // Use LE to convert equal sized loads to zext.
7463 if (ExtVT.bitsLE(VT: Load->getMemoryVT()))
7464 Loads.push_back(Elt: Load);
7465
7466 continue;
7467 }
7468 return false;
7469 }
7470 case ISD::ZERO_EXTEND:
7471 case ISD::AssertZext: {
7472 unsigned ActiveBits = Mask->getAPIntValue().countr_one();
7473 EVT ExtVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: ActiveBits);
7474 EVT VT = Op.getOpcode() == ISD::AssertZext ?
7475 cast<VTSDNode>(Val: Op.getOperand(i: 1))->getVT() :
7476 Op.getOperand(i: 0).getValueType();
7477
7478 // We can accept extending nodes if the mask is wider or an equal
7479 // width to the original type.
7480 if (ExtVT.bitsGE(VT))
7481 continue;
7482 break;
7483 }
7484 case ISD::OR:
7485 case ISD::XOR:
7486 case ISD::AND:
7487 if (!SearchForAndLoads(N: Op.getNode(), Loads, NodesWithConsts, Mask,
7488 NodeToMask))
7489 return false;
7490 continue;
7491 }
7492
7493 // Allow one node which will masked along with any loads found.
7494 if (NodeToMask)
7495 return false;
7496
7497 // Also ensure that the node to be masked only produces one data result.
7498 NodeToMask = Op.getNode();
7499 if (NodeToMask->getNumValues() > 1) {
7500 bool HasValue = false;
7501 for (unsigned i = 0, e = NodeToMask->getNumValues(); i < e; ++i) {
7502 MVT VT = SDValue(NodeToMask, i).getSimpleValueType();
7503 if (VT != MVT::Glue && VT != MVT::Other) {
7504 if (HasValue) {
7505 NodeToMask = nullptr;
7506 return false;
7507 }
7508 HasValue = true;
7509 }
7510 }
7511 assert(HasValue && "Node to be masked has no data result?");
7512 }
7513 }
7514 return true;
7515}
7516
7517bool DAGCombiner::BackwardsPropagateMask(SDNode *N) {
7518 auto *Mask = dyn_cast<ConstantSDNode>(Val: N->getOperand(Num: 1));
7519 if (!Mask)
7520 return false;
7521
7522 if (!Mask->getAPIntValue().isMask())
7523 return false;
7524
7525 // No need to do anything if the and directly uses a load.
7526 if (isa<LoadSDNode>(Val: N->getOperand(Num: 0)))
7527 return false;
7528
7529 SmallVector<LoadSDNode*, 8> Loads;
7530 SmallPtrSet<SDNode*, 2> NodesWithConsts;
7531 SDNode *FixupNode = nullptr;
7532 if (SearchForAndLoads(N, Loads, NodesWithConsts, Mask, NodeToMask&: FixupNode)) {
7533 if (Loads.empty())
7534 return false;
7535
7536 LLVM_DEBUG(dbgs() << "Backwards propagate AND: "; N->dump());
7537 SDValue MaskOp = N->getOperand(Num: 1);
7538
7539 // If it exists, fixup the single node we allow in the tree that needs
7540 // masking.
7541 if (FixupNode) {
7542 LLVM_DEBUG(dbgs() << "First, need to fix up: "; FixupNode->dump());
7543 SDValue And = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(FixupNode),
7544 VT: FixupNode->getValueType(ResNo: 0),
7545 N1: SDValue(FixupNode, 0), N2: MaskOp);
7546 DAG.ReplaceAllUsesOfValueWith(From: SDValue(FixupNode, 0), To: And);
7547 if (And.getOpcode() == ISD ::AND)
7548 DAG.UpdateNodeOperands(N: And.getNode(), Op1: SDValue(FixupNode, 0), Op2: MaskOp);
7549 }
7550
7551 // Narrow any constants that need it.
7552 for (auto *LogicN : NodesWithConsts) {
7553 SDValue Op0 = LogicN->getOperand(Num: 0);
7554 SDValue Op1 = LogicN->getOperand(Num: 1);
7555
7556 // We only need to fix AND if both inputs are constants. And we only need
7557 // to fix one of the constants.
7558 if (LogicN->getOpcode() == ISD::AND &&
7559 (!isa<ConstantSDNode>(Val: Op0) || !isa<ConstantSDNode>(Val: Op1)))
7560 continue;
7561
7562 if (isa<ConstantSDNode>(Val: Op0) && LogicN->getOpcode() != ISD::AND)
7563 Op0 =
7564 DAG.getNode(Opcode: ISD::AND, DL: SDLoc(Op0), VT: Op0.getValueType(), N1: Op0, N2: MaskOp);
7565
7566 if (isa<ConstantSDNode>(Val: Op1))
7567 Op1 =
7568 DAG.getNode(Opcode: ISD::AND, DL: SDLoc(Op1), VT: Op1.getValueType(), N1: Op1, N2: MaskOp);
7569
7570 if (isa<ConstantSDNode>(Val: Op0) && !isa<ConstantSDNode>(Val: Op1))
7571 std::swap(a&: Op0, b&: Op1);
7572
7573 DAG.UpdateNodeOperands(N: LogicN, Op1: Op0, Op2: Op1);
7574 }
7575
7576 // Create narrow loads.
7577 for (auto *Load : Loads) {
7578 LLVM_DEBUG(dbgs() << "Propagate AND back to: "; Load->dump());
7579 SDValue And = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(Load), VT: Load->getValueType(ResNo: 0),
7580 N1: SDValue(Load, 0), N2: MaskOp);
7581 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Load, 0), To: And);
7582 if (And.getOpcode() == ISD ::AND)
7583 And = SDValue(
7584 DAG.UpdateNodeOperands(N: And.getNode(), Op1: SDValue(Load, 0), Op2: MaskOp), 0);
7585 SDValue NewLoad = reduceLoadWidth(N: And.getNode());
7586 assert(NewLoad &&
7587 "Shouldn't be masking the load if it can't be narrowed");
7588 CombineTo(N: Load, Res0: NewLoad, Res1: NewLoad.getValue(R: 1));
7589 }
7590 DAG.ReplaceAllUsesWith(From: N, To: N->getOperand(Num: 0).getNode());
7591 return true;
7592 }
7593 return false;
7594}
7595
7596// Unfold
7597// x & (-1 'logical shift' y)
7598// To
7599// (x 'opposite logical shift' y) 'logical shift' y
7600// if it is better for performance.
7601SDValue DAGCombiner::unfoldExtremeBitClearingToShifts(SDNode *N) {
7602 assert(N->getOpcode() == ISD::AND);
7603
7604 SDValue N0 = N->getOperand(Num: 0);
7605 SDValue N1 = N->getOperand(Num: 1);
7606
7607 // Do we actually prefer shifts over mask?
7608 if (!TLI.shouldFoldMaskToVariableShiftPair(X: N0))
7609 return SDValue();
7610
7611 // Try to match (-1 '[outer] logical shift' y)
7612 unsigned OuterShift;
7613 unsigned InnerShift; // The opposite direction to the OuterShift.
7614 SDValue Y; // Shift amount.
7615 auto matchMask = [&OuterShift, &InnerShift, &Y](SDValue M) -> bool {
7616 if (!M.hasOneUse())
7617 return false;
7618 OuterShift = M->getOpcode();
7619 if (OuterShift == ISD::SHL)
7620 InnerShift = ISD::SRL;
7621 else if (OuterShift == ISD::SRL)
7622 InnerShift = ISD::SHL;
7623 else
7624 return false;
7625 if (!isAllOnesConstant(V: M->getOperand(Num: 0)))
7626 return false;
7627 Y = M->getOperand(Num: 1);
7628 return true;
7629 };
7630
7631 SDValue X;
7632 if (matchMask(N1))
7633 X = N0;
7634 else if (matchMask(N0))
7635 X = N1;
7636 else
7637 return SDValue();
7638
7639 SDLoc DL(N);
7640 EVT VT = N->getValueType(ResNo: 0);
7641
7642 // tmp = x 'opposite logical shift' y
7643 SDValue T0 = DAG.getNode(Opcode: InnerShift, DL, VT, N1: X, N2: Y);
7644 // ret = tmp 'logical shift' y
7645 SDValue T1 = DAG.getNode(Opcode: OuterShift, DL, VT, N1: T0, N2: Y);
7646
7647 return T1;
7648}
7649
7650/// Try to replace shift/logic that tests if a bit is clear with mask + setcc.
7651/// For a target with a bit test, this is expected to become test + set and save
7652/// at least 1 instruction.
7653static SDValue combineShiftAnd1ToBitTest(SDNode *And, SelectionDAG &DAG) {
7654 assert(And->getOpcode() == ISD::AND && "Expected an 'and' op");
7655
7656 // Look through an optional extension.
7657 SDValue And0 = And->getOperand(Num: 0), And1 = And->getOperand(Num: 1);
7658 if (And0.getOpcode() == ISD::ANY_EXTEND && And0.hasOneUse())
7659 And0 = And0.getOperand(i: 0);
7660 if (!isOneConstant(V: And1) || !And0.hasOneUse())
7661 return SDValue();
7662
7663 SDValue Src = And0;
7664
7665 // Attempt to find a 'not' op.
7666 // TODO: Should we favor test+set even without the 'not' op?
7667 bool FoundNot = false;
7668 if (isBitwiseNot(V: Src)) {
7669 FoundNot = true;
7670 Src = Src.getOperand(i: 0);
7671
7672 // Look though an optional truncation. The source operand may not be the
7673 // same type as the original 'and', but that is ok because we are masking
7674 // off everything but the low bit.
7675 if (Src.getOpcode() == ISD::TRUNCATE && Src.hasOneUse())
7676 Src = Src.getOperand(i: 0);
7677 }
7678
7679 // Match a shift-right by constant.
7680 if (Src.getOpcode() != ISD::SRL || !Src.hasOneUse())
7681 return SDValue();
7682
7683 // This is probably not worthwhile without a supported type.
7684 EVT SrcVT = Src.getValueType();
7685 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
7686 if (!TLI.isTypeLegal(VT: SrcVT))
7687 return SDValue();
7688
7689 // We might have looked through casts that make this transform invalid.
7690 unsigned BitWidth = SrcVT.getScalarSizeInBits();
7691 SDValue ShiftAmt = Src.getOperand(i: 1);
7692 auto *ShiftAmtC = dyn_cast<ConstantSDNode>(Val&: ShiftAmt);
7693 if (!ShiftAmtC || !ShiftAmtC->getAPIntValue().ult(RHS: BitWidth))
7694 return SDValue();
7695
7696 // Set source to shift source.
7697 Src = Src.getOperand(i: 0);
7698
7699 // Try again to find a 'not' op.
7700 // TODO: Should we favor test+set even with two 'not' ops?
7701 if (!FoundNot) {
7702 if (!isBitwiseNot(V: Src))
7703 return SDValue();
7704 Src = Src.getOperand(i: 0);
7705 }
7706
7707 if (!TLI.hasBitTest(X: Src, Y: ShiftAmt))
7708 return SDValue();
7709
7710 // Turn this into a bit-test pattern using mask op + setcc:
7711 // and (not (srl X, C)), 1 --> (and X, 1<<C) == 0
7712 // and (srl (not X), C)), 1 --> (and X, 1<<C) == 0
7713 SDLoc DL(And);
7714 SDValue X = DAG.getZExtOrTrunc(Op: Src, DL, VT: SrcVT);
7715 EVT CCVT =
7716 TLI.getSetCCResultType(DL: DAG.getDataLayout(), Context&: *DAG.getContext(), VT: SrcVT);
7717 SDValue Mask = DAG.getConstant(
7718 Val: APInt::getOneBitSet(numBits: BitWidth, BitNo: ShiftAmtC->getZExtValue()), DL, VT: SrcVT);
7719 SDValue NewAnd = DAG.getNode(Opcode: ISD::AND, DL, VT: SrcVT, N1: X, N2: Mask);
7720 SDValue Zero = DAG.getConstant(Val: 0, DL, VT: SrcVT);
7721 SDValue Setcc = DAG.getSetCC(DL, VT: CCVT, LHS: NewAnd, RHS: Zero, Cond: ISD::SETEQ);
7722 return DAG.getZExtOrTrunc(Op: Setcc, DL, VT: And->getValueType(ResNo: 0));
7723}
7724
7725/// For targets that support usubsat, match a bit-hack form of that operation
7726/// that ends in 'and' and convert it.
7727static SDValue foldAndToUsubsat(SDNode *N, SelectionDAG &DAG, const SDLoc &DL) {
7728 EVT VT = N->getValueType(ResNo: 0);
7729 unsigned BitWidth = VT.getScalarSizeInBits();
7730 APInt SignMask = APInt::getSignMask(BitWidth);
7731
7732 // (i8 X ^ 128) & (i8 X s>> 7) --> usubsat X, 128
7733 // (i8 X + 128) & (i8 X s>> 7) --> usubsat X, 128
7734 // xor/add with SMIN (signmask) are logically equivalent.
7735 SDValue X;
7736 if (!sd_match(N, P: m_And(L: m_OneUse(P: m_Xor(L: m_Value(N&: X), R: m_SpecificInt(V: SignMask))),
7737 R: m_OneUse(P: m_Sra(L: m_Deferred(V&: X),
7738 R: m_SpecificInt(V: BitWidth - 1))))) &&
7739 !sd_match(N, P: m_And(L: m_OneUse(P: m_Add(L: m_Value(N&: X), R: m_SpecificInt(V: SignMask))),
7740 R: m_OneUse(P: m_Sra(L: m_Deferred(V&: X),
7741 R: m_SpecificInt(V: BitWidth - 1))))))
7742 return SDValue();
7743
7744 return DAG.getNode(Opcode: ISD::USUBSAT, DL, VT, N1: X,
7745 N2: DAG.getConstant(Val: SignMask, DL, VT));
7746}
7747
7748/// Given a bitwise logic operation N with a matching bitwise logic operand,
7749/// fold a pattern where 2 of the source operands are identically shifted
7750/// values. For example:
7751/// ((X0 << Y) | Z) | (X1 << Y) --> ((X0 | X1) << Y) | Z
7752static SDValue foldLogicOfShifts(SDNode *N, SDValue LogicOp, SDValue ShiftOp,
7753 SelectionDAG &DAG) {
7754 unsigned LogicOpcode = N->getOpcode();
7755 assert(ISD::isBitwiseLogicOp(LogicOpcode) &&
7756 "Expected bitwise logic operation");
7757
7758 if (!LogicOp.hasOneUse() || !ShiftOp.hasOneUse())
7759 return SDValue();
7760
7761 // Match another bitwise logic op and a shift.
7762 unsigned ShiftOpcode = ShiftOp.getOpcode();
7763 if (LogicOp.getOpcode() != LogicOpcode ||
7764 !(ShiftOpcode == ISD::SHL || ShiftOpcode == ISD::SRL ||
7765 ShiftOpcode == ISD::SRA))
7766 return SDValue();
7767
7768 // Match another shift op inside the first logic operand. Handle both commuted
7769 // possibilities.
7770 // LOGIC (LOGIC (SH X0, Y), Z), (SH X1, Y) --> LOGIC (SH (LOGIC X0, X1), Y), Z
7771 // LOGIC (LOGIC Z, (SH X0, Y)), (SH X1, Y) --> LOGIC (SH (LOGIC X0, X1), Y), Z
7772 SDValue X1 = ShiftOp.getOperand(i: 0);
7773 SDValue Y = ShiftOp.getOperand(i: 1);
7774 SDValue X0, Z;
7775 if (LogicOp.getOperand(i: 0).getOpcode() == ShiftOpcode &&
7776 LogicOp.getOperand(i: 0).getOperand(i: 1) == Y) {
7777 X0 = LogicOp.getOperand(i: 0).getOperand(i: 0);
7778 Z = LogicOp.getOperand(i: 1);
7779 } else if (LogicOp.getOperand(i: 1).getOpcode() == ShiftOpcode &&
7780 LogicOp.getOperand(i: 1).getOperand(i: 1) == Y) {
7781 X0 = LogicOp.getOperand(i: 1).getOperand(i: 0);
7782 Z = LogicOp.getOperand(i: 0);
7783 } else {
7784 return SDValue();
7785 }
7786
7787 EVT VT = N->getValueType(ResNo: 0);
7788 SDLoc DL(N);
7789 SDValue LogicX = DAG.getNode(Opcode: LogicOpcode, DL, VT, N1: X0, N2: X1);
7790 SDValue NewShift = DAG.getNode(Opcode: ShiftOpcode, DL, VT, N1: LogicX, N2: Y);
7791 return DAG.getNode(Opcode: LogicOpcode, DL, VT, N1: NewShift, N2: Z);
7792}
7793
7794/// Given a tree of logic operations with shape like
7795/// (LOGIC (LOGIC (X, Y), LOGIC (Z, Y)))
7796/// try to match and fold shift operations with the same shift amount.
7797/// For example:
7798/// LOGIC (LOGIC (SH X0, Y), Z), (LOGIC (SH X1, Y), W) -->
7799/// --> LOGIC (SH (LOGIC X0, X1), Y), (LOGIC Z, W)
7800static SDValue foldLogicTreeOfShifts(SDNode *N, SDValue LeftHand,
7801 SDValue RightHand, SelectionDAG &DAG) {
7802 unsigned LogicOpcode = N->getOpcode();
7803 assert(ISD::isBitwiseLogicOp(LogicOpcode) &&
7804 "Expected bitwise logic operation");
7805 if (LeftHand.getOpcode() != LogicOpcode ||
7806 RightHand.getOpcode() != LogicOpcode)
7807 return SDValue();
7808 if (!LeftHand.hasOneUse() || !RightHand.hasOneUse())
7809 return SDValue();
7810
7811 // Try to match one of following patterns:
7812 // LOGIC (LOGIC (SH X0, Y), Z), (LOGIC (SH X1, Y), W)
7813 // LOGIC (LOGIC (SH X0, Y), Z), (LOGIC W, (SH X1, Y))
7814 // Note that foldLogicOfShifts will handle commuted versions of the left hand
7815 // itself.
7816 SDValue CombinedShifts, W;
7817 SDValue R0 = RightHand.getOperand(i: 0);
7818 SDValue R1 = RightHand.getOperand(i: 1);
7819 if ((CombinedShifts = foldLogicOfShifts(N, LogicOp: LeftHand, ShiftOp: R0, DAG)))
7820 W = R1;
7821 else if ((CombinedShifts = foldLogicOfShifts(N, LogicOp: LeftHand, ShiftOp: R1, DAG)))
7822 W = R0;
7823 else
7824 return SDValue();
7825
7826 EVT VT = N->getValueType(ResNo: 0);
7827 SDLoc DL(N);
7828 return DAG.getNode(Opcode: LogicOpcode, DL, VT, N1: CombinedShifts, N2: W);
7829}
7830
7831/// Fold "masked merge" expressions like `(m & x) | (~m & y)` and its DeMorgan
7832/// variant `(~m | x) & (m | y)` into the equivalent `((x ^ y) & m) ^ y)`
7833/// pattern. This is typically a better representation for targets without a
7834/// fused "and-not" operation.
7835static SDValue foldMaskedMerge(SDNode *Node, SelectionDAG &DAG,
7836 const TargetLowering &TLI, const SDLoc &DL) {
7837 // Note that masked-merge variants using XOR or ADD expressions are
7838 // normalized to OR by InstCombine so we only check for OR or AND.
7839 assert((Node->getOpcode() == ISD::OR || Node->getOpcode() == ISD::AND) &&
7840 "Must be called with ISD::OR or ISD::AND node");
7841
7842 // If the target supports and-not, don't fold this.
7843 if (TLI.hasAndNot(X: SDValue(Node, 0)))
7844 return SDValue();
7845
7846 SDValue M, X, Y;
7847
7848 if (sd_match(N: Node,
7849 P: m_Or(L: m_OneUse(P: m_And(L: m_OneUse(P: m_Not(V: m_Value(N&: M))), R: m_Value(N&: Y))),
7850 R: m_OneUse(P: m_And(L: m_Deferred(V&: M), R: m_Value(N&: X))))) ||
7851 sd_match(N: Node,
7852 P: m_And(L: m_OneUse(P: m_Or(L: m_OneUse(P: m_Not(V: m_Value(N&: M))), R: m_Value(N&: X))),
7853 R: m_OneUse(P: m_Or(L: m_Deferred(V&: M), R: m_Value(N&: Y)))))) {
7854 EVT VT = M.getValueType();
7855 SDValue Xor = DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: X, N2: Y);
7856 SDValue And = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Xor, N2: M);
7857 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: And, N2: Y);
7858 }
7859 return SDValue();
7860}
7861
7862SDValue DAGCombiner::visitAND(SDNode *N) {
7863 SDValue N0 = N->getOperand(Num: 0);
7864 SDValue N1 = N->getOperand(Num: 1);
7865 EVT VT = N1.getValueType();
7866 SDLoc DL(N);
7867
7868 // x & x --> x
7869 if (N0 == N1)
7870 return N0;
7871
7872 // fold (and c1, c2) -> c1&c2
7873 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::AND, DL, VT, Ops: {N0, N1}))
7874 return C;
7875
7876 // canonicalize constant to RHS
7877 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
7878 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
7879 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1, N2: N0);
7880
7881 if (areBitwiseNotOfEachother(Op0: N0, Op1: N1))
7882 return DAG.getConstant(Val: APInt::getZero(numBits: VT.getScalarSizeInBits()), DL, VT);
7883
7884 // fold vector ops
7885 if (VT.isVector()) {
7886 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
7887 return FoldedVOp;
7888
7889 // fold (and x, 0) -> 0, vector edition
7890 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
7891 // do not return N1, because undef node may exist in N1
7892 return DAG.getConstant(Val: APInt::getZero(numBits: N1.getScalarValueSizeInBits()), DL,
7893 VT: N1.getValueType());
7894
7895 // fold (and x, -1) -> x, vector edition
7896 if (ISD::isConstantSplatVectorAllOnes(N: N1.getNode()))
7897 return N0;
7898
7899 // fold (and buildvector(x,0,-1,w), buildvector(0,y,z,w))
7900 // --> buildvector(0,0,z,w)
7901 auto *BV0 = dyn_cast<BuildVectorSDNode>(Val&: N0);
7902 auto *BV1 = dyn_cast<BuildVectorSDNode>(Val&: N1);
7903 if (BV0 && BV1 && !BV0->getSplatValue() && !BV1->getSplatValue() &&
7904 N0.hasOneUse() && N1.hasOneUse() &&
7905 BV0->getOperand(Num: 0).getValueType() ==
7906 BV1->getOperand(Num: 0).getValueType()) {
7907 SmallVector<SDValue> MergedOps;
7908 unsigned NumElts = VT.getVectorNumElements();
7909 EVT EltVT = BV0->getOperand(Num: 0).getValueType();
7910 for (unsigned I = 0; I != NumElts; ++I) {
7911 auto *C0 = dyn_cast<ConstantSDNode>(Val: BV0->getOperand(Num: I));
7912 auto *C1 = dyn_cast<ConstantSDNode>(Val: BV1->getOperand(Num: I));
7913 if (C0 && C1)
7914 MergedOps.push_back(Elt: DAG.getConstant(
7915 Val: C0->getAPIntValue() & C1->getAPIntValue(), DL, VT: EltVT));
7916 else if (C0 && C0->isZero())
7917 MergedOps.push_back(Elt: BV0->getOperand(Num: I));
7918 else if (C1 && C1->isZero())
7919 MergedOps.push_back(Elt: BV1->getOperand(Num: I));
7920 else if (C0 && C0->isAllOnes())
7921 MergedOps.push_back(Elt: BV1->getOperand(Num: I));
7922 else if (C1 && C1->isAllOnes())
7923 MergedOps.push_back(Elt: BV0->getOperand(Num: I));
7924 else if (BV0->getOperand(Num: I) == BV1->getOperand(Num: I))
7925 MergedOps.push_back(Elt: BV0->getOperand(Num: I));
7926 else
7927 break;
7928 }
7929 if (MergedOps.size() == NumElts)
7930 return DAG.getBuildVector(VT, DL, Ops: MergedOps);
7931 }
7932
7933 // fold (and (masked_load) (splat_vec (x, ...))) to zext_masked_load
7934 bool Frozen = N0.getOpcode() == ISD::FREEZE;
7935 auto *MLoad = dyn_cast<MaskedLoadSDNode>(Val: Frozen ? N0.getOperand(i: 0) : N0);
7936 ConstantSDNode *Splat = isConstOrConstSplat(N: N1, AllowUndefs: true, AllowTruncation: true);
7937 if (MLoad && MLoad->getExtensionType() == ISD::EXTLOAD && Splat) {
7938 EVT MemVT = MLoad->getMemoryVT();
7939 if (TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: MLoad->getAlign(),
7940 AddrSpace: MLoad->getAddressSpace(), ExtType: ISD::ZEXTLOAD, Atomic: false)) {
7941 // For this AND to be a zero extension of the masked load the elements
7942 // of the BuildVec must mask the bottom bits of the extended element
7943 // type
7944 if (Splat->getAPIntValue().isMask(numBits: MemVT.getScalarSizeInBits())) {
7945 SDValue NewLoad = DAG.getMaskedLoad(
7946 VT, dl: DL, Chain: MLoad->getChain(), Base: MLoad->getBasePtr(),
7947 Offset: MLoad->getOffset(), Mask: MLoad->getMask(), Src0: MLoad->getPassThru(), MemVT,
7948 MMO: MLoad->getMemOperand(), AM: MLoad->getAddressingMode(), ISD::ZEXTLOAD,
7949 IsExpanding: MLoad->isExpandingLoad());
7950 CombineTo(N, Res: Frozen ? N0 : NewLoad);
7951 CombineTo(N: MLoad, Res0: NewLoad, Res1: NewLoad.getValue(R: 1));
7952 return SDValue(N, 0);
7953 }
7954 }
7955 }
7956 }
7957
7958 // fold (and x, -1) -> x
7959 if (isAllOnesConstant(V: N1))
7960 return N0;
7961
7962 // if (and x, c) is known to be zero, return 0
7963 unsigned BitWidth = VT.getScalarSizeInBits();
7964 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
7965 if (N1C && DAG.MaskedValueIsZero(Op: SDValue(N, 0), Mask: APInt::getAllOnes(numBits: BitWidth)))
7966 return DAG.getConstant(Val: 0, DL, VT);
7967
7968 if (SDValue R = foldAndOrOfSETCC(LogicOp: N, DAG))
7969 return R;
7970
7971 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
7972 return NewSel;
7973
7974 // reassociate and
7975 if (SDValue RAND = reassociateOps(Opc: ISD::AND, DL, N0, N1, Flags: N->getFlags()))
7976 return RAND;
7977
7978 // Fold and(vecreduce(x), vecreduce(y)) -> vecreduce(and(x, y))
7979 if (SDValue SD =
7980 reassociateReduction(RedOpc: ISD::VECREDUCE_AND, Opc: ISD::AND, DL, VT, N0, N1))
7981 return SD;
7982
7983 // fold (and (or x, C), D) -> D if (C & D) == D
7984 auto MatchSubset = [](ConstantSDNode *LHS, ConstantSDNode *RHS) {
7985 return RHS->getAPIntValue().isSubsetOf(RHS: LHS->getAPIntValue());
7986 };
7987 if (N0.getOpcode() == ISD::OR &&
7988 ISD::matchBinaryPredicate(LHS: N0.getOperand(i: 1), RHS: N1, Match: MatchSubset))
7989 return N1;
7990
7991 if (N1C && N0.getOpcode() == ISD::ANY_EXTEND) {
7992 SDValue N0Op0 = N0.getOperand(i: 0);
7993 EVT SrcVT = N0Op0.getValueType();
7994 unsigned SrcBitWidth = SrcVT.getScalarSizeInBits();
7995 APInt Mask = ~N1C->getAPIntValue();
7996 Mask = Mask.trunc(width: SrcBitWidth);
7997
7998 // fold (and (any_ext V), c) -> (zero_ext V) if 'and' only clears top bits.
7999 if (DAG.MaskedValueIsZero(Op: N0Op0, Mask))
8000 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: N0Op0);
8001
8002 // fold (and (any_ext V), c) -> (zero_ext (and V, c)) if profitable, when
8003 // the zext is free or the anyext costs the same as the zext.
8004 if (N1C->getAPIntValue().countLeadingZeros() >= (BitWidth - SrcBitWidth) &&
8005 (TLI.isZExtFree(FromTy: SrcVT, ToTy: VT) || !TLI.isAnyExtFree(FromTy: SrcVT, ToTy: VT)) &&
8006 TLI.isTypeDesirableForOp(ISD::AND, VT: SrcVT) &&
8007 TLI.isNarrowingProfitable(N, SrcVT: VT, DestVT: SrcVT))
8008 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT,
8009 Operand: DAG.getNode(Opcode: ISD::AND, DL, VT: SrcVT, N1: N0Op0,
8010 N2: DAG.getZExtOrTrunc(Op: N1, DL, VT: SrcVT)));
8011 }
8012
8013 // fold (and (ext (and V, c1)), c2) -> (and (ext V), (and c1, (ext c2)))
8014 if (ISD::isExtOpcode(Opcode: N0.getOpcode())) {
8015 unsigned ExtOpc = N0.getOpcode();
8016 SDValue N0Op0 = N0.getOperand(i: 0);
8017 if (N0Op0.getOpcode() == ISD::AND &&
8018 (ExtOpc != ISD::ZERO_EXTEND || !TLI.isZExtFree(Val: N0Op0, VT2: VT)) &&
8019 N0->hasOneUse() && N0Op0->hasOneUse()) {
8020 if (SDValue NewExt = DAG.FoldConstantArithmetic(Opcode: ExtOpc, DL, VT,
8021 Ops: {N0Op0.getOperand(i: 1)})) {
8022 if (SDValue NewMask =
8023 DAG.FoldConstantArithmetic(Opcode: ISD::AND, DL, VT, Ops: {N1, NewExt})) {
8024 return DAG.getNode(Opcode: ISD::AND, DL, VT,
8025 N1: DAG.getNode(Opcode: ExtOpc, DL, VT, Operand: N0Op0.getOperand(i: 0)),
8026 N2: NewMask);
8027 }
8028 }
8029 }
8030 }
8031
8032 // similarly fold (and (X (load ([non_ext|any_ext|zero_ext] V))), c) ->
8033 // (X (load ([non_ext|zero_ext] V))) if 'and' only clears top bits which must
8034 // already be zero by virtue of the width of the base type of the load.
8035 //
8036 // the 'X' node here can either be nothing or an extract_vector_elt to catch
8037 // more cases.
8038 if ((N0.getOpcode() == ISD::EXTRACT_VECTOR_ELT &&
8039 N0.getValueSizeInBits() == N0.getOperand(i: 0).getScalarValueSizeInBits() &&
8040 N0.getOperand(i: 0).getOpcode() == ISD::LOAD &&
8041 N0.getOperand(i: 0).getResNo() == 0) ||
8042 (N0.getOpcode() == ISD::LOAD && N0.getResNo() == 0)) {
8043 auto *Load =
8044 cast<LoadSDNode>(Val: (N0.getOpcode() == ISD::LOAD) ? N0 : N0.getOperand(i: 0));
8045
8046 // Get the constant (if applicable) the zero'th operand is being ANDed with.
8047 // This can be a pure constant or a vector splat, in which case we treat the
8048 // vector as a scalar and use the splat value.
8049 APInt Constant = APInt::getZero(numBits: 1);
8050 if (const ConstantSDNode *C = isConstOrConstSplat(
8051 N: N1, /*AllowUndefs=*/false, /*AllowTruncation=*/true)) {
8052 Constant = C->getAPIntValue();
8053 } else if (BuildVectorSDNode *Vector = dyn_cast<BuildVectorSDNode>(Val&: N1)) {
8054 unsigned EltBitWidth = Vector->getValueType(ResNo: 0).getScalarSizeInBits();
8055 APInt SplatValue, SplatUndef;
8056 unsigned SplatBitSize;
8057 bool HasAnyUndefs;
8058 // Endianness should not matter here. Code below makes sure that we only
8059 // use the result if the SplatBitSize is a multiple of the vector element
8060 // size. And after that we AND all element sized parts of the splat
8061 // together. So the end result should be the same regardless of in which
8062 // order we do those operations.
8063 const bool IsBigEndian = false;
8064 bool IsSplat =
8065 Vector->isConstantSplat(SplatValue, SplatUndef, SplatBitSize,
8066 HasAnyUndefs, MinSplatBits: EltBitWidth, isBigEndian: IsBigEndian);
8067
8068 // Make sure that variable 'Constant' is only set if 'SplatBitSize' is a
8069 // multiple of 'BitWidth'. Otherwise, we could propagate a wrong value.
8070 if (IsSplat && (SplatBitSize % EltBitWidth) == 0) {
8071 // Undef bits can contribute to a possible optimisation if set, so
8072 // set them.
8073 SplatValue |= SplatUndef;
8074
8075 // The splat value may be something like "0x00FFFFFF", which means 0 for
8076 // the first vector value and FF for the rest, repeating. We need a mask
8077 // that will apply equally to all members of the vector, so AND all the
8078 // lanes of the constant together.
8079 Constant = APInt::getAllOnes(numBits: EltBitWidth);
8080 for (unsigned i = 0, n = (SplatBitSize / EltBitWidth); i < n; ++i)
8081 Constant &= SplatValue.extractBits(numBits: EltBitWidth, bitPosition: i * EltBitWidth);
8082 }
8083 }
8084
8085 // If we want to change an EXTLOAD to a ZEXTLOAD, ensure a ZEXTLOAD is
8086 // actually legal and isn't going to get expanded, else this is a false
8087 // optimisation.
8088 bool CanZextLoadProfitably = TLI.isLoadLegal(
8089 ValVT: Load->getValueType(ResNo: 0), MemVT: Load->getMemoryVT(), Alignment: Load->getAlign(),
8090 AddrSpace: Load->getAddressSpace(), ExtType: ISD::ZEXTLOAD, Atomic: false);
8091
8092 // Resize the constant to the same size as the original memory access before
8093 // extension. If it is still the AllOnesValue then this AND is completely
8094 // unneeded.
8095 Constant = Constant.zextOrTrunc(width: Load->getMemoryVT().getScalarSizeInBits());
8096
8097 bool B;
8098 switch (Load->getExtensionType()) {
8099 default: B = false; break;
8100 case ISD::EXTLOAD: B = CanZextLoadProfitably; break;
8101 case ISD::ZEXTLOAD:
8102 case ISD::NON_EXTLOAD: B = true; break;
8103 }
8104
8105 if (B && Constant.isAllOnes()) {
8106 // If the load type was an EXTLOAD, convert to ZEXTLOAD in order to
8107 // preserve semantics once we get rid of the AND.
8108 SDValue NewLoad(Load, 0);
8109
8110 // Fold the AND away. NewLoad may get replaced immediately.
8111 CombineTo(N, Res: (N0.getNode() == Load) ? NewLoad : N0);
8112
8113 if (Load->getExtensionType() == ISD::EXTLOAD) {
8114 NewLoad = DAG.getLoad(AM: Load->getAddressingMode(), ExtType: ISD::ZEXTLOAD,
8115 VT: Load->getValueType(ResNo: 0), dl: SDLoc(Load),
8116 Chain: Load->getChain(), Ptr: Load->getBasePtr(),
8117 Offset: Load->getOffset(), MemVT: Load->getMemoryVT(),
8118 MMO: Load->getMemOperand());
8119 // Replace uses of the EXTLOAD with the new ZEXTLOAD.
8120 if (Load->getNumValues() == 3) {
8121 // PRE/POST_INC loads have 3 values.
8122 SDValue To[] = { NewLoad.getValue(R: 0), NewLoad.getValue(R: 1),
8123 NewLoad.getValue(R: 2) };
8124 CombineTo(N: Load, To, NumTo: 3, AddTo: true);
8125 } else {
8126 CombineTo(N: Load, Res0: NewLoad.getValue(R: 0), Res1: NewLoad.getValue(R: 1));
8127 }
8128 }
8129
8130 return SDValue(N, 0); // Return N so it doesn't get rechecked!
8131 }
8132 }
8133
8134 // Try to convert a constant mask AND into a shuffle clear mask.
8135 if (VT.isVector())
8136 if (SDValue Shuffle = XformToShuffleWithZero(N))
8137 return Shuffle;
8138
8139 if (SDValue Combined = combineCarryDiamond(DAG, TLI, N0, N1, N))
8140 return Combined;
8141
8142 if (N0.getOpcode() == ISD::EXTRACT_SUBVECTOR && N0.hasOneUse() && N1C &&
8143 ISD::isExtOpcode(Opcode: N0.getOperand(i: 0).getOpcode())) {
8144 SDValue Ext = N0.getOperand(i: 0);
8145 EVT ExtVT = Ext->getValueType(ResNo: 0);
8146 SDValue Extendee = Ext->getOperand(Num: 0);
8147
8148 unsigned ScalarWidth = Extendee.getValueType().getScalarSizeInBits();
8149 if (N1C->getAPIntValue().isMask(numBits: ScalarWidth) &&
8150 (!LegalOperations || TLI.isOperationLegal(Op: ISD::ZERO_EXTEND, VT: ExtVT))) {
8151 // (and (extract_subvector (zext|anyext|sext v) _) iN_mask)
8152 // => (extract_subvector (iN_zeroext v))
8153 SDValue ZeroExtExtendee =
8154 DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: ExtVT, Operand: Extendee);
8155
8156 return DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT, N1: ZeroExtExtendee,
8157 N2: N0.getOperand(i: 1));
8158 }
8159 }
8160
8161 // fold (and (masked_gather x)) -> (zext_masked_gather x)
8162 if (auto *GN0 = dyn_cast<MaskedGatherSDNode>(Val&: N0)) {
8163 EVT MemVT = GN0->getMemoryVT();
8164 EVT ScalarVT = MemVT.getScalarType();
8165
8166 if (SDValue(GN0, 0).hasOneUse() &&
8167 isConstantSplatVectorMaskForType(N: N1.getNode(), ScalarTy: ScalarVT) &&
8168 TLI.isVectorLoadExtDesirable(ExtVal: SDValue(N, 0))) {
8169 SDValue Ops[] = {GN0->getChain(), GN0->getPassThru(), GN0->getMask(),
8170 GN0->getBasePtr(), GN0->getIndex(), GN0->getScale()};
8171
8172 SDValue ZExtLoad = DAG.getMaskedGather(
8173 VTs: DAG.getVTList(VT1: VT, VT2: MVT::Other), MemVT, dl: DL, Ops, MMO: GN0->getMemOperand(),
8174 IndexType: GN0->getIndexType(), ExtTy: ISD::ZEXTLOAD);
8175
8176 CombineTo(N, Res: ZExtLoad);
8177 AddToWorklist(N: ZExtLoad.getNode());
8178 // Avoid recheck of N.
8179 return SDValue(N, 0);
8180 }
8181 }
8182
8183 // fold (and (load x), 255) -> (zextload x, i8)
8184 // fold (and (extload x, i16), 255) -> (zextload x, i8)
8185 // fold (and (freeze (load x)), 255) -> (freeze (zextload x, i8))
8186 // fold (and (freeze (extload x, i16)), 255) -> (freeze (zextload x, i8))
8187 if (N1C && !VT.isVector()) {
8188 SDValue Inner = peekThroughFreeze(V: N0);
8189 if (Inner.getOpcode() == ISD::LOAD)
8190 if (SDValue Res = reduceLoadWidth(N))
8191 return Res;
8192 }
8193
8194 if (LegalTypes) {
8195 // Attempt to propagate the AND back up to the leaves which, if they're
8196 // loads, can be combined to narrow loads and the AND node can be removed.
8197 // Perform after legalization so that extend nodes will already be
8198 // combined into the loads.
8199 if (BackwardsPropagateMask(N))
8200 return SDValue(N, 0);
8201 }
8202
8203 if (SDValue Combined = visitANDLike(N0, N1, N))
8204 return Combined;
8205
8206 // Simplify: (and (op x...), (op y...)) -> (op (and x, y))
8207 if (N0.getOpcode() == N1.getOpcode())
8208 if (SDValue V = hoistLogicOpWithSameOpcodeHands(N))
8209 return V;
8210
8211 if (SDValue R = foldLogicOfShifts(N, LogicOp: N0, ShiftOp: N1, DAG))
8212 return R;
8213 if (SDValue R = foldLogicOfShifts(N, LogicOp: N1, ShiftOp: N0, DAG))
8214 return R;
8215
8216 // Fold (and X, (bswap (not Y))) -> (and X, (not (bswap Y)))
8217 // Fold (and X, (bitreverse (not Y))) -> (and X, (not (bitreverse Y)))
8218 SDValue X, Y, Z, NotY;
8219 for (unsigned Opc : {ISD::BSWAP, ISD::BITREVERSE})
8220 if (sd_match(N,
8221 P: m_And(L: m_Value(N&: X), R: m_OneUse(P: m_UnaryOp(Opc, Op: m_Value(N&: NotY))))) &&
8222 sd_match(N: NotY, P: m_Not(V: m_Value(N&: Y))) &&
8223 (TLI.hasAndNot(X: SDValue(N, 0)) || NotY->hasOneUse()))
8224 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: X,
8225 N2: DAG.getNOT(DL, Val: DAG.getNode(Opcode: Opc, DL, VT, Operand: Y), VT));
8226
8227 // Fold (and X, (rot (not Y), Z)) -> (and X, (not (rot Y, Z)))
8228 for (unsigned Opc : {ISD::ROTL, ISD::ROTR})
8229 if (sd_match(N, P: m_And(L: m_Value(N&: X),
8230 R: m_OneUse(P: m_BinOp(Opc, L: m_Value(N&: NotY), R: m_Value(N&: Z))))) &&
8231 sd_match(N: NotY, P: m_Not(V: m_Value(N&: Y))) &&
8232 (TLI.hasAndNot(X: SDValue(N, 0)) || NotY->hasOneUse()))
8233 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: X,
8234 N2: DAG.getNOT(DL, Val: DAG.getNode(Opcode: Opc, DL, VT, N1: Y, N2: Z), VT));
8235
8236 // Fold (and X, (add (not Y), Z)) -> (and X, (not (sub Y, Z)))
8237 // Fold (and X, (sub (not Y), Z)) -> (and X, (not (add Y, Z)))
8238 if (TLI.hasAndNot(X: SDValue(N, 0)))
8239 if (SDValue Folded = foldBitwiseOpWithNeg(N, DL, VT))
8240 return Folded;
8241
8242 // Fold (and (srl X, C), 1) -> (srl X, BW-1) for signbit extraction
8243 // If we are shifting down an extended sign bit, see if we can simplify
8244 // this to shifting the MSB directly to expose further simplifications.
8245 // This pattern often appears after sext_inreg legalization.
8246 APInt Amt;
8247 if (sd_match(N, P: m_And(L: m_Srl(L: m_Value(N&: X), R: m_ConstInt(V&: Amt)), R: m_One())) &&
8248 Amt.ult(RHS: BitWidth - 1) && Amt.uge(RHS: BitWidth - DAG.ComputeNumSignBits(Op: X)))
8249 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: X,
8250 N2: DAG.getShiftAmountConstant(Val: BitWidth - 1, VT, DL));
8251
8252 // Masking the negated extension of a boolean is just the zero-extended
8253 // boolean:
8254 // and (sub 0, zext(bool X)), 1 --> zext(bool X)
8255 // and (sub 0, sext(bool X)), 1 --> zext(bool X)
8256 //
8257 // Note: the SimplifyDemandedBits fold below can make an information-losing
8258 // transform, and then we have no way to find this better fold.
8259 if (sd_match(N, P: m_And(L: m_Sub(L: m_Zero(), R: m_Value(N&: X)), R: m_One()))) {
8260 if (X.getOpcode() == ISD::ZERO_EXTEND &&
8261 X.getOperand(i: 0).getScalarValueSizeInBits() == 1)
8262 return X;
8263 if (X.getOpcode() == ISD::SIGN_EXTEND &&
8264 X.getOperand(i: 0).getScalarValueSizeInBits() == 1)
8265 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: X.getOperand(i: 0));
8266 }
8267
8268 // fold (and (sign_extend_inreg x, i16 to i32), 1) -> (and x, 1)
8269 // fold (and (sra)) -> (and (srl)) when possible.
8270 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
8271 return SDValue(N, 0);
8272
8273 // fold (zext_inreg (extload x)) -> (zextload x)
8274 // fold (zext_inreg (sextload x)) -> (zextload x) iff load has one use
8275 if (ISD::isUNINDEXEDLoad(N: N0.getNode()) &&
8276 (ISD::isEXTLoad(N: N0.getNode()) ||
8277 (ISD::isSEXTLoad(N: N0.getNode()) && N0.hasOneUse()))) {
8278 auto *LN0 = cast<LoadSDNode>(Val&: N0);
8279 EVT MemVT = LN0->getMemoryVT();
8280 // If we zero all the possible extended bits, then we can turn this into
8281 // a zextload if we are running before legalize or the operation is legal.
8282 unsigned ExtBitSize = N1.getScalarValueSizeInBits();
8283 unsigned MemBitSize = MemVT.getScalarSizeInBits();
8284 APInt ExtBits = APInt::getHighBitsSet(numBits: ExtBitSize, hiBitsSet: ExtBitSize - MemBitSize);
8285 if (DAG.MaskedValueIsZero(Op: N1, Mask: ExtBits) &&
8286 ((!LegalOperations && LN0->isSimple()) ||
8287 TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: LN0->getAlign(), AddrSpace: LN0->getAddressSpace(),
8288 ExtType: ISD::ZEXTLOAD, Atomic: false))) {
8289 SDValue ExtLoad =
8290 DAG.getExtLoad(ExtType: ISD::ZEXTLOAD, dl: SDLoc(N0), VT, Chain: LN0->getChain(),
8291 Ptr: LN0->getBasePtr(), MemVT, MMO: LN0->getMemOperand());
8292 AddToWorklist(N);
8293 CombineTo(N: N0.getNode(), Res0: ExtLoad, Res1: ExtLoad.getValue(R: 1));
8294 return SDValue(N, 0); // Return N so it doesn't get rechecked!
8295 }
8296 }
8297
8298 // fold (and (or (srl N, 8), (shl N, 8)), 0xffff) -> (srl (bswap N), const)
8299 if (N1C && N1C->getAPIntValue() == 0xffff && N0.getOpcode() == ISD::OR) {
8300 if (SDValue BSwap = MatchBSwapHWordLow(N: N0.getNode(), N0: N0.getOperand(i: 0),
8301 N1: N0.getOperand(i: 1), DemandHighBits: false))
8302 return BSwap;
8303 }
8304
8305 if (SDValue Shifts = unfoldExtremeBitClearingToShifts(N))
8306 return Shifts;
8307
8308 if (SDValue V = combineShiftAnd1ToBitTest(And: N, DAG))
8309 return V;
8310
8311 // Recognize the following pattern:
8312 //
8313 // AndVT = (and (sign_extend NarrowVT to AndVT) #bitmask)
8314 //
8315 // where bitmask is a mask that clears the upper bits of AndVT. The
8316 // number of bits in bitmask must be a power of two.
8317 auto IsAndZeroExtMask = [](SDValue LHS, SDValue RHS) {
8318 if (LHS->getOpcode() != ISD::SIGN_EXTEND)
8319 return false;
8320
8321 auto *C = isConstOrConstSplat(N: RHS, AllowUndefs: false, AllowTruncation: true);
8322 if (!C)
8323 return false;
8324
8325 if (!C->getAPIntValue().isMask(
8326 numBits: LHS.getOperand(i: 0).getValueType().getScalarSizeInBits()))
8327 return false;
8328
8329 return true;
8330 };
8331
8332 // Replace (and (sign_extend ...) #bitmask) with (zero_extend ...).
8333 if (IsAndZeroExtMask(N0, N1) &&
8334 (!LegalOperations || TLI.isOperationLegal(Op: ISD::ZERO_EXTEND, VT)))
8335 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: N0.getOperand(i: 0));
8336
8337 if (hasOperation(Opcode: ISD::USUBSAT, VT))
8338 if (SDValue V = foldAndToUsubsat(N, DAG, DL))
8339 return V;
8340
8341 // Postpone until legalization completed to avoid interference with bswap
8342 // folding
8343 if (LegalOperations || VT.isVector())
8344 if (SDValue R = foldLogicTreeOfShifts(N, LeftHand: N0, RightHand: N1, DAG))
8345 return R;
8346
8347 if (VT.isScalarInteger() && VT != MVT::i1)
8348 if (SDValue R = foldMaskedMerge(Node: N, DAG, TLI, DL))
8349 return R;
8350
8351 return SDValue();
8352}
8353
8354/// Match (a >> 8) | (a << 8) as (bswap a) >> 16.
8355SDValue DAGCombiner::MatchBSwapHWordLow(SDNode *N, SDValue N0, SDValue N1,
8356 bool DemandHighBits) {
8357 if (!LegalOperations)
8358 return SDValue();
8359
8360 EVT VT = N->getValueType(ResNo: 0);
8361 if (VT != MVT::i64 && VT != MVT::i32 && VT != MVT::i16)
8362 return SDValue();
8363 if (!TLI.isOperationLegalOrCustom(Op: ISD::BSWAP, VT))
8364 return SDValue();
8365
8366 // Recognize (and (shl a, 8), 0xff00), (and (srl a, 8), 0xff)
8367 bool LookPassAnd0 = false;
8368 bool LookPassAnd1 = false;
8369 if (N0.getOpcode() == ISD::AND && N0.getOperand(i: 0).getOpcode() == ISD::SRL)
8370 std::swap(a&: N0, b&: N1);
8371 if (N1.getOpcode() == ISD::AND && N1.getOperand(i: 0).getOpcode() == ISD::SHL)
8372 std::swap(a&: N0, b&: N1);
8373 if (N0.getOpcode() == ISD::AND) {
8374 if (!N0->hasOneUse())
8375 return SDValue();
8376 ConstantSDNode *N01C = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
8377 // Also handle 0xffff since the LHS is guaranteed to have zeros there.
8378 // This is needed for X86.
8379 if (!N01C || (N01C->getZExtValue() != 0xFF00 &&
8380 N01C->getZExtValue() != 0xFFFF))
8381 return SDValue();
8382 N0 = N0.getOperand(i: 0);
8383 LookPassAnd0 = true;
8384 }
8385
8386 if (N1.getOpcode() == ISD::AND) {
8387 if (!N1->hasOneUse())
8388 return SDValue();
8389 ConstantSDNode *N11C = dyn_cast<ConstantSDNode>(Val: N1.getOperand(i: 1));
8390 if (!N11C || N11C->getZExtValue() != 0xFF)
8391 return SDValue();
8392 N1 = N1.getOperand(i: 0);
8393 LookPassAnd1 = true;
8394 }
8395
8396 if (N0.getOpcode() == ISD::SRL && N1.getOpcode() == ISD::SHL)
8397 std::swap(a&: N0, b&: N1);
8398 if (N0.getOpcode() != ISD::SHL || N1.getOpcode() != ISD::SRL)
8399 return SDValue();
8400 if (!N0->hasOneUse() || !N1->hasOneUse())
8401 return SDValue();
8402
8403 ConstantSDNode *N01C = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
8404 ConstantSDNode *N11C = dyn_cast<ConstantSDNode>(Val: N1.getOperand(i: 1));
8405 if (!N01C || !N11C)
8406 return SDValue();
8407 if (N01C->getZExtValue() != 8 || N11C->getZExtValue() != 8)
8408 return SDValue();
8409
8410 // Look for (shl (and a, 0xff), 8), (srl (and a, 0xff00), 8)
8411 SDValue N00 = N0->getOperand(Num: 0);
8412 if (!LookPassAnd0 && N00.getOpcode() == ISD::AND) {
8413 if (!N00->hasOneUse())
8414 return SDValue();
8415 ConstantSDNode *N001C = dyn_cast<ConstantSDNode>(Val: N00.getOperand(i: 1));
8416 if (!N001C || N001C->getZExtValue() != 0xFF)
8417 return SDValue();
8418 N00 = N00.getOperand(i: 0);
8419 LookPassAnd0 = true;
8420 }
8421
8422 SDValue N10 = N1->getOperand(Num: 0);
8423 if (!LookPassAnd1 && N10.getOpcode() == ISD::AND) {
8424 if (!N10->hasOneUse())
8425 return SDValue();
8426 ConstantSDNode *N101C = dyn_cast<ConstantSDNode>(Val: N10.getOperand(i: 1));
8427 // Also allow 0xFFFF since the bits will be shifted out. This is needed
8428 // for X86.
8429 if (!N101C || (N101C->getZExtValue() != 0xFF00 &&
8430 N101C->getZExtValue() != 0xFFFF))
8431 return SDValue();
8432 N10 = N10.getOperand(i: 0);
8433 LookPassAnd1 = true;
8434 }
8435
8436 if (N00 != N10)
8437 return SDValue();
8438
8439 // Make sure everything beyond the low halfword gets set to zero since the SRL
8440 // 16 will clear the top bits.
8441 unsigned OpSizeInBits = VT.getSizeInBits();
8442 if (OpSizeInBits > 16) {
8443 // If the left-shift isn't masked out then the only way this is a bswap is
8444 // if all bits beyond the low 8 are 0. In that case the entire pattern
8445 // reduces to a left shift anyway: leave it for other parts of the combiner.
8446 if (DemandHighBits && !LookPassAnd0)
8447 return SDValue();
8448
8449 // However, if the right shift isn't masked out then it might be because
8450 // it's not needed. See if we can spot that too. If the high bits aren't
8451 // demanded, we only need bits 23:16 to be zero. Otherwise, we need all
8452 // upper bits to be zero.
8453 if (!LookPassAnd1) {
8454 unsigned HighBit = DemandHighBits ? OpSizeInBits : 24;
8455 if (!DAG.MaskedValueIsZero(Op: N10,
8456 Mask: APInt::getBitsSet(numBits: OpSizeInBits, loBit: 16, hiBit: HighBit)))
8457 return SDValue();
8458 }
8459 }
8460
8461 SDValue Res = DAG.getNode(Opcode: ISD::BSWAP, DL: SDLoc(N), VT, Operand: N00);
8462 if (OpSizeInBits > 16) {
8463 SDLoc DL(N);
8464 Res = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: Res,
8465 N2: DAG.getShiftAmountConstant(Val: OpSizeInBits - 16, VT, DL));
8466 }
8467 return Res;
8468}
8469
8470/// Return true if the specified node is an element that makes up a 32-bit
8471/// packed halfword byteswap.
8472/// ((x & 0x000000ff) << 8) |
8473/// ((x & 0x0000ff00) >> 8) |
8474/// ((x & 0x00ff0000) << 8) |
8475/// ((x & 0xff000000) >> 8)
8476static bool isBSwapHWordElement(SDValue N, MutableArrayRef<SDNode *> Parts) {
8477 if (!N->hasOneUse())
8478 return false;
8479
8480 unsigned Opc = N.getOpcode();
8481 if (Opc != ISD::AND && Opc != ISD::SHL && Opc != ISD::SRL)
8482 return false;
8483
8484 SDValue N0 = N.getOperand(i: 0);
8485 unsigned Opc0 = N0.getOpcode();
8486 if (Opc0 != ISD::AND && Opc0 != ISD::SHL && Opc0 != ISD::SRL)
8487 return false;
8488
8489 ConstantSDNode *N1C = nullptr;
8490 // SHL or SRL: look upstream for AND mask operand
8491 if (Opc == ISD::AND)
8492 N1C = dyn_cast<ConstantSDNode>(Val: N.getOperand(i: 1));
8493 else if (Opc0 == ISD::AND)
8494 N1C = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
8495 if (!N1C)
8496 return false;
8497
8498 unsigned MaskByteOffset;
8499 switch (N1C->getZExtValue()) {
8500 default:
8501 return false;
8502 case 0xFF: MaskByteOffset = 0; break;
8503 case 0xFF00: MaskByteOffset = 1; break;
8504 case 0xFFFF:
8505 // In case demanded bits didn't clear the bits that will be shifted out.
8506 // This is needed for X86.
8507 if (Opc == ISD::SRL || (Opc == ISD::AND && Opc0 == ISD::SHL)) {
8508 MaskByteOffset = 1;
8509 break;
8510 }
8511 return false;
8512 case 0xFF0000: MaskByteOffset = 2; break;
8513 case 0xFF000000: MaskByteOffset = 3; break;
8514 }
8515
8516 // Look for (x & 0xff) << 8 as well as ((x << 8) & 0xff00).
8517 if (Opc == ISD::AND) {
8518 if (MaskByteOffset == 0 || MaskByteOffset == 2) {
8519 // (x >> 8) & 0xff
8520 // (x >> 8) & 0xff0000
8521 if (Opc0 != ISD::SRL)
8522 return false;
8523 ConstantSDNode *C = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
8524 if (!C || C->getZExtValue() != 8)
8525 return false;
8526 } else {
8527 // (x << 8) & 0xff00
8528 // (x << 8) & 0xff000000
8529 if (Opc0 != ISD::SHL)
8530 return false;
8531 ConstantSDNode *C = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
8532 if (!C || C->getZExtValue() != 8)
8533 return false;
8534 }
8535 } else if (Opc == ISD::SHL) {
8536 // (x & 0xff) << 8
8537 // (x & 0xff0000) << 8
8538 if (MaskByteOffset != 0 && MaskByteOffset != 2)
8539 return false;
8540 ConstantSDNode *C = dyn_cast<ConstantSDNode>(Val: N.getOperand(i: 1));
8541 if (!C || C->getZExtValue() != 8)
8542 return false;
8543 } else { // Opc == ISD::SRL
8544 // (x & 0xff00) >> 8
8545 // (x & 0xff000000) >> 8
8546 if (MaskByteOffset != 1 && MaskByteOffset != 3)
8547 return false;
8548 ConstantSDNode *C = dyn_cast<ConstantSDNode>(Val: N.getOperand(i: 1));
8549 if (!C || C->getZExtValue() != 8)
8550 return false;
8551 }
8552
8553 if (Parts[MaskByteOffset])
8554 return false;
8555
8556 Parts[MaskByteOffset] = N0.getOperand(i: 0).getNode();
8557 return true;
8558}
8559
8560// Match 2 elements of a packed halfword bswap.
8561static bool isBSwapHWordPair(SDValue N, MutableArrayRef<SDNode *> Parts) {
8562 if (N.getOpcode() == ISD::OR)
8563 return isBSwapHWordElement(N: N.getOperand(i: 0), Parts) &&
8564 isBSwapHWordElement(N: N.getOperand(i: 1), Parts);
8565
8566 if (N.getOpcode() == ISD::SRL && N.getOperand(i: 0).getOpcode() == ISD::BSWAP) {
8567 ConstantSDNode *C = isConstOrConstSplat(N: N.getOperand(i: 1));
8568 if (!C || C->getAPIntValue() != 16)
8569 return false;
8570 Parts[0] = Parts[1] = N.getOperand(i: 0).getOperand(i: 0).getNode();
8571 return true;
8572 }
8573
8574 return false;
8575}
8576
8577// Match this pattern:
8578// (or (and (shl (A, 8)), 0xff00ff00), (and (srl (A, 8)), 0x00ff00ff))
8579// And rewrite this to:
8580// (rotr (bswap A), 16)
8581static SDValue matchBSwapHWordOrAndAnd(const TargetLowering &TLI,
8582 SelectionDAG &DAG, SDNode *N, SDValue N0,
8583 SDValue N1, EVT VT) {
8584 assert(N->getOpcode() == ISD::OR && VT == MVT::i32 &&
8585 "MatchBSwapHWordOrAndAnd: expecting i32");
8586 if (!TLI.isOperationLegalOrCustom(Op: ISD::ROTR, VT))
8587 return SDValue();
8588 if (N0.getOpcode() != ISD::AND || N1.getOpcode() != ISD::AND)
8589 return SDValue();
8590 // TODO: this is too restrictive; lifting this restriction requires more tests
8591 if (!N0->hasOneUse() || !N1->hasOneUse())
8592 return SDValue();
8593 ConstantSDNode *Mask0 = isConstOrConstSplat(N: N0.getOperand(i: 1));
8594 ConstantSDNode *Mask1 = isConstOrConstSplat(N: N1.getOperand(i: 1));
8595 if (!Mask0 || !Mask1)
8596 return SDValue();
8597 if (Mask0->getAPIntValue() != 0xff00ff00 ||
8598 Mask1->getAPIntValue() != 0x00ff00ff)
8599 return SDValue();
8600 SDValue Shift0 = N0.getOperand(i: 0);
8601 SDValue Shift1 = N1.getOperand(i: 0);
8602 if (Shift0.getOpcode() != ISD::SHL || Shift1.getOpcode() != ISD::SRL)
8603 return SDValue();
8604 ConstantSDNode *ShiftAmt0 = isConstOrConstSplat(N: Shift0.getOperand(i: 1));
8605 ConstantSDNode *ShiftAmt1 = isConstOrConstSplat(N: Shift1.getOperand(i: 1));
8606 if (!ShiftAmt0 || !ShiftAmt1)
8607 return SDValue();
8608 if (ShiftAmt0->getAPIntValue() != 8 || ShiftAmt1->getAPIntValue() != 8)
8609 return SDValue();
8610 if (Shift0.getOperand(i: 0) != Shift1.getOperand(i: 0))
8611 return SDValue();
8612
8613 SDLoc DL(N);
8614 SDValue BSwap = DAG.getNode(Opcode: ISD::BSWAP, DL, VT, Operand: Shift0.getOperand(i: 0));
8615 SDValue ShAmt = DAG.getShiftAmountConstant(Val: 16, VT, DL);
8616 return DAG.getNode(Opcode: ISD::ROTR, DL, VT, N1: BSwap, N2: ShAmt);
8617}
8618
8619/// Match a 32-bit packed halfword bswap. That is
8620/// ((x & 0x000000ff) << 8) |
8621/// ((x & 0x0000ff00) >> 8) |
8622/// ((x & 0x00ff0000) << 8) |
8623/// ((x & 0xff000000) >> 8)
8624/// => (rotl (bswap x), 16)
8625SDValue DAGCombiner::MatchBSwapHWord(SDNode *N, SDValue N0, SDValue N1) {
8626 if (!LegalOperations)
8627 return SDValue();
8628
8629 EVT VT = N->getValueType(ResNo: 0);
8630 if (VT != MVT::i32)
8631 return SDValue();
8632 if (!TLI.isOperationLegalOrCustom(Op: ISD::BSWAP, VT))
8633 return SDValue();
8634
8635 if (SDValue BSwap = matchBSwapHWordOrAndAnd(TLI, DAG, N, N0, N1, VT))
8636 return BSwap;
8637
8638 // Try again with commuted operands.
8639 if (SDValue BSwap = matchBSwapHWordOrAndAnd(TLI, DAG, N, N0: N1, N1: N0, VT))
8640 return BSwap;
8641
8642
8643 // Look for either
8644 // (or (bswaphpair), (bswaphpair))
8645 // (or (or (bswaphpair), (and)), (and))
8646 // (or (or (and), (bswaphpair)), (and))
8647 SDNode *Parts[4] = {};
8648
8649 if (isBSwapHWordPair(N: N0, Parts)) {
8650 // (or (or (and), (and)), (or (and), (and)))
8651 if (!isBSwapHWordPair(N: N1, Parts))
8652 return SDValue();
8653 } else if (N0.getOpcode() == ISD::OR) {
8654 // (or (or (or (and), (and)), (and)), (and))
8655 if (!isBSwapHWordElement(N: N1, Parts))
8656 return SDValue();
8657 SDValue N00 = N0.getOperand(i: 0);
8658 SDValue N01 = N0.getOperand(i: 1);
8659 if (!(isBSwapHWordElement(N: N01, Parts) && isBSwapHWordPair(N: N00, Parts)) &&
8660 !(isBSwapHWordElement(N: N00, Parts) && isBSwapHWordPair(N: N01, Parts)))
8661 return SDValue();
8662 } else {
8663 return SDValue();
8664 }
8665
8666 // Make sure the parts are all coming from the same node.
8667 if (Parts[0] != Parts[1] || Parts[0] != Parts[2] || Parts[0] != Parts[3])
8668 return SDValue();
8669
8670 SDLoc DL(N);
8671 SDValue BSwap = DAG.getNode(Opcode: ISD::BSWAP, DL, VT,
8672 Operand: SDValue(Parts[0], 0));
8673
8674 // Result of the bswap should be rotated by 16. If it's not legal, then
8675 // do (x << 16) | (x >> 16).
8676 SDValue ShAmt = DAG.getShiftAmountConstant(Val: 16, VT, DL);
8677 if (TLI.isOperationLegalOrCustom(Op: ISD::ROTL, VT))
8678 return DAG.getNode(Opcode: ISD::ROTL, DL, VT, N1: BSwap, N2: ShAmt);
8679 if (TLI.isOperationLegalOrCustom(Op: ISD::ROTR, VT))
8680 return DAG.getNode(Opcode: ISD::ROTR, DL, VT, N1: BSwap, N2: ShAmt);
8681 return DAG.getNode(Opcode: ISD::OR, DL, VT,
8682 N1: DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: BSwap, N2: ShAmt),
8683 N2: DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: BSwap, N2: ShAmt));
8684}
8685
8686/// This contains all DAGCombine rules which reduce two values combined by
8687/// an Or operation to a single value \see visitANDLike().
8688SDValue DAGCombiner::visitORLike(SDValue N0, SDValue N1, const SDLoc &DL) {
8689 EVT VT = N1.getValueType();
8690
8691 // fold (or x, undef) -> -1
8692 if (!LegalOperations && (N0.isUndef() || N1.isUndef()))
8693 return DAG.getAllOnesConstant(DL, VT);
8694
8695 if (SDValue V = foldLogicOfSetCCs(IsAnd: false, N0, N1, DL))
8696 return V;
8697
8698 // (or (and X, C1), (and Y, C2)) -> (and (or X, Y), C3) if possible.
8699 if (N0.getOpcode() == ISD::AND && N1.getOpcode() == ISD::AND &&
8700 // Don't increase # computations.
8701 (N0->hasOneUse() || N1->hasOneUse())) {
8702 // We can only do this xform if we know that bits from X that are set in C2
8703 // but not in C1 are already zero. Likewise for Y.
8704 if (const ConstantSDNode *N0O1C =
8705 getAsNonOpaqueConstant(N: N0.getOperand(i: 1))) {
8706 if (const ConstantSDNode *N1O1C =
8707 getAsNonOpaqueConstant(N: N1.getOperand(i: 1))) {
8708 // We can only do this xform if we know that bits from X that are set in
8709 // C2 but not in C1 are already zero. Likewise for Y.
8710 const APInt &LHSMask = N0O1C->getAPIntValue();
8711 const APInt &RHSMask = N1O1C->getAPIntValue();
8712
8713 if (DAG.MaskedValueIsZero(Op: N0.getOperand(i: 0), Mask: RHSMask&~LHSMask) &&
8714 DAG.MaskedValueIsZero(Op: N1.getOperand(i: 0), Mask: LHSMask&~RHSMask)) {
8715 SDValue X = DAG.getNode(Opcode: ISD::OR, DL: SDLoc(N0), VT,
8716 N1: N0.getOperand(i: 0), N2: N1.getOperand(i: 0));
8717 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: X,
8718 N2: DAG.getConstant(Val: LHSMask | RHSMask, DL, VT));
8719 }
8720 }
8721 }
8722 }
8723
8724 // (or (and X, M), (and X, N)) -> (and X, (or M, N))
8725 if (N0.getOpcode() == ISD::AND &&
8726 N1.getOpcode() == ISD::AND &&
8727 N0.getOperand(i: 0) == N1.getOperand(i: 0) &&
8728 // Don't increase # computations.
8729 (N0->hasOneUse() || N1->hasOneUse())) {
8730 SDValue X = DAG.getNode(Opcode: ISD::OR, DL: SDLoc(N0), VT,
8731 N1: N0.getOperand(i: 1), N2: N1.getOperand(i: 1));
8732 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N0.getOperand(i: 0), N2: X);
8733 }
8734
8735 return SDValue();
8736}
8737
8738/// OR combines for which the commuted variant will be tried as well.
8739static SDValue visitORCommutative(SelectionDAG &DAG, SDValue N0, SDValue N1,
8740 SDNode *N) {
8741 EVT VT = N0.getValueType();
8742 unsigned BW = VT.getScalarSizeInBits();
8743 SDLoc DL(N);
8744
8745 auto peekThroughResize = [](SDValue V) {
8746 if (V->getOpcode() == ISD::ZERO_EXTEND || V->getOpcode() == ISD::TRUNCATE)
8747 return V->getOperand(Num: 0);
8748 return V;
8749 };
8750
8751 SDValue N0Resized = peekThroughResize(N0);
8752 if (N0Resized.getOpcode() == ISD::AND) {
8753 SDValue N1Resized = peekThroughResize(N1);
8754 SDValue N00 = N0Resized.getOperand(i: 0);
8755 SDValue N01 = N0Resized.getOperand(i: 1);
8756
8757 // fold or (and x, y), x --> x
8758 if (N00 == N1Resized || N01 == N1Resized)
8759 return N1;
8760
8761 // fold (or (and X, (xor Y, -1)), Y) -> (or X, Y)
8762 // TODO: Set AllowUndefs = true.
8763 if (SDValue NotOperand = getBitwiseNotOperand(V: N01, Mask: N00,
8764 /* AllowUndefs */ false)) {
8765 if (peekThroughResize(NotOperand) == N1Resized)
8766 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: DAG.getZExtOrTrunc(Op: N00, DL, VT),
8767 N2: N1);
8768 }
8769
8770 // fold (or (and (xor Y, -1), X), Y) -> (or X, Y)
8771 if (SDValue NotOperand = getBitwiseNotOperand(V: N00, Mask: N01,
8772 /* AllowUndefs */ false)) {
8773 if (peekThroughResize(NotOperand) == N1Resized)
8774 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: DAG.getZExtOrTrunc(Op: N01, DL, VT),
8775 N2: N1);
8776 }
8777 }
8778
8779 SDValue X, Y;
8780
8781 // fold or (xor X, N1), N1 --> or X, N1
8782 if (sd_match(N: N0, P: m_Xor(L: m_Value(N&: X), R: m_Specific(N: N1))))
8783 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: X, N2: N1);
8784
8785 // fold or (xor x, y), (x and/or y) --> or x, y
8786 if (sd_match(N: N0, P: m_Xor(L: m_Value(N&: X), R: m_Value(N&: Y))) &&
8787 (sd_match(N: N1, P: m_And(L: m_Specific(N: X), R: m_Specific(N: Y))) ||
8788 sd_match(N: N1, P: m_Or(L: m_Specific(N: X), R: m_Specific(N: Y)))))
8789 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: X, N2: Y);
8790
8791 if (SDValue R = foldLogicOfShifts(N, LogicOp: N0, ShiftOp: N1, DAG))
8792 return R;
8793
8794 auto peekThroughZext = [](SDValue V) {
8795 if (V->getOpcode() == ISD::ZERO_EXTEND)
8796 return V->getOperand(Num: 0);
8797 return V;
8798 };
8799
8800 if (N0.getOpcode() == ISD::FSHL && N1.getOpcode() == ISD::SHL &&
8801 peekThroughZext(N0.getOperand(i: 2)) == peekThroughZext(N1.getOperand(i: 1))) {
8802 // (fshl X, ?, Y) | (shl X, Y) --> fshl X, ?, Y
8803 if (N0.getOperand(i: 0) == N1.getOperand(i: 0))
8804 return N0;
8805 // (fshl A, X, Y) | (shl X, Y) --> fshl (A|X), X, Y
8806 if (N0.getOperand(i: 1) == N1.getOperand(i: 0) && N0.hasOneUse() &&
8807 N1.hasOneUse()) {
8808 SDValue A = N0.getOperand(i: 0);
8809 SDValue X = N1.getOperand(i: 0);
8810 SDValue NewLHS = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: A, N2: X);
8811 return DAG.getNode(Opcode: ISD::FSHL, DL, VT, N1: NewLHS, N2: X, N3: N0.getOperand(i: 2));
8812 }
8813 }
8814
8815 if (N0.getOpcode() == ISD::FSHR && N1.getOpcode() == ISD::SRL &&
8816 peekThroughZext(N0.getOperand(i: 2)) == peekThroughZext(N1.getOperand(i: 1))) {
8817 // (fshr ?, X, Y) | (srl X, Y) --> fshr ?, X, Y
8818 if (N0.getOperand(i: 1) == N1.getOperand(i: 0))
8819 return N0;
8820 // (fshr X, B, Y) | (srl X, Y) --> fshr X, (X|B), Y
8821 if (N0.getOperand(i: 0) == N1.getOperand(i: 0) && N0.hasOneUse() &&
8822 N1.hasOneUse()) {
8823 SDValue X = N1.getOperand(i: 0);
8824 SDValue B = N0.getOperand(i: 1);
8825 SDValue NewRHS = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: X, N2: B);
8826 return DAG.getNode(Opcode: ISD::FSHR, DL, VT, N1: X, N2: NewRHS, N3: N0.getOperand(i: 2));
8827 }
8828 }
8829
8830 // (fshl A, B, S0) | (fshr C, D, S1) --> fshl (A|C), (B|D), S0
8831 // iff S0 + S1 == bitwidth(S1)
8832 if (N0.getOpcode() == ISD::FSHL && N1.getOpcode() == ISD::FSHR &&
8833 N0.hasOneUse() && N1.hasOneUse()) {
8834 auto *S0 = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 2));
8835 auto *S1 = dyn_cast<ConstantSDNode>(Val: N1.getOperand(i: 2));
8836 if (S0 && S1 && S0->getZExtValue() < BW && S1->getZExtValue() < BW &&
8837 S0->getZExtValue() == (BW - S1->getZExtValue())) {
8838 SDValue A = N0.getOperand(i: 0);
8839 SDValue B = N0.getOperand(i: 1);
8840 SDValue C = N1.getOperand(i: 0);
8841 SDValue D = N1.getOperand(i: 1);
8842 SDValue NewLHS = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: A, N2: C);
8843 SDValue NewRHS = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: B, N2: D);
8844 return DAG.getNode(Opcode: ISD::FSHL, DL, VT, N1: NewLHS, N2: NewRHS, N3: N0.getOperand(i: 2));
8845 }
8846 }
8847
8848 // Attempt to match a legalized build_pair-esque pattern:
8849 // or(shl(aext(Hi),BW/2),zext(Lo))
8850 SDValue Lo, Hi;
8851 if (sd_match(N: N0,
8852 P: m_OneUse(P: m_Shl(L: m_AnyExt(Op: m_Value(N&: Hi)), R: m_SpecificInt(V: BW / 2)))) &&
8853 sd_match(N: N1, P: m_ZExt(Op: m_Value(N&: Lo))) &&
8854 Lo.getScalarValueSizeInBits() == (BW / 2) &&
8855 Lo.getValueType() == Hi.getValueType()) {
8856 // Fold build_pair(not(Lo),not(Hi)) -> not(build_pair(Lo,Hi)).
8857 SDValue NotLo, NotHi;
8858 if (sd_match(N: Lo, P: m_OneUse(P: m_Not(V: m_Value(N&: NotLo)))) &&
8859 sd_match(N: Hi, P: m_OneUse(P: m_Not(V: m_Value(N&: NotHi))))) {
8860 Lo = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: NotLo);
8861 Hi = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT, Operand: NotHi);
8862 Hi = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Hi,
8863 N2: DAG.getShiftAmountConstant(Val: BW / 2, VT, DL));
8864 return DAG.getNOT(DL, Val: DAG.getNode(Opcode: ISD::OR, DL, VT, N1: Lo, N2: Hi), VT);
8865 }
8866 }
8867
8868 return SDValue();
8869}
8870
8871SDValue DAGCombiner::visitOR(SDNode *N) {
8872 SDValue N0 = N->getOperand(Num: 0);
8873 SDValue N1 = N->getOperand(Num: 1);
8874 EVT VT = N1.getValueType();
8875 SDLoc DL(N);
8876
8877 // x | x --> x
8878 if (N0 == N1)
8879 return N0;
8880
8881 // fold (or c1, c2) -> c1|c2
8882 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::OR, DL, VT, Ops: {N0, N1}))
8883 return C;
8884
8885 // canonicalize constant to RHS
8886 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
8887 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
8888 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1, N2: N0);
8889
8890 // fold vector ops
8891 if (VT.isVector()) {
8892 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
8893 return FoldedVOp;
8894
8895 // fold (or x, 0) -> x, vector edition
8896 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
8897 return N0;
8898
8899 // fold (or x, -1) -> -1, vector edition
8900 if (ISD::isConstantSplatVectorAllOnes(N: N1.getNode()))
8901 // do not return N1, because undef node may exist in N1
8902 return DAG.getAllOnesConstant(DL, VT: N1.getValueType());
8903
8904 // fold (or buildvector(x,0,-1,w), buildvector(0,y,z,w))
8905 // --> buildvector(x,y,-1,w)
8906 auto *BV0 = dyn_cast<BuildVectorSDNode>(Val&: N0);
8907 auto *BV1 = dyn_cast<BuildVectorSDNode>(Val&: N1);
8908 if (BV0 && BV1 && !BV0->getSplatValue() && !BV1->getSplatValue() &&
8909 N0.hasOneUse() && N1.hasOneUse() &&
8910 BV0->getOperand(Num: 0).getValueType() ==
8911 BV1->getOperand(Num: 0).getValueType()) {
8912 SmallVector<SDValue> MergedOps;
8913 unsigned NumElts = VT.getVectorNumElements();
8914 EVT EltVT = BV0->getOperand(Num: 0).getValueType();
8915 for (unsigned I = 0; I != NumElts; ++I) {
8916 auto *C0 = dyn_cast<ConstantSDNode>(Val: BV0->getOperand(Num: I));
8917 auto *C1 = dyn_cast<ConstantSDNode>(Val: BV1->getOperand(Num: I));
8918 if (C0 && C1)
8919 MergedOps.push_back(Elt: DAG.getConstant(
8920 Val: C0->getAPIntValue() | C1->getAPIntValue(), DL, VT: EltVT));
8921 else if (C0 && C0->isZero())
8922 MergedOps.push_back(Elt: BV1->getOperand(Num: I));
8923 else if (C1 && C1->isZero())
8924 MergedOps.push_back(Elt: BV0->getOperand(Num: I));
8925 else if (C0 && C0->isAllOnes())
8926 MergedOps.push_back(Elt: BV0->getOperand(Num: I));
8927 else if (C1 && C1->isAllOnes())
8928 MergedOps.push_back(Elt: BV1->getOperand(Num: I));
8929 else if (BV0->getOperand(Num: I) == BV1->getOperand(Num: I))
8930 MergedOps.push_back(Elt: BV0->getOperand(Num: I));
8931 else
8932 break;
8933 }
8934 if (MergedOps.size() == NumElts)
8935 return DAG.getBuildVector(VT, DL, Ops: MergedOps);
8936 }
8937
8938 // fold (or (shuf A, V_0, MA), (shuf B, V_0, MB)) -> (shuf A, B, Mask)
8939 // Do this only if the resulting type / shuffle is legal.
8940 auto *SV0 = dyn_cast<ShuffleVectorSDNode>(Val&: N0);
8941 auto *SV1 = dyn_cast<ShuffleVectorSDNode>(Val&: N1);
8942 if (SV0 && SV1 && TLI.isTypeLegal(VT)) {
8943 bool ZeroN00 = ISD::isBuildVectorAllZeros(N: N0.getOperand(i: 0).getNode());
8944 bool ZeroN01 = ISD::isBuildVectorAllZeros(N: N0.getOperand(i: 1).getNode());
8945 bool ZeroN10 = ISD::isBuildVectorAllZeros(N: N1.getOperand(i: 0).getNode());
8946 bool ZeroN11 = ISD::isBuildVectorAllZeros(N: N1.getOperand(i: 1).getNode());
8947 // Ensure both shuffles have a zero input.
8948 if ((ZeroN00 != ZeroN01) && (ZeroN10 != ZeroN11)) {
8949 assert((!ZeroN00 || !ZeroN01) && "Both inputs zero!");
8950 assert((!ZeroN10 || !ZeroN11) && "Both inputs zero!");
8951 bool CanFold = true;
8952 int NumElts = VT.getVectorNumElements();
8953 SmallVector<int, 4> Mask(NumElts, -1);
8954
8955 for (int i = 0; i != NumElts; ++i) {
8956 int M0 = SV0->getMaskElt(Idx: i);
8957 int M1 = SV1->getMaskElt(Idx: i);
8958
8959 // Determine if either index is pointing to a zero vector.
8960 bool M0Zero = M0 < 0 || (ZeroN00 == (M0 < NumElts));
8961 bool M1Zero = M1 < 0 || (ZeroN10 == (M1 < NumElts));
8962
8963 // If one element is zero and the otherside is undef, keep undef.
8964 // This also handles the case that both are undef.
8965 if ((M0Zero && M1 < 0) || (M1Zero && M0 < 0))
8966 continue;
8967
8968 // Make sure only one of the elements is zero.
8969 if (M0Zero == M1Zero) {
8970 CanFold = false;
8971 break;
8972 }
8973
8974 assert((M0 >= 0 || M1 >= 0) && "Undef index!");
8975
8976 // We have a zero and non-zero element. If the non-zero came from
8977 // SV0 make the index a LHS index. If it came from SV1, make it
8978 // a RHS index. We need to mod by NumElts because we don't care
8979 // which operand it came from in the original shuffles.
8980 Mask[i] = M1Zero ? M0 % NumElts : (M1 % NumElts) + NumElts;
8981 }
8982
8983 if (CanFold) {
8984 SDValue NewLHS = ZeroN00 ? N0.getOperand(i: 1) : N0.getOperand(i: 0);
8985 SDValue NewRHS = ZeroN10 ? N1.getOperand(i: 1) : N1.getOperand(i: 0);
8986 SDValue LegalShuffle =
8987 TLI.buildLegalVectorShuffle(VT, DL, N0: NewLHS, N1: NewRHS, Mask, DAG);
8988 if (LegalShuffle)
8989 return LegalShuffle;
8990 }
8991 }
8992 }
8993 }
8994
8995 // fold (or x, 0) -> x
8996 if (isNullConstant(V: N1))
8997 return N0;
8998
8999 // fold (or x, -1) -> -1
9000 if (isAllOnesConstant(V: N1))
9001 return N1;
9002
9003 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
9004 return NewSel;
9005
9006 // fold (or x, c) -> c iff (x & ~c) == 0
9007 ConstantSDNode *N1C = dyn_cast<ConstantSDNode>(Val&: N1);
9008 if (N1C && DAG.MaskedValueIsZero(Op: N0, Mask: ~N1C->getAPIntValue()))
9009 return N1;
9010
9011 if (SDValue R = foldAndOrOfSETCC(LogicOp: N, DAG))
9012 return R;
9013
9014 if (SDValue Combined = visitORLike(N0, N1, DL))
9015 return Combined;
9016
9017 if (SDValue Combined = combineCarryDiamond(DAG, TLI, N0, N1, N))
9018 return Combined;
9019
9020 if (SDValue Combined = combineOrOfSetCCToUSUBOCarry(N, DAG, TLI))
9021 return Combined;
9022
9023 // Recognize halfword bswaps as (bswap + rotl 16) or (bswap + shl 16)
9024 if (SDValue BSwap = MatchBSwapHWord(N, N0, N1))
9025 return BSwap;
9026 if (SDValue BSwap = MatchBSwapHWordLow(N, N0, N1))
9027 return BSwap;
9028
9029 // reassociate or
9030 if (SDValue ROR = reassociateOps(Opc: ISD::OR, DL, N0, N1, Flags: N->getFlags()))
9031 return ROR;
9032
9033 // Fold or(vecreduce(x), vecreduce(y)) -> vecreduce(or(x, y))
9034 if (SDValue SD =
9035 reassociateReduction(RedOpc: ISD::VECREDUCE_OR, Opc: ISD::OR, DL, VT, N0, N1))
9036 return SD;
9037
9038 // Canonicalize (or (and X, c1), c2) -> (and (or X, c2), c1|c2)
9039 // iff (c1 & c2) != 0 or c1/c2 are undef.
9040 auto MatchIntersect = [](ConstantSDNode *C1, ConstantSDNode *C2) {
9041 return !C1 || !C2 || C1->getAPIntValue().intersects(RHS: C2->getAPIntValue());
9042 };
9043 if (N0.getOpcode() == ISD::AND && N0->hasOneUse() &&
9044 ISD::matchBinaryPredicate(LHS: N0.getOperand(i: 1), RHS: N1, Match: MatchIntersect, AllowUndefs: true)) {
9045 if (SDValue COR = DAG.FoldConstantArithmetic(Opcode: ISD::OR, DL: SDLoc(N1), VT,
9046 Ops: {N1, N0.getOperand(i: 1)})) {
9047 SDValue IOR = DAG.getNode(Opcode: ISD::OR, DL: SDLoc(N0), VT, N1: N0.getOperand(i: 0), N2: N1);
9048 AddToWorklist(N: IOR.getNode());
9049 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: COR, N2: IOR);
9050 }
9051 }
9052
9053 if (SDValue Combined = visitORCommutative(DAG, N0, N1, N))
9054 return Combined;
9055 if (SDValue Combined = visitORCommutative(DAG, N0: N1, N1: N0, N))
9056 return Combined;
9057
9058 // Simplify: (or (op x...), (op y...)) -> (op (or x, y))
9059 if (N0.getOpcode() == N1.getOpcode())
9060 if (SDValue V = hoistLogicOpWithSameOpcodeHands(N))
9061 return V;
9062
9063 // See if this is some rotate idiom.
9064 if (SDValue Rot = MatchRotate(LHS: N0, RHS: N1, DL, /*FromAdd=*/false))
9065 return Rot;
9066
9067 if (SDValue Load = MatchLoadCombine(N))
9068 return Load;
9069
9070 // Simplify the operands using demanded-bits information.
9071 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
9072 return SDValue(N, 0);
9073
9074 // If OR can be rewritten into ADD, try combines based on ADD.
9075 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::ADD, VT)) &&
9076 DAG.isADDLike(Op: SDValue(N, 0)))
9077 if (SDValue Combined = visitADDLike(N))
9078 return Combined;
9079
9080 // Postpone until legalization completed to avoid interference with bswap
9081 // folding
9082 if (LegalOperations || VT.isVector())
9083 if (SDValue R = foldLogicTreeOfShifts(N, LeftHand: N0, RightHand: N1, DAG))
9084 return R;
9085
9086 if (VT.isScalarInteger() && VT != MVT::i1)
9087 if (SDValue R = foldMaskedMerge(Node: N, DAG, TLI, DL))
9088 return R;
9089
9090 return SDValue();
9091}
9092
9093static SDValue stripConstantMask(const SelectionDAG &DAG, SDValue Op,
9094 SDValue &Mask) {
9095 if (Op.getOpcode() == ISD::AND &&
9096 DAG.isConstantIntBuildVectorOrConstantInt(N: Op.getOperand(i: 1))) {
9097 Mask = Op.getOperand(i: 1);
9098 return Op.getOperand(i: 0);
9099 }
9100 return Op;
9101}
9102
9103/// Match "(X shl/srl V1) & V2" where V2 may not be present.
9104static bool matchRotateHalf(const SelectionDAG &DAG, SDValue Op, SDValue &Shift,
9105 SDValue &Mask) {
9106 Op = stripConstantMask(DAG, Op, Mask);
9107 if (Op.getOpcode() == ISD::SRL || Op.getOpcode() == ISD::SHL) {
9108 Shift = Op;
9109 return true;
9110 }
9111 return false;
9112}
9113
9114/// Helper function for visitOR to extract the needed side of a rotate idiom
9115/// from a shl/srl/mul/udiv. This is meant to handle cases where
9116/// InstCombine merged some outside op with one of the shifts from
9117/// the rotate pattern.
9118/// \returns An empty \c SDValue if the needed shift couldn't be extracted.
9119/// Otherwise, returns an expansion of \p ExtractFrom based on the following
9120/// patterns:
9121///
9122/// (or (add v v) (shrl v bitwidth-1)):
9123/// expands (add v v) -> (shl v 1)
9124///
9125/// (or (mul v c0) (shrl (mul v c1) c2)):
9126/// expands (mul v c0) -> (shl (mul v c1) c3)
9127///
9128/// (or (udiv v c0) (shl (udiv v c1) c2)):
9129/// expands (udiv v c0) -> (shrl (udiv v c1) c3)
9130///
9131/// (or (shl v c0) (shrl (shl v c1) c2)):
9132/// expands (shl v c0) -> (shl (shl v c1) c3)
9133///
9134/// (or (shrl v c0) (shl (shrl v c1) c2)):
9135/// expands (shrl v c0) -> (shrl (shrl v c1) c3)
9136///
9137/// Such that in all cases, c3+c2==bitwidth(op v c1).
9138static SDValue extractShiftForRotate(SelectionDAG &DAG, SDValue OppShift,
9139 SDValue ExtractFrom, SDValue &Mask,
9140 const SDLoc &DL) {
9141 assert(OppShift && ExtractFrom && "Empty SDValue");
9142 if (OppShift.getOpcode() != ISD::SHL && OppShift.getOpcode() != ISD::SRL)
9143 return SDValue();
9144
9145 ExtractFrom = stripConstantMask(DAG, Op: ExtractFrom, Mask);
9146
9147 // Value and Type of the shift.
9148 SDValue OppShiftLHS = OppShift.getOperand(i: 0);
9149 EVT ShiftedVT = OppShiftLHS.getValueType();
9150
9151 // Amount of the existing shift.
9152 ConstantSDNode *OppShiftCst = isConstOrConstSplat(N: OppShift.getOperand(i: 1));
9153
9154 // (add v v) -> (shl v 1)
9155 // TODO: Should this be a general DAG canonicalization?
9156 if (OppShift.getOpcode() == ISD::SRL && OppShiftCst &&
9157 ExtractFrom.getOpcode() == ISD::ADD &&
9158 ExtractFrom.getOperand(i: 0) == ExtractFrom.getOperand(i: 1) &&
9159 ExtractFrom.getOperand(i: 0) == OppShiftLHS &&
9160 OppShiftCst->getAPIntValue() == ShiftedVT.getScalarSizeInBits() - 1)
9161 return DAG.getNode(Opcode: ISD::SHL, DL, VT: ShiftedVT, N1: OppShiftLHS,
9162 N2: DAG.getShiftAmountConstant(Val: 1, VT: ShiftedVT, DL));
9163
9164 // Preconditions:
9165 // (or (op0 v c0) (shiftl/r (op0 v c1) c2))
9166 //
9167 // Find opcode of the needed shift to be extracted from (op0 v c0).
9168 unsigned Opcode = ISD::DELETED_NODE;
9169 bool IsMulOrDiv = false;
9170 // Set Opcode and IsMulOrDiv if the extract opcode matches the needed shift
9171 // opcode or its arithmetic (mul or udiv) variant.
9172 auto SelectOpcode = [&](unsigned NeededShift, unsigned MulOrDivVariant) {
9173 IsMulOrDiv = ExtractFrom.getOpcode() == MulOrDivVariant;
9174 if (!IsMulOrDiv && ExtractFrom.getOpcode() != NeededShift)
9175 return false;
9176 Opcode = NeededShift;
9177 return true;
9178 };
9179 // op0 must be either the needed shift opcode or the mul/udiv equivalent
9180 // that the needed shift can be extracted from.
9181 if ((OppShift.getOpcode() != ISD::SRL || !SelectOpcode(ISD::SHL, ISD::MUL)) &&
9182 (OppShift.getOpcode() != ISD::SHL || !SelectOpcode(ISD::SRL, ISD::UDIV)))
9183 return SDValue();
9184
9185 // op0 must be the same opcode on both sides, have the same LHS argument,
9186 // and produce the same value type.
9187 if (OppShiftLHS.getOpcode() != ExtractFrom.getOpcode() ||
9188 OppShiftLHS.getOperand(i: 0) != ExtractFrom.getOperand(i: 0) ||
9189 ShiftedVT != ExtractFrom.getValueType())
9190 return SDValue();
9191
9192 // Constant mul/udiv/shift amount from the RHS of the shift's LHS op.
9193 ConstantSDNode *OppLHSCst = isConstOrConstSplat(N: OppShiftLHS.getOperand(i: 1));
9194 // Constant mul/udiv/shift amount from the RHS of the ExtractFrom op.
9195 ConstantSDNode *ExtractFromCst =
9196 isConstOrConstSplat(N: ExtractFrom.getOperand(i: 1));
9197 // TODO: We should be able to handle non-uniform constant vectors for these values
9198 // Check that we have constant values.
9199 if (!OppShiftCst || !OppShiftCst->getAPIntValue() ||
9200 !OppLHSCst || !OppLHSCst->getAPIntValue() ||
9201 !ExtractFromCst || !ExtractFromCst->getAPIntValue())
9202 return SDValue();
9203
9204 // Compute the shift amount we need to extract to complete the rotate.
9205 const unsigned VTWidth = ShiftedVT.getScalarSizeInBits();
9206 if (OppShiftCst->getAPIntValue().ugt(RHS: VTWidth))
9207 return SDValue();
9208 APInt NeededShiftAmt = VTWidth - OppShiftCst->getAPIntValue();
9209 // Normalize the bitwidth of the two mul/udiv/shift constant operands.
9210 APInt ExtractFromAmt = ExtractFromCst->getAPIntValue();
9211 APInt OppLHSAmt = OppLHSCst->getAPIntValue();
9212 zeroExtendToMatch(LHS&: ExtractFromAmt, RHS&: OppLHSAmt);
9213
9214 // Now try extract the needed shift from the ExtractFrom op and see if the
9215 // result matches up with the existing shift's LHS op.
9216 if (IsMulOrDiv) {
9217 // Op to extract from is a mul or udiv by a constant.
9218 // Check:
9219 // c2 / (1 << (bitwidth(op0 v c0) - c1)) == c0
9220 // c2 % (1 << (bitwidth(op0 v c0) - c1)) == 0
9221 const APInt ExtractDiv = APInt::getOneBitSet(numBits: ExtractFromAmt.getBitWidth(),
9222 BitNo: NeededShiftAmt.getZExtValue());
9223 APInt ResultAmt;
9224 APInt Rem;
9225 APInt::udivrem(LHS: ExtractFromAmt, RHS: ExtractDiv, Quotient&: ResultAmt, Remainder&: Rem);
9226 if (Rem != 0 || ResultAmt != OppLHSAmt)
9227 return SDValue();
9228 } else {
9229 // Op to extract from is a shift by a constant.
9230 // Check:
9231 // c2 - (bitwidth(op0 v c0) - c1) == c0
9232 if (OppLHSAmt != ExtractFromAmt - NeededShiftAmt.zextOrTrunc(
9233 width: ExtractFromAmt.getBitWidth()))
9234 return SDValue();
9235 }
9236
9237 // Return the expanded shift op that should allow a rotate to be formed.
9238 EVT ShiftVT = OppShift.getOperand(i: 1).getValueType();
9239 EVT ResVT = ExtractFrom.getValueType();
9240 SDValue NewShiftNode = DAG.getConstant(Val: NeededShiftAmt, DL, VT: ShiftVT);
9241 return DAG.getNode(Opcode, DL, VT: ResVT, N1: OppShiftLHS, N2: NewShiftNode);
9242}
9243
9244// Return true if we can prove that, whenever Neg and Pos are both in the
9245// range [0, EltSize), Neg == (Pos == 0 ? 0 : EltSize - Pos). This means that
9246// for two opposing shifts shift1 and shift2 and a value X with OpBits bits:
9247//
9248// (or (shift1 X, Neg), (shift2 X, Pos))
9249//
9250// reduces to a rotate in direction shift2 by Pos or (equivalently) a rotate
9251// in direction shift1 by Neg. The range [0, EltSize) means that we only need
9252// to consider shift amounts with defined behavior.
9253//
9254// The IsRotate flag should be set when the LHS of both shifts is the same.
9255// Otherwise if matching a general funnel shift, it should be clear.
9256static bool matchRotateSub(SDValue Pos, SDValue Neg, unsigned EltSize,
9257 SelectionDAG &DAG, bool IsRotate, bool FromAdd) {
9258 const auto &TLI = DAG.getTargetLoweringInfo();
9259 // If EltSize is a power of 2 then:
9260 //
9261 // (a) (Pos == 0 ? 0 : EltSize - Pos) == (EltSize - Pos) & (EltSize - 1)
9262 // (b) Neg == Neg & (EltSize - 1) whenever Neg is in [0, EltSize).
9263 //
9264 // So if EltSize is a power of 2 and Neg is (and Neg', EltSize-1), we check
9265 // for the stronger condition:
9266 //
9267 // Neg & (EltSize - 1) == (EltSize - Pos) & (EltSize - 1) [A]
9268 //
9269 // for all Neg and Pos. Since Neg & (EltSize - 1) == Neg' & (EltSize - 1)
9270 // we can just replace Neg with Neg' for the rest of the function.
9271 //
9272 // In other cases we check for the even stronger condition:
9273 //
9274 // Neg == EltSize - Pos [B]
9275 //
9276 // for all Neg and Pos. Note that the (or ...) then invokes undefined
9277 // behavior if Pos == 0 (and consequently Neg == EltSize).
9278 //
9279 // We could actually use [A] whenever EltSize is a power of 2, but the
9280 // only extra cases that it would match are those uninteresting ones
9281 // where Neg and Pos are never in range at the same time. E.g. for
9282 // EltSize == 32, using [A] would allow a Neg of the form (sub 64, Pos)
9283 // as well as (sub 32, Pos), but:
9284 //
9285 // (or (shift1 X, (sub 64, Pos)), (shift2 X, Pos))
9286 //
9287 // always invokes undefined behavior for 32-bit X.
9288 //
9289 // Below, Mask == EltSize - 1 when using [A] and is all-ones otherwise.
9290 // This allows us to peek through any operations that only affect Mask's
9291 // un-demanded bits.
9292 //
9293 // NOTE: We can only do this when matching operations which won't modify the
9294 // least Log2(EltSize) significant bits and not a general funnel shift.
9295 unsigned MaskLoBits = 0;
9296 if (IsRotate && !FromAdd && isPowerOf2_64(Value: EltSize)) {
9297 unsigned Bits = Log2_64(Value: EltSize);
9298 unsigned NegBits = Neg.getScalarValueSizeInBits();
9299 if (NegBits >= Bits) {
9300 APInt DemandedBits = APInt::getLowBitsSet(numBits: NegBits, loBitsSet: Bits);
9301 if (SDValue Inner =
9302 TLI.SimplifyMultipleUseDemandedBits(Op: Neg, DemandedBits, DAG)) {
9303 Neg = Inner;
9304 MaskLoBits = Bits;
9305 }
9306 }
9307 }
9308
9309 // Check whether Neg has the form (sub NegC, NegOp1) for some NegC and NegOp1.
9310 if (Neg.getOpcode() != ISD::SUB)
9311 return false;
9312 ConstantSDNode *NegC = isConstOrConstSplat(N: Neg.getOperand(i: 0));
9313 if (!NegC)
9314 return false;
9315 SDValue NegOp1 = Neg.getOperand(i: 1);
9316
9317 // On the RHS of [A], if Pos is the result of operation on Pos' that won't
9318 // affect Mask's demanded bits, just replace Pos with Pos'. These operations
9319 // are redundant for the purpose of the equality.
9320 if (MaskLoBits) {
9321 unsigned PosBits = Pos.getScalarValueSizeInBits();
9322 if (PosBits >= MaskLoBits) {
9323 APInt DemandedBits = APInt::getLowBitsSet(numBits: PosBits, loBitsSet: MaskLoBits);
9324 if (SDValue Inner =
9325 TLI.SimplifyMultipleUseDemandedBits(Op: Pos, DemandedBits, DAG)) {
9326 Pos = Inner;
9327 }
9328 }
9329 }
9330
9331 // The condition we need is now:
9332 //
9333 // (NegC - NegOp1) & Mask == (EltSize - Pos) & Mask
9334 //
9335 // If NegOp1 == Pos then we need:
9336 //
9337 // EltSize & Mask == NegC & Mask
9338 //
9339 // (because "x & Mask" is a truncation and distributes through subtraction).
9340 //
9341 // We also need to account for a potential truncation of NegOp1 if the amount
9342 // has already been legalized to a shift amount type.
9343 APInt Width;
9344 if ((Pos == NegOp1) ||
9345 (NegOp1.getOpcode() == ISD::TRUNCATE && Pos == NegOp1.getOperand(i: 0)))
9346 Width = NegC->getAPIntValue();
9347
9348 // Check for cases where Pos has the form (add NegOp1, PosC) for some PosC.
9349 // Then the condition we want to prove becomes:
9350 //
9351 // (NegC - NegOp1) & Mask == (EltSize - (NegOp1 + PosC)) & Mask
9352 //
9353 // which, again because "x & Mask" is a truncation, becomes:
9354 //
9355 // NegC & Mask == (EltSize - PosC) & Mask
9356 // EltSize & Mask == (NegC + PosC) & Mask
9357 else if (Pos.getOpcode() == ISD::ADD && Pos.getOperand(i: 0) == NegOp1) {
9358 if (ConstantSDNode *PosC = isConstOrConstSplat(N: Pos.getOperand(i: 1)))
9359 Width = PosC->getAPIntValue() + NegC->getAPIntValue();
9360 else
9361 return false;
9362 } else
9363 return false;
9364
9365 // Now we just need to check that EltSize & Mask == Width & Mask.
9366 if (MaskLoBits)
9367 // EltSize & Mask is 0 since Mask is EltSize - 1.
9368 return Width.getLoBits(numBits: MaskLoBits) == 0;
9369 return Width == EltSize;
9370}
9371
9372// A subroutine of MatchRotate used once we have found an OR of two opposite
9373// shifts of Shifted. If Neg == <operand size> - Pos then the OR reduces
9374// to both (PosOpcode Shifted, Pos) and (NegOpcode Shifted, Neg), with the
9375// former being preferred if supported. InnerPos and InnerNeg are Pos and
9376// Neg with outer conversions stripped away.
9377SDValue DAGCombiner::MatchRotatePosNeg(SDValue Shifted, SDValue Pos,
9378 SDValue Neg, SDValue InnerPos,
9379 SDValue InnerNeg, bool FromAdd,
9380 bool HasPos, unsigned PosOpcode,
9381 unsigned NegOpcode, const SDLoc &DL) {
9382 // fold (or/add (shl x, (*ext y)),
9383 // (srl x, (*ext (sub 32, y)))) ->
9384 // (rotl x, y) or (rotr x, (sub 32, y))
9385 //
9386 // fold (or/add (shl x, (*ext (sub 32, y))),
9387 // (srl x, (*ext y))) ->
9388 // (rotr x, y) or (rotl x, (sub 32, y))
9389 EVT VT = Shifted.getValueType();
9390 if (matchRotateSub(Pos: InnerPos, Neg: InnerNeg, EltSize: VT.getScalarSizeInBits(), DAG,
9391 /*IsRotate*/ true, FromAdd))
9392 return DAG.getNode(Opcode: HasPos ? PosOpcode : NegOpcode, DL, VT, N1: Shifted,
9393 N2: HasPos ? Pos : Neg);
9394
9395 return SDValue();
9396}
9397
9398// A subroutine of MatchRotate used once we have found an OR of two opposite
9399// shifts of N0 + N1. If Neg == <operand size> - Pos then the OR reduces
9400// to both (PosOpcode N0, N1, Pos) and (NegOpcode N0, N1, Neg), with the
9401// former being preferred if supported. InnerPos and InnerNeg are Pos and
9402// Neg with outer conversions stripped away.
9403// TODO: Merge with MatchRotatePosNeg.
9404SDValue DAGCombiner::MatchFunnelPosNeg(SDValue N0, SDValue N1, SDValue Pos,
9405 SDValue Neg, SDValue InnerPos,
9406 SDValue InnerNeg, bool FromAdd,
9407 bool HasPos, unsigned PosOpcode,
9408 unsigned NegOpcode, const SDLoc &DL) {
9409 EVT VT = N0.getValueType();
9410 unsigned EltBits = VT.getScalarSizeInBits();
9411
9412 // fold (or/add (shl x0, (*ext y)),
9413 // (srl x1, (*ext (sub 32, y)))) ->
9414 // (fshl x0, x1, y) or (fshr x0, x1, (sub 32, y))
9415 //
9416 // fold (or/add (shl x0, (*ext (sub 32, y))),
9417 // (srl x1, (*ext y))) ->
9418 // (fshr x0, x1, y) or (fshl x0, x1, (sub 32, y))
9419 if (matchRotateSub(Pos: InnerPos, Neg: InnerNeg, EltSize: EltBits, DAG, /*IsRotate*/ N0 == N1,
9420 FromAdd))
9421 return DAG.getNode(Opcode: HasPos ? PosOpcode : NegOpcode, DL, VT, N1: N0, N2: N1,
9422 N3: HasPos ? Pos : Neg);
9423
9424 // Matching the shift+xor cases, we can't easily use the xor'd shift amount
9425 // so for now just use the PosOpcode case if its legal.
9426 // TODO: When can we use the NegOpcode case?
9427 if (PosOpcode == ISD::FSHL && isPowerOf2_32(Value: EltBits)) {
9428 SDValue X;
9429 // fold (or/add (shl x0, y), (srl (srl x1, 1), (xor y, 31)))
9430 // -> (fshl x0, x1, y)
9431 if (sd_match(N: N1, P: m_Srl(L: m_Value(N&: X), R: m_One())) &&
9432 sd_match(N: InnerNeg,
9433 P: m_Xor(L: m_Specific(N: InnerPos), R: m_SpecificInt(V: EltBits - 1))) &&
9434 TLI.isOperationLegalOrCustom(Op: ISD::FSHL, VT)) {
9435 return DAG.getNode(Opcode: ISD::FSHL, DL, VT, N1: N0, N2: X, N3: Pos);
9436 }
9437
9438 // fold (or/add (shl (shl x0, 1), (xor y, 31)), (srl x1, y))
9439 // -> (fshr x0, x1, y)
9440 if (sd_match(N: N0, P: m_Shl(L: m_Value(N&: X), R: m_One())) &&
9441 sd_match(N: InnerPos,
9442 P: m_Xor(L: m_Specific(N: InnerNeg), R: m_SpecificInt(V: EltBits - 1))) &&
9443 TLI.isOperationLegalOrCustom(Op: ISD::FSHR, VT)) {
9444 return DAG.getNode(Opcode: ISD::FSHR, DL, VT, N1: X, N2: N1, N3: Neg);
9445 }
9446
9447 // fold (or/add (shl (add x0, x0), (xor y, 31)), (srl x1, y))
9448 // -> (fshr x0, x1, y)
9449 // TODO: Should add(x,x) -> shl(x,1) be a general DAG canonicalization?
9450 if (sd_match(N: N0, P: m_Add(L: m_Value(N&: X), R: m_Deferred(V&: X))) &&
9451 sd_match(N: InnerPos,
9452 P: m_Xor(L: m_Specific(N: InnerNeg), R: m_SpecificInt(V: EltBits - 1))) &&
9453 TLI.isOperationLegalOrCustom(Op: ISD::FSHR, VT)) {
9454 return DAG.getNode(Opcode: ISD::FSHR, DL, VT, N1: X, N2: N1, N3: Neg);
9455 }
9456 }
9457
9458 return SDValue();
9459}
9460
9461// MatchRotate - Handle an 'or' or 'add' of two operands. If this is one of the
9462// many idioms for rotate, and if the target supports rotation instructions,
9463// generate a rot[lr]. This also matches funnel shift patterns, similar to
9464// rotation but with different shifted sources.
9465SDValue DAGCombiner::MatchRotate(SDValue LHS, SDValue RHS, const SDLoc &DL,
9466 bool FromAdd) {
9467 EVT VT = LHS.getValueType();
9468
9469 // The target must have at least one rotate/funnel flavor.
9470 // We still try to match rotate by constant pre-legalization.
9471 // TODO: Support pre-legalization funnel-shift by constant.
9472 bool HasROTL = hasOperation(Opcode: ISD::ROTL, VT);
9473 bool HasROTR = hasOperation(Opcode: ISD::ROTR, VT);
9474 bool HasFSHL = hasOperation(Opcode: ISD::FSHL, VT);
9475 bool HasFSHR = hasOperation(Opcode: ISD::FSHR, VT);
9476
9477 // If the type is going to be promoted and the target has enabled custom
9478 // lowering for rotate, allow matching rotate by non-constants. Only allow
9479 // this for scalar types.
9480 if (VT.isScalarInteger() && TLI.getTypeAction(Context&: *DAG.getContext(), VT) ==
9481 TargetLowering::TypePromoteInteger) {
9482 HasROTL |= TLI.getOperationAction(Op: ISD::ROTL, VT) == TargetLowering::Custom;
9483 HasROTR |= TLI.getOperationAction(Op: ISD::ROTR, VT) == TargetLowering::Custom;
9484 }
9485
9486 if (LegalOperations && !HasROTL && !HasROTR && !HasFSHL && !HasFSHR)
9487 return SDValue();
9488
9489 // Check for truncated rotate.
9490 if (LHS.getOpcode() == ISD::TRUNCATE && RHS.getOpcode() == ISD::TRUNCATE &&
9491 LHS.getOperand(i: 0).getValueType() == RHS.getOperand(i: 0).getValueType()) {
9492 assert(LHS.getValueType() == RHS.getValueType());
9493 if (SDValue Rot =
9494 MatchRotate(LHS: LHS.getOperand(i: 0), RHS: RHS.getOperand(i: 0), DL, FromAdd))
9495 return DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(LHS), VT: LHS.getValueType(), Operand: Rot);
9496 }
9497
9498 // Match "(X shl/srl V1) & V2" where V2 may not be present.
9499 SDValue LHSShift; // The shift.
9500 SDValue LHSMask; // AND value if any.
9501 matchRotateHalf(DAG, Op: LHS, Shift&: LHSShift, Mask&: LHSMask);
9502
9503 SDValue RHSShift; // The shift.
9504 SDValue RHSMask; // AND value if any.
9505 matchRotateHalf(DAG, Op: RHS, Shift&: RHSShift, Mask&: RHSMask);
9506
9507 // If neither side matched a rotate half, bail
9508 if (!LHSShift && !RHSShift)
9509 return SDValue();
9510
9511 // InstCombine may have combined a constant shl, srl, mul, or udiv with one
9512 // side of the rotate, so try to handle that here. In all cases we need to
9513 // pass the matched shift from the opposite side to compute the opcode and
9514 // needed shift amount to extract. We still want to do this if both sides
9515 // matched a rotate half because one half may be a potential overshift that
9516 // can be broken down (ie if InstCombine merged two shl or srl ops into a
9517 // single one).
9518
9519 // Have LHS side of the rotate, try to extract the needed shift from the RHS.
9520 if (LHSShift)
9521 if (SDValue NewRHSShift =
9522 extractShiftForRotate(DAG, OppShift: LHSShift, ExtractFrom: RHS, Mask&: RHSMask, DL))
9523 RHSShift = NewRHSShift;
9524 // Have RHS side of the rotate, try to extract the needed shift from the LHS.
9525 if (RHSShift)
9526 if (SDValue NewLHSShift =
9527 extractShiftForRotate(DAG, OppShift: RHSShift, ExtractFrom: LHS, Mask&: LHSMask, DL))
9528 LHSShift = NewLHSShift;
9529
9530 // If a side is still missing, nothing else we can do.
9531 if (!RHSShift || !LHSShift)
9532 return SDValue();
9533
9534 // At this point we've matched or extracted a shift op on each side.
9535
9536 if (LHSShift.getOpcode() == RHSShift.getOpcode())
9537 return SDValue(); // Shifts must disagree.
9538
9539 // Canonicalize shl to left side in a shl/srl pair.
9540 if (RHSShift.getOpcode() == ISD::SHL) {
9541 std::swap(a&: LHS, b&: RHS);
9542 std::swap(a&: LHSShift, b&: RHSShift);
9543 std::swap(a&: LHSMask, b&: RHSMask);
9544 }
9545
9546 // Something has gone wrong - we've lost the shl/srl pair - bail.
9547 if (LHSShift.getOpcode() != ISD::SHL || RHSShift.getOpcode() != ISD::SRL)
9548 return SDValue();
9549
9550 unsigned EltSizeInBits = VT.getScalarSizeInBits();
9551 SDValue LHSShiftArg = LHSShift.getOperand(i: 0);
9552 SDValue LHSShiftAmt = LHSShift.getOperand(i: 1);
9553 SDValue RHSShiftArg = RHSShift.getOperand(i: 0);
9554 SDValue RHSShiftAmt = RHSShift.getOperand(i: 1);
9555
9556 auto MatchRotateSum = [EltSizeInBits](ConstantSDNode *LHS,
9557 ConstantSDNode *RHS) {
9558 return (LHS->getAPIntValue() + RHS->getAPIntValue()) == EltSizeInBits;
9559 };
9560
9561 auto ApplyMasks = [&](SDValue Res) {
9562 // If there is an AND of either shifted operand, apply it to the result.
9563 if (LHSMask.getNode() || RHSMask.getNode()) {
9564 SDValue AllOnes = DAG.getAllOnesConstant(DL, VT);
9565 SDValue Mask = AllOnes;
9566
9567 if (LHSMask.getNode()) {
9568 SDValue RHSBits = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: AllOnes, N2: RHSShiftAmt);
9569 Mask = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Mask,
9570 N2: DAG.getNode(Opcode: ISD::OR, DL, VT, N1: LHSMask, N2: RHSBits));
9571 }
9572 if (RHSMask.getNode()) {
9573 SDValue LHSBits = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: AllOnes, N2: LHSShiftAmt);
9574 Mask = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Mask,
9575 N2: DAG.getNode(Opcode: ISD::OR, DL, VT, N1: RHSMask, N2: LHSBits));
9576 }
9577
9578 Res = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Res, N2: Mask);
9579 }
9580
9581 return Res;
9582 };
9583
9584 // TODO: Support pre-legalization funnel-shift by constant.
9585 bool IsRotate = LHSShiftArg == RHSShiftArg;
9586 if (!IsRotate && !(HasFSHL || HasFSHR)) {
9587 if (TLI.isTypeLegal(VT) && LHS.hasOneUse() && RHS.hasOneUse() &&
9588 ISD::matchBinaryPredicate(LHS: LHSShiftAmt, RHS: RHSShiftAmt, Match: MatchRotateSum)) {
9589 // Look for a disguised rotate by constant.
9590 // The common shifted operand X may be hidden inside another 'or'.
9591 SDValue X, Y;
9592 auto matchOr = [&X, &Y](SDValue Or, SDValue CommonOp) {
9593 if (!Or.hasOneUse() || Or.getOpcode() != ISD::OR)
9594 return false;
9595 if (CommonOp == Or.getOperand(i: 0)) {
9596 X = CommonOp;
9597 Y = Or.getOperand(i: 1);
9598 return true;
9599 }
9600 if (CommonOp == Or.getOperand(i: 1)) {
9601 X = CommonOp;
9602 Y = Or.getOperand(i: 0);
9603 return true;
9604 }
9605 return false;
9606 };
9607
9608 SDValue Res;
9609 if (matchOr(LHSShiftArg, RHSShiftArg)) {
9610 // (shl (X | Y), C1) | (srl X, C2) --> (rotl X, C1) | (shl Y, C1)
9611 SDValue RotX = DAG.getNode(Opcode: ISD::ROTL, DL, VT, N1: X, N2: LHSShiftAmt);
9612 SDValue ShlY = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Y, N2: LHSShiftAmt);
9613 Res = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: RotX, N2: ShlY);
9614 } else if (matchOr(RHSShiftArg, LHSShiftArg)) {
9615 // (shl X, C1) | (srl (X | Y), C2) --> (rotl X, C1) | (srl Y, C2)
9616 SDValue RotX = DAG.getNode(Opcode: ISD::ROTL, DL, VT, N1: X, N2: LHSShiftAmt);
9617 SDValue SrlY = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: Y, N2: RHSShiftAmt);
9618 Res = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: RotX, N2: SrlY);
9619 } else {
9620 return SDValue();
9621 }
9622
9623 return ApplyMasks(Res);
9624 }
9625
9626 return SDValue(); // Requires funnel shift support.
9627 }
9628
9629 // fold (or/add (shl x, C1), (srl x, C2)) -> (rotl x, C1)
9630 // fold (or/add (shl x, C1), (srl x, C2)) -> (rotr x, C2)
9631 // fold (or/add (shl x, C1), (srl y, C2)) -> (fshl x, y, C1)
9632 // fold (or/add (shl x, C1), (srl y, C2)) -> (fshr x, y, C2)
9633 // iff C1+C2 == EltSizeInBits
9634 if (ISD::matchBinaryPredicate(LHS: LHSShiftAmt, RHS: RHSShiftAmt, Match: MatchRotateSum)) {
9635 SDValue Res;
9636 if (IsRotate && (HasROTL || HasROTR || !(HasFSHL || HasFSHR))) {
9637 bool UseROTL = !LegalOperations || HasROTL;
9638 Res = DAG.getNode(Opcode: UseROTL ? ISD::ROTL : ISD::ROTR, DL, VT, N1: LHSShiftArg,
9639 N2: UseROTL ? LHSShiftAmt : RHSShiftAmt);
9640 } else {
9641 bool UseFSHL = !LegalOperations || HasFSHL;
9642 Res = DAG.getNode(Opcode: UseFSHL ? ISD::FSHL : ISD::FSHR, DL, VT, N1: LHSShiftArg,
9643 N2: RHSShiftArg, N3: UseFSHL ? LHSShiftAmt : RHSShiftAmt);
9644 }
9645
9646 return ApplyMasks(Res);
9647 }
9648
9649 // Even pre-legalization, we can't easily rotate/funnel-shift by a variable
9650 // shift.
9651 if (!HasROTL && !HasROTR && !HasFSHL && !HasFSHR)
9652 return SDValue();
9653
9654 // If there is a mask here, and we have a variable shift, we can't be sure
9655 // that we're masking out the right stuff.
9656 if (LHSMask.getNode() || RHSMask.getNode())
9657 return SDValue();
9658
9659 // If the shift amount is sign/zext/any-extended just peel it off.
9660 SDValue LExtOp0 = LHSShiftAmt;
9661 SDValue RExtOp0 = RHSShiftAmt;
9662 if ((LHSShiftAmt.getOpcode() == ISD::SIGN_EXTEND ||
9663 LHSShiftAmt.getOpcode() == ISD::ZERO_EXTEND ||
9664 LHSShiftAmt.getOpcode() == ISD::ANY_EXTEND ||
9665 LHSShiftAmt.getOpcode() == ISD::TRUNCATE) &&
9666 (RHSShiftAmt.getOpcode() == ISD::SIGN_EXTEND ||
9667 RHSShiftAmt.getOpcode() == ISD::ZERO_EXTEND ||
9668 RHSShiftAmt.getOpcode() == ISD::ANY_EXTEND ||
9669 RHSShiftAmt.getOpcode() == ISD::TRUNCATE)) {
9670 LExtOp0 = LHSShiftAmt.getOperand(i: 0);
9671 RExtOp0 = RHSShiftAmt.getOperand(i: 0);
9672 }
9673
9674 if (IsRotate && (HasROTL || HasROTR)) {
9675 if (SDValue TryL = MatchRotatePosNeg(Shifted: LHSShiftArg, Pos: LHSShiftAmt, Neg: RHSShiftAmt,
9676 InnerPos: LExtOp0, InnerNeg: RExtOp0, FromAdd, HasPos: HasROTL,
9677 PosOpcode: ISD::ROTL, NegOpcode: ISD::ROTR, DL))
9678 return TryL;
9679
9680 if (SDValue TryR = MatchRotatePosNeg(Shifted: RHSShiftArg, Pos: RHSShiftAmt, Neg: LHSShiftAmt,
9681 InnerPos: RExtOp0, InnerNeg: LExtOp0, FromAdd, HasPos: HasROTR,
9682 PosOpcode: ISD::ROTR, NegOpcode: ISD::ROTL, DL))
9683 return TryR;
9684 }
9685
9686 if (SDValue TryL = MatchFunnelPosNeg(N0: LHSShiftArg, N1: RHSShiftArg, Pos: LHSShiftAmt,
9687 Neg: RHSShiftAmt, InnerPos: LExtOp0, InnerNeg: RExtOp0, FromAdd,
9688 HasPos: HasFSHL, PosOpcode: ISD::FSHL, NegOpcode: ISD::FSHR, DL))
9689 return TryL;
9690
9691 if (SDValue TryR = MatchFunnelPosNeg(N0: LHSShiftArg, N1: RHSShiftArg, Pos: RHSShiftAmt,
9692 Neg: LHSShiftAmt, InnerPos: RExtOp0, InnerNeg: LExtOp0, FromAdd,
9693 HasPos: HasFSHR, PosOpcode: ISD::FSHR, NegOpcode: ISD::FSHL, DL))
9694 return TryR;
9695
9696 return SDValue();
9697}
9698
9699/// Recursively traverses the expression calculating the origin of the requested
9700/// byte of the given value. Returns std::nullopt if the provider can't be
9701/// calculated.
9702///
9703/// For all the values except the root of the expression, we verify that the
9704/// value has exactly one use and if not then return std::nullopt. This way if
9705/// the origin of the byte is returned it's guaranteed that the values which
9706/// contribute to the byte are not used outside of this expression.
9707
9708/// However, there is a special case when dealing with vector loads -- we allow
9709/// more than one use if the load is a vector type. Since the values that
9710/// contribute to the byte ultimately come from the ExtractVectorElements of the
9711/// Load, we don't care if the Load has uses other than ExtractVectorElements,
9712/// because those operations are independent from the pattern to be combined.
9713/// For vector loads, we simply care that the ByteProviders are adjacent
9714/// positions of the same vector, and their index matches the byte that is being
9715/// provided. This is captured by the \p VectorIndex algorithm. \p VectorIndex
9716/// is the index used in an ExtractVectorElement, and \p StartingIndex is the
9717/// byte position we are trying to provide for the LoadCombine. If these do
9718/// not match, then we can not combine the vector loads. \p Index uses the
9719/// byte position we are trying to provide for and is matched against the
9720/// shl and load size. The \p Index algorithm ensures the requested byte is
9721/// provided for by the pattern, and the pattern does not over provide bytes.
9722///
9723///
9724/// The supported LoadCombine pattern for vector loads is as follows
9725/// or
9726/// / \
9727/// or shl
9728/// / \ |
9729/// or shl zext
9730/// / \ | |
9731/// shl zext zext EVE*
9732/// | | | |
9733/// zext EVE* EVE* LOAD
9734/// | | |
9735/// EVE* LOAD LOAD
9736/// |
9737/// LOAD
9738///
9739/// *ExtractVectorElement
9740using SDByteProvider = ByteProvider<SDNode *>;
9741
9742static std::optional<SDByteProvider>
9743calculateByteProvider(SDValue Op, unsigned Index, unsigned Depth,
9744 std::optional<uint64_t> VectorIndex,
9745 unsigned StartingIndex = 0,
9746 MutableArrayRef<uint8_t> ByteMask = {}) {
9747
9748 // Typical i64 by i8 pattern requires recursion up to 8 calls depth
9749 if (Depth == 10)
9750 return std::nullopt;
9751
9752 // Only allow multiple uses if the instruction is a vector load (in which
9753 // case we will use the load for every ExtractVectorElement)
9754 if (Depth && !Op.hasOneUse() &&
9755 (Op.getOpcode() != ISD::LOAD || !Op.getValueType().isVector()))
9756 return std::nullopt;
9757
9758 // Fail to combine if we have encountered anything but a LOAD after handling
9759 // an ExtractVectorElement.
9760 if (Op.getOpcode() != ISD::LOAD && VectorIndex.has_value())
9761 return std::nullopt;
9762
9763 unsigned BitWidth = Op.getScalarValueSizeInBits();
9764 if (BitWidth % 8 != 0)
9765 return std::nullopt;
9766 unsigned ByteWidth = BitWidth / 8;
9767 assert(Index < ByteWidth && "invalid index requested");
9768 (void) ByteWidth;
9769
9770 switch (Op.getOpcode()) {
9771 case ISD::OR: {
9772 auto LHS = calculateByteProvider(Op: Op->getOperand(Num: 0), Index, Depth: Depth + 1,
9773 VectorIndex, StartingIndex, ByteMask);
9774 if (!LHS)
9775 return std::nullopt;
9776 auto RHS = calculateByteProvider(Op: Op->getOperand(Num: 1), Index, Depth: Depth + 1,
9777 VectorIndex, StartingIndex, ByteMask);
9778 if (!RHS)
9779 return std::nullopt;
9780
9781 if (LHS->isConstantZero())
9782 return RHS;
9783 if (RHS->isConstantZero())
9784 return LHS;
9785 return std::nullopt;
9786 }
9787 case ISD::SHL: {
9788 auto ShiftOp = dyn_cast<ConstantSDNode>(Val: Op->getOperand(Num: 1));
9789 if (!ShiftOp)
9790 return std::nullopt;
9791
9792 uint64_t BitShift = ShiftOp->getZExtValue();
9793
9794 if (BitShift % 8 != 0)
9795 return std::nullopt;
9796 uint64_t ByteShift = BitShift / 8;
9797
9798 // If we are shifting by an amount greater than the index we are trying to
9799 // provide, then do not provide anything. Otherwise, subtract the index by
9800 // the amount we shifted by.
9801 return Index < ByteShift
9802 ? SDByteProvider::getConstantZero()
9803 : calculateByteProvider(Op: Op->getOperand(Num: 0), Index: Index - ByteShift,
9804 Depth: Depth + 1, VectorIndex, StartingIndex: Index, ByteMask);
9805 }
9806 case ISD::ANY_EXTEND:
9807 case ISD::SIGN_EXTEND:
9808 case ISD::ZERO_EXTEND: {
9809 SDValue NarrowOp = Op->getOperand(Num: 0);
9810 unsigned NarrowBitWidth = NarrowOp.getScalarValueSizeInBits();
9811 if (NarrowBitWidth % 8 != 0)
9812 return std::nullopt;
9813 uint64_t NarrowByteWidth = NarrowBitWidth / 8;
9814
9815 if (Index >= NarrowByteWidth)
9816 return Op.getOpcode() == ISD::ZERO_EXTEND
9817 ? std::optional<SDByteProvider>(
9818 SDByteProvider::getConstantZero())
9819 : std::nullopt;
9820 return calculateByteProvider(Op: NarrowOp, Index, Depth: Depth + 1, VectorIndex,
9821 StartingIndex, ByteMask);
9822 }
9823 case ISD::BSWAP:
9824 return calculateByteProvider(Op: Op->getOperand(Num: 0), Index: ByteWidth - Index - 1,
9825 Depth: Depth + 1, VectorIndex, StartingIndex,
9826 ByteMask);
9827 case ISD::AND: {
9828 // Constants are canonicalized to the RHS of AND, so only operand 1 needs
9829 // to be checked.
9830 auto *MaskOp = dyn_cast<ConstantSDNode>(Val: Op->getOperand(Num: 1));
9831 if (!MaskOp)
9832 return std::nullopt;
9833
9834 uint8_t MaskByte =
9835 MaskOp->getAPIntValue().extractBitsAsZExtValue(numBits: 8, bitPosition: Index * 8);
9836
9837 if (MaskByte == 0x00)
9838 return SDByteProvider::getConstantZero();
9839
9840 auto Result = calculateByteProvider(Op: Op->getOperand(Num: 0), Index, Depth: Depth + 1,
9841 VectorIndex, StartingIndex, ByteMask);
9842 if (!Result)
9843 return std::nullopt;
9844
9845 // Only record the mask if this byte is actually provided (not zero).
9846 // A ConstantZero result may be discarded by the OR handler in favor of
9847 // the other operand, so writing the mask here would corrupt ByteMask.
9848 if (MaskByte != 0xFF && !ByteMask.empty() && !Result->isConstantZero())
9849 ByteMask[StartingIndex] &= MaskByte;
9850
9851 return Result;
9852 }
9853 case ISD::EXTRACT_VECTOR_ELT: {
9854 auto OffsetOp = dyn_cast<ConstantSDNode>(Val: Op->getOperand(Num: 1));
9855 if (!OffsetOp)
9856 return std::nullopt;
9857
9858 VectorIndex = OffsetOp->getZExtValue();
9859
9860 SDValue NarrowOp = Op->getOperand(Num: 0);
9861 unsigned NarrowBitWidth = NarrowOp.getScalarValueSizeInBits();
9862 if (NarrowBitWidth % 8 != 0)
9863 return std::nullopt;
9864 uint64_t NarrowByteWidth = NarrowBitWidth / 8;
9865 // EXTRACT_VECTOR_ELT can extend the element type to the width of the return
9866 // type, leaving the high bits undefined.
9867 if (Index >= NarrowByteWidth)
9868 return std::nullopt;
9869
9870 // Check to see if the position of the element in the vector corresponds
9871 // with the byte we are trying to provide for. In the case of a vector of
9872 // i8, this simply means the VectorIndex == StartingIndex. For non i8 cases,
9873 // the element will provide a range of bytes. For example, if we have a
9874 // vector of i16s, each element provides two bytes (V[1] provides byte 2 and
9875 // 3).
9876 if (*VectorIndex * NarrowByteWidth > StartingIndex)
9877 return std::nullopt;
9878 if ((*VectorIndex + 1) * NarrowByteWidth <= StartingIndex)
9879 return std::nullopt;
9880
9881 return calculateByteProvider(Op: Op->getOperand(Num: 0), Index, Depth: Depth + 1,
9882 VectorIndex, StartingIndex, ByteMask);
9883 }
9884 case ISD::LOAD: {
9885 auto L = cast<LoadSDNode>(Val: Op.getNode());
9886 if (!L->isSimple() || L->isIndexed())
9887 return std::nullopt;
9888
9889 unsigned NarrowBitWidth = L->getMemoryVT().getScalarSizeInBits();
9890 if (NarrowBitWidth % 8 != 0)
9891 return std::nullopt;
9892 uint64_t NarrowByteWidth = NarrowBitWidth / 8;
9893
9894 // If the width of the load does not reach byte we are trying to provide for
9895 // and it is not a ZEXTLOAD, then the load does not provide for the byte in
9896 // question
9897 if (Index >= NarrowByteWidth)
9898 return L->getExtensionType() == ISD::ZEXTLOAD
9899 ? std::optional<SDByteProvider>(
9900 SDByteProvider::getConstantZero())
9901 : std::nullopt;
9902
9903 unsigned BPVectorIndex = VectorIndex.value_or(u: 0U);
9904 return SDByteProvider::getSrc(Val: L, ByteOffset: Index, VectorOffset: BPVectorIndex);
9905 }
9906 }
9907
9908 return std::nullopt;
9909}
9910
9911static unsigned littleEndianByteAt(unsigned BW, unsigned i) {
9912 return i;
9913}
9914
9915static unsigned bigEndianByteAt(unsigned BW, unsigned i) {
9916 return BW - i - 1;
9917}
9918
9919// Check if the bytes offsets we are looking at match with either big or
9920// little endian value loaded. Return true for big endian, false for little
9921// endian, and std::nullopt if match failed.
9922static std::optional<bool> isBigEndian(ArrayRef<int64_t> ByteOffsets,
9923 int64_t FirstOffset) {
9924 // The endian can be decided only when it is 2 bytes at least.
9925 unsigned Width = ByteOffsets.size();
9926 if (Width < 2)
9927 return std::nullopt;
9928
9929 bool BigEndian = true, LittleEndian = true;
9930 for (unsigned i = 0; i < Width; i++) {
9931 int64_t CurrentByteOffset = ByteOffsets[i] - FirstOffset;
9932 LittleEndian &= CurrentByteOffset == littleEndianByteAt(BW: Width, i);
9933 BigEndian &= CurrentByteOffset == bigEndianByteAt(BW: Width, i);
9934 if (!BigEndian && !LittleEndian)
9935 return std::nullopt;
9936 }
9937
9938 assert((BigEndian != LittleEndian) && "It should be either big endian or"
9939 "little endian");
9940 return BigEndian;
9941}
9942
9943// Look through one layer of truncate or extend.
9944static SDValue stripTruncAndExt(SDValue Value) {
9945 switch (Value.getOpcode()) {
9946 case ISD::TRUNCATE:
9947 case ISD::ZERO_EXTEND:
9948 case ISD::SIGN_EXTEND:
9949 case ISD::ANY_EXTEND:
9950 return Value.getOperand(i: 0);
9951 }
9952 return SDValue();
9953}
9954
9955/// Match a pattern where a wide type scalar value is stored by several narrow
9956/// stores. Fold it into a single store or a BSWAP and a store if the targets
9957/// supports it.
9958///
9959/// Assuming little endian target:
9960/// i8 *p = ...
9961/// i32 val = ...
9962/// p[0] = (val >> 0) & 0xFF;
9963/// p[1] = (val >> 8) & 0xFF;
9964/// p[2] = (val >> 16) & 0xFF;
9965/// p[3] = (val >> 24) & 0xFF;
9966/// =>
9967/// *((i32)p) = val;
9968///
9969/// i8 *p = ...
9970/// i32 val = ...
9971/// p[0] = (val >> 24) & 0xFF;
9972/// p[1] = (val >> 16) & 0xFF;
9973/// p[2] = (val >> 8) & 0xFF;
9974/// p[3] = (val >> 0) & 0xFF;
9975/// =>
9976/// *((i32)p) = BSWAP(val);
9977SDValue DAGCombiner::mergeTruncStores(StoreSDNode *N) {
9978 // The matching looks for "store (trunc x)" patterns that appear early but are
9979 // likely to be replaced by truncating store nodes during combining.
9980 // TODO: If there is evidence that running this later would help, this
9981 // limitation could be removed. Legality checks may need to be added
9982 // for the created store and optional bswap/rotate.
9983 if (LegalOperations || OptLevel == CodeGenOptLevel::None)
9984 return SDValue();
9985
9986 // We only handle merging simple stores of 1-4 bytes.
9987 // TODO: Allow unordered atomics when wider type is legal (see D66309)
9988 EVT MemVT = N->getMemoryVT();
9989 if (!(MemVT == MVT::i8 || MemVT == MVT::i16 || MemVT == MVT::i32) ||
9990 !N->isSimple() || N->isIndexed())
9991 return SDValue();
9992
9993 // Collect all of the stores in the chain, upto the maximum store width (i64).
9994 SDValue Chain = N->getChain();
9995 SmallVector<StoreSDNode *, 8> Stores = {N};
9996 unsigned NarrowNumBits = MemVT.getScalarSizeInBits();
9997 unsigned MaxWideNumBits = 64;
9998 unsigned MaxStores = MaxWideNumBits / NarrowNumBits;
9999 while (auto *Store = dyn_cast<StoreSDNode>(Val&: Chain)) {
10000 // All stores must be the same size to ensure that we are writing all of the
10001 // bytes in the wide value.
10002 // This store should have exactly one use as a chain operand for another
10003 // store in the merging set. If there are other chain uses, then the
10004 // transform may not be safe because order of loads/stores outside of this
10005 // set may not be preserved.
10006 // TODO: We could allow multiple sizes by tracking each stored byte.
10007 if (Store->getMemoryVT() != MemVT || !Store->isSimple() ||
10008 Store->isIndexed() || !Store->hasOneUse())
10009 return SDValue();
10010 Stores.push_back(Elt: Store);
10011 Chain = Store->getChain();
10012 if (MaxStores < Stores.size())
10013 return SDValue();
10014 }
10015 // There is no reason to continue if we do not have at least a pair of stores.
10016 if (Stores.size() < 2)
10017 return SDValue();
10018
10019 // Handle simple types only.
10020 LLVMContext &Context = *DAG.getContext();
10021 unsigned NumStores = Stores.size();
10022 unsigned WideNumBits = NumStores * NarrowNumBits;
10023 if (WideNumBits != 16 && WideNumBits != 32 && WideNumBits != 64)
10024 return SDValue();
10025
10026 // Check if all bytes of the source value that we are looking at are stored
10027 // to the same base address. Collect offsets from Base address into OffsetMap.
10028 SDValue SourceValue;
10029 SmallVector<int64_t, 8> OffsetMap(NumStores, INT64_MAX);
10030 int64_t FirstOffset = INT64_MAX;
10031 StoreSDNode *FirstStore = nullptr;
10032 std::optional<BaseIndexOffset> Base;
10033 for (auto *Store : Stores) {
10034 // All the stores store different parts of the CombinedValue. A truncate is
10035 // required to get the partial value.
10036 SDValue Trunc = Store->getValue();
10037 if (Trunc.getOpcode() != ISD::TRUNCATE)
10038 return SDValue();
10039 // Other than the first/last part, a shift operation is required to get the
10040 // offset.
10041 int64_t Offset = 0;
10042 SDValue WideVal = Trunc.getOperand(i: 0);
10043 if ((WideVal.getOpcode() == ISD::SRL || WideVal.getOpcode() == ISD::SRA) &&
10044 isa<ConstantSDNode>(Val: WideVal.getOperand(i: 1))) {
10045 // The shift amount must be a constant multiple of the narrow type.
10046 // It is translated to the offset address in the wide source value "y".
10047 //
10048 // x = srl y, ShiftAmtC
10049 // i8 z = trunc x
10050 // store z, ...
10051 uint64_t ShiftAmtC = WideVal.getConstantOperandVal(i: 1);
10052 if (ShiftAmtC % NarrowNumBits != 0)
10053 return SDValue();
10054
10055 // Make sure we aren't reading bits that are shifted in.
10056 if (ShiftAmtC > WideVal.getScalarValueSizeInBits() - NarrowNumBits)
10057 return SDValue();
10058
10059 Offset = ShiftAmtC / NarrowNumBits;
10060 WideVal = WideVal.getOperand(i: 0);
10061 }
10062
10063 // Stores must share the same source value with different offsets.
10064 if (!SourceValue)
10065 SourceValue = WideVal;
10066 else if (SourceValue != WideVal) {
10067 // Truncate and extends can be stripped to see if the values are related.
10068 if (stripTruncAndExt(Value: SourceValue) != WideVal &&
10069 stripTruncAndExt(Value: WideVal) != SourceValue)
10070 return SDValue();
10071
10072 if (WideVal.getScalarValueSizeInBits() >
10073 SourceValue.getScalarValueSizeInBits())
10074 SourceValue = WideVal;
10075
10076 // Give up if the source value type is smaller than the store size.
10077 if (SourceValue.getScalarValueSizeInBits() < WideNumBits)
10078 return SDValue();
10079 }
10080
10081 // Stores must share the same base address.
10082 BaseIndexOffset Ptr = BaseIndexOffset::match(N: Store, DAG);
10083 int64_t ByteOffsetFromBase = 0;
10084 if (!Base)
10085 Base = Ptr;
10086 else if (!Base->equalBaseIndex(Other: Ptr, DAG, Off&: ByteOffsetFromBase))
10087 return SDValue();
10088
10089 // Remember the first store.
10090 if (ByteOffsetFromBase < FirstOffset) {
10091 FirstStore = Store;
10092 FirstOffset = ByteOffsetFromBase;
10093 }
10094 // Map the offset in the store and the offset in the combined value, and
10095 // early return if it has been set before.
10096 if (Offset < 0 || Offset >= NumStores || OffsetMap[Offset] != INT64_MAX)
10097 return SDValue();
10098 OffsetMap[Offset] = ByteOffsetFromBase;
10099 }
10100
10101 EVT WideVT = EVT::getIntegerVT(Context, BitWidth: WideNumBits);
10102
10103 assert(FirstOffset != INT64_MAX && "First byte offset must be set");
10104 assert(FirstStore && "First store must be set");
10105
10106 // Check that a store of the wide type is both allowed and fast on the target
10107 const DataLayout &Layout = DAG.getDataLayout();
10108 unsigned Fast = 0;
10109 bool Allowed = TLI.allowsMemoryAccess(Context, DL: Layout, VT: WideVT,
10110 MMO: *FirstStore->getMemOperand(), Fast: &Fast);
10111 if (!Allowed || !Fast)
10112 return SDValue();
10113
10114 // Check if the pieces of the value are going to the expected places in memory
10115 // to merge the stores.
10116 auto checkOffsets = [&](bool MatchLittleEndian) {
10117 if (MatchLittleEndian) {
10118 for (unsigned i = 0; i != NumStores; ++i)
10119 if (OffsetMap[i] != i * (NarrowNumBits / 8) + FirstOffset)
10120 return false;
10121 } else { // MatchBigEndian by reversing loop counter.
10122 for (unsigned i = 0, j = NumStores - 1; i != NumStores; ++i, --j)
10123 if (OffsetMap[j] != i * (NarrowNumBits / 8) + FirstOffset)
10124 return false;
10125 }
10126 return true;
10127 };
10128
10129 // Check if the offsets line up for the native data layout of this target.
10130 bool NeedBswap = false;
10131 bool NeedRotate = false;
10132 if (!checkOffsets(Layout.isLittleEndian())) {
10133 // Special-case: check if byte offsets line up for the opposite endian.
10134 if (NarrowNumBits == 8 && checkOffsets(Layout.isBigEndian()))
10135 NeedBswap = true;
10136 else if (NumStores == 2 && checkOffsets(Layout.isBigEndian()))
10137 NeedRotate = true;
10138 else
10139 return SDValue();
10140 }
10141
10142 SDLoc DL(N);
10143 if (WideVT != SourceValue.getValueType()) {
10144 assert(SourceValue.getValueType().getScalarSizeInBits() > WideNumBits &&
10145 "Unexpected store value to merge");
10146 SourceValue = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: WideVT, Operand: SourceValue);
10147 }
10148
10149 // Before legalize we can introduce illegal bswaps/rotates which will be later
10150 // converted to an explicit bswap sequence. This way we end up with a single
10151 // store and byte shuffling instead of several stores and byte shuffling.
10152 if (NeedBswap) {
10153 SourceValue = DAG.getNode(Opcode: ISD::BSWAP, DL, VT: WideVT, Operand: SourceValue);
10154 } else if (NeedRotate) {
10155 assert(WideNumBits % 2 == 0 && "Unexpected type for rotate");
10156 SDValue RotAmt = DAG.getConstant(Val: WideNumBits / 2, DL, VT: WideVT);
10157 SourceValue = DAG.getNode(Opcode: ISD::ROTR, DL, VT: WideVT, N1: SourceValue, N2: RotAmt);
10158 }
10159
10160 SDValue NewStore =
10161 DAG.getStore(Chain, dl: DL, Val: SourceValue, Ptr: FirstStore->getBasePtr(),
10162 PtrInfo: FirstStore->getPointerInfo(), Alignment: FirstStore->getAlign());
10163
10164 // Rely on other DAG combine rules to remove the other individual stores.
10165 DAG.ReplaceAllUsesWith(From: N, To: NewStore.getNode());
10166 return NewStore;
10167}
10168
10169/// Match a pattern where a wide type scalar value is loaded by several narrow
10170/// loads and combined by shifts and ors. Fold it into a single load or a load
10171/// and a BSWAP if the targets supports it.
10172///
10173/// Assuming little endian target:
10174/// i8 *a = ...
10175/// i32 val = a[0] | (a[1] << 8) | (a[2] << 16) | (a[3] << 24)
10176/// =>
10177/// i32 val = *((i32)a)
10178///
10179/// i8 *a = ...
10180/// i32 val = (a[0] << 24) | (a[1] << 16) | (a[2] << 8) | a[3]
10181/// =>
10182/// i32 val = BSWAP(*((i32)a))
10183///
10184/// TODO: This rule matches complex patterns with OR node roots and doesn't
10185/// interact well with the worklist mechanism. When a part of the pattern is
10186/// updated (e.g. one of the loads) its direct users are put into the worklist,
10187/// but the root node of the pattern which triggers the load combine is not
10188/// necessarily a direct user of the changed node. For example, once the address
10189/// of t28 load is reassociated load combine won't be triggered:
10190/// t25: i32 = add t4, Constant:i32<2>
10191/// t26: i64 = sign_extend t25
10192/// t27: i64 = add t2, t26
10193/// t28: i8,ch = load<LD1[%tmp9]> t0, t27, undef:i64
10194/// t29: i32 = zero_extend t28
10195/// t32: i32 = shl t29, Constant:i8<8>
10196/// t33: i32 = or t23, t32
10197/// As a possible fix visitLoad can check if the load can be a part of a load
10198/// combine pattern and add corresponding OR roots to the worklist.
10199SDValue DAGCombiner::MatchLoadCombine(SDNode *N) {
10200 assert(N->getOpcode() == ISD::OR &&
10201 "Can only match load combining against OR nodes");
10202
10203 // Handles simple types only
10204 EVT VT = N->getValueType(ResNo: 0);
10205 if (VT != MVT::i16 && VT != MVT::i32 && VT != MVT::i64)
10206 return SDValue();
10207 unsigned ByteWidth = VT.getSizeInBits() / 8;
10208
10209 bool IsBigEndianTarget = DAG.getDataLayout().isBigEndian();
10210 auto MemoryByteOffset = [&](SDByteProvider P) {
10211 assert(P.hasSrc() && "Must be a memory byte provider");
10212 auto *Load = cast<LoadSDNode>(Val: P.Src.value());
10213
10214 unsigned LoadBitWidth = Load->getMemoryVT().getScalarSizeInBits();
10215
10216 assert(LoadBitWidth % 8 == 0 &&
10217 "can only analyze providers for individual bytes not bit");
10218 unsigned LoadByteWidth = LoadBitWidth / 8;
10219 return IsBigEndianTarget ? bigEndianByteAt(BW: LoadByteWidth, i: P.DestOffset)
10220 : littleEndianByteAt(BW: LoadByteWidth, i: P.DestOffset);
10221 };
10222
10223 std::optional<BaseIndexOffset> Base;
10224 SDValue Chain;
10225
10226 SmallPtrSet<LoadSDNode *, 8> Loads;
10227 std::optional<SDByteProvider> FirstByteProvider;
10228 int64_t FirstOffset = INT64_MAX;
10229
10230 // Check if all the bytes of the OR we are looking at are loaded from the same
10231 // base address. Collect bytes offsets from Base address in ByteOffsets.
10232 SmallVector<int64_t, 8> ByteOffsets(ByteWidth);
10233 SmallVector<uint8_t, 8> ByteMasks(ByteWidth, 0xFF);
10234 unsigned ZeroExtendedBytes = 0;
10235 for (int i = ByteWidth - 1; i >= 0; --i) {
10236 auto P =
10237 calculateByteProvider(Op: SDValue(N, 0), Index: i, Depth: 0, /*VectorIndex*/ std::nullopt,
10238 /*StartingIndex*/ i, ByteMask: ByteMasks);
10239 if (!P)
10240 return SDValue();
10241
10242 if (P->isConstantZero()) {
10243 // It's OK for the N most significant bytes to be 0, we can just
10244 // zero-extend the load.
10245 if (++ZeroExtendedBytes != (ByteWidth - static_cast<unsigned>(i)))
10246 return SDValue();
10247 continue;
10248 }
10249 assert(P->hasSrc() && "provenance should either be memory or zero");
10250 auto *L = cast<LoadSDNode>(Val: P->Src.value());
10251
10252 // All loads must share the same chain
10253 SDValue LChain = L->getChain();
10254 if (!Chain)
10255 Chain = LChain;
10256 else if (Chain != LChain)
10257 return SDValue();
10258
10259 // Loads must share the same base address
10260 BaseIndexOffset Ptr = BaseIndexOffset::match(N: L, DAG);
10261 int64_t ByteOffsetFromBase = 0;
10262
10263 // For vector loads, the expected load combine pattern will have an
10264 // ExtractElement for each index in the vector. While each of these
10265 // ExtractElements will be accessing the same base address as determined
10266 // by the load instruction, the actual bytes they interact with will differ
10267 // due to different ExtractElement indices. To accurately determine the
10268 // byte position of an ExtractElement, we offset the base load ptr with
10269 // the index multiplied by the byte size of each element in the vector.
10270 if (L->getMemoryVT().isVector()) {
10271 unsigned LoadWidthInBit = L->getMemoryVT().getScalarSizeInBits();
10272 if (LoadWidthInBit % 8 != 0)
10273 return SDValue();
10274 unsigned ByteOffsetFromVector = P->SrcOffset * LoadWidthInBit / 8;
10275 Ptr.addToOffset(VectorOff: ByteOffsetFromVector);
10276 }
10277
10278 if (!Base)
10279 Base = Ptr;
10280
10281 else if (!Base->equalBaseIndex(Other: Ptr, DAG, Off&: ByteOffsetFromBase))
10282 return SDValue();
10283
10284 // Calculate the offset of the current byte from the base address
10285 ByteOffsetFromBase += MemoryByteOffset(*P);
10286 ByteOffsets[i] = ByteOffsetFromBase;
10287
10288 // Remember the first byte load
10289 if (ByteOffsetFromBase < FirstOffset) {
10290 FirstByteProvider = P;
10291 FirstOffset = ByteOffsetFromBase;
10292 }
10293
10294 Loads.insert(Ptr: L);
10295 }
10296
10297 assert(!Loads.empty() && "All the bytes of the value must be loaded from "
10298 "memory, so there must be at least one load which produces the value");
10299 assert(Base && "Base address of the accessed memory location must be set");
10300 assert(FirstOffset != INT64_MAX && "First byte offset must be set");
10301
10302 bool NeedsZext = ZeroExtendedBytes > 0;
10303
10304 EVT MemVT =
10305 EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: (ByteWidth - ZeroExtendedBytes) * 8);
10306
10307 if (!MemVT.isSimple())
10308 return SDValue();
10309
10310 // Check if the bytes of the OR we are looking at match with either big or
10311 // little endian value load
10312 std::optional<bool> IsBigEndian = isBigEndian(
10313 ByteOffsets: ArrayRef(ByteOffsets).drop_back(N: ZeroExtendedBytes), FirstOffset);
10314 if (!IsBigEndian)
10315 return SDValue();
10316
10317 assert(FirstByteProvider && "must be set");
10318
10319 // Ensure that the first byte is loaded from zero offset of the first load.
10320 // So the combined value can be loaded from the first load address.
10321 if (MemoryByteOffset(*FirstByteProvider) != 0)
10322 return SDValue();
10323 auto *FirstLoad = cast<LoadSDNode>(Val: FirstByteProvider->Src.value());
10324
10325 // Before legalization we allow introducing loads that are wider than legal,
10326 // which will later be split into legally sized loads. This enables us to
10327 // combine, for example, i8 loads forming an i64 into an i64 load, which get
10328 // then gets split up into couple of i32 loads on 32 bit targets.
10329 if (LegalOperations &&
10330 !TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: FirstLoad->getAlign(),
10331 AddrSpace: FirstLoad->getAddressSpace(),
10332 ExtType: NeedsZext ? ISD::ZEXTLOAD : ISD::NON_EXTLOAD, Atomic: false))
10333 return SDValue();
10334
10335 // The node we are looking at matches with the pattern, check if we can
10336 // replace it with a single (possibly zero-extended) load and bswap + shift if
10337 // needed.
10338
10339 // If the load needs byte swap check if the target supports it
10340 bool NeedsBswap = IsBigEndianTarget != *IsBigEndian;
10341
10342 // Before legalize we can introduce illegal bswaps which will be later
10343 // converted to an explicit bswap sequence. This way we end up with a single
10344 // load and byte shuffling instead of several loads and byte shuffling.
10345 // We do not introduce illegal bswaps when zero-extending as this tends to
10346 // introduce too many arithmetic instructions.
10347 if (NeedsBswap && (LegalOperations || NeedsZext) &&
10348 !TLI.isOperationLegal(Op: ISD::BSWAP, VT))
10349 return SDValue();
10350
10351 // If we need to bswap and zero extend, we have to insert a shift. Check that
10352 // it is legal.
10353 if (NeedsBswap && NeedsZext && LegalOperations &&
10354 !TLI.isOperationLegal(Op: ISD::SHL, VT))
10355 return SDValue();
10356
10357 // Check that a load of the wide type is both allowed and fast on the target
10358 unsigned Fast = 0;
10359 bool Allowed =
10360 TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT: MemVT,
10361 MMO: *FirstLoad->getMemOperand(), Fast: &Fast);
10362 if (!Allowed || !Fast)
10363 return SDValue();
10364
10365 SDValue NewLoad =
10366 DAG.getExtLoad(ExtType: NeedsZext ? ISD::ZEXTLOAD : ISD::NON_EXTLOAD, dl: SDLoc(N), VT,
10367 Chain, Ptr: FirstLoad->getBasePtr(),
10368 PtrInfo: FirstLoad->getPointerInfo(), MemVT, Alignment: FirstLoad->getAlign());
10369
10370 // Transfer chain users from old loads to the new load.
10371 for (LoadSDNode *L : Loads)
10372 DAG.makeEquivalentMemoryOrdering(OldLoad: L, NewMemOp: NewLoad);
10373
10374 // Apply combined mask if any bytes were partially masked by AND operations.
10375 bool HasPartialMask = false;
10376 uint64_t CombinedMask = 0;
10377 for (unsigned i = 0; i < ByteWidth; ++i) {
10378 CombinedMask |= (uint64_t)ByteMasks[i] << (i * 8);
10379 if (ByteMasks[i] != 0xFF)
10380 HasPartialMask = true;
10381 }
10382
10383 if (!NeedsBswap) {
10384 if (HasPartialMask)
10385 NewLoad = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(N), VT, N1: NewLoad,
10386 N2: DAG.getConstant(Val: CombinedMask, DL: SDLoc(N), VT));
10387 return NewLoad;
10388 }
10389
10390 SDValue ShiftedLoad =
10391 NeedsZext ? DAG.getNode(Opcode: ISD::SHL, DL: SDLoc(N), VT, N1: NewLoad,
10392 N2: DAG.getShiftAmountConstant(Val: ZeroExtendedBytes * 8,
10393 VT, DL: SDLoc(N)))
10394 : NewLoad;
10395 SDValue Result = DAG.getNode(Opcode: ISD::BSWAP, DL: SDLoc(N), VT, Operand: ShiftedLoad);
10396
10397 // The mask is built in final-result byte order (ByteMasks[i] corresponds to
10398 // byte i of the result), so it is correct to apply after the bswap.
10399 if (HasPartialMask)
10400 Result = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(N), VT, N1: Result,
10401 N2: DAG.getConstant(Val: CombinedMask, DL: SDLoc(N), VT));
10402
10403 return Result;
10404}
10405
10406// If the target has andn, bsl, or a similar bit-select instruction,
10407// we want to unfold masked merge, with canonical pattern of:
10408// | A | |B|
10409// ((x ^ y) & m) ^ y
10410// | D |
10411// Into:
10412// (x & m) | (y & ~m)
10413// If y is a constant, m is not a 'not', and the 'andn' does not work with
10414// immediates, we unfold into a different pattern:
10415// ~(~x & m) & (m | y)
10416// If x is a constant, m is a 'not', and the 'andn' does not work with
10417// immediates, we unfold into a different pattern:
10418// (x | ~m) & ~(~m & ~y)
10419// NOTE: we don't unfold the pattern if 'xor' is actually a 'not', because at
10420// the very least that breaks andnpd / andnps patterns, and because those
10421// patterns are simplified in IR and shouldn't be created in the DAG
10422SDValue DAGCombiner::unfoldMaskedMerge(SDNode *N) {
10423 assert(N->getOpcode() == ISD::XOR);
10424
10425 // Don't touch 'not' (i.e. where y = -1).
10426 if (isAllOnesOrAllOnesSplat(V: N->getOperand(Num: 1)))
10427 return SDValue();
10428
10429 EVT VT = N->getValueType(ResNo: 0);
10430
10431 // There are 3 commutable operators in the pattern,
10432 // so we have to deal with 8 possible variants of the basic pattern.
10433 SDValue X, Y, M;
10434 auto matchAndXor = [&X, &Y, &M](SDValue And, unsigned XorIdx, SDValue Other) {
10435 if (And.getOpcode() != ISD::AND || !And.hasOneUse())
10436 return false;
10437 SDValue Xor = And.getOperand(i: XorIdx);
10438 if (Xor.getOpcode() != ISD::XOR || !Xor.hasOneUse())
10439 return false;
10440 SDValue Xor0 = Xor.getOperand(i: 0);
10441 SDValue Xor1 = Xor.getOperand(i: 1);
10442 // Don't touch 'not' (i.e. where y = -1).
10443 if (isAllOnesOrAllOnesSplat(V: Xor1))
10444 return false;
10445 if (Other == Xor0)
10446 std::swap(a&: Xor0, b&: Xor1);
10447 if (Other != Xor1)
10448 return false;
10449 X = Xor0;
10450 Y = Xor1;
10451 M = And.getOperand(i: XorIdx ? 0 : 1);
10452 return true;
10453 };
10454
10455 SDValue N0 = N->getOperand(Num: 0);
10456 SDValue N1 = N->getOperand(Num: 1);
10457 if (!matchAndXor(N0, 0, N1) && !matchAndXor(N0, 1, N1) &&
10458 !matchAndXor(N1, 0, N0) && !matchAndXor(N1, 1, N0))
10459 return SDValue();
10460
10461 // Don't do anything if the mask is constant. This should not be reachable.
10462 // InstCombine should have already unfolded this pattern, and DAGCombiner
10463 // probably shouldn't produce it, too.
10464 if (isa<ConstantSDNode>(Val: M.getNode()))
10465 return SDValue();
10466
10467 // We can transform if the target has AndNot
10468 if (!TLI.hasAndNot(X: M))
10469 return SDValue();
10470
10471 SDLoc DL(N);
10472
10473 // If Y is a constant, check that 'andn' works with immediates. Unless M is
10474 // a bitwise not that would already allow ANDN to be used.
10475 if (!TLI.hasAndNot(X: Y) && !isBitwiseNot(V: M)) {
10476 assert(TLI.hasAndNot(X) && "Only mask is a variable? Unreachable.");
10477 // If not, we need to do a bit more work to make sure andn is still used.
10478 SDValue NotX = DAG.getNOT(DL, Val: X, VT);
10479 SDValue LHS = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: NotX, N2: M);
10480 SDValue NotLHS = DAG.getNOT(DL, Val: LHS, VT);
10481 SDValue RHS = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: M, N2: Y);
10482 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: NotLHS, N2: RHS);
10483 }
10484
10485 // If X is a constant and M is a bitwise not, check that 'andn' works with
10486 // immediates.
10487 if (!TLI.hasAndNot(X) && isBitwiseNot(V: M)) {
10488 assert(TLI.hasAndNot(Y) && "Only mask is a variable? Unreachable.");
10489 // If not, we need to do a bit more work to make sure andn is still used.
10490 SDValue NotM = M.getOperand(i: 0);
10491 SDValue LHS = DAG.getNode(Opcode: ISD::OR, DL, VT, N1: X, N2: NotM);
10492 SDValue NotY = DAG.getNOT(DL, Val: Y, VT);
10493 SDValue RHS = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: NotM, N2: NotY);
10494 SDValue NotRHS = DAG.getNOT(DL, Val: RHS, VT);
10495 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: LHS, N2: NotRHS);
10496 }
10497
10498 SDValue LHS = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: X, N2: M);
10499 SDValue NotM = DAG.getNOT(DL, Val: M, VT);
10500 SDValue RHS = DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Y, N2: NotM);
10501
10502 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: LHS, N2: RHS);
10503}
10504
10505SDValue DAGCombiner::visitXOR(SDNode *N) {
10506 SDValue N0 = N->getOperand(Num: 0);
10507 SDValue N1 = N->getOperand(Num: 1);
10508 EVT VT = N0.getValueType();
10509 SDLoc DL(N);
10510
10511 // fold (xor undef, undef) -> 0. This is a common idiom (misuse).
10512 if (N0.isUndef() && N1.isUndef())
10513 return DAG.getConstant(Val: 0, DL, VT);
10514
10515 // fold (xor x, undef) -> undef
10516 if (N0.isUndef())
10517 return N0;
10518 if (N1.isUndef())
10519 return N1;
10520
10521 // fold (xor c1, c2) -> c1^c2
10522 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::XOR, DL, VT, Ops: {N0, N1}))
10523 return C;
10524
10525 // canonicalize constant to RHS
10526 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
10527 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
10528 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1, N2: N0);
10529
10530 // fold vector ops
10531 if (VT.isVector()) {
10532 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
10533 return FoldedVOp;
10534
10535 // fold (xor x, 0) -> x, vector edition
10536 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
10537 return N0;
10538 }
10539
10540 // fold (xor x, 0) -> x
10541 if (isNullConstant(V: N1))
10542 return N0;
10543
10544 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
10545 return NewSel;
10546
10547 // reassociate xor
10548 if (SDValue RXOR = reassociateOps(Opc: ISD::XOR, DL, N0, N1, Flags: N->getFlags()))
10549 return RXOR;
10550
10551 // Fold xor(vecreduce(x), vecreduce(y)) -> vecreduce(xor(x, y))
10552 if (SDValue SD =
10553 reassociateReduction(RedOpc: ISD::VECREDUCE_XOR, Opc: ISD::XOR, DL, VT, N0, N1))
10554 return SD;
10555
10556 // fold (a^b) -> (a|b) iff a and b share no bits.
10557 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::OR, VT)) &&
10558 DAG.haveNoCommonBitsSet(A: N0, B: N1))
10559 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: N0, N2: N1, Flags: SDNodeFlags::Disjoint);
10560
10561 // look for 'add-like' folds:
10562 // XOR(N0,MIN_SIGNED_VALUE) == ADD(N0,MIN_SIGNED_VALUE)
10563 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::ADD, VT)) &&
10564 isMinSignedConstant(V: N1))
10565 if (SDValue Combined = visitADDLike(N))
10566 return Combined;
10567
10568 // fold not (setcc x, y, cc) -> setcc x y !cc
10569 // Avoid breaking: and (not(setcc x, y, cc), z) -> andn for vec
10570 unsigned N0Opcode = N0.getOpcode();
10571 SDValue LHS, RHS, CC;
10572 if (TLI.isConstTrueVal(N: N1) &&
10573 isSetCCEquivalent(N: N0, LHS, RHS, CC, /*MatchStrict*/ true) &&
10574 !(VT.isVector() && TLI.hasAndNot(X: SDValue(N, 0)) && N->hasOneUse() &&
10575 N->use_begin()->getUser()->getOpcode() == ISD::AND)) {
10576 ISD::CondCode NotCC = ISD::getSetCCInverse(Operation: cast<CondCodeSDNode>(Val&: CC)->get(),
10577 Type: LHS.getValueType());
10578 if (!LegalOperations ||
10579 TLI.isCondCodeLegal(CC: NotCC, VT: LHS.getSimpleValueType())) {
10580 // Propagate fast-math-flags.
10581 SDNodeFlags Flags = N0->getFlags();
10582 switch (N0Opcode) {
10583 default:
10584 llvm_unreachable("Unhandled SetCC Equivalent!");
10585 case ISD::SETCC:
10586 return DAG.getSetCC(DL: SDLoc(N0), VT, LHS, RHS, Cond: NotCC, Chain: SDValue(),
10587 /*IsSignaling=*/false, Flags);
10588 case ISD::SELECT_CC:
10589 return DAG.getSelectCC(DL: SDLoc(N0), LHS, RHS, True: N0.getOperand(i: 2),
10590 False: N0.getOperand(i: 3), Cond: NotCC, Flags);
10591 case ISD::STRICT_FSETCC:
10592 case ISD::STRICT_FSETCCS: {
10593 if (N0.hasOneUse()) {
10594 // FIXME Can we handle multiple uses? Could we token factor the chain
10595 // results from the new/old setcc?
10596 SDValue SetCC =
10597 DAG.getSetCC(DL: SDLoc(N0), VT, LHS, RHS, Cond: NotCC, Chain: N0.getOperand(i: 0),
10598 IsSignaling: N0Opcode == ISD::STRICT_FSETCCS, Flags);
10599 CombineTo(N, Res: SetCC);
10600 DAG.ReplaceAllUsesOfValueWith(From: N0.getValue(R: 1), To: SetCC.getValue(R: 1));
10601 recursivelyDeleteUnusedNodes(N: N0.getNode());
10602 return SDValue(N, 0); // Return N so it doesn't get rechecked!
10603 }
10604 break;
10605 }
10606 }
10607 }
10608 }
10609
10610 // fold (not (zext (setcc x, y))) -> (zext (not (setcc x, y)))
10611 if (isOneConstant(V: N1) && N0Opcode == ISD::ZERO_EXTEND && N0.hasOneUse() &&
10612 isSetCCEquivalent(N: N0.getOperand(i: 0), LHS, RHS, CC)){
10613 SDValue V = N0.getOperand(i: 0);
10614 SDLoc DL0(N0);
10615 V = DAG.getNode(Opcode: ISD::XOR, DL: DL0, VT: V.getValueType(), N1: V,
10616 N2: DAG.getConstant(Val: 1, DL: DL0, VT: V.getValueType()));
10617 AddToWorklist(N: V.getNode());
10618 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: V);
10619 }
10620
10621 // fold (not (or x, y)) -> (and (not x), (not y)) iff x or y are setcc
10622 // fold (not (and x, y)) -> (or (not x), (not y)) iff x or y are setcc
10623 if (isOneConstant(V: N1) && VT == MVT::i1 && N0.hasOneUse() &&
10624 (N0Opcode == ISD::OR || N0Opcode == ISD::AND)) {
10625 SDValue N00 = N0.getOperand(i: 0), N01 = N0.getOperand(i: 1);
10626 if (isOneUseSetCC(N: N01) || isOneUseSetCC(N: N00)) {
10627 unsigned NewOpcode = N0Opcode == ISD::AND ? ISD::OR : ISD::AND;
10628 N00 = DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N00), VT, N1: N00, N2: N1); // N00 = ~N00
10629 N01 = DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N01), VT, N1: N01, N2: N1); // N01 = ~N01
10630 AddToWorklist(N: N00.getNode()); AddToWorklist(N: N01.getNode());
10631 return DAG.getNode(Opcode: NewOpcode, DL, VT, N1: N00, N2: N01);
10632 }
10633 }
10634 // fold (not (or x, y)) -> (and (not x), (not y)) iff x or y are constants
10635 // fold (not (and x, y)) -> (or (not x), (not y)) iff x or y are constants
10636 if (isAllOnesConstant(V: N1) && N0.hasOneUse() &&
10637 (N0Opcode == ISD::OR || N0Opcode == ISD::AND)) {
10638 SDValue N00 = N0.getOperand(i: 0), N01 = N0.getOperand(i: 1);
10639 if (isa<ConstantSDNode>(Val: N01) || isa<ConstantSDNode>(Val: N00)) {
10640 unsigned NewOpcode = N0Opcode == ISD::AND ? ISD::OR : ISD::AND;
10641 N00 = DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N00), VT, N1: N00, N2: N1); // N00 = ~N00
10642 N01 = DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N01), VT, N1: N01, N2: N1); // N01 = ~N01
10643 AddToWorklist(N: N00.getNode()); AddToWorklist(N: N01.getNode());
10644 return DAG.getNode(Opcode: NewOpcode, DL, VT, N1: N00, N2: N01);
10645 }
10646 }
10647
10648 // fold (not (sub Y, X)) -> (add X, ~Y) if Y is a constant
10649 if (N0.getOpcode() == ISD::SUB && isAllOnesConstant(V: N1)) {
10650 SDValue Y = N0.getOperand(i: 0);
10651 SDValue X = N0.getOperand(i: 1);
10652
10653 if (auto *YConst = dyn_cast<ConstantSDNode>(Val&: Y)) {
10654 APInt NotYValue = ~YConst->getAPIntValue();
10655 SDValue NotY = DAG.getConstant(Val: NotYValue, DL, VT);
10656 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: X, N2: NotY, Flags: N->getFlags());
10657 }
10658 }
10659
10660 // fold (not (add X, -1)) -> (neg X)
10661 if (N0.getOpcode() == ISD::ADD && N0.hasOneUse() && isAllOnesConstant(V: N1) &&
10662 isAllOnesOrAllOnesSplat(V: N0.getOperand(i: 1))) {
10663 return DAG.getNegative(Val: N0.getOperand(i: 0), DL, VT);
10664 }
10665
10666 // fold (xor (and x, y), y) -> (and (not x), y)
10667 if (N0Opcode == ISD::AND && N0.hasOneUse() && N0->getOperand(Num: 1) == N1) {
10668 SDValue X = N0.getOperand(i: 0);
10669 SDValue NotX = DAG.getNOT(DL: SDLoc(X), Val: X, VT);
10670 AddToWorklist(N: NotX.getNode());
10671 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: NotX, N2: N1);
10672 }
10673
10674 // fold Y = sra (X, size(X)-1); xor (add (X, Y), Y) -> (abs X)
10675 if (!LegalOperations || hasOperation(Opcode: ISD::ABS, VT)) {
10676 SDValue A = N0Opcode == ISD::ADD ? N0 : N1;
10677 SDValue S = N0Opcode == ISD::SRA ? N0 : N1;
10678 if (A.getOpcode() == ISD::ADD && S.getOpcode() == ISD::SRA) {
10679 SDValue A0 = A.getOperand(i: 0), A1 = A.getOperand(i: 1);
10680 SDValue S0 = S.getOperand(i: 0);
10681 if ((A0 == S && A1 == S0) || (A1 == S && A0 == S0))
10682 if (ConstantSDNode *C = isConstOrConstSplat(N: S.getOperand(i: 1)))
10683 if (C->getAPIntValue() == (VT.getScalarSizeInBits() - 1))
10684 return DAG.getNode(Opcode: ISD::ABS, DL, VT, Operand: S0);
10685 }
10686 }
10687
10688 // fold (xor x, x) -> 0
10689 if (N0 == N1)
10690 return tryFoldToZero(DL, TLI, VT, DAG, LegalOperations);
10691
10692 // fold (xor (shl 1, x), -1) -> (rotl ~1, x)
10693 // Here is a concrete example of this equivalence:
10694 // i16 x == 14
10695 // i16 shl == 1 << 14 == 16384 == 0b0100000000000000
10696 // i16 xor == ~(1 << 14) == 49151 == 0b1011111111111111
10697 //
10698 // =>
10699 //
10700 // i16 ~1 == 0b1111111111111110
10701 // i16 rol(~1, 14) == 0b1011111111111111
10702 //
10703 // Some additional tips to help conceptualize this transform:
10704 // - Try to see the operation as placing a single zero in a value of all ones.
10705 // - There exists no value for x which would allow the result to contain zero.
10706 // - Values of x larger than the bitwidth are undefined and do not require a
10707 // consistent result.
10708 // - Pushing the zero left requires shifting one bits in from the right.
10709 // A rotate left of ~1 is a nice way of achieving the desired result.
10710 if (TLI.isOperationLegalOrCustom(Op: ISD::ROTL, VT) && N0Opcode == ISD::SHL &&
10711 isAllOnesConstant(V: N1) && isOneConstant(V: N0.getOperand(i: 0))) {
10712 return DAG.getNode(Opcode: ISD::ROTL, DL, VT, N1: DAG.getSignedConstant(Val: ~1, DL, VT),
10713 N2: N0.getOperand(i: 1));
10714 }
10715
10716 // Simplify: xor (op x...), (op y...) -> (op (xor x, y))
10717 if (N0Opcode == N1.getOpcode())
10718 if (SDValue V = hoistLogicOpWithSameOpcodeHands(N))
10719 return V;
10720
10721 if (SDValue R = foldLogicOfShifts(N, LogicOp: N0, ShiftOp: N1, DAG))
10722 return R;
10723 if (SDValue R = foldLogicOfShifts(N, LogicOp: N1, ShiftOp: N0, DAG))
10724 return R;
10725 if (SDValue R = foldLogicTreeOfShifts(N, LeftHand: N0, RightHand: N1, DAG))
10726 return R;
10727
10728 // Unfold ((x ^ y) & m) ^ y into (x & m) | (y & ~m) if profitable
10729 if (SDValue MM = unfoldMaskedMerge(N))
10730 return MM;
10731
10732 // Simplify the expression using non-local knowledge.
10733 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
10734 return SDValue(N, 0);
10735
10736 if (SDValue Combined = combineCarryDiamond(DAG, TLI, N0, N1, N))
10737 return Combined;
10738
10739 // fold (xor (smin(x, C), C)) -> select (x < C), xor(x, C), 0
10740 // fold (xor (smax(x, C), C)) -> select (x > C), xor(x, C), 0
10741 // fold (xor (umin(x, C), C)) -> select (x < C), xor(x, C), 0
10742 // fold (xor (umax(x, C), C)) -> select (x > C), xor(x, C), 0
10743 SDValue Op0;
10744 if (sd_match(N: N0, P: m_OneUse(P: m_AnyOf(preds: m_SMin(L: m_Value(N&: Op0), R: m_Specific(N: N1)),
10745 preds: m_SMax(L: m_Value(N&: Op0), R: m_Specific(N: N1)),
10746 preds: m_UMin(L: m_Value(N&: Op0), R: m_Specific(N: N1)),
10747 preds: m_UMax(L: m_Value(N&: Op0), R: m_Specific(N: N1)))))) {
10748
10749 if (isa<ConstantSDNode>(Val: N1) ||
10750 ISD::isBuildVectorOfConstantSDNodes(N: N1.getNode())) {
10751 // For vectors, only optimize when the constant is zero or all-ones to
10752 // avoid generating more instructions
10753 if (VT.isVector()) {
10754 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
10755 if (!N1C || (!N1C->isZero() && !N1C->isAllOnes()))
10756 return SDValue();
10757 }
10758
10759 // Avoid the fold if the minmax operation is legal and select is expensive
10760 if (TLI.isOperationLegal(Op: N0.getOpcode(), VT) &&
10761 TLI.isPredictableSelectExpensive())
10762 return SDValue();
10763
10764 EVT CCVT = getSetCCResultType(VT);
10765 ISD::CondCode CC;
10766 switch (N0.getOpcode()) {
10767 case ISD::SMIN:
10768 CC = ISD::SETLT;
10769 break;
10770 case ISD::SMAX:
10771 CC = ISD::SETGT;
10772 break;
10773 case ISD::UMIN:
10774 CC = ISD::SETULT;
10775 break;
10776 case ISD::UMAX:
10777 CC = ISD::SETUGT;
10778 break;
10779 }
10780 SDValue FN1 = DAG.getFreeze(V: N1);
10781 SDValue Cmp = DAG.getSetCC(DL, VT: CCVT, LHS: Op0, RHS: FN1, Cond: CC);
10782 SDValue XorXC = DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: Op0, N2: FN1);
10783 SDValue Zero = DAG.getConstant(Val: 0, DL, VT);
10784 return DAG.getSelect(DL, VT, Cond: Cmp, LHS: XorXC, RHS: Zero);
10785 }
10786 }
10787
10788 return SDValue();
10789}
10790
10791/// If we have a shift-by-constant of a bitwise logic op that itself has a
10792/// shift-by-constant operand with identical opcode, we may be able to convert
10793/// that into 2 independent shifts followed by the logic op. This is a
10794/// throughput improvement.
10795static SDValue combineShiftOfShiftedLogic(SDNode *Shift, SelectionDAG &DAG) {
10796 // Match a one-use bitwise logic op.
10797 SDValue LogicOp = Shift->getOperand(Num: 0);
10798 if (!LogicOp.hasOneUse())
10799 return SDValue();
10800
10801 unsigned LogicOpcode = LogicOp.getOpcode();
10802 if (LogicOpcode != ISD::AND && LogicOpcode != ISD::OR &&
10803 LogicOpcode != ISD::XOR)
10804 return SDValue();
10805
10806 // Find a matching one-use shift by constant.
10807 unsigned ShiftOpcode = Shift->getOpcode();
10808 SDValue C1 = Shift->getOperand(Num: 1);
10809 ConstantSDNode *C1Node = isConstOrConstSplat(N: C1);
10810 assert(C1Node && "Expected a shift with constant operand");
10811 const APInt &C1Val = C1Node->getAPIntValue();
10812 auto matchFirstShift = [&](SDValue V, SDValue &ShiftOp,
10813 const APInt *&ShiftAmtVal) {
10814 if (V.getOpcode() != ShiftOpcode || !V.hasOneUse())
10815 return false;
10816
10817 ConstantSDNode *ShiftCNode = isConstOrConstSplat(N: V.getOperand(i: 1));
10818 if (!ShiftCNode)
10819 return false;
10820
10821 // Capture the shifted operand and shift amount value.
10822 ShiftOp = V.getOperand(i: 0);
10823 ShiftAmtVal = &ShiftCNode->getAPIntValue();
10824
10825 // Shift amount types do not have to match their operand type, so check that
10826 // the constants are the same width.
10827 if (ShiftAmtVal->getBitWidth() != C1Val.getBitWidth())
10828 return false;
10829
10830 // The fold is not valid if the sum of the shift values doesn't fit in the
10831 // given shift amount type.
10832 bool Overflow = false;
10833 APInt NewShiftAmt = C1Val.uadd_ov(RHS: *ShiftAmtVal, Overflow);
10834 if (Overflow)
10835 return false;
10836
10837 // The fold is not valid if the sum of the shift values exceeds bitwidth.
10838 if (NewShiftAmt.uge(RHS: V.getScalarValueSizeInBits()))
10839 return false;
10840
10841 return true;
10842 };
10843
10844 // Logic ops are commutative, so check each operand for a match.
10845 SDValue X, Y;
10846 const APInt *C0Val;
10847 if (matchFirstShift(LogicOp.getOperand(i: 0), X, C0Val))
10848 Y = LogicOp.getOperand(i: 1);
10849 else if (matchFirstShift(LogicOp.getOperand(i: 1), X, C0Val))
10850 Y = LogicOp.getOperand(i: 0);
10851 else
10852 return SDValue();
10853
10854 // shift (logic (shift X, C0), Y), C1 -> logic (shift X, C0+C1), (shift Y, C1)
10855 SDLoc DL(Shift);
10856 EVT VT = Shift->getValueType(ResNo: 0);
10857 EVT ShiftAmtVT = Shift->getOperand(Num: 1).getValueType();
10858 SDValue ShiftSumC = DAG.getConstant(Val: *C0Val + C1Val, DL, VT: ShiftAmtVT);
10859 SDValue NewShift1 = DAG.getNode(Opcode: ShiftOpcode, DL, VT, N1: X, N2: ShiftSumC);
10860 SDValue NewShift2 = DAG.getNode(Opcode: ShiftOpcode, DL, VT, N1: Y, N2: C1);
10861 return DAG.getNode(Opcode: LogicOpcode, DL, VT, N1: NewShift1, N2: NewShift2,
10862 Flags: LogicOp->getFlags());
10863}
10864
10865/// Handle transforms common to the three shifts, when the shift amount is a
10866/// constant.
10867/// We are looking for: (shift being one of shl/sra/srl)
10868/// shift (binop X, C0), C1
10869/// And want to transform into:
10870/// binop (shift X, C1), (shift C0, C1)
10871SDValue DAGCombiner::visitShiftByConstant(SDNode *N) {
10872 assert(isConstOrConstSplat(N->getOperand(1)) && "Expected constant operand");
10873
10874 // Do not turn a 'not' into a regular xor.
10875 if (isBitwiseNot(V: N->getOperand(Num: 0)))
10876 return SDValue();
10877
10878 // The inner binop must be one-use, since we want to replace it.
10879 SDValue LHS = N->getOperand(Num: 0);
10880 if (!LHS.hasOneUse() || !TLI.isDesirableToCommuteWithShift(N, Level))
10881 return SDValue();
10882
10883 // Fold shift(bitop(shift(x,c1),y), c2) -> bitop(shift(x,c1+c2),shift(y,c2)).
10884 if (SDValue R = combineShiftOfShiftedLogic(Shift: N, DAG))
10885 return R;
10886
10887 // We want to pull some binops through shifts, so that we have (and (shift))
10888 // instead of (shift (and)), likewise for add, or, xor, etc. This sort of
10889 // thing happens with address calculations, so it's important to canonicalize
10890 // it.
10891 switch (LHS.getOpcode()) {
10892 default:
10893 return SDValue();
10894 case ISD::OR:
10895 case ISD::XOR:
10896 case ISD::AND:
10897 break;
10898 case ISD::ADD:
10899 if (N->getOpcode() != ISD::SHL)
10900 return SDValue(); // only shl(add) not sr[al](add).
10901 break;
10902 }
10903
10904 // FIXME: disable this unless the input to the binop is a shift by a constant
10905 // or is copy/select. Enable this in other cases when figure out it's exactly
10906 // profitable.
10907 SDValue BinOpLHSVal = LHS.getOperand(i: 0);
10908 bool IsShiftByConstant = (BinOpLHSVal.getOpcode() == ISD::SHL ||
10909 BinOpLHSVal.getOpcode() == ISD::SRA ||
10910 BinOpLHSVal.getOpcode() == ISD::SRL) &&
10911 isa<ConstantSDNode>(Val: BinOpLHSVal.getOperand(i: 1));
10912 bool IsCopyOrSelect = BinOpLHSVal.getOpcode() == ISD::CopyFromReg ||
10913 BinOpLHSVal.getOpcode() == ISD::SELECT;
10914
10915 if (!IsShiftByConstant && !IsCopyOrSelect)
10916 return SDValue();
10917
10918 if (IsCopyOrSelect && N->hasOneUse())
10919 return SDValue();
10920
10921 // Attempt to fold the constants, shifting the binop RHS by the shift amount.
10922 SDLoc DL(N);
10923 EVT VT = N->getValueType(ResNo: 0);
10924 if (SDValue NewRHS = DAG.FoldConstantArithmetic(
10925 Opcode: N->getOpcode(), DL, VT, Ops: {LHS.getOperand(i: 1), N->getOperand(Num: 1)})) {
10926 SDValue NewShift = DAG.getNode(Opcode: N->getOpcode(), DL, VT, N1: LHS.getOperand(i: 0),
10927 N2: N->getOperand(Num: 1));
10928 return DAG.getNode(Opcode: LHS.getOpcode(), DL, VT, N1: NewShift, N2: NewRHS);
10929 }
10930
10931 return SDValue();
10932}
10933
10934SDValue DAGCombiner::distributeTruncateThroughAnd(SDNode *N) {
10935 assert(N->getOpcode() == ISD::TRUNCATE);
10936 assert(N->getOperand(0).getOpcode() == ISD::AND);
10937
10938 // (truncate:TruncVT (and N00, N01C)) -> (and (truncate:TruncVT N00), TruncC)
10939 EVT TruncVT = N->getValueType(ResNo: 0);
10940 if (N->hasOneUse() && N->getOperand(Num: 0).hasOneUse() &&
10941 TLI.isTypeDesirableForOp(ISD::AND, VT: TruncVT)) {
10942 SDValue N01 = N->getOperand(Num: 0).getOperand(i: 1);
10943 if (isConstantOrConstantVector(N: N01, /* NoOpaques */ true)) {
10944 SDLoc DL(N);
10945 SDValue N00 = N->getOperand(Num: 0).getOperand(i: 0);
10946 SDValue Trunc00 = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: TruncVT, Operand: N00);
10947 SDValue Trunc01 = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: TruncVT, Operand: N01);
10948 AddToWorklist(N: Trunc00.getNode());
10949 AddToWorklist(N: Trunc01.getNode());
10950 return DAG.getNode(Opcode: ISD::AND, DL, VT: TruncVT, N1: Trunc00, N2: Trunc01);
10951 }
10952 }
10953
10954 return SDValue();
10955}
10956
10957SDValue DAGCombiner::visitRotate(SDNode *N) {
10958 SDLoc dl(N);
10959 SDValue N0 = N->getOperand(Num: 0);
10960 SDValue N1 = N->getOperand(Num: 1);
10961 EVT VT = N->getValueType(ResNo: 0);
10962 unsigned Bitsize = VT.getScalarSizeInBits();
10963
10964 // fold (rot x, 0) -> x
10965 if (isNullOrNullSplat(V: N1))
10966 return N0;
10967
10968 // fold (rot x, c) -> x iff (c % BitSize) == 0
10969 if (isPowerOf2_32(Value: Bitsize) && Bitsize > 1) {
10970 APInt ModuloMask(N1.getScalarValueSizeInBits(), Bitsize - 1);
10971 if (DAG.MaskedValueIsZero(Op: N1, Mask: ModuloMask))
10972 return N0;
10973 }
10974
10975 // fold (rot x, c) -> (rot x, c % BitSize)
10976 bool OutOfRange = false;
10977 auto MatchOutOfRange = [Bitsize, &OutOfRange](ConstantSDNode *C) {
10978 OutOfRange |= C->getAPIntValue().uge(RHS: Bitsize);
10979 return true;
10980 };
10981 if (ISD::matchUnaryPredicate(Op: N1, Match: MatchOutOfRange) && OutOfRange) {
10982 EVT AmtVT = N1.getValueType();
10983 SDValue Bits = DAG.getConstant(Val: Bitsize, DL: dl, VT: AmtVT);
10984 if (SDValue Amt =
10985 DAG.FoldConstantArithmetic(Opcode: ISD::UREM, DL: dl, VT: AmtVT, Ops: {N1, Bits}))
10986 return DAG.getNode(Opcode: N->getOpcode(), DL: dl, VT, N1: N0, N2: Amt);
10987 }
10988
10989 // rot i16 X, 8 --> bswap X
10990 auto *RotAmtC = isConstOrConstSplat(N: N1);
10991 if (RotAmtC && RotAmtC->getAPIntValue() == 8 &&
10992 VT.getScalarSizeInBits() == 16 && hasOperation(Opcode: ISD::BSWAP, VT))
10993 return DAG.getNode(Opcode: ISD::BSWAP, DL: dl, VT, Operand: N0);
10994
10995 // Simplify the operands using demanded-bits information.
10996 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
10997 return SDValue(N, 0);
10998
10999 // fold (rot* x, (trunc (and y, c))) -> (rot* x, (and (trunc y), (trunc c))).
11000 if (N1.getOpcode() == ISD::TRUNCATE &&
11001 N1.getOperand(i: 0).getOpcode() == ISD::AND) {
11002 if (SDValue NewOp1 = distributeTruncateThroughAnd(N: N1.getNode()))
11003 return DAG.getNode(Opcode: N->getOpcode(), DL: dl, VT, N1: N0, N2: NewOp1);
11004 }
11005
11006 unsigned NextOp = N0.getOpcode();
11007
11008 // fold (rot* (rot* x, c2), c1)
11009 // -> (rot* x, ((c1 % bitsize) +- (c2 % bitsize) + bitsize) % bitsize)
11010 if (NextOp == ISD::ROTL || NextOp == ISD::ROTR) {
11011 bool C1 = DAG.isConstantIntBuildVectorOrConstantInt(N: N1);
11012 bool C2 = DAG.isConstantIntBuildVectorOrConstantInt(N: N0.getOperand(i: 1));
11013 if (C1 && C2 && N1.getValueType() == N0.getOperand(i: 1).getValueType()) {
11014 EVT ShiftVT = N1.getValueType();
11015 bool SameSide = (N->getOpcode() == NextOp);
11016 unsigned CombineOp = SameSide ? ISD::ADD : ISD::SUB;
11017 SDValue BitsizeC = DAG.getConstant(Val: Bitsize, DL: dl, VT: ShiftVT);
11018 SDValue Norm1 = DAG.FoldConstantArithmetic(Opcode: ISD::UREM, DL: dl, VT: ShiftVT,
11019 Ops: {N1, BitsizeC});
11020 SDValue Norm2 = DAG.FoldConstantArithmetic(Opcode: ISD::UREM, DL: dl, VT: ShiftVT,
11021 Ops: {N0.getOperand(i: 1), BitsizeC});
11022 if (Norm1 && Norm2)
11023 if (SDValue CombinedShift = DAG.FoldConstantArithmetic(
11024 Opcode: CombineOp, DL: dl, VT: ShiftVT, Ops: {Norm1, Norm2})) {
11025 CombinedShift = DAG.FoldConstantArithmetic(Opcode: ISD::ADD, DL: dl, VT: ShiftVT,
11026 Ops: {CombinedShift, BitsizeC});
11027 SDValue CombinedShiftNorm = DAG.FoldConstantArithmetic(
11028 Opcode: ISD::UREM, DL: dl, VT: ShiftVT, Ops: {CombinedShift, BitsizeC});
11029 return DAG.getNode(Opcode: N->getOpcode(), DL: dl, VT, N1: N0->getOperand(Num: 0),
11030 N2: CombinedShiftNorm);
11031 }
11032 }
11033 }
11034 return SDValue();
11035}
11036
11037SDValue DAGCombiner::visitSHL(SDNode *N) {
11038 SDValue N0 = N->getOperand(Num: 0);
11039 SDValue N1 = N->getOperand(Num: 1);
11040 if (SDValue V = DAG.simplifyShift(X: N0, Y: N1))
11041 return V;
11042
11043 SDLoc DL(N);
11044 EVT VT = N0.getValueType();
11045 EVT ShiftVT = N1.getValueType();
11046 unsigned OpSizeInBits = VT.getScalarSizeInBits();
11047
11048 // fold (shl c1, c2) -> c1<<c2
11049 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::SHL, DL, VT, Ops: {N0, N1}))
11050 return C;
11051
11052 // fold vector ops
11053 if (VT.isVector()) {
11054 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
11055 return FoldedVOp;
11056
11057 BuildVectorSDNode *N1CV = dyn_cast<BuildVectorSDNode>(Val&: N1);
11058 // If setcc produces all-one true value then:
11059 // (shl (and (setcc) N01CV) N1CV) -> (and (setcc) N01CV<<N1CV)
11060 if (N1CV && N1CV->isConstant()) {
11061 if (N0.getOpcode() == ISD::AND) {
11062 SDValue N00 = N0->getOperand(Num: 0);
11063 SDValue N01 = N0->getOperand(Num: 1);
11064 BuildVectorSDNode *N01CV = dyn_cast<BuildVectorSDNode>(Val&: N01);
11065
11066 if (N01CV && N01CV->isConstant() && N00.getOpcode() == ISD::SETCC &&
11067 TLI.getBooleanContents(Type: N00.getOperand(i: 0).getValueType()) ==
11068 TargetLowering::ZeroOrNegativeOneBooleanContent) {
11069 if (SDValue C =
11070 DAG.FoldConstantArithmetic(Opcode: ISD::SHL, DL, VT, Ops: {N01, N1}))
11071 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N00, N2: C);
11072 }
11073 }
11074 }
11075 }
11076
11077 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
11078 return NewSel;
11079
11080 // if (shl x, c) is known to be zero, return 0
11081 if (DAG.MaskedValueIsZero(Op: SDValue(N, 0), Mask: APInt::getAllOnes(numBits: OpSizeInBits)))
11082 return DAG.getConstant(Val: 0, DL, VT);
11083
11084 // fold (shl x, (trunc (and y, c))) -> (shl x, (and (trunc y), (trunc c))).
11085 if (N1.getOpcode() == ISD::TRUNCATE &&
11086 N1.getOperand(i: 0).getOpcode() == ISD::AND) {
11087 if (SDValue NewOp1 = distributeTruncateThroughAnd(N: N1.getNode()))
11088 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0, N2: NewOp1);
11089 }
11090
11091 // fold (shl (shl x, c1), c2) -> 0 or (shl x, (add c1, c2))
11092 if (N0.getOpcode() == ISD::SHL) {
11093 auto MatchOutOfRange = [OpSizeInBits](ConstantSDNode *LHS,
11094 ConstantSDNode *RHS) {
11095 APInt c1 = LHS->getAPIntValue();
11096 APInt c2 = RHS->getAPIntValue();
11097 zeroExtendToMatch(LHS&: c1, RHS&: c2, Offset: 1 /* Overflow Bit */);
11098 return (c1 + c2).uge(RHS: OpSizeInBits);
11099 };
11100 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: MatchOutOfRange))
11101 return DAG.getConstant(Val: 0, DL, VT);
11102
11103 auto MatchInRange = [OpSizeInBits](ConstantSDNode *LHS,
11104 ConstantSDNode *RHS) {
11105 APInt c1 = LHS->getAPIntValue();
11106 APInt c2 = RHS->getAPIntValue();
11107 zeroExtendToMatch(LHS&: c1, RHS&: c2, Offset: 1 /* Overflow Bit */);
11108 return (c1 + c2).ult(RHS: OpSizeInBits);
11109 };
11110 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: MatchInRange)) {
11111 SDValue Sum = DAG.getNode(Opcode: ISD::ADD, DL, VT: ShiftVT, N1, N2: N0.getOperand(i: 1));
11112 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0.getOperand(i: 0), N2: Sum);
11113 }
11114 }
11115
11116 // fold (shl (ext (shl x, c1)), c2) -> (shl (ext x), (add c1, c2))
11117 // For this to be valid, the second form must not preserve any of the bits
11118 // that are shifted out by the inner shift in the first form. This means
11119 // the outer shift size must be >= the number of bits added by the ext.
11120 // As a corollary, we don't care what kind of ext it is.
11121 if ((N0.getOpcode() == ISD::ZERO_EXTEND ||
11122 N0.getOpcode() == ISD::ANY_EXTEND ||
11123 N0.getOpcode() == ISD::SIGN_EXTEND) &&
11124 N0.getOperand(i: 0).getOpcode() == ISD::SHL) {
11125 SDValue N0Op0 = N0.getOperand(i: 0);
11126 SDValue InnerShiftAmt = N0Op0.getOperand(i: 1);
11127 EVT InnerVT = N0Op0.getValueType();
11128 uint64_t InnerBitwidth = InnerVT.getScalarSizeInBits();
11129
11130 auto MatchOutOfRange = [OpSizeInBits, InnerBitwidth](ConstantSDNode *LHS,
11131 ConstantSDNode *RHS) {
11132 APInt c1 = LHS->getAPIntValue();
11133 APInt c2 = RHS->getAPIntValue();
11134 zeroExtendToMatch(LHS&: c1, RHS&: c2, Offset: 1 /* Overflow Bit */);
11135 return c2.uge(RHS: OpSizeInBits - InnerBitwidth) &&
11136 (c1 + c2).uge(RHS: OpSizeInBits);
11137 };
11138 if (ISD::matchBinaryPredicate(LHS: InnerShiftAmt, RHS: N1, Match: MatchOutOfRange,
11139 /*AllowUndefs*/ false,
11140 /*AllowTypeMismatch*/ true))
11141 return DAG.getConstant(Val: 0, DL, VT);
11142
11143 auto MatchInRange = [OpSizeInBits, InnerBitwidth](ConstantSDNode *LHS,
11144 ConstantSDNode *RHS) {
11145 APInt c1 = LHS->getAPIntValue();
11146 APInt c2 = RHS->getAPIntValue();
11147 zeroExtendToMatch(LHS&: c1, RHS&: c2, Offset: 1 /* Overflow Bit */);
11148 return c2.uge(RHS: OpSizeInBits - InnerBitwidth) &&
11149 (c1 + c2).ult(RHS: OpSizeInBits);
11150 };
11151 if (ISD::matchBinaryPredicate(LHS: InnerShiftAmt, RHS: N1, Match: MatchInRange,
11152 /*AllowUndefs*/ false,
11153 /*AllowTypeMismatch*/ true)) {
11154 SDValue Ext = DAG.getNode(Opcode: N0.getOpcode(), DL, VT, Operand: N0Op0.getOperand(i: 0));
11155 SDValue Sum = DAG.getZExtOrTrunc(Op: InnerShiftAmt, DL, VT: ShiftVT);
11156 Sum = DAG.getNode(Opcode: ISD::ADD, DL, VT: ShiftVT, N1: Sum, N2: N1);
11157 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Ext, N2: Sum);
11158 }
11159 }
11160
11161 // fold (shl (zext (srl x, C)), C) -> (zext (shl (srl x, C), C))
11162 // Only fold this if the inner zext has no other uses to avoid increasing
11163 // the total number of instructions.
11164 if (N0.getOpcode() == ISD::ZERO_EXTEND && N0.hasOneUse() &&
11165 N0.getOperand(i: 0).getOpcode() == ISD::SRL) {
11166 SDValue N0Op0 = N0.getOperand(i: 0);
11167 SDValue InnerShiftAmt = N0Op0.getOperand(i: 1);
11168
11169 auto MatchEqual = [VT](ConstantSDNode *LHS, ConstantSDNode *RHS) {
11170 APInt c1 = LHS->getAPIntValue();
11171 APInt c2 = RHS->getAPIntValue();
11172 zeroExtendToMatch(LHS&: c1, RHS&: c2);
11173 return c1.ult(RHS: VT.getScalarSizeInBits()) && (c1 == c2);
11174 };
11175 if (ISD::matchBinaryPredicate(LHS: InnerShiftAmt, RHS: N1, Match: MatchEqual,
11176 /*AllowUndefs*/ false,
11177 /*AllowTypeMismatch*/ true)) {
11178 EVT InnerShiftAmtVT = N0Op0.getOperand(i: 1).getValueType();
11179 SDValue NewSHL = DAG.getZExtOrTrunc(Op: N1, DL, VT: InnerShiftAmtVT);
11180 NewSHL = DAG.getNode(Opcode: ISD::SHL, DL, VT: N0Op0.getValueType(), N1: N0Op0, N2: NewSHL);
11181 AddToWorklist(N: NewSHL.getNode());
11182 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL: SDLoc(N0), VT, Operand: NewSHL);
11183 }
11184 }
11185
11186 if (N0.getOpcode() == ISD::SRL || N0.getOpcode() == ISD::SRA) {
11187 auto MatchShiftAmount = [OpSizeInBits](ConstantSDNode *LHS,
11188 ConstantSDNode *RHS) {
11189 const APInt &LHSC = LHS->getAPIntValue();
11190 const APInt &RHSC = RHS->getAPIntValue();
11191 return LHSC.ult(RHS: OpSizeInBits) && RHSC.ult(RHS: OpSizeInBits) &&
11192 LHSC.getZExtValue() <= RHSC.getZExtValue();
11193 };
11194
11195 // fold (shl (sr[la] exact X, C1), C2) -> (shl X, (C2-C1)) if C1 <= C2
11196 // fold (shl (sr[la] exact X, C1), C2) -> (sr[la] X, (C2-C1)) if C1 >= C2
11197 if (N0->getFlags().hasExact()) {
11198 if (ISD::matchBinaryPredicate(LHS: N0.getOperand(i: 1), RHS: N1, Match: MatchShiftAmount,
11199 /*AllowUndefs*/ false,
11200 /*AllowTypeMismatch*/ true)) {
11201 SDValue N01 = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 1), DL, VT: ShiftVT);
11202 SDValue Diff = DAG.getNode(Opcode: ISD::SUB, DL, VT: ShiftVT, N1, N2: N01);
11203 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0.getOperand(i: 0), N2: Diff);
11204 }
11205 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: MatchShiftAmount,
11206 /*AllowUndefs*/ false,
11207 /*AllowTypeMismatch*/ true)) {
11208 SDValue N01 = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 1), DL, VT: ShiftVT);
11209 SDValue Diff = DAG.getNode(Opcode: ISD::SUB, DL, VT: ShiftVT, N1: N01, N2: N1);
11210 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, N1: N0.getOperand(i: 0), N2: Diff);
11211 }
11212 }
11213
11214 // fold (shl (srl x, c1), c2) -> (and (shl x, (sub c2, c1), MASK) or
11215 // (and (srl x, (sub c1, c2), MASK)
11216 // Only fold this if the inner shift has no other uses -- if it does,
11217 // folding this will increase the total number of instructions.
11218 if (N0.getOpcode() == ISD::SRL &&
11219 (N0.getOperand(i: 1) == N1 || N0.hasOneUse()) &&
11220 TLI.shouldFoldConstantShiftPairToMask(N)) {
11221 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: MatchShiftAmount,
11222 /*AllowUndefs*/ false,
11223 /*AllowTypeMismatch*/ true)) {
11224 SDValue N01 = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 1), DL, VT: ShiftVT);
11225 SDValue Diff = DAG.getNode(Opcode: ISD::SUB, DL, VT: ShiftVT, N1: N01, N2: N1);
11226 SDValue Mask = DAG.getAllOnesConstant(DL, VT);
11227 Mask = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Mask, N2: N01);
11228 Mask = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: Mask, N2: Diff);
11229 SDValue Shift = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0.getOperand(i: 0), N2: Diff);
11230 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Shift, N2: Mask);
11231 }
11232 if (ISD::matchBinaryPredicate(LHS: N0.getOperand(i: 1), RHS: N1, Match: MatchShiftAmount,
11233 /*AllowUndefs*/ false,
11234 /*AllowTypeMismatch*/ true)) {
11235 SDValue N01 = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 1), DL, VT: ShiftVT);
11236 SDValue Diff = DAG.getNode(Opcode: ISD::SUB, DL, VT: ShiftVT, N1, N2: N01);
11237 SDValue Mask = DAG.getAllOnesConstant(DL, VT);
11238 Mask = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Mask, N2: N1);
11239 SDValue Shift = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0.getOperand(i: 0), N2: Diff);
11240 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Shift, N2: Mask);
11241 }
11242 }
11243 }
11244
11245 // fold (shl (sra x, c1), c1) -> (and x, (shl -1, c1))
11246 if (N0.getOpcode() == ISD::SRA && N1 == N0.getOperand(i: 1) &&
11247 isConstantOrConstantVector(N: N1, /* No Opaques */ NoOpaques: true)) {
11248 SDValue AllBits = DAG.getAllOnesConstant(DL, VT);
11249 SDValue HiBitsMask = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: AllBits, N2: N1);
11250 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: N0.getOperand(i: 0), N2: HiBitsMask);
11251 }
11252
11253 // fold (shl (add x, c1), c2) -> (add (shl x, c2), c1 << c2)
11254 // fold (shl (or x, c1), c2) -> (or (shl x, c2), c1 << c2)
11255 // Variant of version done on multiply, except mul by a power of 2 is turned
11256 // into a shift.
11257 if ((N0.getOpcode() == ISD::ADD || N0.getOpcode() == ISD::OR) &&
11258 TLI.isDesirableToCommuteWithShift(N, Level)) {
11259 SDValue N01 = N0.getOperand(i: 1);
11260 if (SDValue Shl1 =
11261 DAG.FoldConstantArithmetic(Opcode: ISD::SHL, DL: SDLoc(N1), VT, Ops: {N01, N1})) {
11262 SDValue Shl0 = DAG.getNode(Opcode: ISD::SHL, DL: SDLoc(N0), VT, N1: N0.getOperand(i: 0), N2: N1);
11263 AddToWorklist(N: Shl0.getNode());
11264 SDNodeFlags Flags;
11265 // Preserve the disjoint flag for Or.
11266 if (N0.getOpcode() == ISD::OR && N0->getFlags().hasDisjoint())
11267 Flags |= SDNodeFlags::Disjoint;
11268 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, N1: Shl0, N2: Shl1, Flags);
11269 }
11270 }
11271
11272 // fold (shl (sext (add_nsw x, c1)), c2) -> (add (shl (sext x), c2), c1 << c2)
11273 // TODO: Add zext/add_nuw variant with suitable test coverage
11274 // TODO: Should we limit this with isLegalAddImmediate?
11275 if (N0.getOpcode() == ISD::SIGN_EXTEND &&
11276 N0.getOperand(i: 0).getOpcode() == ISD::ADD &&
11277 N0.getOperand(i: 0)->getFlags().hasNoSignedWrap() &&
11278 TLI.isDesirableToCommuteWithShift(N, Level)) {
11279 SDValue Add = N0.getOperand(i: 0);
11280 SDLoc DL(N0);
11281 if (SDValue ExtC = DAG.FoldConstantArithmetic(Opcode: N0.getOpcode(), DL, VT,
11282 Ops: {Add.getOperand(i: 1)})) {
11283 if (SDValue ShlC =
11284 DAG.FoldConstantArithmetic(Opcode: ISD::SHL, DL, VT, Ops: {ExtC, N1})) {
11285 SDValue ExtX = DAG.getNode(Opcode: N0.getOpcode(), DL, VT, Operand: Add.getOperand(i: 0));
11286 SDValue ShlX = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: ExtX, N2: N1);
11287 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: ShlX, N2: ShlC);
11288 }
11289 }
11290 }
11291
11292 // fold (shl (mul x, c1), c2) -> (mul x, c1 << c2)
11293 if (N0.getOpcode() == ISD::MUL && N0->hasOneUse()) {
11294 SDValue N01 = N0.getOperand(i: 1);
11295 if (SDValue Shl =
11296 DAG.FoldConstantArithmetic(Opcode: ISD::SHL, DL: SDLoc(N1), VT, Ops: {N01, N1}))
11297 return DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: N0.getOperand(i: 0), N2: Shl);
11298 }
11299
11300 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
11301 if (N1C && !N1C->isOpaque())
11302 if (SDValue NewSHL = visitShiftByConstant(N))
11303 return NewSHL;
11304
11305 // fold (shl X, cttz(Y)) -> (mul (Y & -Y), X) if cttz is unsupported on the
11306 // target.
11307 if (((N1.getOpcode() == ISD::CTTZ &&
11308 VT.getScalarSizeInBits() <= ShiftVT.getScalarSizeInBits()) ||
11309 N1.getOpcode() == ISD::CTTZ_ZERO_POISON) &&
11310 N1.hasOneUse() && !TLI.isOperationLegalOrCustom(Op: ISD::CTTZ, VT: ShiftVT) &&
11311 TLI.isOperationLegalOrCustom(Op: ISD::MUL, VT)) {
11312 SDValue Y = N1.getOperand(i: 0);
11313 SDLoc DL(N);
11314 SDValue NegY = DAG.getNegative(Val: Y, DL, VT: ShiftVT);
11315 SDValue And =
11316 DAG.getZExtOrTrunc(Op: DAG.getNode(Opcode: ISD::AND, DL, VT: ShiftVT, N1: Y, N2: NegY), DL, VT);
11317 return DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: And, N2: N0);
11318 }
11319
11320 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
11321 return SDValue(N, 0);
11322
11323 // Fold (shl (vscale * C0), C1) to (vscale * (C0 << C1)).
11324 if (N0.getOpcode() == ISD::VSCALE && N1C) {
11325 const APInt &C0 = N0.getConstantOperandAPInt(i: 0);
11326 const APInt &C1 = N1C->getAPIntValue();
11327 return DAG.getVScale(DL, VT, MulImm: C0 << C1);
11328 }
11329
11330 SDValue X;
11331 APInt VS0;
11332
11333 // fold (shl (X * vscale(VS0)), C1) -> (X * vscale(VS0 << C1))
11334 if (N1C && sd_match(N: N0, P: m_Mul(L: m_Value(N&: X), R: m_VScale(Op: m_ConstInt(V&: VS0))))) {
11335 SDNodeFlags Flags;
11336 Flags.setNoUnsignedWrap(N->getFlags().hasNoUnsignedWrap() &&
11337 N0->getFlags().hasNoUnsignedWrap());
11338
11339 SDValue VScale = DAG.getVScale(DL, VT, MulImm: VS0 << N1C->getAPIntValue());
11340 return DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: X, N2: VScale, Flags);
11341 }
11342
11343 // Fold (shl step_vector(C0), C1) to (step_vector(C0 << C1)).
11344 APInt ShlVal;
11345 if (N0.getOpcode() == ISD::STEP_VECTOR &&
11346 ISD::isConstantSplatVector(N: N1.getNode(), SplatValue&: ShlVal)) {
11347 const APInt &C0 = N0.getConstantOperandAPInt(i: 0);
11348 if (ShlVal.ult(RHS: C0.getBitWidth())) {
11349 APInt NewStep = C0 << ShlVal;
11350 return DAG.getStepVector(DL, ResVT: VT, StepVal: NewStep);
11351 }
11352 }
11353
11354 return SDValue();
11355}
11356
11357// Transform a right shift of a multiply into a multiply-high.
11358// Examples:
11359// (srl (mul (zext i32:$a to i64), (zext i32:$a to i64)), 32) -> (mulhu $a, $b)
11360// (sra (mul (sext i32:$a to i64), (sext i32:$a to i64)), 32) -> (mulhs $a, $b)
11361static SDValue combineShiftToMULH(SDNode *N, const SDLoc &DL, SelectionDAG &DAG,
11362 const TargetLowering &TLI) {
11363 assert((N->getOpcode() == ISD::SRL || N->getOpcode() == ISD::SRA) &&
11364 "SRL or SRA node is required here!");
11365
11366 // Check the shift amount. Proceed with the transformation if the shift
11367 // amount is constant.
11368 ConstantSDNode *ShiftAmtSrc = isConstOrConstSplat(N: N->getOperand(Num: 1));
11369 if (!ShiftAmtSrc)
11370 return SDValue();
11371
11372 // The operation feeding into the shift must be a multiply.
11373 SDValue ShiftOperand = N->getOperand(Num: 0);
11374 if (ShiftOperand.getOpcode() != ISD::MUL)
11375 return SDValue();
11376
11377 // Both operands must be equivalent extend nodes.
11378 SDValue LeftOp = ShiftOperand.getOperand(i: 0);
11379 SDValue RightOp = ShiftOperand.getOperand(i: 1);
11380
11381 if (LeftOp.getOpcode() != ISD::SIGN_EXTEND &&
11382 LeftOp.getOpcode() != ISD::ZERO_EXTEND)
11383 std::swap(a&: LeftOp, b&: RightOp);
11384
11385 bool IsSignExt = LeftOp.getOpcode() == ISD::SIGN_EXTEND;
11386 bool IsZeroExt = LeftOp.getOpcode() == ISD::ZERO_EXTEND;
11387
11388 if (!IsSignExt && !IsZeroExt)
11389 return SDValue();
11390
11391 EVT NarrowVT = LeftOp.getOperand(i: 0).getValueType();
11392 unsigned NarrowVTSize = NarrowVT.getScalarSizeInBits();
11393
11394 // return true if U may use the lower bits of its operands
11395 auto UserOfLowerBits = [NarrowVTSize](SDNode *U) {
11396 if (U->getOpcode() != ISD::SRL && U->getOpcode() != ISD::SRA) {
11397 return true;
11398 }
11399 ConstantSDNode *UShiftAmtSrc = isConstOrConstSplat(N: U->getOperand(Num: 1));
11400 if (!UShiftAmtSrc) {
11401 return true;
11402 }
11403 unsigned UShiftAmt = UShiftAmtSrc->getZExtValue();
11404 return UShiftAmt < NarrowVTSize;
11405 };
11406
11407 // If the lower part of the MUL is also used and MUL_LOHI is supported
11408 // do not introduce the MULH in favor of MUL_LOHI
11409 unsigned MulLoHiOp = IsSignExt ? ISD::SMUL_LOHI : ISD::UMUL_LOHI;
11410 if (!ShiftOperand.hasOneUse() &&
11411 TLI.isOperationLegalOrCustom(Op: MulLoHiOp, VT: NarrowVT) &&
11412 llvm::any_of(Range: ShiftOperand->users(), P: UserOfLowerBits)) {
11413 return SDValue();
11414 }
11415
11416 SDValue MulhRightOp;
11417 if (LeftOp.getOpcode() != RightOp.getOpcode()) {
11418 if (IsZeroExt && ShiftOperand.hasOneUse() &&
11419 DAG.computeKnownBits(Op: RightOp).countMaxActiveBits() <= NarrowVTSize) {
11420 MulhRightOp = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: NarrowVT, Operand: RightOp);
11421 } else if (IsSignExt && ShiftOperand.hasOneUse() &&
11422 DAG.ComputeMaxSignificantBits(Op: RightOp) <= NarrowVTSize) {
11423 MulhRightOp = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: NarrowVT, Operand: RightOp);
11424 } else {
11425 return SDValue();
11426 }
11427 } else {
11428 // Check that the two extend nodes are the same type.
11429 if (NarrowVT != RightOp.getOperand(i: 0).getValueType())
11430 return SDValue();
11431 MulhRightOp = RightOp.getOperand(i: 0);
11432 }
11433
11434 EVT WideVT = LeftOp.getValueType();
11435 // Proceed with the transformation if the wide types match.
11436 assert((WideVT == RightOp.getValueType()) &&
11437 "Cannot have a multiply node with two different operand types.");
11438
11439 // Proceed with the transformation if the wide type is twice as large
11440 // as the narrow type.
11441 if (WideVT.getScalarSizeInBits() != 2 * NarrowVTSize)
11442 return SDValue();
11443
11444 // Check the shift amount with the narrow type size.
11445 // Proceed with the transformation if the shift amount is the width
11446 // of the narrow type.
11447 unsigned ShiftAmt = ShiftAmtSrc->getZExtValue();
11448 if (ShiftAmt != NarrowVTSize)
11449 return SDValue();
11450
11451 // If the operation feeding into the MUL is a sign extend (sext),
11452 // we use mulhs. Othewise, zero extends (zext) use mulhu.
11453 unsigned MulhOpcode = IsSignExt ? ISD::MULHS : ISD::MULHU;
11454
11455 // Combine to mulh if mulh is legal/custom for the narrow type on the target
11456 // or if it is a vector type then we could transform to an acceptable type and
11457 // rely on legalization to split/combine the result.
11458 EVT TransformVT = NarrowVT;
11459 if (NarrowVT.isVector()) {
11460 TransformVT = TLI.getLegalTypeToTransformTo(Context&: *DAG.getContext(), VT: NarrowVT);
11461 if (TransformVT.getScalarType() != NarrowVT.getScalarType())
11462 return SDValue();
11463 }
11464 if (!TLI.isOperationLegalOrCustom(Op: MulhOpcode, VT: TransformVT))
11465 return SDValue();
11466
11467 SDValue Result =
11468 DAG.getNode(Opcode: MulhOpcode, DL, VT: NarrowVT, N1: LeftOp.getOperand(i: 0), N2: MulhRightOp);
11469 bool IsSigned = N->getOpcode() == ISD::SRA;
11470 return DAG.getExtOrTrunc(IsSigned, Op: Result, DL, VT: WideVT);
11471}
11472
11473// fold (bswap (logic_op(bswap(x),y))) -> logic_op(x,bswap(y))
11474// This helper function accept SDNode with opcode ISD::BSWAP and ISD::BITREVERSE
11475static SDValue foldBitOrderCrossLogicOp(SDNode *N, SelectionDAG &DAG) {
11476 unsigned Opcode = N->getOpcode();
11477 if (Opcode != ISD::BSWAP && Opcode != ISD::BITREVERSE)
11478 return SDValue();
11479
11480 SDValue N0 = N->getOperand(Num: 0);
11481 EVT VT = N->getValueType(ResNo: 0);
11482 SDLoc DL(N);
11483 SDValue X, Y;
11484
11485 // If both operands are bswap/bitreverse, ignore the multiuse
11486 if (sd_match(N: N0, P: m_OneUse(P: m_BitwiseLogic(L: m_UnaryOp(Opc: Opcode, Op: m_Value(N&: X)),
11487 R: m_UnaryOp(Opc: Opcode, Op: m_Value(N&: Y))))))
11488 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, N1: X, N2: Y);
11489
11490 // Otherwise need to ensure logic_op and bswap/bitreverse(x) have one use.
11491 if (sd_match(N: N0, P: m_OneUse(P: m_BitwiseLogic(
11492 L: m_OneUse(P: m_UnaryOp(Opc: Opcode, Op: m_Value(N&: X))), R: m_Value(N&: Y))))) {
11493 SDValue NewBitReorder = DAG.getNode(Opcode, DL, VT, Operand: Y);
11494 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, N1: X, N2: NewBitReorder);
11495 }
11496
11497 return SDValue();
11498}
11499
11500SDValue DAGCombiner::visitSRA(SDNode *N) {
11501 SDValue N0 = N->getOperand(Num: 0);
11502 SDValue N1 = N->getOperand(Num: 1);
11503 if (SDValue V = DAG.simplifyShift(X: N0, Y: N1))
11504 return V;
11505
11506 SDLoc DL(N);
11507 EVT VT = N0.getValueType();
11508 unsigned OpSizeInBits = VT.getScalarSizeInBits();
11509
11510 // fold (sra c1, c2) -> (sra c1, c2)
11511 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::SRA, DL, VT, Ops: {N0, N1}))
11512 return C;
11513
11514 // Arithmetic shifting an all-sign-bit value is a no-op.
11515 // fold (sra 0, x) -> 0
11516 // fold (sra -1, x) -> -1
11517 if (DAG.ComputeNumSignBits(Op: N0) == OpSizeInBits)
11518 return N0;
11519
11520 // fold vector ops
11521 if (VT.isVector())
11522 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
11523 return FoldedVOp;
11524
11525 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
11526 return NewSel;
11527
11528 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
11529
11530 // fold (sra (sra x, c1), c2) -> (sra x, (add c1, c2))
11531 // clamp (add c1, c2) to max shift.
11532 if (N0.getOpcode() == ISD::SRA) {
11533 EVT ShiftVT = N1.getValueType();
11534 EVT ShiftSVT = ShiftVT.getScalarType();
11535 SmallVector<SDValue, 16> ShiftValues;
11536
11537 auto SumOfShifts = [&](ConstantSDNode *LHS, ConstantSDNode *RHS) {
11538 APInt c1 = LHS->getAPIntValue();
11539 APInt c2 = RHS->getAPIntValue();
11540 zeroExtendToMatch(LHS&: c1, RHS&: c2, Offset: 1 /* Overflow Bit */);
11541 APInt Sum = c1 + c2;
11542 unsigned ShiftSum =
11543 Sum.uge(RHS: OpSizeInBits) ? (OpSizeInBits - 1) : Sum.getZExtValue();
11544 ShiftValues.push_back(Elt: DAG.getConstant(Val: ShiftSum, DL, VT: ShiftSVT));
11545 return true;
11546 };
11547 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: SumOfShifts)) {
11548 SDValue ShiftValue;
11549 if (N1.getOpcode() == ISD::BUILD_VECTOR)
11550 ShiftValue = DAG.getBuildVector(VT: ShiftVT, DL, Ops: ShiftValues);
11551 else if (N1.getOpcode() == ISD::SPLAT_VECTOR) {
11552 assert(ShiftValues.size() == 1 &&
11553 "Expected matchBinaryPredicate to return one element for "
11554 "SPLAT_VECTORs");
11555 ShiftValue = DAG.getSplatVector(VT: ShiftVT, DL, Op: ShiftValues[0]);
11556 } else
11557 ShiftValue = ShiftValues[0];
11558 return DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N0.getOperand(i: 0), N2: ShiftValue);
11559 }
11560 }
11561
11562 // fold (sra (xor (sra x, c1), -1), c2) -> (xor (sra x, c3), -1)
11563 // This allows merging two arithmetic shifts even when there's a NOT in
11564 // between.
11565 SDValue X;
11566 APInt C1;
11567 if (N1C && sd_match(N: N0, P: m_OneUse(P: m_Not(
11568 V: m_OneUse(P: m_Sra(L: m_Value(N&: X), R: m_ConstInt(V&: C1))))))) {
11569 APInt C2 = N1C->getAPIntValue();
11570 zeroExtendToMatch(LHS&: C1, RHS&: C2, Offset: 1 /* Overflow Bit */);
11571 APInt Sum = C1 + C2;
11572 unsigned ShiftSum = Sum.getLimitedValue(Limit: OpSizeInBits - 1);
11573 SDValue NewShift = DAG.getNode(
11574 Opcode: ISD::SRA, DL, VT, N1: X, N2: DAG.getShiftAmountConstant(Val: ShiftSum, VT, DL));
11575 return DAG.getNOT(DL, Val: NewShift, VT);
11576 }
11577
11578 // fold (sra (shl X, m), (sub result_size, n))
11579 // -> (sign_extend (trunc (shl X, (sub (sub result_size, n), m)))) for
11580 // result_size - n != m.
11581 // If truncate is free for the target sext(shl) is likely to result in better
11582 // code.
11583 if (N0.getOpcode() == ISD::SHL && N1C) {
11584 // Get the two constants of the shifts, CN0 = m, CN = n.
11585 const ConstantSDNode *N01C = isConstOrConstSplat(N: N0.getOperand(i: 1));
11586 if (N01C) {
11587 LLVMContext &Ctx = *DAG.getContext();
11588 // Determine what the truncate's result bitsize and type would be.
11589 EVT TruncVT = VT.changeElementType(
11590 Context&: Ctx, EltVT: EVT::getIntegerVT(Context&: Ctx, BitWidth: OpSizeInBits - N1C->getZExtValue()));
11591
11592 // Determine the residual right-shift amount.
11593 int ShiftAmt = N1C->getZExtValue() - N01C->getZExtValue();
11594
11595 // If the shift is not a no-op (in which case this should be just a sign
11596 // extend already), the truncated to type is legal, sign_extend is legal
11597 // on that type, and the truncate to that type is both legal and free,
11598 // perform the transform.
11599 if ((ShiftAmt > 0) &&
11600 TLI.isOperationLegalOrCustom(Op: ISD::SIGN_EXTEND, VT: TruncVT) &&
11601 TLI.isOperationLegalOrCustom(Op: ISD::TRUNCATE, VT) &&
11602 TLI.isTruncateFree(FromVT: VT, ToVT: TruncVT)) {
11603 SDValue Amt = DAG.getShiftAmountConstant(Val: ShiftAmt, VT, DL);
11604 SDValue Shift = DAG.getNode(Opcode: ISD::SRL, DL, VT,
11605 N1: N0.getOperand(i: 0), N2: Amt);
11606 SDValue Trunc = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: TruncVT,
11607 Operand: Shift);
11608 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL,
11609 VT: N->getValueType(ResNo: 0), Operand: Trunc);
11610 }
11611 }
11612 }
11613
11614 // We convert trunc/ext to opposing shifts in IR, but casts may be cheaper.
11615 // sra (add (shl X, N1C), AddC), N1C -->
11616 // sext (add (trunc X to (width - N1C)), AddC')
11617 // sra (sub AddC, (shl X, N1C)), N1C -->
11618 // sext (sub AddC1',(trunc X to (width - N1C)))
11619 if ((N0.getOpcode() == ISD::ADD || N0.getOpcode() == ISD::SUB) && N1C &&
11620 N0.hasOneUse()) {
11621 bool IsAdd = N0.getOpcode() == ISD::ADD;
11622 SDValue Shl = N0.getOperand(i: IsAdd ? 0 : 1);
11623 if (Shl.getOpcode() == ISD::SHL && Shl.getOperand(i: 1) == N1 &&
11624 Shl.hasOneUse()) {
11625 // TODO: AddC does not need to be a splat.
11626 if (ConstantSDNode *AddC =
11627 isConstOrConstSplat(N: N0.getOperand(i: IsAdd ? 1 : 0))) {
11628 // Determine what the truncate's type would be and ask the target if
11629 // that is a free operation.
11630 LLVMContext &Ctx = *DAG.getContext();
11631 unsigned ShiftAmt = N1C->getZExtValue();
11632 EVT TruncVT = VT.changeElementType(
11633 Context&: Ctx, EltVT: EVT::getIntegerVT(Context&: Ctx, BitWidth: OpSizeInBits - ShiftAmt));
11634
11635 // TODO: The simple type check probably belongs in the default hook
11636 // implementation and/or target-specific overrides (because
11637 // non-simple types likely require masking when legalized), but
11638 // that restriction may conflict with other transforms.
11639 if (TruncVT.isSimple() && isTypeLegal(VT: TruncVT) &&
11640 TLI.isTruncateFree(FromVT: VT, ToVT: TruncVT)) {
11641 SDValue Trunc = DAG.getZExtOrTrunc(Op: Shl.getOperand(i: 0), DL, VT: TruncVT);
11642 SDValue ShiftC =
11643 DAG.getConstant(Val: AddC->getAPIntValue().lshr(shiftAmt: ShiftAmt).trunc(
11644 width: TruncVT.getScalarSizeInBits()),
11645 DL, VT: TruncVT);
11646 SDValue Add;
11647 if (IsAdd)
11648 Add = DAG.getNode(Opcode: ISD::ADD, DL, VT: TruncVT, N1: Trunc, N2: ShiftC);
11649 else
11650 Add = DAG.getNode(Opcode: ISD::SUB, DL, VT: TruncVT, N1: ShiftC, N2: Trunc);
11651 return DAG.getSExtOrTrunc(Op: Add, DL, VT);
11652 }
11653 }
11654 }
11655 }
11656
11657 // fold (sra x, (trunc (and y, c))) -> (sra x, (and (trunc y), (trunc c))).
11658 if (N1.getOpcode() == ISD::TRUNCATE &&
11659 N1.getOperand(i: 0).getOpcode() == ISD::AND) {
11660 if (SDValue NewOp1 = distributeTruncateThroughAnd(N: N1.getNode()))
11661 return DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N0, N2: NewOp1);
11662 }
11663
11664 // fold (sra (trunc (sra x, c1)), c2) -> (trunc (sra x, c1 + c2))
11665 // fold (sra (trunc (srl x, c1)), c2) -> (trunc (sra x, c1 + c2))
11666 // if c1 is equal to the number of bits the trunc removes
11667 // TODO - support non-uniform vector shift amounts.
11668 if (N0.getOpcode() == ISD::TRUNCATE &&
11669 (N0.getOperand(i: 0).getOpcode() == ISD::SRL ||
11670 N0.getOperand(i: 0).getOpcode() == ISD::SRA) &&
11671 N0.getOperand(i: 0).hasOneUse() &&
11672 N0.getOperand(i: 0).getOperand(i: 1).hasOneUse() && N1C) {
11673 SDValue N0Op0 = N0.getOperand(i: 0);
11674 if (ConstantSDNode *LargeShift = isConstOrConstSplat(N: N0Op0.getOperand(i: 1))) {
11675 EVT LargeVT = N0Op0.getValueType();
11676 unsigned TruncBits = LargeVT.getScalarSizeInBits() - OpSizeInBits;
11677 if (LargeShift->getAPIntValue() == TruncBits) {
11678 EVT LargeShiftVT = getShiftAmountTy(LHSTy: LargeVT);
11679 SDValue Amt = DAG.getZExtOrTrunc(Op: N1, DL, VT: LargeShiftVT);
11680 Amt = DAG.getNode(Opcode: ISD::ADD, DL, VT: LargeShiftVT, N1: Amt,
11681 N2: DAG.getConstant(Val: TruncBits, DL, VT: LargeShiftVT));
11682 SDValue SRA =
11683 DAG.getNode(Opcode: ISD::SRA, DL, VT: LargeVT, N1: N0Op0.getOperand(i: 0), N2: Amt);
11684 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: SRA);
11685 }
11686 }
11687 }
11688
11689 // fold (sra (add nsw X, C), D) -> (add nsw (sra X, D), C s>> D)
11690 // when C has D trailing zeros (so C s>> D is exact).
11691 if (N1C && N0.hasOneUse() && N0.getOpcode() == ISD::ADD &&
11692 N0->getFlags().hasNoSignedWrap()) {
11693 if (ConstantSDNode *AddC = isConstOrConstSplat(N: N0.getOperand(i: 1))) {
11694 const APInt &ShAmt = N1C->getAPIntValue();
11695 const APInt &AddVal = AddC->getAPIntValue();
11696 if (ShAmt.ult(RHS: AddVal.countr_zero())) {
11697 SDNodeFlags ShiftFlags = N->getFlags();
11698 SDValue NewSra =
11699 DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N0.getOperand(i: 0), N2: N1, Flags: ShiftFlags);
11700 SDValue NewC = DAG.getConstant(Val: AddVal.ashr(ShiftAmt: ShAmt), DL, VT);
11701 SDNodeFlags AddFlags = N0->getFlags();
11702 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: NewSra, N2: NewC, Flags: AddFlags);
11703 }
11704 }
11705 }
11706
11707 // Simplify, based on bits shifted out of the LHS.
11708 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
11709 return SDValue(N, 0);
11710
11711 // If the sign bit is known to be zero, switch this to a SRL.
11712 if (DAG.SignBitIsZero(Op: N0))
11713 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0, N2: N1);
11714
11715 if (N1C && !N1C->isOpaque())
11716 if (SDValue NewSRA = visitShiftByConstant(N))
11717 return NewSRA;
11718
11719 // Try to transform this shift into a multiply-high if
11720 // it matches the appropriate pattern detected in combineShiftToMULH.
11721 if (SDValue MULH = combineShiftToMULH(N, DL, DAG, TLI))
11722 return MULH;
11723
11724 // Attempt to convert a sra of a load into a narrower sign-extending load.
11725 if (SDValue NarrowLoad = reduceLoadWidth(N))
11726 return NarrowLoad;
11727
11728 if (SDValue AVG = foldShiftToAvg(N, DL))
11729 return AVG;
11730
11731 return SDValue();
11732}
11733
11734SDValue DAGCombiner::visitSRL(SDNode *N) {
11735 SDValue N0 = N->getOperand(Num: 0);
11736 SDValue N1 = N->getOperand(Num: 1);
11737 if (SDValue V = DAG.simplifyShift(X: N0, Y: N1))
11738 return V;
11739
11740 SDLoc DL(N);
11741 EVT VT = N0.getValueType();
11742 EVT ShiftVT = N1.getValueType();
11743 unsigned OpSizeInBits = VT.getScalarSizeInBits();
11744
11745 // fold (srl c1, c2) -> c1 >>u c2
11746 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::SRL, DL, VT, Ops: {N0, N1}))
11747 return C;
11748
11749 // fold vector ops
11750 if (VT.isVector())
11751 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
11752 return FoldedVOp;
11753
11754 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
11755 return NewSel;
11756
11757 // if (srl x, c) is known to be zero, return 0
11758 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
11759 if (N1C &&
11760 DAG.MaskedValueIsZero(Op: SDValue(N, 0), Mask: APInt::getAllOnes(numBits: OpSizeInBits)))
11761 return DAG.getConstant(Val: 0, DL, VT);
11762
11763 // fold (srl (srl x, c1), c2) -> 0 or (srl x, (add c1, c2))
11764 if (N0.getOpcode() == ISD::SRL) {
11765 auto MatchOutOfRange = [OpSizeInBits](ConstantSDNode *LHS,
11766 ConstantSDNode *RHS) {
11767 APInt c1 = LHS->getAPIntValue();
11768 APInt c2 = RHS->getAPIntValue();
11769 zeroExtendToMatch(LHS&: c1, RHS&: c2, Offset: 1 /* Overflow Bit */);
11770 return (c1 + c2).uge(RHS: OpSizeInBits);
11771 };
11772 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: MatchOutOfRange))
11773 return DAG.getConstant(Val: 0, DL, VT);
11774
11775 auto MatchInRange = [OpSizeInBits](ConstantSDNode *LHS,
11776 ConstantSDNode *RHS) {
11777 APInt c1 = LHS->getAPIntValue();
11778 APInt c2 = RHS->getAPIntValue();
11779 zeroExtendToMatch(LHS&: c1, RHS&: c2, Offset: 1 /* Overflow Bit */);
11780 return (c1 + c2).ult(RHS: OpSizeInBits);
11781 };
11782 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: MatchInRange)) {
11783 SDValue Sum = DAG.getNode(Opcode: ISD::ADD, DL, VT: ShiftVT, N1, N2: N0.getOperand(i: 1));
11784 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0.getOperand(i: 0), N2: Sum);
11785 }
11786 }
11787
11788 if (N1C && N0.getOpcode() == ISD::TRUNCATE &&
11789 N0.getOperand(i: 0).getOpcode() == ISD::SRL) {
11790 SDValue InnerShift = N0.getOperand(i: 0);
11791 // TODO - support non-uniform vector shift amounts.
11792 if (auto *N001C = isConstOrConstSplat(N: InnerShift.getOperand(i: 1))) {
11793 uint64_t c1 = N001C->getZExtValue();
11794 uint64_t c2 = N1C->getZExtValue();
11795 EVT InnerShiftVT = InnerShift.getValueType();
11796 EVT ShiftAmtVT = InnerShift.getOperand(i: 1).getValueType();
11797 uint64_t InnerShiftSize = InnerShiftVT.getScalarSizeInBits();
11798 // srl (trunc (srl x, c1)), c2 --> 0 or (trunc (srl x, (add c1, c2)))
11799 // This is only valid if the OpSizeInBits + c1 = size of inner shift.
11800 if (c1 + OpSizeInBits == InnerShiftSize) {
11801 if (c1 + c2 >= InnerShiftSize)
11802 return DAG.getConstant(Val: 0, DL, VT);
11803 SDValue NewShiftAmt = DAG.getConstant(Val: c1 + c2, DL, VT: ShiftAmtVT);
11804 SDValue NewShift = DAG.getNode(Opcode: ISD::SRL, DL, VT: InnerShiftVT,
11805 N1: InnerShift.getOperand(i: 0), N2: NewShiftAmt);
11806 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: NewShift);
11807 }
11808 // In the more general case, we can clear the high bits after the shift:
11809 // srl (trunc (srl x, c1)), c2 --> trunc (and (srl x, (c1+c2)), Mask)
11810 if (N0.hasOneUse() && InnerShift.hasOneUse() &&
11811 c1 + c2 < InnerShiftSize) {
11812 SDValue NewShiftAmt = DAG.getConstant(Val: c1 + c2, DL, VT: ShiftAmtVT);
11813 SDValue NewShift = DAG.getNode(Opcode: ISD::SRL, DL, VT: InnerShiftVT,
11814 N1: InnerShift.getOperand(i: 0), N2: NewShiftAmt);
11815 SDValue Mask = DAG.getConstant(Val: APInt::getLowBitsSet(numBits: InnerShiftSize,
11816 loBitsSet: OpSizeInBits - c2),
11817 DL, VT: InnerShiftVT);
11818 SDValue And = DAG.getNode(Opcode: ISD::AND, DL, VT: InnerShiftVT, N1: NewShift, N2: Mask);
11819 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: And);
11820 }
11821 }
11822 }
11823
11824 if (N0.getOpcode() == ISD::SHL) {
11825 // fold (srl (shl nuw x, c), c) -> x
11826 if (N0.getOperand(i: 1) == N1 && N0->getFlags().hasNoUnsignedWrap())
11827 return N0.getOperand(i: 0);
11828
11829 // fold (srl (shl x, c1), c2) -> (and (shl x, (sub c1, c2), MASK) or
11830 // (and (srl x, (sub c2, c1), MASK)
11831 if ((N0.getOperand(i: 1) == N1 || N0->hasOneUse()) &&
11832 TLI.shouldFoldConstantShiftPairToMask(N)) {
11833 auto MatchShiftAmount = [OpSizeInBits](ConstantSDNode *LHS,
11834 ConstantSDNode *RHS) {
11835 const APInt &LHSC = LHS->getAPIntValue();
11836 const APInt &RHSC = RHS->getAPIntValue();
11837 return LHSC.ult(RHS: OpSizeInBits) && RHSC.ult(RHS: OpSizeInBits) &&
11838 LHSC.getZExtValue() <= RHSC.getZExtValue();
11839 };
11840 if (ISD::matchBinaryPredicate(LHS: N1, RHS: N0.getOperand(i: 1), Match: MatchShiftAmount,
11841 /*AllowUndefs*/ false,
11842 /*AllowTypeMismatch*/ true)) {
11843 SDValue N01 = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 1), DL, VT: ShiftVT);
11844 SDValue Diff = DAG.getNode(Opcode: ISD::SUB, DL, VT: ShiftVT, N1: N01, N2: N1);
11845 SDValue Mask = DAG.getAllOnesConstant(DL, VT);
11846 Mask = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: Mask, N2: N01);
11847 Mask = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Mask, N2: Diff);
11848 SDValue Shift = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0.getOperand(i: 0), N2: Diff);
11849 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Shift, N2: Mask);
11850 }
11851 if (ISD::matchBinaryPredicate(LHS: N0.getOperand(i: 1), RHS: N1, Match: MatchShiftAmount,
11852 /*AllowUndefs*/ false,
11853 /*AllowTypeMismatch*/ true)) {
11854 SDValue N01 = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 1), DL, VT: ShiftVT);
11855 SDValue Diff = DAG.getNode(Opcode: ISD::SUB, DL, VT: ShiftVT, N1, N2: N01);
11856 SDValue Mask = DAG.getAllOnesConstant(DL, VT);
11857 Mask = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: Mask, N2: N1);
11858 SDValue Shift = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0.getOperand(i: 0), N2: Diff);
11859 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Shift, N2: Mask);
11860 }
11861 }
11862 }
11863
11864 // fold (srl (anyextend x), c) -> (and (anyextend (srl x, c)), mask)
11865 // TODO - support non-uniform vector shift amounts.
11866 if (N1C && N0.getOpcode() == ISD::ANY_EXTEND) {
11867 // Shifting in all undef bits?
11868 EVT SmallVT = N0.getOperand(i: 0).getValueType();
11869 unsigned BitSize = SmallVT.getScalarSizeInBits();
11870 if (N1C->getAPIntValue().uge(RHS: BitSize))
11871 return DAG.getUNDEF(VT);
11872
11873 if (!LegalTypes || TLI.isTypeDesirableForOp(ISD::SRL, VT: SmallVT)) {
11874 uint64_t ShiftAmt = N1C->getZExtValue();
11875 SDLoc DL0(N0);
11876 SDValue SmallShift =
11877 DAG.getNode(Opcode: ISD::SRL, DL: DL0, VT: SmallVT, N1: N0.getOperand(i: 0),
11878 N2: DAG.getShiftAmountConstant(Val: ShiftAmt, VT: SmallVT, DL: DL0));
11879 AddToWorklist(N: SmallShift.getNode());
11880 APInt Mask = APInt::getLowBitsSet(numBits: OpSizeInBits, loBitsSet: OpSizeInBits - ShiftAmt);
11881 return DAG.getNode(Opcode: ISD::AND, DL, VT,
11882 N1: DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT, Operand: SmallShift),
11883 N2: DAG.getConstant(Val: Mask, DL, VT));
11884 }
11885 }
11886
11887 // fold (srl (sra X, Y), 31) -> (srl X, 31). This srl only looks at the sign
11888 // bit, which is unmodified by sra.
11889 if (N1C && N1C->getAPIntValue() == (OpSizeInBits - 1)) {
11890 if (N0.getOpcode() == ISD::SRA)
11891 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0.getOperand(i: 0), N2: N1);
11892 }
11893
11894 // fold (srl (ctlz x), "5") -> x iff x has one bit set (the low bit), and x has a power
11895 // of two bitwidth. The "5" represents (log2 (bitwidth x)).
11896 if (N1C && N0.getOpcode() == ISD::CTLZ &&
11897 isPowerOf2_32(Value: OpSizeInBits) &&
11898 N1C->getAPIntValue() == Log2_32(Value: OpSizeInBits)) {
11899 KnownBits Known = DAG.computeKnownBits(Op: N0.getOperand(i: 0));
11900
11901 // If any of the input bits are KnownOne, then the input couldn't be all
11902 // zeros, thus the result of the srl will always be zero.
11903 if (Known.One.getBoolValue()) return DAG.getConstant(Val: 0, DL: SDLoc(N0), VT);
11904
11905 // If all of the bits input the to ctlz node are known to be zero, then
11906 // the result of the ctlz is "32" and the result of the shift is one.
11907 APInt UnknownBits = ~Known.Zero;
11908 if (UnknownBits == 0) return DAG.getConstant(Val: 1, DL: SDLoc(N0), VT);
11909
11910 // Otherwise, check to see if there is exactly one bit input to the ctlz.
11911 if (UnknownBits.isPowerOf2()) {
11912 // Okay, we know that only that the single bit specified by UnknownBits
11913 // could be set on input to the CTLZ node. If this bit is set, the SRL
11914 // will return 0, if it is clear, it returns 1. Change the CTLZ/SRL pair
11915 // to an SRL/XOR pair, which is likely to simplify more.
11916 unsigned ShAmt = UnknownBits.countr_zero();
11917 SDValue Op = N0.getOperand(i: 0);
11918
11919 if (ShAmt) {
11920 SDLoc DL(N0);
11921 Op = DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: Op,
11922 N2: DAG.getShiftAmountConstant(Val: ShAmt, VT, DL));
11923 AddToWorklist(N: Op.getNode());
11924 }
11925 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: Op, N2: DAG.getConstant(Val: 1, DL, VT));
11926 }
11927 }
11928
11929 // fold (srl x, (trunc (and y, c))) -> (srl x, (and (trunc y), (trunc c))).
11930 if (N1.getOpcode() == ISD::TRUNCATE &&
11931 N1.getOperand(i: 0).getOpcode() == ISD::AND) {
11932 if (SDValue NewOp1 = distributeTruncateThroughAnd(N: N1.getNode()))
11933 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0, N2: NewOp1);
11934 }
11935
11936 // fold (srl (logic_op x, (shl (zext y), c1)), c1)
11937 // -> (logic_op (srl x, c1), (zext y))
11938 // c1 <= leadingzeros(zext(y))
11939 // TODO: Replace c1 with valuetracking?
11940 SDValue X, ZExtY;
11941 if (sd_match(
11942 N: N0,
11943 P: m_OneUse(P: m_BitwiseLogic(
11944 L: m_Value(N&: X),
11945 R: m_OneUse(P: m_Shl(L: m_Value(N&: ZExtY, P: m_SpecificOpc<ISD::ZERO_EXTEND>()),
11946 R: m_Specific(N: N1))))))) {
11947 unsigned NumLeadingZeros = ZExtY.getScalarValueSizeInBits() -
11948 ZExtY.getOperand(i: 0).getScalarValueSizeInBits();
11949 if (N1C && N1C->getZExtValue() <= NumLeadingZeros)
11950 return DAG.getNode(Opcode: N0.getOpcode(), DL: SDLoc(N0), VT,
11951 N1: DAG.getNode(Opcode: ISD::SRL, DL: SDLoc(N0), VT, N1: X, N2: N1), N2: ZExtY);
11952 }
11953
11954 // fold (srl (bitcast (build_vector e1, ..., eN)), (N-1) * eltsize)
11955 // -> (zext eN)
11956 if (N1C && VT.isScalarInteger() && DAG.getDataLayout().isLittleEndian()) {
11957 SDValue BV = peekThroughBitcasts(V: N0);
11958 if (BV.getOpcode() == ISD::BUILD_VECTOR) {
11959 EVT BVVT = BV.getValueType();
11960 unsigned EltSizeInBits = BVVT.getScalarSizeInBits();
11961 unsigned NumElts = BVVT.getVectorNumElements();
11962 if (N1C->getZExtValue() == (NumElts - 1) * EltSizeInBits) {
11963 SDValue LastElt = BV.getOperand(i: NumElts - 1);
11964 assert(LastElt.getScalarValueSizeInBits() >= EltSizeInBits &&
11965 "Expected BUILD_VECTOR operand as wide as element type");
11966 EVT IntEltVT = LastElt.getValueType().changeTypeToInteger();
11967 if (!LegalTypes || TLI.isTypeLegal(VT: IntEltVT)) {
11968 LastElt = DAG.getBitcast(VT: IntEltVT, V: LastElt);
11969 SDValue Ext = DAG.getZExtOrTrunc(Op: LastElt, DL, VT);
11970 APInt Mask = APInt::getLowBitsSet(numBits: VT.getSizeInBits(), loBitsSet: EltSizeInBits);
11971 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Ext,
11972 N2: DAG.getConstant(Val: Mask, DL, VT));
11973 }
11974 }
11975 }
11976 }
11977
11978 // fold (srl (add nuw X, C), D) -> (add nuw (srl X, D), C u>> D)
11979 // when C has D trailing zeros (so C >> D is exact).
11980 if (N1C && N0.hasOneUse() && N0.getOpcode() == ISD::ADD &&
11981 N0->getFlags().hasNoUnsignedWrap()) {
11982 if (ConstantSDNode *AddC = isConstOrConstSplat(N: N0.getOperand(i: 1))) {
11983 const APInt &ShAmt = N1C->getAPIntValue();
11984 const APInt &AddVal = AddC->getAPIntValue();
11985 if (ShAmt.ult(RHS: AddVal.countr_zero())) {
11986 SDNodeFlags ShiftFlags = N->getFlags();
11987 SDValue NewSrl =
11988 DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: N0.getOperand(i: 0), N2: N1, Flags: ShiftFlags);
11989 SDValue NewC = DAG.getConstant(Val: AddVal.lshr(ShiftAmt: ShAmt), DL, VT);
11990 SDNodeFlags AddFlags = N0->getFlags();
11991 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: NewSrl, N2: NewC, Flags: AddFlags);
11992 }
11993 }
11994 }
11995
11996 // fold operands of srl based on knowledge that the low bits are not
11997 // demanded.
11998 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
11999 return SDValue(N, 0);
12000
12001 if (N1C && !N1C->isOpaque())
12002 if (SDValue NewSRL = visitShiftByConstant(N))
12003 return NewSRL;
12004
12005 // Attempt to convert a srl of a load into a narrower zero-extending load.
12006 if (SDValue NarrowLoad = reduceLoadWidth(N))
12007 return NarrowLoad;
12008
12009 // Here is a common situation. We want to optimize:
12010 //
12011 // %a = ...
12012 // %b = and i32 %a, 2
12013 // %c = srl i32 %b, 1
12014 // brcond i32 %c ...
12015 //
12016 // into
12017 //
12018 // %a = ...
12019 // %b = and %a, 2
12020 // %c = setcc eq %b, 0
12021 // brcond %c ...
12022 //
12023 // However when after the source operand of SRL is optimized into AND, the SRL
12024 // itself may not be optimized further. Look for it and add the BRCOND into
12025 // the worklist.
12026 //
12027 // The also tends to happen for binary operations when SimplifyDemandedBits
12028 // is involved.
12029 //
12030 // FIXME: This is unecessary if we process the DAG in topological order,
12031 // which we plan to do. This workaround can be removed once the DAG is
12032 // processed in topological order.
12033 if (N->hasOneUse()) {
12034 SDNode *User = *N->user_begin();
12035
12036 // Look pass the truncate.
12037 if (User->getOpcode() == ISD::TRUNCATE && User->hasOneUse())
12038 User = *User->user_begin();
12039
12040 if (User->getOpcode() == ISD::BRCOND || User->getOpcode() == ISD::AND ||
12041 User->getOpcode() == ISD::OR || User->getOpcode() == ISD::XOR)
12042 AddToWorklist(N: User);
12043 }
12044
12045 // Try to transform this shift into a multiply-high if
12046 // it matches the appropriate pattern detected in combineShiftToMULH.
12047 if (SDValue MULH = combineShiftToMULH(N, DL, DAG, TLI))
12048 return MULH;
12049
12050 if (SDValue AVG = foldShiftToAvg(N, DL))
12051 return AVG;
12052
12053 SDValue Y;
12054 if (VT.getScalarSizeInBits() % 2 == 0 && N1C) {
12055 // Fold clmul(zext(x), zext(y)) >> (BW - 1 | BW) -> clmul(r|h)(x, y).
12056 unsigned HalfBW = VT.getScalarSizeInBits() / 2;
12057 if (sd_match(N: N0, P: m_Clmul(L: m_ZExt(Op: m_Value(N&: X)), R: m_ZExt(Op: m_Value(N&: Y)))) &&
12058 X.getScalarValueSizeInBits() == HalfBW &&
12059 Y.getScalarValueSizeInBits() == HalfBW) {
12060 if (N1C->getZExtValue() == HalfBW - 1 &&
12061 (!LegalOperations ||
12062 TLI.isOperationLegalOrCustom(Op: ISD::CLMULR, VT: X.getValueType())))
12063 return DAG.getNode(
12064 Opcode: ISD::ZERO_EXTEND, DL, VT,
12065 Operand: DAG.getNode(Opcode: ISD::CLMULR, DL, VT: X.getValueType(), N1: X, N2: Y));
12066 if (N1C->getZExtValue() == HalfBW &&
12067 (!LegalOperations ||
12068 TLI.isOperationLegalOrCustom(Op: ISD::CLMULH, VT: X.getValueType())))
12069 return DAG.getNode(
12070 Opcode: ISD::ZERO_EXTEND, DL, VT,
12071 Operand: DAG.getNode(Opcode: ISD::CLMULH, DL, VT: X.getValueType(), N1: X, N2: Y));
12072 }
12073 }
12074
12075 // Fold bitreverse(clmul(bitreverse(x), bitreverse(y))) >> 1 ->
12076 // clmulh(x, y).
12077 if (N1C && N1C->getZExtValue() == 1 &&
12078 sd_match(N: N0, P: m_BitReverse(Op: m_Clmul(L: m_BitReverse(Op: m_Value(N&: X)),
12079 R: m_BitReverse(Op: m_Value(N&: Y))))))
12080 return DAG.getNode(Opcode: ISD::CLMULH, DL, VT, N1: X, N2: Y);
12081
12082 return SDValue();
12083}
12084
12085SDValue DAGCombiner::visitFunnelShift(SDNode *N) {
12086 EVT VT = N->getValueType(ResNo: 0);
12087 SDValue N0 = N->getOperand(Num: 0);
12088 SDValue N1 = N->getOperand(Num: 1);
12089 SDValue N2 = N->getOperand(Num: 2);
12090 bool IsFSHL = N->getOpcode() == ISD::FSHL;
12091 unsigned BitWidth = VT.getScalarSizeInBits();
12092 SDLoc DL(N);
12093
12094 // fold (fshl/fshr C0, C1, C2) -> C3
12095 if (SDValue C =
12096 DAG.FoldConstantArithmetic(Opcode: N->getOpcode(), DL, VT, Ops: {N0, N1, N2}))
12097 return C;
12098
12099 // fold (fshl N0, N1, 0) -> N0
12100 // fold (fshr N0, N1, 0) -> N1
12101 if (isPowerOf2_32(Value: BitWidth))
12102 if (DAG.MaskedValueIsZero(
12103 Op: N2, Mask: APInt(N2.getScalarValueSizeInBits(), BitWidth - 1)))
12104 return IsFSHL ? N0 : N1;
12105
12106 auto IsUndefOrZero = [](SDValue V) {
12107 return V.isUndef() || isNullOrNullSplat(V, /*AllowUndefs*/ true);
12108 };
12109
12110 // TODO - support non-uniform vector shift amounts.
12111 if (ConstantSDNode *Cst = isConstOrConstSplat(N: N2)) {
12112 EVT ShAmtTy = N2.getValueType();
12113
12114 // fold (fsh* N0, N1, c) -> (fsh* N0, N1, c % BitWidth)
12115 if (Cst->getAPIntValue().uge(RHS: BitWidth)) {
12116 uint64_t RotAmt = Cst->getAPIntValue().urem(RHS: BitWidth);
12117 return DAG.getNode(Opcode: N->getOpcode(), DL, VT, N1: N0, N2: N1,
12118 N3: DAG.getConstant(Val: RotAmt, DL, VT: ShAmtTy));
12119 }
12120
12121 unsigned ShAmt = Cst->getZExtValue();
12122 if (ShAmt == 0)
12123 return IsFSHL ? N0 : N1;
12124
12125 // fold fshl(undef_or_zero, N1, C) -> lshr(N1, BW-C)
12126 // fold fshr(undef_or_zero, N1, C) -> lshr(N1, C)
12127 // fold fshl(N0, undef_or_zero, C) -> shl(N0, C)
12128 // fold fshr(N0, undef_or_zero, C) -> shl(N0, BW-C)
12129 if (IsUndefOrZero(N0))
12130 return DAG.getNode(
12131 Opcode: ISD::SRL, DL, VT, N1,
12132 N2: DAG.getConstant(Val: IsFSHL ? BitWidth - ShAmt : ShAmt, DL, VT: ShAmtTy));
12133 if (IsUndefOrZero(N1))
12134 return DAG.getNode(
12135 Opcode: ISD::SHL, DL, VT, N1: N0,
12136 N2: DAG.getConstant(Val: IsFSHL ? ShAmt : BitWidth - ShAmt, DL, VT: ShAmtTy));
12137
12138 // fold fshl(N0, N1, c) -> x and fshr(N0, N1, c) -> x
12139 // where N0 is any node that contributes "x >> C0" to the result:
12140 // lshr(x, C0) | fshr(_, x, C0) | fshl(_, x, C1)
12141 // and N1 is any node that contributes "x << C1" to the result:
12142 // shl(x, C1) | fshl(x, _, C1) | fshr(x, _, C0)
12143 // with C0 = IsFSHL ? amnt : BW-amnt, C1 = BW - C0
12144
12145 // ShAmt == 0 was handled above; uge(BitWidth) was reduced via modulo above.
12146 assert(ShAmt >= 1 && ShAmt < BitWidth &&
12147 "ShAmt must be in [1, BW-1] for the identity fold to be valid");
12148 SDValue Val;
12149 unsigned C0Expected = IsFSHL ? ShAmt : BitWidth - ShAmt;
12150 unsigned C1Expected = IsFSHL ? BitWidth - ShAmt : ShAmt;
12151
12152 if ((sd_match(N: N0, P: m_Srl(L: m_Value(N&: Val), R: m_SpecificInt(V: C0Expected))) ||
12153 sd_match(N: N0,
12154 P: m_FShR(Op0: m_Value(), Op1: m_Value(N&: Val), Op2: m_SpecificInt(V: C0Expected))) ||
12155 sd_match(
12156 N: N0, P: m_FShL(Op0: m_Value(), Op1: m_Value(N&: Val), Op2: m_SpecificInt(V: C1Expected)))) &&
12157 (sd_match(N: N1, P: m_Shl(L: m_Specific(N: Val), R: m_SpecificInt(V: C1Expected))) ||
12158 sd_match(N: N1, P: m_FShL(Op0: m_Specific(N: Val), Op1: m_Value(),
12159 Op2: m_SpecificInt(V: C1Expected))) ||
12160 sd_match(N: N1, P: m_FShR(Op0: m_Specific(N: Val), Op1: m_Value(),
12161 Op2: m_SpecificInt(V: C0Expected)))))
12162 return Val;
12163
12164 // fold (fshl ld1, ld0, c) -> (ld0[ofs]) iff ld0 and ld1 are consecutive.
12165 // fold (fshr ld1, ld0, c) -> (ld0[ofs]) iff ld0 and ld1 are consecutive.
12166 // TODO - bigendian support once we have test coverage.
12167 // TODO - can we merge this with CombineConseutiveLoads/MatchLoadCombine?
12168 // TODO - permit LHS EXTLOAD if extensions are shifted out.
12169 if ((BitWidth % 8) == 0 && (ShAmt % 8) == 0 && !VT.isVector() &&
12170 !DAG.getDataLayout().isBigEndian()) {
12171 auto *LHS = dyn_cast<LoadSDNode>(Val&: N0);
12172 auto *RHS = dyn_cast<LoadSDNode>(Val&: N1);
12173 if (LHS && RHS && LHS->isSimple() && RHS->isSimple() &&
12174 LHS->getAddressSpace() == RHS->getAddressSpace() &&
12175 (LHS->hasNUsesOfValue(NUses: 1, Value: 0) || RHS->hasNUsesOfValue(NUses: 1, Value: 0)) &&
12176 ISD::isNON_EXTLoad(N: RHS) && ISD::isNON_EXTLoad(N: LHS)) {
12177 if (DAG.areNonVolatileConsecutiveLoads(LD: LHS, Base: RHS, Bytes: BitWidth / 8, Dist: 1)) {
12178 SDLoc DL(RHS);
12179 uint64_t PtrOff =
12180 IsFSHL ? (((BitWidth - ShAmt) % BitWidth) / 8) : (ShAmt / 8);
12181 Align NewAlign = commonAlignment(A: RHS->getAlign(), Offset: PtrOff);
12182 unsigned Fast = 0;
12183 if (TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT,
12184 AddrSpace: RHS->getAddressSpace(), Alignment: NewAlign,
12185 Flags: RHS->getMemOperand()->getFlags(), Fast: &Fast) &&
12186 Fast) {
12187 SDValue NewPtr = DAG.getMemBasePlusOffset(
12188 Base: RHS->getBasePtr(), Offset: TypeSize::getFixed(ExactSize: PtrOff), DL);
12189 AddToWorklist(N: NewPtr.getNode());
12190 SDValue Load = DAG.getLoad(
12191 VT, dl: DL, Chain: RHS->getChain(), Ptr: NewPtr,
12192 PtrInfo: RHS->getPointerInfo().getWithOffset(O: PtrOff), Alignment: NewAlign,
12193 MMOFlags: RHS->getMemOperand()->getFlags(), Metadata: RHS->getAAInfo());
12194 DAG.makeEquivalentMemoryOrdering(OldLoad: LHS, NewMemOp: Load.getValue(R: 1));
12195 DAG.makeEquivalentMemoryOrdering(OldLoad: RHS, NewMemOp: Load.getValue(R: 1));
12196 return Load;
12197 }
12198 }
12199 }
12200 }
12201 }
12202
12203 // fold fshr(undef_or_zero, N1, N2) -> lshr(N1, N2)
12204 // fold fshl(N0, undef_or_zero, N2) -> shl(N0, N2)
12205 // iff We know the shift amount is in range.
12206 // TODO: when is it worth doing SUB(BW, N2) as well?
12207 if (isPowerOf2_32(Value: BitWidth)) {
12208 APInt ModuloBits(N2.getScalarValueSizeInBits(), BitWidth - 1);
12209 if (IsUndefOrZero(N0) && !IsFSHL && DAG.MaskedValueIsZero(Op: N2, Mask: ~ModuloBits))
12210 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1, N2);
12211 if (IsUndefOrZero(N1) && IsFSHL && DAG.MaskedValueIsZero(Op: N2, Mask: ~ModuloBits))
12212 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0, N2);
12213 }
12214
12215 // fold (fshl N0, N0, N2) -> (rotl N0, N2)
12216 // fold (fshr N0, N0, N2) -> (rotr N0, N2)
12217 // TODO: Investigate flipping this rotate if only one is legal.
12218 // If funnel shift is legal as well we might be better off avoiding
12219 // non-constant (BW - N2).
12220 unsigned RotOpc = IsFSHL ? ISD::ROTL : ISD::ROTR;
12221 if (N0 == N1 && hasOperation(Opcode: RotOpc, VT))
12222 return DAG.getNode(Opcode: RotOpc, DL, VT, N1: N0, N2);
12223
12224 // Simplify, based on bits shifted out of N0/N1.
12225 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
12226 return SDValue(N, 0);
12227
12228 return SDValue();
12229}
12230
12231SDValue DAGCombiner::visitSHLSAT(SDNode *N) {
12232 SDValue N0 = N->getOperand(Num: 0);
12233 SDValue N1 = N->getOperand(Num: 1);
12234 if (SDValue V = DAG.simplifyShift(X: N0, Y: N1))
12235 return V;
12236
12237 SDLoc DL(N);
12238 EVT VT = N0.getValueType();
12239
12240 // fold (*shlsat c1, c2) -> c1<<c2
12241 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: N->getOpcode(), DL, VT, Ops: {N0, N1}))
12242 return C;
12243
12244 ConstantSDNode *N1C = isConstOrConstSplat(N: N1);
12245
12246 if (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::SHL, VT)) {
12247 // fold (sshlsat x, c) -> (shl x, c)
12248 if (N->getOpcode() == ISD::SSHLSAT && N1C &&
12249 N1C->getAPIntValue().ult(RHS: DAG.ComputeNumSignBits(Op: N0)))
12250 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0, N2: N1);
12251
12252 // fold (ushlsat x, c) -> (shl x, c)
12253 if (N->getOpcode() == ISD::USHLSAT && N1C &&
12254 N1C->getAPIntValue().ule(
12255 RHS: DAG.computeKnownBits(Op: N0).countMinLeadingZeros()))
12256 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0, N2: N1);
12257 }
12258
12259 return SDValue();
12260}
12261
12262// Given a ABS node, detect the following patterns:
12263// (ABS (SUB (EXTEND a), (EXTEND b))).
12264// (TRUNC (ABS (SUB (EXTEND a), (EXTEND b)))).
12265// Generates UABD/SABD instruction.
12266SDValue DAGCombiner::foldABSToABD(SDNode *N, const SDLoc &DL) {
12267 EVT SrcVT = N->getValueType(ResNo: 0);
12268
12269 if (N->getOpcode() == ISD::TRUNCATE)
12270 N = N->getOperand(Num: 0).getNode();
12271
12272 EVT VT = N->getValueType(ResNo: 0);
12273 SDValue Op0, Op1;
12274
12275 if (!sd_match(N, P: m_Abs(Op: m_AnyOf(preds: m_Sub(L: m_Value(N&: Op0), R: m_Value(N&: Op1)),
12276 preds: m_Add(L: m_Value(N&: Op0), R: m_Value(N&: Op1))))))
12277 return SDValue();
12278
12279 SDValue AbsOp0 = N->getOperand(Num: 0);
12280 bool IsAdd = AbsOp0.getOpcode() == ISD::ADD;
12281 // Make sure (abs B) is positive.
12282 if (IsAdd) {
12283 // Elements of Op1 must be constant and != VT.minSignedValue() (or undef)
12284 auto IsNotMinSignedInt = [VT](ConstantSDNode *C) {
12285 if (C == nullptr)
12286 return true;
12287 return !C->getAPIntValue()
12288 .trunc(width: VT.getScalarSizeInBits())
12289 .isMinSignedValue();
12290 };
12291
12292 if (!ISD::matchUnaryPredicate(Op: Op1, Match: IsNotMinSignedInt, /*AllowUndefs=*/true,
12293 /*AllowTruncation=*/true))
12294 return SDValue();
12295 }
12296
12297 unsigned Opc0 = Op0.getOpcode();
12298
12299 // Check if the operands of the sub are (zero|sign)-extended, otherwise
12300 // fallback to ValueTracking.
12301 if (Opc0 != Op1.getOpcode() ||
12302 (Opc0 != ISD::ZERO_EXTEND && Opc0 != ISD::SIGN_EXTEND &&
12303 Opc0 != ISD::SIGN_EXTEND_INREG)) {
12304
12305 auto CreateZextedAbd = [&](unsigned AbdOpc) {
12306 if (IsAdd)
12307 Op1 = DAG.getNegative(Val: Op1, DL: SDLoc(Op1), VT);
12308 SDValue ABD = DAG.getNode(Opcode: AbdOpc, DL, VT, N1: Op0, N2: Op1);
12309 return DAG.getZExtOrTrunc(Op: ABD, DL, VT: SrcVT);
12310 };
12311
12312 // fold (abs (sub nsw x, y)) -> abds(x, y)
12313 // fold (abs (add nsw x, -y)) -> abds(x, y)
12314 bool AbsOpWillNSW =
12315 AbsOp0->getFlags().hasNoSignedWrap() ||
12316 (IsAdd ? DAG.willNotOverflowAdd(/*IsSigned=*/true, N0: Op0, N1: Op1)
12317 : DAG.willNotOverflowSub(/*IsSigned=*/true, N0: Op0, N1: Op1));
12318
12319 // Don't fold this for unsupported types as we lose the NSW handling.
12320 if (hasOperation(Opcode: ISD::ABDS, VT) && TLI.preferABDSToABSWithNSW(VT) &&
12321 AbsOpWillNSW)
12322 return CreateZextedAbd(ISD::ABDS);
12323
12324 // fold (abs (sub x, y)) -> abdu(x, y)
12325 bool AbsOpWillNUW =
12326 !IsAdd && DAG.SignBitIsZero(Op: Op0) && DAG.SignBitIsZero(Op: Op1);
12327
12328 if (hasOperation(Opcode: ISD::ABDU, VT) && AbsOpWillNUW)
12329 return CreateZextedAbd(ISD::ABDU);
12330
12331 return SDValue();
12332 }
12333
12334 // The IsAdd case explicitly checks for const/bv-of-const. This implies either
12335 // (Opc0 != Op1.getOpcode() || Opc0 is not in {zext/sext/sign_ext_inreg}. This
12336 // implies it was alrady handled by the above if statement.
12337 assert(!IsAdd && "Unexpected abs(add(x,y)) pattern");
12338
12339 EVT VT0, VT1;
12340 if (Opc0 == ISD::SIGN_EXTEND_INREG) {
12341 VT0 = cast<VTSDNode>(Val: Op0.getOperand(i: 1))->getVT();
12342 VT1 = cast<VTSDNode>(Val: Op1.getOperand(i: 1))->getVT();
12343 } else {
12344 VT0 = Op0.getOperand(i: 0).getValueType();
12345 VT1 = Op1.getOperand(i: 0).getValueType();
12346 }
12347 unsigned ABDOpcode = (Opc0 == ISD::ZERO_EXTEND) ? ISD::ABDU : ISD::ABDS;
12348
12349 // fold abs(sext(x) - sext(y)) -> zext(abds(x, y))
12350 // fold abs(zext(x) - zext(y)) -> zext(abdu(x, y))
12351 EVT MaxVT = VT0.bitsGT(VT: VT1) ? VT0 : VT1;
12352 if ((VT0 == MaxVT || Op0->hasOneUse()) &&
12353 (VT1 == MaxVT || Op1->hasOneUse()) &&
12354 (!LegalTypes || hasOperation(Opcode: ABDOpcode, VT: MaxVT))) {
12355 SDValue ABD = DAG.getNode(Opcode: ABDOpcode, DL, VT: MaxVT,
12356 N1: DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: MaxVT, Operand: Op0),
12357 N2: DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: MaxVT, Operand: Op1));
12358 ABD = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: ABD);
12359 return DAG.getZExtOrTrunc(Op: ABD, DL, VT: SrcVT);
12360 }
12361
12362 // fold abs(sext(x) - sext(y)) -> abds(sext(x), sext(y))
12363 // fold abs(zext(x) - zext(y)) -> abdu(zext(x), zext(y))
12364 if (!LegalOperations || hasOperation(Opcode: ABDOpcode, VT)) {
12365 SDValue ABD = DAG.getNode(Opcode: ABDOpcode, DL, VT, N1: Op0, N2: Op1);
12366 return DAG.getZExtOrTrunc(Op: ABD, DL, VT: SrcVT);
12367 }
12368
12369 return SDValue();
12370}
12371
12372SDValue DAGCombiner::visitABS(SDNode *N) {
12373 SDValue N0 = N->getOperand(Num: 0);
12374 EVT VT = N->getValueType(ResNo: 0);
12375 SDLoc DL(N);
12376
12377 // fold (abs c1) -> c2
12378 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::ABS, DL, VT, Ops: {N0}))
12379 return C;
12380 // fold (abs (abs x)) -> (abs x)
12381 // fold (abs (abs_min_poison x)) -> (abs_min_poison x)
12382 if (ISD::isAbsOpcode(Opcode: N0.getOpcode()))
12383 return N0;
12384 // fold (abs x) -> x iff not-negative
12385 if (DAG.SignBitIsZero(Op: N0))
12386 return N0;
12387
12388 if (SDValue ABD = foldABSToABD(N, DL))
12389 return ABD;
12390
12391 // fold (abs (sign_extend_inreg x)) -> (zero_extend (abs (truncate x)))
12392 // iff zero_extend/truncate are free.
12393 if (N0.getOpcode() == ISD::SIGN_EXTEND_INREG) {
12394 EVT ExtVT = cast<VTSDNode>(Val: N0.getOperand(i: 1))->getVT();
12395 if (TLI.isTruncateFree(FromVT: VT, ToVT: ExtVT) && TLI.isZExtFree(FromTy: ExtVT, ToTy: VT) &&
12396 TLI.isTypeDesirableForOp(ISD::ABS, VT: ExtVT) &&
12397 hasOperation(Opcode: ISD::ABS, VT: ExtVT)) {
12398 return DAG.getNode(
12399 Opcode: ISD::ZERO_EXTEND, DL, VT,
12400 Operand: DAG.getNode(Opcode: ISD::ABS, DL, VT: ExtVT,
12401 Operand: DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: ExtVT, Operand: N0.getOperand(i: 0))));
12402 }
12403 }
12404
12405 return SDValue();
12406}
12407
12408SDValue DAGCombiner::visitABS_MIN_POISON(SDNode *N) {
12409 SDValue N0 = N->getOperand(Num: 0);
12410 EVT VT = N->getValueType(ResNo: 0);
12411 SDLoc DL(N);
12412
12413 // fold (abs_min_poison c1) -> c2 (or poison if c1 == INT_MIN)
12414 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::ABS_MIN_POISON, DL, VT, Ops: {N0}))
12415 return C;
12416 // fold (abs_min_poison (abs_min_poison x)) -> (abs_min_poison x)
12417 // fold (abs_min_poison (abs x)) -> (abs x)
12418 // fold (abs_min_poison (freeze (abs x))) -> (freeze (abs x))
12419 // fold (abs_min_poison (freeze (abs_min_poison x))) ->
12420 // (freeze (abs_min_poison x))
12421 //
12422 // Freeze case is valid because: for x != INT_MIN both sides equal abs(x);
12423 // for x == INT_MIN both forms produce a non-deterministic but well-defined
12424 // value since freeze already consumed the poison.
12425 if (ISD::isAbsOpcode(Opcode: peekThroughFreeze(V: N0).getOpcode()))
12426 return N0;
12427 // fold (abs_min_poison x) -> x iff not-negative
12428 if (DAG.SignBitIsZero(Op: N0))
12429 return N0;
12430
12431 if (SDValue ABD = foldABSToABD(N, DL))
12432 return ABD;
12433
12434 // fold (abs_min_poison (sign_extend_inreg x)) ->
12435 // (zero_extend (abs (truncate x)))
12436 // iff zero_extend/truncate are free. The sign_extend_inreg keeps the value
12437 // in the narrow type's range, so the wide abs_min_poison is never actually
12438 // poison.
12439 if (N0.getOpcode() == ISD::SIGN_EXTEND_INREG) {
12440 EVT ExtVT = cast<VTSDNode>(Val: N0.getOperand(i: 1))->getVT();
12441 if (TLI.isTruncateFree(FromVT: VT, ToVT: ExtVT) && TLI.isZExtFree(FromTy: ExtVT, ToTy: VT) &&
12442 TLI.isTypeDesirableForOp(ISD::ABS, VT: ExtVT) &&
12443 hasOperation(Opcode: ISD::ABS, VT: ExtVT)) {
12444 return DAG.getNode(
12445 Opcode: ISD::ZERO_EXTEND, DL, VT,
12446 Operand: DAG.getNode(Opcode: ISD::ABS, DL, VT: ExtVT,
12447 Operand: DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: ExtVT, Operand: N0.getOperand(i: 0))));
12448 }
12449 }
12450
12451 return SDValue();
12452}
12453
12454SDValue DAGCombiner::visitCLMUL(SDNode *N) {
12455 unsigned Opcode = N->getOpcode();
12456 SDValue N0 = N->getOperand(Num: 0);
12457 SDValue N1 = N->getOperand(Num: 1);
12458 EVT VT = N->getValueType(ResNo: 0);
12459 SDLoc DL(N);
12460
12461 // fold (clmul c1, c2)
12462 if (SDValue C = DAG.FoldConstantArithmetic(Opcode, DL, VT, Ops: {N0, N1}))
12463 return C;
12464
12465 // canonicalize constant to RHS
12466 if (DAG.isConstantIntBuildVectorOrConstantInt(N: N0) &&
12467 !DAG.isConstantIntBuildVectorOrConstantInt(N: N1))
12468 return DAG.getNode(Opcode, DL, VT, N1, N2: N0);
12469
12470 // fold (clmul x, 0) -> 0
12471 if (isNullConstant(V: N1) || ISD::isConstantSplatVectorAllZeros(N: N1.getNode()))
12472 return DAG.getConstant(Val: 0, DL, VT);
12473
12474 // fold (clmul x, c_pow2) -> (shl x, log2(c_pow2))
12475 // This also handles (clmul x, 1) -> x since (shl x, 0) simplifies to x.
12476 if (Opcode == ISD::CLMUL) {
12477 if (ConstantSDNode *C = isConstOrConstSplat(N: N1)) {
12478 APInt CV = C->getAPIntValue().trunc(width: VT.getScalarSizeInBits());
12479 if (CV.isPowerOf2() &&
12480 (!LegalOperations || TLI.isOperationLegal(Op: ISD::SHL, VT)))
12481 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: N0,
12482 N2: DAG.getShiftAmountConstant(Val: CV.logBase2(), VT, DL));
12483 }
12484 }
12485
12486 return SDValue();
12487}
12488
12489SDValue DAGCombiner::visitPEXT(SDNode *N) {
12490 EVT VT = N->getValueType(ResNo: 0);
12491 SDValue N0 = N->getOperand(Num: 0);
12492 SDValue N1 = N->getOperand(Num: 1);
12493 SDLoc DL(N);
12494
12495 // pext(x, 0) -> 0
12496 if (isNullOrNullSplat(V: N1))
12497 return DAG.getConstant(Val: 0, DL, VT);
12498 // pext(x, -1) -> x (all bits selected, packed into low positions = x)
12499 if (isAllOnesOrAllOnesSplat(V: N1))
12500 return N0;
12501 // fold pext(c1, c2) -> c3
12502 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::PEXT, DL, VT, Ops: {N0, N1}))
12503 return C;
12504 return SDValue();
12505}
12506
12507SDValue DAGCombiner::visitPDEP(SDNode *N) {
12508 EVT VT = N->getValueType(ResNo: 0);
12509 SDValue N0 = N->getOperand(Num: 0);
12510 SDValue N1 = N->getOperand(Num: 1);
12511 SDLoc DL(N);
12512
12513 // pdep(x, 0) -> 0
12514 if (isNullOrNullSplat(V: N1))
12515 return DAG.getConstant(Val: 0, DL, VT);
12516
12517 // pdep(x, -1) -> x (all positions selected, bits deposited at identity)
12518 if (isAllOnesOrAllOnesSplat(V: N1))
12519 return N0;
12520
12521 // fold pdep(c1, c2) -> c3
12522 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::PDEP, DL, VT, Ops: {N0, N1}))
12523 return C;
12524
12525 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
12526 return SDValue(N, 0);
12527
12528 return SDValue();
12529}
12530
12531SDValue DAGCombiner::visitBSWAP(SDNode *N) {
12532 SDValue N0 = N->getOperand(Num: 0);
12533 EVT VT = N->getValueType(ResNo: 0);
12534 SDLoc DL(N);
12535
12536 // fold (bswap c1) -> c2
12537 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::BSWAP, DL, VT, Ops: {N0}))
12538 return C;
12539 // fold (bswap (bswap x)) -> x
12540 if (N0.getOpcode() == ISD::BSWAP)
12541 return N0.getOperand(i: 0);
12542
12543 // Canonicalize bswap(bitreverse(x)) -> bitreverse(bswap(x)). If bitreverse
12544 // isn't supported, it will be expanded to bswap followed by a manual reversal
12545 // of bits in each byte. By placing bswaps before bitreverse, we can remove
12546 // the two bswaps if the bitreverse gets expanded.
12547 if (N0.getOpcode() == ISD::BITREVERSE && N0.hasOneUse()) {
12548 SDValue BSwap = DAG.getNode(Opcode: ISD::BSWAP, DL, VT, Operand: N0.getOperand(i: 0));
12549 return DAG.getNode(Opcode: ISD::BITREVERSE, DL, VT, Operand: BSwap);
12550 }
12551
12552 unsigned BW = VT.getScalarSizeInBits();
12553 // fold (bswap shl(x,c)) -> (zext(bswap(trunc(shl(x,sub(c,bw/2))))))
12554 // iff x >= bw/2 (i.e. lower half is known zero)
12555 if (BW >= 32 && N0.getOpcode() == ISD::SHL && N0.hasOneUse()) {
12556 auto *ShAmt = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
12557 EVT HalfVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: BW / 2);
12558 if (ShAmt && ShAmt->getAPIntValue().ult(RHS: BW) &&
12559 ShAmt->getZExtValue() >= (BW / 2) && (ShAmt->getZExtValue() % 8) == 0 &&
12560 TLI.isTypeLegal(VT: HalfVT) && TLI.isTruncateFree(FromVT: VT, ToVT: HalfVT) &&
12561 (!LegalOperations || hasOperation(Opcode: ISD::BSWAP, VT: HalfVT))) {
12562 SDValue Res = N0.getOperand(i: 0);
12563 if (uint64_t NewShAmt = (ShAmt->getZExtValue() - (BW / 2)))
12564 Res = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Res,
12565 N2: DAG.getShiftAmountConstant(Val: NewShAmt, VT, DL));
12566 Res = DAG.getZExtOrTrunc(Op: Res, DL, VT: HalfVT);
12567 Res = DAG.getNode(Opcode: ISD::BSWAP, DL, VT: HalfVT, Operand: Res);
12568 return DAG.getZExtOrTrunc(Op: Res, DL, VT);
12569 }
12570 }
12571
12572 // Try to canonicalize bswap-of-logical-shift-by-8-bit-multiple as
12573 // inverse-shift-of-bswap:
12574 // bswap (X u<< C) --> (bswap X) u>> C
12575 // bswap (X u>> C) --> (bswap X) u<< C
12576 if ((N0.getOpcode() == ISD::SHL || N0.getOpcode() == ISD::SRL) &&
12577 N0.hasOneUse()) {
12578 auto *ShAmt = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
12579 if (ShAmt && ShAmt->getAPIntValue().ult(RHS: BW) &&
12580 ShAmt->getZExtValue() % 8 == 0) {
12581 SDValue NewSwap = DAG.getNode(Opcode: ISD::BSWAP, DL, VT, Operand: N0.getOperand(i: 0));
12582 unsigned InverseShift = N0.getOpcode() == ISD::SHL ? ISD::SRL : ISD::SHL;
12583 return DAG.getNode(Opcode: InverseShift, DL, VT, N1: NewSwap, N2: N0.getOperand(i: 1));
12584 }
12585 }
12586
12587 if (SDValue V = foldBitOrderCrossLogicOp(N, DAG))
12588 return V;
12589
12590 // Folds that depend on computeKnownBits of the operand.
12591 KnownBits Known = DAG.computeKnownBits(Op: N0);
12592 // bswap(0) = 0. Catch cases that computeKnownBits can prove are zero but
12593 // that structural combines haven't simplified to a constant yet
12594 // (e.g. and of disjoint byte masks).
12595 if (Known.isZero())
12596 return DAG.getConstant(Val: 0, DL, VT);
12597 // If only one byte of the operand may be nonzero, bswap becomes a shift
12598 // to the mirror byte.
12599 unsigned TZ = alignDown(Value: Known.countMinTrailingZeros(), Align: 8);
12600 unsigned LZ = alignDown(Value: Known.countMinLeadingZeros(), Align: 8);
12601 if (BW - (LZ + TZ) == 8) {
12602 unsigned Opc = LZ > TZ ? ISD::SHL : ISD::SRL;
12603 // Skip if the target would re-expand the produced shift post-legalize.
12604 // Targets that custom-lower byte-multiple shifts via bswap (e.g. MSP430
12605 // for shl i16) would loop with this combine.
12606 if (!LegalOperations || hasOperation(Opcode: Opc, VT)) {
12607 unsigned Amt = AbsoluteDifference(X: LZ, Y: TZ);
12608 SDNodeFlags Flags =
12609 Opc == ISD::SHL ? SDNodeFlags::NoUnsignedWrap : SDNodeFlags::Exact;
12610 return DAG.getNode(Opcode: Opc, DL, VT, N1: N0,
12611 N2: DAG.getShiftAmountConstant(Val: Amt, VT, DL), Flags);
12612 }
12613 }
12614
12615 return SDValue();
12616}
12617
12618SDValue DAGCombiner::visitBITREVERSE(SDNode *N) {
12619 SDValue N0 = N->getOperand(Num: 0);
12620 EVT VT = N->getValueType(ResNo: 0);
12621 SDLoc DL(N);
12622
12623 // fold (bitreverse c1) -> c2
12624 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::BITREVERSE, DL, VT, Ops: {N0}))
12625 return C;
12626
12627 // fold (bitreverse (bitreverse x)) -> x
12628 if (N0.getOpcode() == ISD::BITREVERSE)
12629 return N0.getOperand(i: 0);
12630
12631 SDValue X, Y;
12632
12633 // fold (bitreverse (lshr (bitreverse x), y)) -> (shl x, y)
12634 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::SHL, VT)) &&
12635 sd_match(N: N0, P: m_Srl(L: m_BitReverse(Op: m_Value(N&: X)), R: m_Value(N&: Y))))
12636 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: X, N2: Y);
12637
12638 // fold (bitreverse (shl (bitreverse x), y)) -> (lshr x, y)
12639 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::SRL, VT)) &&
12640 sd_match(N: N0, P: m_Shl(L: m_BitReverse(Op: m_Value(N&: X)), R: m_Value(N&: Y))))
12641 return DAG.getNode(Opcode: ISD::SRL, DL, VT, N1: X, N2: Y);
12642
12643 // fold bitreverse(clmul(bitreverse(x), bitreverse(y))) -> clmulr(x, y)
12644 if ((!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::CLMULR, VT)) &&
12645 sd_match(N: N0, P: m_Clmul(L: m_BitReverse(Op: m_Value(N&: X)), R: m_BitReverse(Op: m_Value(N&: Y)))))
12646 return DAG.getNode(Opcode: ISD::CLMULR, DL, VT, N1: X, N2: Y);
12647
12648 return SDValue();
12649}
12650
12651// Fold (ctlz (xor x, (sra x, bitwidth-1))) -> (add (ctls x), 1).
12652// Fold (ctlz (or (shl (xor x, (sra x, bitwidth-1)), 1), 1) -> (ctls x)
12653SDValue DAGCombiner::foldCTLZToCTLS(SDValue Src, const SDLoc &DL) {
12654 EVT VT = Src.getValueType();
12655
12656 auto LK = TLI.getTypeConversion(Context&: *DAG.getContext(), VT);
12657 if ((LK.first != TargetLoweringBase::TypeLegal &&
12658 LK.first != TargetLoweringBase::TypePromoteInteger) ||
12659 !TLI.isOperationLegalOrCustom(Op: ISD::CTLS, VT: LK.second))
12660 return SDValue();
12661
12662 unsigned BitWidth = VT.getScalarSizeInBits();
12663
12664 bool NeedAdd = true;
12665
12666 SDValue X;
12667 if (sd_match(N: Src,
12668 P: m_OneUse(P: m_Or(L: m_OneUse(P: m_Shl(L: m_Value(N&: X), R: m_One())), R: m_One())))) {
12669 NeedAdd = false;
12670 Src = X;
12671 }
12672
12673 if (!sd_match(N: Src,
12674 P: m_OneUse(P: m_Xor(L: m_Value(N&: X),
12675 R: m_OneUse(P: m_Sra(L: m_Deferred(V&: X),
12676 R: m_SpecificInt(V: BitWidth - 1)))))))
12677 return SDValue();
12678
12679 SDValue Res = DAG.getNode(Opcode: ISD::CTLS, DL, VT, Operand: X);
12680 if (!NeedAdd)
12681 return Res;
12682
12683 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Res, N2: DAG.getConstant(Val: 1, DL, VT));
12684}
12685
12686SDValue DAGCombiner::visitCTLZ(SDNode *N) {
12687 SDValue N0 = N->getOperand(Num: 0);
12688 EVT VT = N->getValueType(ResNo: 0);
12689 SDLoc DL(N);
12690
12691 // fold (ctlz c1) -> c2
12692 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::CTLZ, DL, VT, Ops: {N0}))
12693 return C;
12694
12695 // If the value is known never to be zero, switch to the poison version.
12696 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::CTLZ_ZERO_POISON, VT))
12697 if (DAG.isKnownNeverZero(Op: N0))
12698 return DAG.getNode(Opcode: ISD::CTLZ_ZERO_POISON, DL, VT, Operand: N0);
12699
12700 if (SDValue V = foldCTLZToCTLS(Src: N0, DL))
12701 return V;
12702
12703 return SDValue();
12704}
12705
12706SDValue DAGCombiner::visitCTLZ_ZERO_POISON(SDNode *N) {
12707 SDValue N0 = N->getOperand(Num: 0);
12708 EVT VT = N->getValueType(ResNo: 0);
12709 SDLoc DL(N);
12710
12711 // fold (ctlz_zero_poison c1) -> c2
12712 if (SDValue C =
12713 DAG.FoldConstantArithmetic(Opcode: ISD::CTLZ_ZERO_POISON, DL, VT, Ops: {N0}))
12714 return C;
12715
12716 if (SDValue V = foldCTLZToCTLS(Src: N0, DL))
12717 return V;
12718
12719 return SDValue();
12720}
12721
12722SDValue DAGCombiner::visitCTTZ(SDNode *N) {
12723 SDValue N0 = N->getOperand(Num: 0);
12724 EVT VT = N->getValueType(ResNo: 0);
12725 SDLoc DL(N);
12726
12727 // fold (cttz c1) -> c2
12728 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::CTTZ, DL, VT, Ops: {N0}))
12729 return C;
12730
12731 // If the value is known never to be zero, switch to the poison version.
12732 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::CTTZ_ZERO_POISON, VT))
12733 if (DAG.isKnownNeverZero(Op: N0))
12734 return DAG.getNode(Opcode: ISD::CTTZ_ZERO_POISON, DL, VT, Operand: N0);
12735
12736 return SDValue();
12737}
12738
12739SDValue DAGCombiner::visitCTTZ_ZERO_POISON(SDNode *N) {
12740 SDValue N0 = N->getOperand(Num: 0);
12741 EVT VT = N->getValueType(ResNo: 0);
12742 SDLoc DL(N);
12743
12744 // fold (cttz_zero_poison c1) -> c2
12745 if (SDValue C =
12746 DAG.FoldConstantArithmetic(Opcode: ISD::CTTZ_ZERO_POISON, DL, VT, Ops: {N0}))
12747 return C;
12748 return SDValue();
12749}
12750
12751SDValue DAGCombiner::visitPARITY(SDNode *N) {
12752 SDValue N0 = N->getOperand(Num: 0);
12753 EVT VT = N->getValueType(ResNo: 0);
12754
12755 // fold (parity c1) -> c2
12756 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::PARITY, DL: SDLoc(N), VT, Ops: {N0}))
12757 return C;
12758
12759 return SDValue();
12760}
12761
12762SDValue DAGCombiner::visitCTPOP(SDNode *N) {
12763 SDValue N0 = N->getOperand(Num: 0);
12764 EVT VT = N->getValueType(ResNo: 0);
12765 unsigned NumBits = VT.getScalarSizeInBits();
12766 SDLoc DL(N);
12767
12768 // fold (ctpop c1) -> c2
12769 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::CTPOP, DL, VT, Ops: {N0}))
12770 return C;
12771
12772 // If the source is being shifted, but doesn't affect any active bits,
12773 // then we can call CTPOP on the shift source directly.
12774 if (N0.getOpcode() == ISD::SRL || N0.getOpcode() == ISD::SHL) {
12775 if (ConstantSDNode *AmtC = isConstOrConstSplat(N: N0.getOperand(i: 1))) {
12776 const APInt &Amt = AmtC->getAPIntValue();
12777 if (Amt.ult(RHS: NumBits)) {
12778 KnownBits KnownSrc = DAG.computeKnownBits(Op: N0.getOperand(i: 0));
12779 if ((N0.getOpcode() == ISD::SRL &&
12780 Amt.ule(RHS: KnownSrc.countMinTrailingZeros())) ||
12781 (N0.getOpcode() == ISD::SHL &&
12782 Amt.ule(RHS: KnownSrc.countMinLeadingZeros()))) {
12783 return DAG.getNode(Opcode: ISD::CTPOP, DL, VT, Operand: N0.getOperand(i: 0));
12784 }
12785 }
12786 }
12787 }
12788
12789 // If the upper bits are known to be zero, then see if its profitable to
12790 // only count the lower bits.
12791 if (VT.isScalarInteger() && NumBits > 8 && (NumBits & 1) == 0) {
12792 EVT HalfVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: NumBits / 2);
12793 if (hasOperation(Opcode: ISD::CTPOP, VT: HalfVT) &&
12794 TLI.isTypeDesirableForOp(ISD::CTPOP, VT: HalfVT) &&
12795 TLI.isTruncateFree(Val: N0, VT2: HalfVT) && TLI.isZExtFree(FromTy: HalfVT, ToTy: VT)) {
12796 APInt UpperBits = APInt::getHighBitsSet(numBits: NumBits, hiBitsSet: NumBits / 2);
12797 if (DAG.MaskedValueIsZero(Op: N0, Mask: UpperBits)) {
12798 SDValue PopCnt = DAG.getNode(Opcode: ISD::CTPOP, DL, VT: HalfVT,
12799 Operand: DAG.getZExtOrTrunc(Op: N0, DL, VT: HalfVT));
12800 return DAG.getZExtOrTrunc(Op: PopCnt, DL, VT);
12801 }
12802 }
12803 }
12804
12805 return SDValue();
12806}
12807
12808static bool isLegalToCombineMinNumMaxNum(SelectionDAG &DAG, SDValue LHS,
12809 SDValue RHS,
12810 const SDNodeFlags SelectFlags,
12811 const SDNodeFlags CmpFlags,
12812 const TargetLowering &TLI) {
12813 EVT VT = LHS.getValueType();
12814 if (!VT.isFloatingPoint())
12815 return false;
12816
12817 return SelectFlags.hasNoSignedZeros() &&
12818 TLI.isProfitableToCombineMinNumMaxNum(VT) &&
12819 (SelectFlags.hasNoNaNs() || CmpFlags.hasNoNaNs() ||
12820 (DAG.isKnownNeverNaN(Op: RHS) && DAG.isKnownNeverNaN(Op: LHS)));
12821}
12822
12823static SDValue combineMinNumMaxNumImpl(const SDLoc &DL, EVT VT, SDValue LHS,
12824 SDValue RHS, SDValue True, SDValue False,
12825 ISD::CondCode CC,
12826 const TargetLowering &TLI,
12827 SelectionDAG &DAG) {
12828 EVT TransformVT = TLI.getLegalTypeToTransformTo(Context&: *DAG.getContext(), VT);
12829
12830 // We have checked nnan and nsz as pre-conditions for the transform.
12831 SDNodeFlags Flags = SDNodeFlags::NoNaNs | SDNodeFlags::NoSignedZeros;
12832
12833 switch (CC) {
12834 case ISD::SETOLT:
12835 case ISD::SETOLE:
12836 case ISD::SETLT:
12837 case ISD::SETLE:
12838 case ISD::SETULT:
12839 case ISD::SETULE: {
12840 // Since it's known never nan to get here already, either fminnum or
12841 // fminnum_ieee are OK. Try the ieee version first, since it's fminnum is
12842 // expanded in terms of it.
12843 unsigned IEEEOpcode = (LHS == True) ? ISD::FMINNUM_IEEE : ISD::FMAXNUM_IEEE;
12844 if (TLI.isOperationLegalOrCustom(Op: IEEEOpcode, VT))
12845 return DAG.getNode(Opcode: IEEEOpcode, DL, VT, N1: LHS, N2: RHS, Flags);
12846
12847 unsigned Opcode = (LHS == True) ? ISD::FMINNUM : ISD::FMAXNUM;
12848 if (TLI.isOperationLegalOrCustom(Op: Opcode, VT: TransformVT))
12849 return DAG.getNode(Opcode, DL, VT, N1: LHS, N2: RHS, Flags);
12850 return SDValue();
12851 }
12852 case ISD::SETOGT:
12853 case ISD::SETOGE:
12854 case ISD::SETGT:
12855 case ISD::SETGE:
12856 case ISD::SETUGT:
12857 case ISD::SETUGE: {
12858 unsigned IEEEOpcode = (LHS == True) ? ISD::FMAXNUM_IEEE : ISD::FMINNUM_IEEE;
12859 if (TLI.isOperationLegalOrCustom(Op: IEEEOpcode, VT))
12860 return DAG.getNode(Opcode: IEEEOpcode, DL, VT, N1: LHS, N2: RHS, Flags);
12861
12862 unsigned Opcode = (LHS == True) ? ISD::FMAXNUM : ISD::FMINNUM;
12863 if (TLI.isOperationLegalOrCustom(Op: Opcode, VT: TransformVT))
12864 return DAG.getNode(Opcode, DL, VT, N1: LHS, N2: RHS, Flags);
12865 return SDValue();
12866 }
12867 default:
12868 return SDValue();
12869 }
12870}
12871
12872// Convert (sr[al] (add n[su]w x, y)) -> (avgfloor[su] x, y)
12873SDValue DAGCombiner::foldShiftToAvg(SDNode *N, const SDLoc &DL) {
12874 const unsigned Opcode = N->getOpcode();
12875 if (Opcode != ISD::SRA && Opcode != ISD::SRL)
12876 return SDValue();
12877
12878 EVT VT = N->getValueType(ResNo: 0);
12879 bool IsUnsigned = Opcode == ISD::SRL;
12880
12881 // Captured values.
12882 SDValue A, B;
12883
12884 // Match floor average as it is common to both floor/ceil avgs, ensure the add
12885 // doesn't wrap.
12886 SDNodeFlags Flags =
12887 IsUnsigned ? SDNodeFlags::NoUnsignedWrap : SDNodeFlags::NoSignedWrap;
12888 if (sd_match(N, P: m_BinOp(Opc: Opcode,
12889 L: m_c_BinOp(Opc: ISD::ADD, L: m_Value(N&: A), R: m_Value(N&: B), Flgs: Flags),
12890 R: m_One()))) {
12891 // Decide whether signed or unsigned.
12892 unsigned FloorISD = IsUnsigned ? ISD::AVGFLOORU : ISD::AVGFLOORS;
12893 if (hasOperation(Opcode: FloorISD, VT))
12894 return DAG.getNode(Opcode: FloorISD, DL, VT, Ops: {A, B});
12895 }
12896
12897 return SDValue();
12898}
12899
12900SDValue DAGCombiner::foldBitwiseOpWithNeg(SDNode *N, const SDLoc &DL, EVT VT) {
12901 unsigned Opc = N->getOpcode();
12902 SDValue X, Y, Z;
12903 if (sd_match(
12904 N, P: m_BitwiseLogic(L: m_Value(N&: X), R: m_Add(L: m_Not(V: m_Value(N&: Y)), R: m_Value(N&: Z)))))
12905 return DAG.getNode(Opcode: Opc, DL, VT, N1: X,
12906 N2: DAG.getNOT(DL, Val: DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Y, N2: Z), VT));
12907
12908 if (sd_match(N, P: m_BitwiseLogic(L: m_Value(N&: X), R: m_Sub(L: m_OneUse(P: m_Not(V: m_Value(N&: Y))),
12909 R: m_Value(N&: Z)))))
12910 return DAG.getNode(Opcode: Opc, DL, VT, N1: X,
12911 N2: DAG.getNOT(DL, Val: DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Y, N2: Z), VT));
12912
12913 return SDValue();
12914}
12915
12916/// Generate Min/Max node
12917SDValue DAGCombiner::combineMinNumMaxNum(const SDLoc &DL, EVT VT, SDValue LHS,
12918 SDValue RHS, SDValue True,
12919 SDValue False, ISD::CondCode CC) {
12920 if ((LHS == True && RHS == False) || (LHS == False && RHS == True))
12921 return combineMinNumMaxNumImpl(DL, VT, LHS, RHS, True, False, CC, TLI, DAG);
12922
12923 // If we can't directly match this, try to see if we can pull an fneg out of
12924 // the select.
12925 SDValue NegTrue = TLI.getCheaperOrNeutralNegatedExpression(
12926 Op: True, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize);
12927 if (!NegTrue)
12928 return SDValue();
12929
12930 HandleSDNode NegTrueHandle(NegTrue);
12931
12932 // Try to unfold an fneg from the select if we are comparing the negated
12933 // constant.
12934 //
12935 // select (setcc x, K) (fneg x), -K -> fneg(minnum(x, K))
12936 //
12937 // TODO: Handle fabs
12938 if (LHS == NegTrue) {
12939 // If we can't directly match this, try to see if we can pull an fneg out of
12940 // the select.
12941 SDValue NegRHS = TLI.getCheaperOrNeutralNegatedExpression(
12942 Op: RHS, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize);
12943 if (NegRHS) {
12944 HandleSDNode NegRHSHandle(NegRHS);
12945 if (NegRHS == False) {
12946 SDValue Combined = combineMinNumMaxNumImpl(DL, VT, LHS, RHS, True: NegTrue,
12947 False, CC, TLI, DAG);
12948 if (Combined)
12949 return DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: Combined);
12950 }
12951 }
12952 }
12953
12954 return SDValue();
12955}
12956
12957/// If a (v)select has a condition value that is a sign-bit test, try to smear
12958/// the condition operand sign-bit across the value width and use it as a mask.
12959static SDValue foldSelectOfConstantsUsingSra(SDNode *N, const SDLoc &DL,
12960 SelectionDAG &DAG) {
12961 SDValue Cond = N->getOperand(Num: 0);
12962 SDValue C1 = N->getOperand(Num: 1);
12963 SDValue C2 = N->getOperand(Num: 2);
12964 if (!isConstantOrConstantVector(N: C1) || !isConstantOrConstantVector(N: C2))
12965 return SDValue();
12966
12967 EVT VT = N->getValueType(ResNo: 0);
12968 if (Cond.getOpcode() != ISD::SETCC || !Cond.hasOneUse() ||
12969 VT != Cond.getOperand(i: 0).getValueType())
12970 return SDValue();
12971
12972 // The inverted-condition + commuted-select variants of these patterns are
12973 // canonicalized to these forms in IR.
12974 SDValue X = Cond.getOperand(i: 0);
12975 SDValue CondC = Cond.getOperand(i: 1);
12976 ISD::CondCode CC = cast<CondCodeSDNode>(Val: Cond.getOperand(i: 2))->get();
12977 if (CC == ISD::SETGT && isAllOnesOrAllOnesSplat(V: CondC) &&
12978 isAllOnesOrAllOnesSplat(V: C2)) {
12979 // i32 X > -1 ? C1 : -1 --> (X >>s 31) | C1
12980 SDValue ShAmtC = DAG.getConstant(Val: X.getScalarValueSizeInBits() - 1, DL, VT);
12981 SDValue Sra = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: X, N2: ShAmtC);
12982 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: Sra, N2: C1);
12983 }
12984 if (CC == ISD::SETLT && isNullOrNullSplat(V: CondC) && isNullOrNullSplat(V: C2)) {
12985 // i8 X < 0 ? C1 : 0 --> (X >>s 7) & C1
12986 SDValue ShAmtC = DAG.getConstant(Val: X.getScalarValueSizeInBits() - 1, DL, VT);
12987 SDValue Sra = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: X, N2: ShAmtC);
12988 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Sra, N2: C1);
12989 }
12990 return SDValue();
12991}
12992
12993static bool shouldConvertSelectOfConstantsToMath(const SDValue &Cond, EVT VT,
12994 const TargetLowering &TLI) {
12995 if (!TLI.convertSelectOfConstantsToMath(VT))
12996 return false;
12997
12998 if (Cond.getOpcode() != ISD::SETCC || !Cond->hasOneUse())
12999 return true;
13000 if (!TLI.isOperationLegalOrCustom(Op: ISD::SELECT_CC, VT))
13001 return true;
13002
13003 ISD::CondCode CC = cast<CondCodeSDNode>(Val: Cond.getOperand(i: 2))->get();
13004 if (CC == ISD::SETLT && isNullOrNullSplat(V: Cond.getOperand(i: 1)))
13005 return true;
13006 if (CC == ISD::SETGT && isAllOnesOrAllOnesSplat(V: Cond.getOperand(i: 1)))
13007 return true;
13008
13009 return false;
13010}
13011
13012SDValue DAGCombiner::foldSelectOfConstants(SDNode *N) {
13013 SDValue Cond = N->getOperand(Num: 0);
13014 SDValue N1 = N->getOperand(Num: 1);
13015 SDValue N2 = N->getOperand(Num: 2);
13016 EVT VT = N->getValueType(ResNo: 0);
13017 EVT CondVT = Cond.getValueType();
13018 SDLoc DL(N);
13019
13020 if (!VT.isInteger())
13021 return SDValue();
13022
13023 auto *C1 = dyn_cast<ConstantSDNode>(Val&: N1);
13024 auto *C2 = dyn_cast<ConstantSDNode>(Val&: N2);
13025 if (!C1 || !C2)
13026 return SDValue();
13027
13028 if (CondVT != MVT::i1 || LegalOperations) {
13029 // We can't do this reliably if integer based booleans have different contents
13030 // to floating point based booleans. This is because we can't tell whether we
13031 // have an integer-based boolean or a floating-point-based boolean unless we
13032 // can find the SETCC that produced it and inspect its operands. This is
13033 // fairly easy if C is the SETCC node, but it can potentially be
13034 // undiscoverable (or not reasonably discoverable). For example, it could be
13035 // in another basic block or it could require searching a complicated
13036 // expression.
13037 if (CondVT.isInteger() &&
13038 TLI.getBooleanContents(/*isVec*/false, /*isFloat*/true) ==
13039 TargetLowering::ZeroOrOneBooleanContent &&
13040 TLI.getBooleanContents(/*isVec*/false, /*isFloat*/false) ==
13041 TargetLowering::ZeroOrOneBooleanContent) {
13042 // fold (select Cond, 0, 1) -> (xor Cond, 1)
13043 if (C1->isZero() && C2->isOne()) {
13044 SDValue NotCond = DAG.getNode(Opcode: ISD::XOR, DL, VT: CondVT, N1: Cond,
13045 N2: DAG.getConstant(Val: 1, DL, VT: CondVT));
13046 if (VT.bitsEq(VT: CondVT))
13047 return NotCond;
13048 return DAG.getZExtOrTrunc(Op: NotCond, DL, VT);
13049 }
13050
13051 // fold (select Cond, 1, 0) -> Cond
13052 if (C1->isOne() && C2->isZero() && CondVT == VT)
13053 return Cond;
13054 }
13055
13056 return SDValue();
13057 }
13058
13059 // Only do this before legalization to avoid conflicting with target-specific
13060 // transforms in the other direction (create a select from a zext/sext). There
13061 // is also a target-independent combine here in DAGCombiner in the other
13062 // direction for (select Cond, -1, 0) when the condition is not i1.
13063 assert(CondVT == MVT::i1 && !LegalOperations);
13064
13065 // select Cond, 1, 0 --> zext (Cond)
13066 if (C1->isOne() && C2->isZero())
13067 return DAG.getZExtOrTrunc(Op: Cond, DL, VT);
13068
13069 // select Cond, -1, 0 --> sext (Cond)
13070 if (C1->isAllOnes() && C2->isZero())
13071 return DAG.getSExtOrTrunc(Op: Cond, DL, VT);
13072
13073 // select Cond, 0, 1 --> zext (!Cond)
13074 if (C1->isZero() && C2->isOne()) {
13075 SDValue NotCond = DAG.getNOT(DL, Val: Cond, VT: MVT::i1);
13076 NotCond = DAG.getZExtOrTrunc(Op: NotCond, DL, VT);
13077 return NotCond;
13078 }
13079
13080 // select Cond, 0, -1 --> sext (!Cond)
13081 if (C1->isZero() && C2->isAllOnes()) {
13082 SDValue NotCond = DAG.getNOT(DL, Val: Cond, VT: MVT::i1);
13083 NotCond = DAG.getSExtOrTrunc(Op: NotCond, DL, VT);
13084 return NotCond;
13085 }
13086
13087 // Use a target hook because some targets may prefer to transform in the
13088 // other direction.
13089 if (!shouldConvertSelectOfConstantsToMath(Cond, VT, TLI))
13090 return SDValue();
13091
13092 // For any constants that differ by 1, we can transform the select into
13093 // an extend and add.
13094 const APInt &C1Val = C1->getAPIntValue();
13095 const APInt &C2Val = C2->getAPIntValue();
13096
13097 // select Cond, C1, C1-1 --> add (zext Cond), C1-1
13098 if (C1Val - 1 == C2Val) {
13099 Cond = DAG.getZExtOrTrunc(Op: Cond, DL, VT);
13100 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Cond, N2);
13101 }
13102
13103 // select Cond, C1, C1+1 --> add (sext Cond), C1+1
13104 if (C1Val + 1 == C2Val) {
13105 Cond = DAG.getSExtOrTrunc(Op: Cond, DL, VT);
13106 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Cond, N2);
13107 }
13108
13109 // select Cond, Pow2, 0 --> (zext Cond) << log2(Pow2)
13110 if (C1Val.isPowerOf2() && C2Val.isZero()) {
13111 Cond = DAG.getZExtOrTrunc(Op: Cond, DL, VT);
13112 SDValue ShAmtC =
13113 DAG.getShiftAmountConstant(Val: C1Val.exactLogBase2(), VT, DL);
13114 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Cond, N2: ShAmtC);
13115 }
13116
13117 // select Cond, -1, C --> or (sext Cond), C
13118 if (C1->isAllOnes()) {
13119 Cond = DAG.getSExtOrTrunc(Op: Cond, DL, VT);
13120 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: Cond, N2);
13121 }
13122
13123 // select Cond, C, -1 --> or (sext (not Cond)), C
13124 if (C2->isAllOnes()) {
13125 SDValue NotCond = DAG.getNOT(DL, Val: Cond, VT: MVT::i1);
13126 NotCond = DAG.getSExtOrTrunc(Op: NotCond, DL, VT);
13127 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: NotCond, N2: N1);
13128 }
13129
13130 if (SDValue V = foldSelectOfConstantsUsingSra(N, DL, DAG))
13131 return V;
13132
13133 return SDValue();
13134}
13135
13136static SDValue foldBoolSelectToLogic(SDNode *N, const SDLoc &DL,
13137 SelectionDAG &DAG) {
13138 assert((N->getOpcode() == ISD::SELECT || N->getOpcode() == ISD::VSELECT) &&
13139 "Expected a (v)select");
13140 SDValue Cond = N->getOperand(Num: 0);
13141 SDValue T = N->getOperand(Num: 1), F = N->getOperand(Num: 2);
13142 EVT VT = N->getValueType(ResNo: 0);
13143
13144 if (VT != Cond.getValueType() || VT.getScalarSizeInBits() != 1)
13145 return SDValue();
13146
13147 // select Cond, Cond, F --> or Cond, freeze(F)
13148 // select Cond, 1, F --> or Cond, freeze(F)
13149 if (Cond == T || isOneOrOneSplat(V: T, /* AllowUndefs */ true))
13150 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: Cond, N2: DAG.getFreeze(V: F));
13151
13152 // select Cond, T, Cond --> and Cond, freeze(T)
13153 // select Cond, T, 0 --> and Cond, freeze(T)
13154 if (Cond == F || isNullOrNullSplat(V: F, /* AllowUndefs */ true))
13155 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Cond, N2: DAG.getFreeze(V: T));
13156
13157 // select Cond, T, 1 --> or (not Cond), freeze(T)
13158 if (isOneOrOneSplat(V: F, /* AllowUndefs */ true)) {
13159 SDValue NotCond =
13160 DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: Cond, N2: DAG.getAllOnesConstant(DL, VT));
13161 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: NotCond, N2: DAG.getFreeze(V: T));
13162 }
13163
13164 // select Cond, 0, F --> and (not Cond), freeze(F)
13165 if (isNullOrNullSplat(V: T, /* AllowUndefs */ true)) {
13166 SDValue NotCond =
13167 DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: Cond, N2: DAG.getAllOnesConstant(DL, VT));
13168 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: NotCond, N2: DAG.getFreeze(V: F));
13169 }
13170
13171 return SDValue();
13172}
13173
13174static SDValue foldVSelectToSignBitSplatMask(SDNode *N, SelectionDAG &DAG) {
13175 SDValue N0 = N->getOperand(Num: 0);
13176 SDValue N1 = N->getOperand(Num: 1);
13177 SDValue N2 = N->getOperand(Num: 2);
13178 EVT VT = N->getValueType(ResNo: 0);
13179 unsigned EltSizeInBits = VT.getScalarSizeInBits();
13180
13181 SDValue Cond0, Cond1;
13182 ISD::CondCode CC;
13183 if (!sd_match(N: N0, P: m_OneUse(P: m_SetCC(CC, LHS: m_Value(N&: Cond0), RHS: m_Value(N&: Cond1)))) ||
13184 VT != Cond0.getValueType())
13185 return SDValue();
13186
13187 // Match a signbit check of Cond0 as "Cond0 s<0". Swap select operands if the
13188 // compare is inverted from that pattern ("Cond0 s> -1").
13189 if (CC == ISD::SETLT && isNullOrNullSplat(V: Cond1))
13190 ; // This is the pattern we are looking for.
13191 else if (CC == ISD::SETGT && isAllOnesOrAllOnesSplat(V: Cond1))
13192 std::swap(a&: N1, b&: N2);
13193 else
13194 return SDValue();
13195
13196 // (Cond0 s< 0) ? N1 : 0 --> (Cond0 s>> BW-1) & freeze(N1)
13197 if (isNullOrNullSplat(V: N2)) {
13198 SDLoc DL(N);
13199 SDValue ShiftAmt = DAG.getShiftAmountConstant(Val: EltSizeInBits - 1, VT, DL);
13200 SDValue Sra = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: Cond0, N2: ShiftAmt);
13201 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Sra, N2: DAG.getFreeze(V: N1));
13202 }
13203
13204 // (Cond0 s< 0) ? -1 : N2 --> (Cond0 s>> BW-1) | freeze(N2)
13205 if (isAllOnesOrAllOnesSplat(V: N1)) {
13206 SDLoc DL(N);
13207 SDValue ShiftAmt = DAG.getShiftAmountConstant(Val: EltSizeInBits - 1, VT, DL);
13208 SDValue Sra = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: Cond0, N2: ShiftAmt);
13209 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: Sra, N2: DAG.getFreeze(V: N2));
13210 }
13211
13212 // If we have to invert the sign bit mask, only do that transform if the
13213 // target has a bitwise 'and not' instruction (the invert is free).
13214 // (Cond0 s< -0) ? 0 : N2 --> ~(Cond0 s>> BW-1) & freeze(N2)
13215 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
13216 if (isNullOrNullSplat(V: N1) && TLI.hasAndNot(X: N1)) {
13217 SDLoc DL(N);
13218 SDValue ShiftAmt = DAG.getShiftAmountConstant(Val: EltSizeInBits - 1, VT, DL);
13219 SDValue Sra = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: Cond0, N2: ShiftAmt);
13220 SDValue Not = DAG.getNOT(DL, Val: Sra, VT);
13221 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Not, N2: DAG.getFreeze(V: N2));
13222 }
13223
13224 // TODO: There's another pattern in this family, but it may require
13225 // implementing hasOrNot() to check for profitability:
13226 // (Cond0 s> -1) ? -1 : N2 --> ~(Cond0 s>> BW-1) | freeze(N2)
13227
13228 return SDValue();
13229}
13230
13231// Match SELECTs with absolute difference patterns.
13232// (select (setcc a, b, set?gt), (sub a, b), (sub b, a)) --> (abd? a, b)
13233// (select (setcc a, b, set?ge), (sub a, b), (sub b, a)) --> (abd? a, b)
13234// (select (setcc a, b, set?lt), (sub b, a), (sub a, b)) --> (abd? a, b)
13235// (select (setcc a, b, set?le), (sub b, a), (sub a, b)) --> (abd? a, b)
13236SDValue DAGCombiner::foldSelectToABD(SDValue LHS, SDValue RHS, SDValue True,
13237 SDValue False, ISD::CondCode CC,
13238 const SDLoc &DL) {
13239 bool IsSigned = isSignedIntSetCC(Code: CC);
13240 unsigned ABDOpc = IsSigned ? ISD::ABDS : ISD::ABDU;
13241 EVT VT = LHS.getValueType();
13242
13243 if (LegalOperations && !hasOperation(Opcode: ABDOpc, VT))
13244 return SDValue();
13245
13246 // (setcc 0, b set???) --> (setcc b, 0, set???)
13247 if (isZeroOrZeroSplat(N: LHS)) {
13248 std::swap(a&: LHS, b&: RHS);
13249 CC = ISD::getSetCCSwappedOperands(Operation: CC);
13250 }
13251
13252 // (setcc (add nsw A, Const), 0, sets??) --> (setcc A, -Const, sets??)
13253 SDValue A, B;
13254 if (ISD::isSignedIntSetCC(Code: CC) && LHS->getFlags().hasNoSignedWrap() &&
13255 isZeroOrZeroSplat(N: RHS) && sd_match(N: LHS, P: m_Add(L: m_Value(N&: A), R: m_Value(N&: B))) &&
13256 DAG.isConstantIntBuildVectorOrConstantInt(N: B)) {
13257 RHS = DAG.getNegative(Val: B, DL: LHS, VT: B.getValueType());
13258 LHS = A;
13259 }
13260
13261 bool IsTypeLegalOrPromote =
13262 TLI.isTypeLegal(VT) || TLI.getTypeAction(Context&: *DAG.getContext(), VT) ==
13263 TargetLowering::TypePromoteInteger;
13264
13265 switch (CC) {
13266 case ISD::SETGT:
13267 case ISD::SETGE:
13268 case ISD::SETUGT:
13269 case ISD::SETUGE:
13270 if (sd_match(N: True, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: LHS), R: m_Specific(N: RHS)),
13271 preds: m_Add(L: m_Specific(N: LHS), R: m_SpecificNeg(V: RHS)))) &&
13272 sd_match(N: False, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: RHS), R: m_Specific(N: LHS)),
13273 preds: m_Add(L: m_Specific(N: RHS), R: m_SpecificNeg(V: LHS)))))
13274 return DAG.getNode(Opcode: ABDOpc, DL, VT, N1: LHS, N2: RHS);
13275 if (sd_match(N: True, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: RHS), R: m_Specific(N: LHS)),
13276 preds: m_Add(L: m_Specific(N: RHS), R: m_SpecificNeg(V: LHS)))) &&
13277 sd_match(N: False, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: LHS), R: m_Specific(N: RHS)),
13278 preds: m_Add(L: m_Specific(N: LHS), R: m_SpecificNeg(V: RHS)))) &&
13279 IsTypeLegalOrPromote)
13280 return DAG.getNegative(Val: DAG.getNode(Opcode: ABDOpc, DL, VT, N1: LHS, N2: RHS), DL, VT);
13281 break;
13282 case ISD::SETLT:
13283 case ISD::SETLE:
13284 case ISD::SETULT:
13285 case ISD::SETULE:
13286 if (sd_match(N: True, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: RHS), R: m_Specific(N: LHS)),
13287 preds: m_Add(L: m_Specific(N: RHS), R: m_SpecificNeg(V: LHS)))) &&
13288 sd_match(N: False, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: LHS), R: m_Specific(N: RHS)),
13289 preds: m_Add(L: m_Specific(N: LHS), R: m_SpecificNeg(V: RHS)))))
13290 return DAG.getNode(Opcode: ABDOpc, DL, VT, N1: LHS, N2: RHS);
13291 if (sd_match(N: True, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: LHS), R: m_Specific(N: RHS)),
13292 preds: m_Add(L: m_Specific(N: LHS), R: m_SpecificNeg(V: RHS)))) &&
13293 sd_match(N: False, P: m_AnyOf(preds: m_Sub(L: m_Specific(N: RHS), R: m_Specific(N: LHS)),
13294 preds: m_Add(L: m_Specific(N: RHS), R: m_SpecificNeg(V: LHS)))) &&
13295 IsTypeLegalOrPromote)
13296 return DAG.getNegative(Val: DAG.getNode(Opcode: ABDOpc, DL, VT, N1: LHS, N2: RHS), DL, VT);
13297 break;
13298 default:
13299 break;
13300 }
13301
13302 return SDValue();
13303}
13304
13305// ([v]select (ugt x, C), (add x, ~C), x) -> (umin (add x, ~C), x)
13306// ([v]select (ult x, C), x, (add x, -C)) -> (umin x, (add x, -C))
13307SDValue DAGCombiner::foldSelectToUMin(SDValue LHS, SDValue RHS, SDValue True,
13308 SDValue False, ISD::CondCode CC,
13309 const SDLoc &DL) {
13310 APInt C;
13311 EVT VT = True.getValueType();
13312 if (sd_match(N: RHS, P: m_ConstInt(V&: C)) && hasUMin(VT)) {
13313 if (CC == ISD::SETUGT && LHS == False &&
13314 sd_match(N: True, P: m_Add(L: m_Specific(N: False), R: m_SpecificInt(V: ~C)))) {
13315 SDValue AddC = DAG.getConstant(Val: ~C, DL, VT);
13316 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: False, N2: AddC);
13317 return DAG.getNode(Opcode: ISD::UMIN, DL, VT, N1: Add, N2: False);
13318 }
13319 if (CC == ISD::SETULT && LHS == True &&
13320 sd_match(N: False, P: m_Add(L: m_Specific(N: True), R: m_SpecificInt(V: -C)))) {
13321 SDValue AddC = DAG.getConstant(Val: -C, DL, VT);
13322 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: True, N2: AddC);
13323 return DAG.getNode(Opcode: ISD::UMIN, DL, VT, N1: True, N2: Add);
13324 }
13325 }
13326 return SDValue();
13327}
13328
13329// Combine x olt y ? x : y to pseudo_fmin and x ogt y ? x : y to pseudo_fmax.
13330// Op0/Op1 are the setcc operands, LHS/RHS are the select operands, Flags are
13331// from the select.
13332// The return value is the opcode and its operands.
13333static std::tuple<unsigned, SDValue, SDValue> combineSelectCCToPseudoMinMax(
13334 SelectionDAG &DAG, const SDLoc &DL, ISD::CondCode CC, SDValue Op0,
13335 SDValue Op1, SDValue LHS, SDValue RHS, SDNodeFlags Flags, bool IsStrict) {
13336 std::tuple<unsigned, SDValue, SDValue> Invalid(0, {}, {});
13337 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
13338 EVT VT = LHS.getValueType();
13339 if (!VT.isFloatingPoint())
13340 return Invalid;
13341
13342 // Check for x CC y ? x : y.
13343 if (!DAG.isEqualTo(A: LHS, B: Op0) || !DAG.isEqualTo(A: RHS, B: Op1)) {
13344 if (!DAG.isEqualTo(A: LHS, B: Op1) || !DAG.isEqualTo(A: RHS, B: Op0))
13345 return Invalid;
13346
13347 // Convert x CC y ? y : x to x inv(CC) y ? x : y.
13348 CC = ISD::getSetCCInverse(Operation: CC, Type: VT);
13349 std::swap(a&: LHS, b&: RHS);
13350 }
13351
13352 // Convert x CC y ? x : y to y swap(inv(CC)) x ? y : x
13353 // to convert an unordered into an ordered comparison.
13354 if (ISD::getUnorderedFlavor(Cond: CC) == 1) {
13355 CC = ISD::getSetCCSwappedOperands(Operation: ISD::getSetCCInverse(Operation: CC, Type: VT));
13356 std::swap(a&: LHS, b&: RHS);
13357 }
13358
13359 unsigned Opcode = 0;
13360 switch (CC) {
13361 default:
13362 break;
13363 case ISD::SETOLE:
13364 // Converting this to a min would handle comparisons between positive
13365 // and negative zero incorrectly.
13366 if (!Flags.hasNoSignedZeros() && !DAG.isKnownNeverLogicalZero(Op: LHS) &&
13367 !DAG.isKnownNeverLogicalZero(Op: RHS))
13368 break;
13369 Opcode = ISD::PSEUDO_FMIN;
13370 break;
13371 case ISD::SETLE:
13372 // Convert setle to setlt via inv+swap.
13373 std::swap(a&: LHS, b&: RHS);
13374 [[fallthrough]];
13375 case ISD::SETOLT:
13376 case ISD::SETLT:
13377 Opcode = ISD::PSEUDO_FMIN;
13378 break;
13379
13380 case ISD::SETOGE:
13381 // Converting this to a max would handle comparisons between positive
13382 // and negative zero incorrectly.
13383 if (!Flags.hasNoSignedZeros() && !DAG.isKnownNeverLogicalZero(Op: LHS) &&
13384 !DAG.isKnownNeverLogicalZero(Op: RHS))
13385 break;
13386 Opcode = ISD::PSEUDO_FMAX;
13387 break;
13388 case ISD::SETGE:
13389 // Convert setge to setgt via inv+swap.
13390 std::swap(a&: LHS, b&: RHS);
13391 [[fallthrough]];
13392 case ISD::SETOGT:
13393 case ISD::SETGT:
13394 Opcode = ISD::PSEUDO_FMAX;
13395 break;
13396 }
13397
13398 if (!Opcode)
13399 return Invalid;
13400
13401 if (IsStrict)
13402 Opcode = Opcode == ISD::PSEUDO_FMIN ? ISD::STRICT_PSEUDO_FMIN
13403 : ISD::STRICT_PSEUDO_FMAX;
13404 if (!TLI.isOperationLegalOrCustom(Op: Opcode, VT))
13405 return Invalid;
13406
13407 return {Opcode, LHS, RHS};
13408}
13409
13410static SDValue combineSelectToPseudoMinMax(SelectionDAG &DAG, SDNode *N) {
13411 SDLoc DL(N);
13412 SDValue Cond = N->getOperand(Num: 0);
13413 SDValue LHS = N->getOperand(Num: 1);
13414 SDValue RHS = N->getOperand(Num: 2);
13415 EVT VT = LHS.getValueType();
13416 if ((Cond.getOpcode() != ISD::SETCC &&
13417 Cond.getOpcode() != ISD::STRICT_FSETCCS))
13418 return SDValue();
13419
13420 bool IsStrict = Cond->isStrictFPOpcode();
13421 ISD::CondCode CC =
13422 cast<CondCodeSDNode>(Val: Cond.getOperand(i: IsStrict ? 3 : 2))->get();
13423 SDValue Op0 = Cond.getOperand(i: IsStrict ? 1 : 0);
13424 SDValue Op1 = Cond.getOperand(i: IsStrict ? 2 : 1);
13425 auto [Opcode, NewLHS, NewRHS] = combineSelectCCToPseudoMinMax(
13426 DAG, DL, CC, Op0, Op1, LHS, RHS, Flags: N->getFlags(), IsStrict);
13427 if (!Opcode)
13428 return SDValue();
13429
13430 // Propagate fast-math-flags.
13431 SelectionDAG::FlagInserter FlagsInserter(DAG, N->getFlags());
13432 if (IsStrict) {
13433 SDValue Ret = DAG.getNode(Opcode, DL, ResultTys: {VT, MVT::Other},
13434 Ops: {Cond.getOperand(i: 0), NewLHS, NewRHS});
13435 DAG.ReplaceAllUsesOfValueWith(From: Cond.getValue(R: 1), To: Ret.getValue(R: 1));
13436 return Ret;
13437 }
13438 return DAG.getNode(Opcode, DL, VT, N1: NewLHS, N2: NewRHS);
13439}
13440
13441/// Fold:
13442/// select_cc (select C, TV, FV), CmpC, TrueV, FalseV, seteq
13443/// -> select C, TrueV, FalseV
13444/// select_cc (select C, TV, FV), CmpC, TrueV, FalseV, setne
13445/// -> select C, FalseV, TrueV
13446/// and the same with CmpC on the LHS of the comparison. TV and FV must be
13447/// distinct integer constants. Also used for select (setcc ...).
13448static SDValue foldSelectOfSelectCmp(SDValue LHS, SDValue RHS, ISD::CondCode CC,
13449 SDValue TrueV, SDValue FalseV,
13450 const SDLoc &DL, EVT VT, SelectionDAG &DAG,
13451 SDNodeFlags Flags) {
13452 if (CC != ISD::SETEQ && CC != ISD::SETNE)
13453 return SDValue();
13454
13455 SDValue InnerSel;
13456 SDValue CmpC;
13457 if (LHS.getOpcode() == ISD::SELECT) {
13458 InnerSel = LHS;
13459 CmpC = RHS;
13460 } else if (RHS.getOpcode() == ISD::SELECT) {
13461 InnerSel = RHS;
13462 CmpC = LHS;
13463 } else
13464 return SDValue();
13465
13466 SDValue Cond = InnerSel.getOperand(i: 0);
13467 SDValue InnerTV = InnerSel.getOperand(i: 1);
13468 SDValue InnerFV = InnerSel.getOperand(i: 2);
13469
13470 auto *CTV = dyn_cast<ConstantSDNode>(Val&: InnerTV);
13471 auto *CFV = dyn_cast<ConstantSDNode>(Val&: InnerFV);
13472 auto *CCnst = dyn_cast<ConstantSDNode>(Val&: CmpC);
13473 if (!CTV || !CFV || !CCnst)
13474 return SDValue();
13475
13476 // If one of the constants is opaque, the SDNodes may differ while the values
13477 // are the same. Check APInt to avoid miscompiles.
13478 if (CTV->getAPIntValue() == CFV->getAPIntValue())
13479 return SDValue();
13480
13481 const APInt &CmpVal = CCnst->getAPIntValue();
13482 bool MatchesTV = CmpVal == CTV->getAPIntValue();
13483 bool MatchesFV = CmpVal == CFV->getAPIntValue();
13484 if (!MatchesTV && !MatchesFV)
13485 return SDValue();
13486
13487 SDValue SelTrueV = TrueV;
13488 SDValue SelFalseV = FalseV;
13489 if (CC == ISD::SETEQ) {
13490 if (MatchesFV)
13491 std::swap(a&: SelTrueV, b&: SelFalseV);
13492 } else {
13493 if (MatchesTV)
13494 std::swap(a&: SelTrueV, b&: SelFalseV);
13495 }
13496
13497 return DAG.getSelect(DL, VT, Cond, LHS: SelTrueV, RHS: SelFalseV, Flags);
13498}
13499
13500static SDValue foldSelectCCOfSelect(SDNode *N, SelectionDAG &DAG) {
13501 ISD::CondCode CC = cast<CondCodeSDNode>(Val: N->getOperand(Num: 4))->get();
13502 return foldSelectOfSelectCmp(LHS: N->getOperand(Num: 0), RHS: N->getOperand(Num: 1), CC,
13503 TrueV: N->getOperand(Num: 2), FalseV: N->getOperand(Num: 3), DL: SDLoc(N),
13504 VT: N->getValueType(ResNo: 0), DAG, Flags: N->getFlags());
13505}
13506
13507SDValue DAGCombiner::visitSELECT(SDNode *N) {
13508 SDValue N0 = N->getOperand(Num: 0);
13509 SDValue N1 = N->getOperand(Num: 1);
13510 SDValue N2 = N->getOperand(Num: 2);
13511 EVT VT = N->getValueType(ResNo: 0);
13512 EVT VT0 = N0.getValueType();
13513 SDLoc DL(N);
13514 SDNodeFlags Flags = N->getFlags();
13515
13516 if (SDValue V = DAG.simplifySelect(Cond: N0, TVal: N1, FVal: N2))
13517 return V;
13518
13519 if (SDValue V = foldBoolSelectToLogic(N, DL, DAG))
13520 return V;
13521
13522 // select (not Cond), N1, N2 -> select Cond, N2, N1
13523 if (SDValue F = extractBooleanFlip(V: N0, DAG, TLI, Force: false))
13524 return DAG.getSelect(DL, VT, Cond: F, LHS: N2, RHS: N1, Flags);
13525
13526 if (SDValue V = foldSelectOfConstants(N))
13527 return V;
13528
13529 // select (setcc (select C, TV, FV), CmpC, cc), TrueV, FalseV
13530 // -> select C, TrueV, FalseV (or swapped FalseV/TrueV)
13531 if (N0.getOpcode() == ISD::SETCC) {
13532 ISD::CondCode CC = cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get();
13533 if (SDValue R = foldSelectOfSelectCmp(LHS: N0.getOperand(i: 0), RHS: N0.getOperand(i: 1),
13534 CC, TrueV: N1, FalseV: N2, DL, VT, DAG, Flags))
13535 return R;
13536 }
13537
13538 // If we can fold this based on the true/false value, do so.
13539 if (SimplifySelectOps(SELECT: N, LHS: N1, RHS: N2))
13540 return SDValue(N, 0); // Don't revisit N.
13541
13542 if (VT0 == MVT::i1) {
13543 // The code in this block deals with the following 2 equivalences:
13544 // select(C0|C1, x, y) <=> select(C0, x, select(C1, x, y))
13545 // select(C0&C1, x, y) <=> select(C0, select(C1, x, y), y)
13546 // The target can specify its preferred form with the
13547 // shouldNormalizeToSelectSequence() callback. However we always transform
13548 // to the right anyway if we find the inner select exists in the DAG anyway
13549 // and we always transform to the left side if we know that we can further
13550 // optimize the combination of the conditions.
13551 bool normalizeToSequence =
13552 TLI.shouldNormalizeToSelectSequence(Context&: *DAG.getContext(), VT, CCVT: VT0);
13553 // select (and Cond0, Cond1), X, Y
13554 // -> select Cond0, (select Cond1, X, Y), Y
13555 if (N0->getOpcode() == ISD::AND && N0->hasOneUse()) {
13556 SDValue Cond0 = N0->getOperand(Num: 0);
13557 SDValue Cond1 = N0->getOperand(Num: 1);
13558 SDValue InnerSelect =
13559 DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(), N1: Cond1, N2: N1, N3: N2, Flags);
13560 if (normalizeToSequence || !InnerSelect.use_empty())
13561 return DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(), N1: Cond0,
13562 N2: InnerSelect, N3: N2, Flags);
13563 // Cleanup on failure.
13564 if (InnerSelect.use_empty())
13565 recursivelyDeleteUnusedNodes(N: InnerSelect.getNode());
13566 }
13567 // select (or Cond0, Cond1), X, Y -> select Cond0, X, (select Cond1, X, Y)
13568 if (N0->getOpcode() == ISD::OR && N0->hasOneUse()) {
13569 SDValue Cond0 = N0->getOperand(Num: 0);
13570 SDValue Cond1 = N0->getOperand(Num: 1);
13571 SDValue InnerSelect = DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(),
13572 N1: Cond1, N2: N1, N3: N2, Flags);
13573 if (normalizeToSequence || !InnerSelect.use_empty())
13574 return DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(), N1: Cond0, N2: N1,
13575 N3: InnerSelect, Flags);
13576 // Cleanup on failure.
13577 if (InnerSelect.use_empty())
13578 recursivelyDeleteUnusedNodes(N: InnerSelect.getNode());
13579 }
13580
13581 // select Cond0, (select Cond1, X, Y), Y -> select (and Cond0, Cond1), X, Y
13582 if (N1->getOpcode() == ISD::SELECT && N1->hasOneUse()) {
13583 SDValue N1_0 = N1->getOperand(Num: 0);
13584 SDValue N1_1 = N1->getOperand(Num: 1);
13585 SDValue N1_2 = N1->getOperand(Num: 2);
13586 if (N1_2 == N2 && N0.getValueType() == N1_0.getValueType()) {
13587 // Create the actual and node if we can generate good code for it.
13588 if (!normalizeToSequence) {
13589 SDValue And = DAG.getNode(Opcode: ISD::AND, DL, VT: N0.getValueType(), N1: N0, N2: N1_0);
13590 return DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(), N1: And, N2: N1_1,
13591 N3: N2, Flags);
13592 }
13593 // Otherwise see if we can optimize the "and" to a better pattern.
13594 if (SDValue Combined = visitANDLike(N0, N1: N1_0, N)) {
13595 return DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(), N1: Combined, N2: N1_1,
13596 N3: N2, Flags);
13597 }
13598 }
13599 }
13600 // select Cond0, X, (select Cond1, X, Y) -> select (or Cond0, Cond1), X, Y
13601 if (N2->getOpcode() == ISD::SELECT && N2->hasOneUse()) {
13602 SDValue N2_0 = N2->getOperand(Num: 0);
13603 SDValue N2_1 = N2->getOperand(Num: 1);
13604 SDValue N2_2 = N2->getOperand(Num: 2);
13605 if (N2_1 == N1 && N0.getValueType() == N2_0.getValueType()) {
13606 // Create the actual or node if we can generate good code for it.
13607 if (!normalizeToSequence) {
13608 SDValue Or = DAG.getNode(Opcode: ISD::OR, DL, VT: N0.getValueType(), N1: N0, N2: N2_0);
13609 return DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(), N1: Or, N2: N1,
13610 N3: N2_2, Flags);
13611 }
13612 // Otherwise see if we can optimize to a better pattern.
13613 if (SDValue Combined = visitORLike(N0, N1: N2_0, DL))
13614 return DAG.getNode(Opcode: ISD::SELECT, DL, VT: N1.getValueType(), N1: Combined, N2: N1,
13615 N3: N2_2, Flags);
13616 }
13617 }
13618
13619 // select usubo(x, y).overflow, (sub y, x), (usubo x, y) -> abdu(x, y)
13620 if (N0.getOpcode() == ISD::USUBO && N0.getResNo() == 1 &&
13621 N2.getNode() == N0.getNode() && N2.getResNo() == 0 &&
13622 N1.getOpcode() == ISD::SUB && N2.getOperand(i: 0) == N1.getOperand(i: 1) &&
13623 N2.getOperand(i: 1) == N1.getOperand(i: 0) &&
13624 (!LegalOperations || TLI.isOperationLegal(Op: ISD::ABDU, VT)))
13625 return DAG.getNode(Opcode: ISD::ABDU, DL, VT, N1: N0.getOperand(i: 0), N2: N0.getOperand(i: 1));
13626
13627 // select usubo(x, y).overflow, (usubo x, y), (sub y, x) -> neg (abdu x, y)
13628 if (N0.getOpcode() == ISD::USUBO && N0.getResNo() == 1 &&
13629 N1.getNode() == N0.getNode() && N1.getResNo() == 0 &&
13630 N2.getOpcode() == ISD::SUB && N2.getOperand(i: 0) == N1.getOperand(i: 1) &&
13631 N2.getOperand(i: 1) == N1.getOperand(i: 0) &&
13632 (!LegalOperations || TLI.isOperationLegal(Op: ISD::ABDU, VT)))
13633 return DAG.getNegative(
13634 Val: DAG.getNode(Opcode: ISD::ABDU, DL, VT, N1: N0.getOperand(i: 0), N2: N0.getOperand(i: 1)),
13635 DL, VT);
13636 }
13637
13638 // Fold selects based on a setcc into other things, such as min/max/abs.
13639 if (N0.getOpcode() == ISD::SETCC) {
13640 SDValue Cond0 = N0.getOperand(i: 0), Cond1 = N0.getOperand(i: 1);
13641 ISD::CondCode CC = cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get();
13642
13643 // select (fcmp lt x, y), x, y -> fminnum x, y
13644 // select (fcmp gt x, y), x, y -> fmaxnum x, y
13645 //
13646 // This is OK if we don't care what happens if either operand is a NaN.
13647 if (N0.hasOneUse() &&
13648 isLegalToCombineMinNumMaxNum(DAG, LHS: N1, RHS: N2, SelectFlags: Flags, CmpFlags: N0->getFlags(), TLI))
13649 if (SDValue FMinMax =
13650 combineMinNumMaxNum(DL, VT, LHS: Cond0, RHS: Cond1, True: N1, False: N2, CC))
13651 return FMinMax;
13652
13653 // Use 'unsigned add with overflow' to optimize an unsigned saturating add.
13654 // This is conservatively limited to pre-legal-operations to give targets
13655 // a chance to reverse the transform if they want to do that. Also, it is
13656 // unlikely that the pattern would be formed late, so it's probably not
13657 // worth going through the other checks.
13658 if (!LegalOperations && TLI.isOperationLegalOrCustom(Op: ISD::UADDO, VT) &&
13659 CC == ISD::SETUGT && N0.hasOneUse() && isAllOnesConstant(V: N1) &&
13660 N2.getOpcode() == ISD::ADD && Cond0 == N2.getOperand(i: 0)) {
13661 auto *C = dyn_cast<ConstantSDNode>(Val: N2.getOperand(i: 1));
13662 auto *NotC = dyn_cast<ConstantSDNode>(Val&: Cond1);
13663 if (C && NotC && C->getAPIntValue() == ~NotC->getAPIntValue()) {
13664 // select (setcc Cond0, ~C, ugt), -1, (add Cond0, C) -->
13665 // uaddo Cond0, C; select uaddo.1, -1, uaddo.0
13666 //
13667 // The IR equivalent of this transform would have this form:
13668 // %a = add %x, C
13669 // %c = icmp ugt %x, ~C
13670 // %r = select %c, -1, %a
13671 // =>
13672 // %u = call {iN,i1} llvm.uadd.with.overflow(%x, C)
13673 // %u0 = extractvalue %u, 0
13674 // %u1 = extractvalue %u, 1
13675 // %r = select %u1, -1, %u0
13676 SDVTList VTs = DAG.getVTList(VT1: VT, VT2: VT0);
13677 SDValue UAO = DAG.getNode(Opcode: ISD::UADDO, DL, VTList: VTs, N1: Cond0, N2: N2.getOperand(i: 1));
13678 return DAG.getSelect(DL, VT, Cond: UAO.getValue(R: 1), LHS: N1, RHS: UAO.getValue(R: 0));
13679 }
13680 }
13681
13682 if (SDValue S = performNanGuardFpToSatCombine(N, DAG))
13683 return S;
13684
13685 if (TLI.isOperationLegal(Op: ISD::SELECT_CC, VT) ||
13686 (!LegalOperations &&
13687 TLI.isOperationLegalOrCustom(Op: ISD::SELECT_CC, VT))) {
13688 // Any flags available in a select/setcc fold will be on the setcc as they
13689 // migrated from fcmp
13690 return DAG.getNode(Opcode: ISD::SELECT_CC, DL, VT, N1: Cond0, N2: Cond1, N3: N1, N4: N2,
13691 N5: N0.getOperand(i: 2), Flags: N0->getFlags());
13692 }
13693
13694 if (SDValue ABD = foldSelectToABD(LHS: Cond0, RHS: Cond1, True: N1, False: N2, CC, DL))
13695 return ABD;
13696
13697 if (SDValue NewSel = SimplifySelect(DL, N0, N1, N2))
13698 return NewSel;
13699
13700 // (select (ugt x, C), (add x, ~C), x) -> (umin (add x, ~C), x)
13701 // (select (ult x, C), x, (add x, -C)) -> (umin x, (add x, -C))
13702 if (SDValue UMin = foldSelectToUMin(LHS: Cond0, RHS: Cond1, True: N1, False: N2, CC, DL))
13703 return UMin;
13704 }
13705
13706 if (!VT.isVector())
13707 if (SDValue BinOp = foldSelectOfBinops(N))
13708 return BinOp;
13709
13710 if (SDValue R = combineSelectAsExtAnd(Cond: N0, T: N1, F: N2, DL, DAG))
13711 return R;
13712
13713 if (SDValue R = combineSelectToPseudoMinMax(DAG, N))
13714 return R;
13715
13716 return SDValue();
13717}
13718
13719// This function assumes all the vselect's arguments are CONCAT_VECTOR
13720// nodes and that the condition is a BV of ConstantSDNodes (or undefs).
13721static SDValue ConvertSelectToConcatVector(SDNode *N, SelectionDAG &DAG) {
13722 SDLoc DL(N);
13723 SDValue Cond = N->getOperand(Num: 0);
13724 SDValue LHS = N->getOperand(Num: 1);
13725 SDValue RHS = N->getOperand(Num: 2);
13726 EVT VT = N->getValueType(ResNo: 0);
13727 int NumElems = VT.getVectorNumElements();
13728 assert(LHS.getOpcode() == ISD::CONCAT_VECTORS &&
13729 RHS.getOpcode() == ISD::CONCAT_VECTORS &&
13730 Cond.getOpcode() == ISD::BUILD_VECTOR);
13731
13732 // CONCAT_VECTOR can take an arbitrary number of arguments. We only care about
13733 // binary ones here.
13734 if (LHS->getNumOperands() != 2 || RHS->getNumOperands() != 2)
13735 return SDValue();
13736
13737 // We're sure we have an even number of elements due to the
13738 // concat_vectors we have as arguments to vselect.
13739 // Skip BV elements until we find one that's not an UNDEF
13740 // After we find an UNDEF element, keep looping until we get to half the
13741 // length of the BV and see if all the non-undef nodes are the same.
13742 ConstantSDNode *BottomHalf = nullptr;
13743 for (int i = 0; i < NumElems / 2; ++i) {
13744 if (Cond->getOperand(Num: i)->isUndef())
13745 continue;
13746
13747 if (BottomHalf == nullptr)
13748 BottomHalf = cast<ConstantSDNode>(Val: Cond.getOperand(i));
13749 else if (Cond->getOperand(Num: i).getNode() != BottomHalf)
13750 return SDValue();
13751 }
13752
13753 // Do the same for the second half of the BuildVector
13754 ConstantSDNode *TopHalf = nullptr;
13755 for (int i = NumElems / 2; i < NumElems; ++i) {
13756 if (Cond->getOperand(Num: i)->isUndef())
13757 continue;
13758
13759 if (TopHalf == nullptr)
13760 TopHalf = cast<ConstantSDNode>(Val: Cond.getOperand(i));
13761 else if (Cond->getOperand(Num: i).getNode() != TopHalf)
13762 return SDValue();
13763 }
13764
13765 assert(TopHalf && BottomHalf &&
13766 "One half of the selector was all UNDEFs and the other was all the "
13767 "same value. This should have been addressed before this function.");
13768 return DAG.getNode(
13769 Opcode: ISD::CONCAT_VECTORS, DL, VT,
13770 N1: BottomHalf->isZero() ? RHS->getOperand(Num: 0) : LHS->getOperand(Num: 0),
13771 N2: TopHalf->isZero() ? RHS->getOperand(Num: 1) : LHS->getOperand(Num: 1));
13772}
13773
13774bool refineUniformBase(SDValue &BasePtr, SDValue &Index, bool IndexIsScaled,
13775 SelectionDAG &DAG, const SDLoc &DL) {
13776
13777 // Only perform the transformation when existing operands can be reused.
13778 if (IndexIsScaled)
13779 return false;
13780
13781 if (!isNullConstant(V: BasePtr) && !Index.hasOneUse())
13782 return false;
13783
13784 EVT VT = BasePtr.getValueType();
13785
13786 if (SDValue SplatVal = DAG.getSplatValue(V: Index);
13787 SplatVal && !isNullConstant(V: SplatVal) &&
13788 SplatVal.getValueType() == VT) {
13789 BasePtr = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: BasePtr, N2: SplatVal);
13790 Index = DAG.getSplat(VT: Index.getValueType(), DL, Op: DAG.getConstant(Val: 0, DL, VT));
13791 return true;
13792 }
13793
13794 if (Index.getOpcode() != ISD::ADD)
13795 return false;
13796
13797 if (SDValue SplatVal = DAG.getSplatValue(V: Index.getOperand(i: 0));
13798 SplatVal && SplatVal.getValueType() == VT) {
13799 BasePtr = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: BasePtr, N2: SplatVal);
13800 Index = Index.getOperand(i: 1);
13801 return true;
13802 }
13803 if (SDValue SplatVal = DAG.getSplatValue(V: Index.getOperand(i: 1));
13804 SplatVal && SplatVal.getValueType() == VT) {
13805 BasePtr = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: BasePtr, N2: SplatVal);
13806 Index = Index.getOperand(i: 0);
13807 return true;
13808 }
13809 return false;
13810}
13811
13812// Fold sext/zext of index into index type.
13813bool refineIndexType(SDValue &Index, ISD::MemIndexType &IndexType, EVT DataVT,
13814 SelectionDAG &DAG) {
13815 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
13816
13817 // It's always safe to look through zero extends.
13818 if (Index.getOpcode() == ISD::ZERO_EXTEND) {
13819 if (TLI.shouldRemoveExtendFromGSIndex(Extend: Index, DataVT)) {
13820 IndexType = ISD::UNSIGNED_SCALED;
13821 Index = Index.getOperand(i: 0);
13822 return true;
13823 }
13824 if (ISD::isIndexTypeSigned(IndexType)) {
13825 IndexType = ISD::UNSIGNED_SCALED;
13826 return true;
13827 }
13828 }
13829
13830 // It's only safe to look through sign extends when Index is signed.
13831 if (Index.getOpcode() == ISD::SIGN_EXTEND &&
13832 ISD::isIndexTypeSigned(IndexType) &&
13833 TLI.shouldRemoveExtendFromGSIndex(Extend: Index, DataVT)) {
13834 Index = Index.getOperand(i: 0);
13835 return true;
13836 }
13837
13838 return false;
13839}
13840
13841SDValue DAGCombiner::visitVPSCATTER(SDNode *N) {
13842 VPScatterSDNode *MSC = cast<VPScatterSDNode>(Val: N);
13843 SDValue Mask = MSC->getMask();
13844 SDValue Chain = MSC->getChain();
13845 SDValue Index = MSC->getIndex();
13846 SDValue Scale = MSC->getScale();
13847 SDValue StoreVal = MSC->getValue();
13848 SDValue BasePtr = MSC->getBasePtr();
13849 SDValue VL = MSC->getVectorLength();
13850 ISD::MemIndexType IndexType = MSC->getIndexType();
13851 SDLoc DL(N);
13852
13853 // Zap scatters with a zero mask.
13854 if (ISD::isConstantSplatVectorAllZeros(N: Mask.getNode()))
13855 return Chain;
13856
13857 if (refineUniformBase(BasePtr, Index, IndexIsScaled: MSC->isIndexScaled(), DAG, DL)) {
13858 SDValue Ops[] = {Chain, StoreVal, BasePtr, Index, Scale, Mask, VL};
13859 return DAG.getScatterVP(VTs: DAG.getVTList(VT: MVT::Other), VT: MSC->getMemoryVT(),
13860 dl: DL, Ops, MMO: MSC->getMemOperand(), IndexType);
13861 }
13862
13863 if (refineIndexType(Index, IndexType, DataVT: StoreVal.getValueType(), DAG)) {
13864 SDValue Ops[] = {Chain, StoreVal, BasePtr, Index, Scale, Mask, VL};
13865 return DAG.getScatterVP(VTs: DAG.getVTList(VT: MVT::Other), VT: MSC->getMemoryVT(),
13866 dl: DL, Ops, MMO: MSC->getMemOperand(), IndexType);
13867 }
13868
13869 return SDValue();
13870}
13871
13872SDValue DAGCombiner::visitMSCATTER(SDNode *N) {
13873 MaskedScatterSDNode *MSC = cast<MaskedScatterSDNode>(Val: N);
13874 SDValue Mask = MSC->getMask();
13875 SDValue Chain = MSC->getChain();
13876 SDValue Index = MSC->getIndex();
13877 SDValue Scale = MSC->getScale();
13878 SDValue StoreVal = MSC->getValue();
13879 SDValue BasePtr = MSC->getBasePtr();
13880 ISD::MemIndexType IndexType = MSC->getIndexType();
13881 SDLoc DL(N);
13882
13883 // Zap scatters with a zero mask.
13884 if (ISD::isConstantSplatVectorAllZeros(N: Mask.getNode()))
13885 return Chain;
13886
13887 if (refineUniformBase(BasePtr, Index, IndexIsScaled: MSC->isIndexScaled(), DAG, DL)) {
13888 SDValue Ops[] = {Chain, StoreVal, Mask, BasePtr, Index, Scale};
13889 return DAG.getMaskedScatter(VTs: DAG.getVTList(VT: MVT::Other), MemVT: MSC->getMemoryVT(),
13890 dl: DL, Ops, MMO: MSC->getMemOperand(), IndexType,
13891 IsTruncating: MSC->isTruncatingStore());
13892 }
13893
13894 if (refineIndexType(Index, IndexType, DataVT: StoreVal.getValueType(), DAG)) {
13895 SDValue Ops[] = {Chain, StoreVal, Mask, BasePtr, Index, Scale};
13896 return DAG.getMaskedScatter(VTs: DAG.getVTList(VT: MVT::Other), MemVT: MSC->getMemoryVT(),
13897 dl: DL, Ops, MMO: MSC->getMemOperand(), IndexType,
13898 IsTruncating: MSC->isTruncatingStore());
13899 }
13900
13901 return SDValue();
13902}
13903
13904/// Check if Mask defines a known constant set of enabled lanes, where only the
13905/// first N lanes are enabled. N is returned if so.
13906static uint64_t calculateConstantLowMaskLanes(SDValue Mask) {
13907 // We expect masks for masked load/store to be i1 predicates.
13908 if (Mask.getValueType().getScalarSizeInBits() != 1)
13909 return 0;
13910
13911 if (Mask.getOpcode() == ISD::BUILD_VECTOR) {
13912 unsigned NumOnes = 0;
13913 auto Op = Mask->op_begin();
13914 for (; Op != Mask->op_end(); ++Op) {
13915 if (!isOneConstant(V: *Op))
13916 break;
13917 ++NumOnes;
13918 }
13919
13920 if (NumOnes == 0)
13921 return 0;
13922
13923 for (; Op != Mask->op_end(); ++Op)
13924 if (!isNullConstant(V: *Op))
13925 return 0;
13926 return NumOnes;
13927 }
13928
13929 if (Mask.getOpcode() == ISD::GET_ACTIVE_LANE_MASK &&
13930 isNullConstant(V: Mask.getOperand(i: 0)) &&
13931 isa<ConstantSDNode>(Val: Mask.getOperand(i: 1)) &&
13932 Mask.getOperand(i: 1).getValueType().getSizeInBits() <= 64)
13933 return Mask.getConstantOperandVal(i: 1);
13934
13935 return 0;
13936}
13937
13938SDValue DAGCombiner::visitMSTORE(SDNode *N) {
13939 MaskedStoreSDNode *MST = cast<MaskedStoreSDNode>(Val: N);
13940 SDValue Mask = MST->getMask();
13941 SDValue Chain = MST->getChain();
13942 SDValue Value = MST->getValue();
13943 SDValue Ptr = MST->getBasePtr();
13944 EVT VT = Value.getValueType();
13945
13946 // Zap masked stores with a zero mask.
13947 if (ISD::isConstantSplatVectorAllZeros(N: Mask.getNode()))
13948 return Chain;
13949
13950 // Remove a masked store if base pointers and masks are equal.
13951 if (MaskedStoreSDNode *MST1 = dyn_cast<MaskedStoreSDNode>(Val&: Chain)) {
13952 if (MST->isUnindexed() && MST->isSimple() && MST1->isUnindexed() &&
13953 MST1->isSimple() && MST1->getBasePtr() == Ptr &&
13954 !MST->getBasePtr().isUndef() &&
13955 ((Mask == MST1->getMask() && MST->getMemoryVT().getStoreSize() ==
13956 MST1->getMemoryVT().getStoreSize()) ||
13957 ISD::isConstantSplatVectorAllOnes(N: Mask.getNode())) &&
13958 TypeSize::isKnownLE(LHS: MST1->getMemoryVT().getStoreSize(),
13959 RHS: MST->getMemoryVT().getStoreSize())) {
13960 CombineTo(N: MST1, Res: MST1->getChain());
13961 if (N->getOpcode() != ISD::DELETED_NODE)
13962 AddToWorklist(N);
13963 return SDValue(N, 0);
13964 }
13965 }
13966
13967 // If this is a masked store with an all ones mask, we can use a unmasked
13968 // store.
13969 // FIXME: Can we do this for indexed, compressing, or truncating stores?
13970 if (MST->isUnindexed() && !MST->isCompressingStore() &&
13971 !MST->isTruncatingStore()) {
13972 if (ISD::isConstantSplatVectorAllOnes(N: Mask.getNode()))
13973 return DAG.getStore(Chain: MST->getChain(), dl: SDLoc(N), Val: MST->getValue(),
13974 Ptr: MST->getBasePtr(), PtrInfo: MST->getPointerInfo(),
13975 Alignment: MST->getBaseAlign(), MMOFlags: MST->getMemOperand()->getFlags(),
13976 Metadata: MST->getAAInfo());
13977
13978 // Convert a masked_store with constant getactivelanemask input mask to a
13979 // standard store.
13980 if (uint64_t Lanes = calculateConstantLowMaskLanes(Mask)) {
13981 if (Lanes < VT.getVectorMinNumElements() && isPowerOf2_32(Value: Lanes)) {
13982 EVT SubVT = EVT::getVectorVT(Context&: *DAG.getContext(),
13983 VT: VT.getVectorElementType(), NumElements: Lanes);
13984 unsigned IsFast = 0;
13985 if ((!LegalTypes || isTypeLegal(VT: SubVT)) &&
13986 TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(),
13987 VT: SubVT, AddrSpace: MST->getAddressSpace(),
13988 Alignment: MST->getBaseAlign(),
13989 Flags: MST->getMemOperand()->getFlags(), Fast: &IsFast) &&
13990 IsFast) {
13991 SDLoc DL(N);
13992 SDValue Ext = DAG.getExtractSubvector(DL, VT: SubVT, Vec: MST->getValue(), Idx: 0);
13993 return DAG.getStore(Chain: MST->getChain(), dl: DL, Val: Ext, Ptr: MST->getBasePtr(),
13994 PtrInfo: MST->getPointerInfo(), Alignment: MST->getBaseAlign(),
13995 MMOFlags: MST->getMemOperand()->getFlags(),
13996 Metadata: MST->getAAInfo());
13997 }
13998 }
13999 }
14000 }
14001
14002 // Try transforming N to an indexed store.
14003 if (CombineToPreIndexedLoadStore(N) || CombineToPostIndexedLoadStore(N))
14004 return SDValue(N, 0);
14005
14006 if (MST->isTruncatingStore() && MST->isUnindexed() && VT.isInteger() &&
14007 (!isa<ConstantSDNode>(Val: Value) ||
14008 !cast<ConstantSDNode>(Val&: Value)->isOpaque())) {
14009 APInt TruncDemandedBits =
14010 APInt::getLowBitsSet(numBits: Value.getScalarValueSizeInBits(),
14011 loBitsSet: MST->getMemoryVT().getScalarSizeInBits());
14012
14013 // See if we can simplify the operation with
14014 // SimplifyDemandedBits, which only works if the value has a single use.
14015 if (SimplifyDemandedBits(Op: Value, DemandedBits: TruncDemandedBits)) {
14016 // Re-visit the store if anything changed and the store hasn't been merged
14017 // with another node (N is deleted) SimplifyDemandedBits will add Value's
14018 // node back to the worklist if necessary, but we also need to re-visit
14019 // the Store node itself.
14020 if (N->getOpcode() != ISD::DELETED_NODE)
14021 AddToWorklist(N);
14022 return SDValue(N, 0);
14023 }
14024 }
14025
14026 // If this is a TRUNC followed by a masked store, fold this into a masked
14027 // truncating store. We can do this even if this is already a masked
14028 // truncstore.
14029 // TODO: Try combine to masked compress store if possiable.
14030 if ((Value.getOpcode() == ISD::TRUNCATE) && Value->hasOneUse() &&
14031 MST->isUnindexed() && !MST->isCompressingStore() &&
14032 TLI.canCombineTruncStore(ValVT: Value.getOperand(i: 0).getValueType(),
14033 MemVT: MST->getMemoryVT(), Alignment: MST->getAlign(),
14034 AddrSpace: MST->getAddressSpace(), LegalOnly: LegalOperations)) {
14035 auto Mask = TLI.promoteTargetBoolean(DAG, Bool: MST->getMask(),
14036 ValVT: Value.getOperand(i: 0).getValueType());
14037 return DAG.getMaskedStore(Chain, dl: SDLoc(N), Val: Value.getOperand(i: 0), Base: Ptr,
14038 Offset: MST->getOffset(), Mask, MemVT: MST->getMemoryVT(),
14039 MMO: MST->getMemOperand(), AM: MST->getAddressingMode(),
14040 /*IsTruncating=*/true);
14041 }
14042
14043 return SDValue();
14044}
14045
14046SDValue DAGCombiner::visitVP_STRIDED_STORE(SDNode *N) {
14047 auto *SST = cast<VPStridedStoreSDNode>(Val: N);
14048 EVT EltVT = SST->getValue().getValueType().getVectorElementType();
14049 // Combine strided stores with unit-stride to a regular VP store.
14050 if (auto *CStride = dyn_cast<ConstantSDNode>(Val: SST->getStride());
14051 CStride && CStride->getZExtValue() == EltVT.getStoreSize()) {
14052 return DAG.getStoreVP(Chain: SST->getChain(), dl: SDLoc(N), Val: SST->getValue(),
14053 Ptr: SST->getBasePtr(), Offset: SST->getOffset(), Mask: SST->getMask(),
14054 EVL: SST->getVectorLength(), MemVT: SST->getMemoryVT(),
14055 MMO: SST->getMemOperand(), AM: SST->getAddressingMode(),
14056 IsTruncating: SST->isTruncatingStore(), IsCompressing: SST->isCompressingStore());
14057 }
14058 return SDValue();
14059}
14060
14061SDValue DAGCombiner::visitVECTOR_COMPRESS(SDNode *N) {
14062 SDLoc DL(N);
14063 SDValue Vec = N->getOperand(Num: 0);
14064 SDValue Mask = N->getOperand(Num: 1);
14065 SDValue Passthru = N->getOperand(Num: 2);
14066 EVT VecVT = Vec.getValueType();
14067
14068 bool HasPassthru = !Passthru.isUndef();
14069
14070 APInt SplatVal;
14071 if (ISD::isConstantSplatVector(N: Mask.getNode(), SplatValue&: SplatVal))
14072 return TLI.isConstTrueVal(N: Mask) ? Vec : Passthru;
14073
14074 if (Vec.isUndef() || Mask.isUndef())
14075 return Passthru;
14076
14077 // No need for potentially expensive compress if the mask is constant.
14078 if (ISD::isBuildVectorOfConstantSDNodes(N: Mask.getNode())) {
14079 SmallVector<SDValue, 16> Ops;
14080 EVT ScalarVT = VecVT.getVectorElementType();
14081 unsigned NumSelected = 0;
14082 unsigned NumElmts = VecVT.getVectorNumElements();
14083 for (unsigned I = 0; I < NumElmts; ++I) {
14084 SDValue MaskI = Mask.getOperand(i: I);
14085 // We treat undef mask entries as "false".
14086 if (MaskI.isUndef())
14087 continue;
14088
14089 if (TLI.isConstTrueVal(N: MaskI)) {
14090 SDValue VecI = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: ScalarVT, N1: Vec,
14091 N2: DAG.getVectorIdxConstant(Val: I, DL));
14092 Ops.push_back(Elt: VecI);
14093 NumSelected++;
14094 }
14095 }
14096 for (unsigned Rest = NumSelected; Rest < NumElmts; ++Rest) {
14097 SDValue Val =
14098 HasPassthru
14099 ? DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: ScalarVT, N1: Passthru,
14100 N2: DAG.getVectorIdxConstant(Val: Rest, DL))
14101 : DAG.getUNDEF(VT: ScalarVT);
14102 Ops.push_back(Elt: Val);
14103 }
14104 return DAG.getBuildVector(VT: VecVT, DL, Ops);
14105 }
14106
14107 return SDValue();
14108}
14109
14110SDValue DAGCombiner::visitVPGATHER(SDNode *N) {
14111 VPGatherSDNode *MGT = cast<VPGatherSDNode>(Val: N);
14112 SDValue Mask = MGT->getMask();
14113 SDValue Chain = MGT->getChain();
14114 SDValue Index = MGT->getIndex();
14115 SDValue Scale = MGT->getScale();
14116 SDValue BasePtr = MGT->getBasePtr();
14117 SDValue VL = MGT->getVectorLength();
14118 ISD::MemIndexType IndexType = MGT->getIndexType();
14119 SDLoc DL(N);
14120
14121 if (refineUniformBase(BasePtr, Index, IndexIsScaled: MGT->isIndexScaled(), DAG, DL)) {
14122 SDValue Ops[] = {Chain, BasePtr, Index, Scale, Mask, VL};
14123 return DAG.getGatherVP(
14124 VTs: DAG.getVTList(VT1: N->getValueType(ResNo: 0), VT2: MVT::Other), VT: MGT->getMemoryVT(), dl: DL,
14125 Ops, MMO: MGT->getMemOperand(), IndexType);
14126 }
14127
14128 if (refineIndexType(Index, IndexType, DataVT: N->getValueType(ResNo: 0), DAG)) {
14129 SDValue Ops[] = {Chain, BasePtr, Index, Scale, Mask, VL};
14130 return DAG.getGatherVP(
14131 VTs: DAG.getVTList(VT1: N->getValueType(ResNo: 0), VT2: MVT::Other), VT: MGT->getMemoryVT(), dl: DL,
14132 Ops, MMO: MGT->getMemOperand(), IndexType);
14133 }
14134
14135 return SDValue();
14136}
14137
14138SDValue DAGCombiner::visitMGATHER(SDNode *N) {
14139 MaskedGatherSDNode *MGT = cast<MaskedGatherSDNode>(Val: N);
14140 SDValue Mask = MGT->getMask();
14141 SDValue Chain = MGT->getChain();
14142 SDValue Index = MGT->getIndex();
14143 SDValue Scale = MGT->getScale();
14144 SDValue PassThru = MGT->getPassThru();
14145 SDValue BasePtr = MGT->getBasePtr();
14146 ISD::MemIndexType IndexType = MGT->getIndexType();
14147 SDLoc DL(N);
14148
14149 // Zap gathers with a zero mask.
14150 if (ISD::isConstantSplatVectorAllZeros(N: Mask.getNode()))
14151 return CombineTo(N, Res0: PassThru, Res1: MGT->getChain());
14152
14153 if (refineUniformBase(BasePtr, Index, IndexIsScaled: MGT->isIndexScaled(), DAG, DL)) {
14154 SDValue Ops[] = {Chain, PassThru, Mask, BasePtr, Index, Scale};
14155 return DAG.getMaskedGather(
14156 VTs: DAG.getVTList(VT1: N->getValueType(ResNo: 0), VT2: MVT::Other), MemVT: MGT->getMemoryVT(), dl: DL,
14157 Ops, MMO: MGT->getMemOperand(), IndexType, ExtTy: MGT->getExtensionType());
14158 }
14159
14160 if (refineIndexType(Index, IndexType, DataVT: N->getValueType(ResNo: 0), DAG)) {
14161 SDValue Ops[] = {Chain, PassThru, Mask, BasePtr, Index, Scale};
14162 return DAG.getMaskedGather(
14163 VTs: DAG.getVTList(VT1: N->getValueType(ResNo: 0), VT2: MVT::Other), MemVT: MGT->getMemoryVT(), dl: DL,
14164 Ops, MMO: MGT->getMemOperand(), IndexType, ExtTy: MGT->getExtensionType());
14165 }
14166
14167 return SDValue();
14168}
14169
14170SDValue DAGCombiner::visitMLOAD(SDNode *N) {
14171 MaskedLoadSDNode *MLD = cast<MaskedLoadSDNode>(Val: N);
14172 SDValue Mask = MLD->getMask();
14173
14174 // Zap masked loads with a zero mask.
14175 if (ISD::isConstantSplatVectorAllZeros(N: Mask.getNode()))
14176 return CombineTo(N, Res0: MLD->getPassThru(), Res1: MLD->getChain());
14177
14178 // If this is a masked load with an all ones mask, we can use a unmasked load.
14179 // FIXME: Can we do this for indexed, expanding, or extending loads?
14180 if (ISD::isConstantSplatVectorAllOnes(N: Mask.getNode()) && MLD->isUnindexed() &&
14181 !MLD->isExpandingLoad() && MLD->getExtensionType() == ISD::NON_EXTLOAD) {
14182 SDValue NewLd =
14183 DAG.getLoad(VT: N->getValueType(ResNo: 0), dl: SDLoc(N), Chain: MLD->getChain(),
14184 Ptr: MLD->getBasePtr(), PtrInfo: MLD->getPointerInfo(),
14185 Alignment: MLD->getBaseAlign(), MMOFlags: MLD->getMemOperand()->getFlags(),
14186 Metadata: MMOMetadata(MLD->getAAInfo(), MLD->getRanges()));
14187 return CombineTo(N, Res0: NewLd, Res1: NewLd.getValue(R: 1));
14188 }
14189
14190 // Try transforming N to an indexed load.
14191 if (CombineToPreIndexedLoadStore(N) || CombineToPostIndexedLoadStore(N))
14192 return SDValue(N, 0);
14193
14194 return SDValue();
14195}
14196
14197SDValue DAGCombiner::visitMHISTOGRAM(SDNode *N) {
14198 MaskedHistogramSDNode *HG = cast<MaskedHistogramSDNode>(Val: N);
14199 SDValue Chain = HG->getChain();
14200 SDValue Inc = HG->getInc();
14201 SDValue Mask = HG->getMask();
14202 SDValue BasePtr = HG->getBasePtr();
14203 SDValue Index = HG->getIndex();
14204 SDLoc DL(HG);
14205
14206 EVT MemVT = HG->getMemoryVT();
14207 EVT DataVT = Index.getValueType();
14208 MachineMemOperand *MMO = HG->getMemOperand();
14209 ISD::MemIndexType IndexType = HG->getIndexType();
14210
14211 if (ISD::isConstantSplatVectorAllZeros(N: Mask.getNode()))
14212 return Chain;
14213
14214 if (refineUniformBase(BasePtr, Index, IndexIsScaled: HG->isIndexScaled(), DAG, DL) ||
14215 refineIndexType(Index, IndexType, DataVT, DAG)) {
14216 SDValue Ops[] = {Chain, Inc, Mask, BasePtr, Index,
14217 HG->getScale(), HG->getIntID()};
14218 return DAG.getMaskedHistogram(VTs: DAG.getVTList(VT: MVT::Other), MemVT, dl: DL, Ops,
14219 MMO, IndexType);
14220 }
14221
14222 return SDValue();
14223}
14224
14225SDValue DAGCombiner::visitPARTIAL_REDUCE_MLA(SDNode *N) {
14226 if (SDValue Res = foldPartialReduceMLAMulOp(N))
14227 return Res;
14228 if (SDValue Res = foldPartialReduceAdd(N))
14229 return Res;
14230 return SDValue();
14231}
14232
14233// partial_reduce_*mla(acc, mul(*ext(a), *ext(b)), splat(1))
14234// -> partial_reduce_*mla(acc, a, b)
14235//
14236// partial_reduce_*mla(acc, mul(*ext(x), splat(C)), splat(1))
14237// -> partial_reduce_*mla(acc, x, splat(C))
14238//
14239// partial_reduce_*mla(acc, sel(p, mul(*ext(a), *ext(b)), splat(0)), splat(1))
14240// -> partial_reduce_*mla(acc, sel(p, a, splat(0)), b)
14241//
14242// partial_reduce_*mla(acc, sel(p, mul(*ext(a), splat(C)), splat(0)), splat(1))
14243// -> partial_reduce_*mla(acc, sel(p, a, splat(0)), splat(C))
14244//
14245// `sel` could either be VSELECT or VP_MERGE.
14246SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
14247 SDLoc DL(N);
14248 auto *Context = DAG.getContext();
14249 SDValue Tmp;
14250 SDValue Acc = N->getOperand(Num: 0);
14251 SDValue Op1 = N->getOperand(Num: 1);
14252 SDValue OrigOp1 = Op1;
14253 SDValue Op2 = N->getOperand(Num: 2);
14254 unsigned Opc = Op1->getOpcode();
14255
14256 // Handle predication by moving the VSELECT / VP_MERGE into the operand of the
14257 // MUL.
14258 SDValue Pred;
14259 if ((Opc == ISD::VSELECT || Opc == ISD::VP_MERGE) &&
14260 (isZeroOrZeroSplat(N: Op1->getOperand(Num: 2)) ||
14261 isZeroOrZeroSplatFP(N: Op1->getOperand(Num: 2)))) {
14262 Pred = Op1->getOperand(Num: 0);
14263 Op1 = Op1->getOperand(Num: 1);
14264 Opc = Op1->getOpcode();
14265 }
14266
14267 // Handle negation (sub-reduction).
14268 bool IsMLS = false;
14269 if (sd_match(N: Op1, P: m_Neg(V: m_Value(N&: Tmp)))) {
14270 Op1 = Tmp;
14271 Opc = Op1->getOpcode();
14272 IsMLS = true;
14273 }
14274
14275 if (Opc != ISD::MUL && Opc != ISD::FMUL && Opc != ISD::SHL)
14276 return SDValue();
14277
14278 SDValue LHS = Op1->getOperand(Num: 0);
14279 SDValue RHS = Op1->getOperand(Num: 1);
14280
14281 // After instcombine, negation for FP operations is on the RHS, so implement:
14282 // fmul(fpext(a), fneg(fpext(b)))
14283 //-> fmul(fpext(a), fpext(fneg(b)))
14284 if (sd_match(N: RHS, P: m_FNeg(Op: m_Value(N&: Tmp)))) {
14285 RHS = Tmp;
14286 IsMLS = true;
14287 }
14288
14289 // Try to treat (shl %a, %c) as (mul %a, (1 << %c)) for constant %c.
14290 if (Opc == ISD::SHL) {
14291 APInt C;
14292 if (!ISD::isConstantSplatVector(N: RHS.getNode(), SplatValue&: C))
14293 return SDValue();
14294
14295 RHS =
14296 DAG.getSplatVector(VT: RHS.getValueType(), DL,
14297 Op: DAG.getConstant(Val: APInt(C.getBitWidth(), 1).shl(ShiftAmt: C), DL,
14298 VT: RHS.getValueType().getScalarType()));
14299 Opc = ISD::MUL;
14300 }
14301
14302 if (!(Opc == ISD::MUL && llvm::isOneOrOneSplat(V: Op2)) &&
14303 !(Opc == ISD::FMUL && llvm::isOneOrOneSplatFP(V: Op2)))
14304 return SDValue();
14305
14306 auto IsIntOrFPExtOpcode = [](unsigned int Opcode) {
14307 return (ISD::isExtOpcode(Opcode) || Opcode == ISD::FP_EXTEND);
14308 };
14309
14310 unsigned LHSOpcode = LHS->getOpcode();
14311 if (!IsIntOrFPExtOpcode(LHSOpcode))
14312 return SDValue();
14313
14314 SDValue LHSExtOp = LHS->getOperand(Num: 0);
14315 EVT LHSExtOpVT = LHSExtOp.getValueType();
14316
14317 // When Pred is non-zero, set Op = select(Pred, Op, splat(0)) and freeze
14318 // OtherOp to keep the same semantics when moving the selects into the MUL
14319 // operands.
14320 auto ApplyPredicate = [&](SDValue &Op, SDValue &OtherOp) {
14321 if (Pred) {
14322 EVT OpVT = Op.getValueType();
14323 SDValue Zero = OpVT.isFloatingPoint() ? DAG.getConstantFP(Val: 0.0, DL, VT: OpVT)
14324 : DAG.getConstant(Val: 0, DL, VT: OpVT);
14325 if (OrigOp1.getOpcode() == ISD::VP_MERGE)
14326 Op = DAG.getNode(Opcode: ISD::VP_MERGE, DL, VT: OpVT, N1: Pred, N2: Op, N3: Zero,
14327 N4: OrigOp1.getOperand(i: 3));
14328 else
14329 Op = DAG.getSelect(DL, VT: OpVT, Cond: Pred, LHS: Op, RHS: Zero);
14330 OtherOp = DAG.getFreeze(V: OtherOp);
14331 }
14332 };
14333
14334 // Generate an MLA or MLS.
14335 auto GetMLA = [&](unsigned Opc, SDValue Acc, SDValue LHS,
14336 SDValue RHS) -> SDValue {
14337 EVT AccVT = Acc.getValueType();
14338 return IsMLS ? DAG.getPartialReduceMLS(Opc, DL, Acc, LHS, RHS)
14339 : DAG.getNode(Opcode: Opc, DL, VT: AccVT, N1: Acc, N2: LHS, N3: RHS);
14340 };
14341
14342 // partial_reduce_*mla(acc, mul(ext(x), splat(C)), splat(1))
14343 // -> partial_reduce_*mla(acc, x, C)
14344 APInt C;
14345 if (ISD::isConstantSplatVector(N: RHS.getNode(), SplatValue&: C)) {
14346 // TODO: Make use of partial_reduce_sumla here
14347 APInt CTrunc = C.trunc(width: LHSExtOpVT.getScalarSizeInBits());
14348 unsigned LHSBits = LHS.getValueType().getScalarSizeInBits();
14349 if ((LHSOpcode != ISD::ZERO_EXTEND || CTrunc.zext(width: LHSBits) != C) &&
14350 (LHSOpcode != ISD::SIGN_EXTEND || CTrunc.sext(width: LHSBits) != C))
14351 return SDValue();
14352
14353 unsigned NewOpcode = LHSOpcode == ISD::SIGN_EXTEND
14354 ? ISD::PARTIAL_REDUCE_SMLA
14355 : ISD::PARTIAL_REDUCE_UMLA;
14356
14357 // Only perform these combines if the target supports folding
14358 // the extends into the operation.
14359 if (!TLI.isPartialReduceMLALegalOrCustom(
14360 Opc: NewOpcode, AccVT: TLI.getTypeToTransformTo(Context&: *Context, VT: N->getValueType(ResNo: 0)),
14361 InputVT: TLI.getTypeToTransformTo(Context&: *Context, VT: LHSExtOpVT)))
14362 return SDValue();
14363
14364 SDValue C = DAG.getConstant(Val: CTrunc, DL, VT: LHSExtOpVT);
14365 ApplyPredicate(C, LHSExtOp);
14366 return GetMLA(NewOpcode, Acc, LHSExtOp, C);
14367 }
14368
14369 unsigned RHSOpcode = RHS->getOpcode();
14370 if (!IsIntOrFPExtOpcode(RHSOpcode))
14371 return SDValue();
14372
14373 SDValue RHSExtOp = RHS->getOperand(Num: 0);
14374 if (LHSExtOpVT != RHSExtOp.getValueType())
14375 return SDValue();
14376
14377 unsigned NewOpc;
14378 if (LHSOpcode == ISD::SIGN_EXTEND && RHSOpcode == ISD::SIGN_EXTEND)
14379 NewOpc = ISD::PARTIAL_REDUCE_SMLA;
14380 else if (LHSOpcode == ISD::ZERO_EXTEND && RHSOpcode == ISD::ZERO_EXTEND)
14381 NewOpc = ISD::PARTIAL_REDUCE_UMLA;
14382 else if (LHSOpcode == ISD::SIGN_EXTEND && RHSOpcode == ISD::ZERO_EXTEND)
14383 NewOpc = ISD::PARTIAL_REDUCE_SUMLA;
14384 else if (LHSOpcode == ISD::ZERO_EXTEND && RHSOpcode == ISD::SIGN_EXTEND) {
14385 NewOpc = ISD::PARTIAL_REDUCE_SUMLA;
14386 std::swap(a&: LHSExtOp, b&: RHSExtOp);
14387 } else if (LHSOpcode == ISD::FP_EXTEND && RHSOpcode == ISD::FP_EXTEND) {
14388 NewOpc = ISD::PARTIAL_REDUCE_FMLA;
14389 } else
14390 return SDValue();
14391 // For a 2-stage extend the signedness of both of the extends must match
14392 // If the mul has the same type, there is no outer extend, and thus we
14393 // can simply use the inner extends to pick the result node.
14394 // TODO: extend to handle nonneg zext as sext
14395 EVT AccElemVT = Acc.getValueType().getVectorElementType();
14396 if (Op1.getValueType().getVectorElementType() != AccElemVT &&
14397 NewOpc != N->getOpcode())
14398 return SDValue();
14399
14400 // Only perform these combines if the target supports folding
14401 // the extends into the operation.
14402 if (!TLI.isPartialReduceMLALegalOrCustom(
14403 Opc: NewOpc, AccVT: TLI.getTypeToTransformTo(Context&: *Context, VT: N->getValueType(ResNo: 0)),
14404 InputVT: TLI.getTypeToTransformTo(Context&: *Context, VT: LHSExtOpVT)))
14405 return SDValue();
14406
14407 ApplyPredicate(RHSExtOp, LHSExtOp);
14408 return GetMLA(NewOpc, Acc, LHSExtOp, RHSExtOp);
14409}
14410
14411// partial.reduce.*mla(acc, *ext(op), splat(1))
14412// -> partial.reduce.*mla(acc, op, splat(trunc(1)))
14413// partial.reduce.sumla(acc, sext(op), splat(1))
14414// -> partial.reduce.smla(acc, op, splat(trunc(1)))
14415//
14416// partial.reduce.*mla(acc, sel(p, *ext(op), splat(0)), splat(1))
14417// -> partial.reduce.*mla(acc, sel(p, op, splat(0)), splat(trunc(1)))
14418SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) {
14419 SDLoc DL(N);
14420 SDValue Tmp;
14421 SDValue Acc = N->getOperand(Num: 0);
14422 SDValue Op1 = N->getOperand(Num: 1);
14423 SDValue Op2 = N->getOperand(Num: 2);
14424
14425 if (!llvm::isOneOrOneSplat(V: Op2) && !llvm::isOneOrOneSplatFP(V: Op2))
14426 return SDValue();
14427
14428 SDValue Pred;
14429 unsigned Op1Opcode = Op1.getOpcode();
14430 if (Op1Opcode == ISD::VSELECT && (isZeroOrZeroSplat(N: Op1->getOperand(Num: 2)) ||
14431 isZeroOrZeroSplatFP(N: Op1->getOperand(Num: 2)))) {
14432 Pred = Op1->getOperand(Num: 0);
14433 Op1 = Op1->getOperand(Num: 1);
14434 Op1Opcode = Op1->getOpcode();
14435 }
14436
14437 // Handle negation (sub-reduction).
14438 bool IsMLS = false;
14439 if (sd_match(N: Op1, P: m_AnyOf(preds: m_Neg(V: m_Value(N&: Tmp)), preds: m_FNeg(Op: m_Value(N&: Tmp))))) {
14440 Op1 = Tmp;
14441 Op1Opcode = Op1.getOpcode();
14442 IsMLS = true;
14443 }
14444
14445 if (!ISD::isExtOpcode(Opcode: Op1Opcode) && Op1Opcode != ISD::FP_EXTEND)
14446 return SDValue();
14447
14448 bool Op1IsSigned =
14449 Op1Opcode == ISD::SIGN_EXTEND || Op1Opcode == ISD::FP_EXTEND;
14450 bool NodeIsSigned = N->getOpcode() != ISD::PARTIAL_REDUCE_UMLA;
14451 EVT AccElemVT = Acc.getValueType().getVectorElementType();
14452 if (Op1IsSigned != NodeIsSigned &&
14453 Op1.getValueType().getVectorElementType() != AccElemVT)
14454 return SDValue();
14455
14456 unsigned NewOpcode = N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA
14457 ? ISD::PARTIAL_REDUCE_FMLA
14458 : Op1IsSigned ? ISD::PARTIAL_REDUCE_SMLA
14459 : ISD::PARTIAL_REDUCE_UMLA;
14460
14461 SDValue UnextOp1 = Op1.getOperand(i: 0);
14462 EVT UnextOp1VT = UnextOp1.getValueType();
14463 auto *Context = DAG.getContext();
14464 EVT PromOp1VT = TLI.getTypeToTransformTo(Context&: *Context, VT: UnextOp1VT);
14465 if (!TLI.isPartialReduceMLALegalOrCustom(
14466 Opc: NewOpcode, AccVT: TLI.getTypeToTransformTo(Context&: *Context, VT: N->getValueType(ResNo: 0)),
14467 InputVT: PromOp1VT))
14468 return SDValue();
14469
14470 // The multiplier below is built at the operand type, where a splat of 1 in i1
14471 // sign extends to -1. Extend i1 masks to the promoted type first.
14472 if (Op1IsSigned && UnextOp1VT.getVectorElementType() == MVT::i1) {
14473 if (PromOp1VT == UnextOp1VT)
14474 return SDValue();
14475 UnextOp1VT = PromOp1VT;
14476 UnextOp1 = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT: UnextOp1VT, Operand: UnextOp1);
14477 }
14478
14479 SDValue Constant = N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA
14480 ? DAG.getConstantFP(Val: 1, DL, VT: UnextOp1VT)
14481 : DAG.getConstant(Val: 1, DL, VT: UnextOp1VT);
14482
14483 if (Pred) {
14484 SDValue Zero = N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA
14485 ? DAG.getConstantFP(Val: 0, DL, VT: UnextOp1VT)
14486 : DAG.getConstant(Val: 0, DL, VT: UnextOp1VT);
14487 Constant = DAG.getSelect(DL, VT: UnextOp1VT, Cond: Pred, LHS: Constant, RHS: Zero);
14488 }
14489 EVT AccVT = Acc.getValueType();
14490 return IsMLS ? DAG.getPartialReduceMLS(Opc: NewOpcode, DL, Acc, LHS: UnextOp1, RHS: Constant)
14491 : DAG.getNode(Opcode: NewOpcode, DL, VT: AccVT, N1: Acc, N2: UnextOp1, N3: Constant);
14492}
14493
14494SDValue DAGCombiner::visitLOOP_DEPENDENCE_MASK(SDNode *N) {
14495 SDLoc DL(N);
14496 EVT VT = N->getValueType(ResNo: 0);
14497 unsigned LaneOffset = N->getConstantOperandVal(Num: 3);
14498
14499 // The first lane is always active, so v1i1 => true.
14500 if (LaneOffset == 0 &&
14501 VT.getVectorElementCount() == ElementCount::getFixed(MinVal: 1))
14502 return DAG.getBoolConstant(V: true, DL, VT, OpVT: VT);
14503
14504 return SDValue();
14505}
14506
14507SDValue DAGCombiner::visitVP_STRIDED_LOAD(SDNode *N) {
14508 auto *SLD = cast<VPStridedLoadSDNode>(Val: N);
14509 EVT EltVT = SLD->getValueType(ResNo: 0).getVectorElementType();
14510 // Combine strided loads with unit-stride to a regular VP load.
14511 if (auto *CStride = dyn_cast<ConstantSDNode>(Val: SLD->getStride());
14512 CStride && CStride->getZExtValue() == EltVT.getStoreSize()) {
14513 SDValue NewLd = DAG.getLoadVP(
14514 AM: SLD->getAddressingMode(), ExtType: SLD->getExtensionType(), VT: SLD->getValueType(ResNo: 0),
14515 dl: SDLoc(N), Chain: SLD->getChain(), Ptr: SLD->getBasePtr(), Offset: SLD->getOffset(),
14516 Mask: SLD->getMask(), EVL: SLD->getVectorLength(), MemVT: SLD->getMemoryVT(),
14517 MMO: SLD->getMemOperand(), IsExpanding: SLD->isExpandingLoad());
14518 return CombineTo(N, Res0: NewLd, Res1: NewLd.getValue(R: 1));
14519 }
14520 return SDValue();
14521}
14522
14523/// A vector select of 2 constant vectors can be simplified to math/logic to
14524/// avoid a variable select instruction and possibly avoid constant loads.
14525SDValue DAGCombiner::foldVSelectOfConstants(SDNode *N) {
14526 SDValue Cond = N->getOperand(Num: 0);
14527 SDValue N1 = N->getOperand(Num: 1);
14528 SDValue N2 = N->getOperand(Num: 2);
14529 EVT VT = N->getValueType(ResNo: 0);
14530 if (!Cond.hasOneUse() || Cond.getScalarValueSizeInBits() != 1 ||
14531 !shouldConvertSelectOfConstantsToMath(Cond, VT, TLI) ||
14532 !ISD::isBuildVectorOfConstantSDNodes(N: N1.getNode()) ||
14533 !ISD::isBuildVectorOfConstantSDNodes(N: N2.getNode()))
14534 return SDValue();
14535
14536 // Check if we can use the condition value to increment/decrement a single
14537 // constant value. This simplifies a select to an add and removes a constant
14538 // load/materialization from the general case.
14539 bool AllAddOne = true;
14540 bool AllSubOne = true;
14541 unsigned Elts = VT.getVectorNumElements();
14542 for (unsigned i = 0; i != Elts; ++i) {
14543 SDValue N1Elt = N1.getOperand(i);
14544 SDValue N2Elt = N2.getOperand(i);
14545 if (N1Elt.isUndef())
14546 continue;
14547 // N2 should not contain undef values since it will be reused in the fold.
14548 if (N2Elt.isUndef() || N1Elt.getValueType() != N2Elt.getValueType()) {
14549 AllAddOne = false;
14550 AllSubOne = false;
14551 break;
14552 }
14553
14554 const APInt &C1 = N1Elt->getAsAPIntVal();
14555 const APInt &C2 = N2Elt->getAsAPIntVal();
14556 if (C1 != C2 + 1)
14557 AllAddOne = false;
14558 if (C1 != C2 - 1)
14559 AllSubOne = false;
14560 }
14561
14562 // Further simplifications for the extra-special cases where the constants are
14563 // all 0 or all -1 should be implemented as folds of these patterns.
14564 SDLoc DL(N);
14565 if (AllAddOne || AllSubOne) {
14566 // vselect <N x i1> Cond, C+1, C --> add (zext Cond), C
14567 // vselect <N x i1> Cond, C-1, C --> add (sext Cond), C
14568 auto ExtendOpcode = AllAddOne ? ISD::ZERO_EXTEND : ISD::SIGN_EXTEND;
14569 SDValue ExtendedCond = DAG.getNode(Opcode: ExtendOpcode, DL, VT, Operand: Cond);
14570 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: ExtendedCond, N2);
14571 }
14572
14573 // select Cond, Pow2C, 0 --> (zext Cond) << log2(Pow2C)
14574 APInt Pow2C;
14575 if (ISD::isConstantSplatVector(N: N1.getNode(), SplatValue&: Pow2C) && Pow2C.isPowerOf2() &&
14576 isNullOrNullSplat(V: N2)) {
14577 SDValue ZextCond = DAG.getZExtOrTrunc(Op: Cond, DL, VT);
14578 SDValue ShAmtC = DAG.getConstant(Val: Pow2C.exactLogBase2(), DL, VT);
14579 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: ZextCond, N2: ShAmtC);
14580 }
14581
14582 if (SDValue V = foldSelectOfConstantsUsingSra(N, DL, DAG))
14583 return V;
14584
14585 // The general case for select-of-constants:
14586 // vselect <N x i1> Cond, C1, C2 --> xor (and (sext Cond), (C1^C2)), C2
14587 // ...but that only makes sense if a vselect is slower than 2 logic ops, so
14588 // leave that to a machine-specific pass.
14589 return SDValue();
14590}
14591
14592static SDValue combineVSelectWithAllOnesOrZeros(SDValue Cond, SDValue TVal,
14593 SDValue FVal,
14594 const TargetLowering &TLI,
14595 SelectionDAG &DAG,
14596 const SDLoc &DL) {
14597 EVT VT = TVal.getValueType();
14598 if (!TLI.isTypeLegal(VT))
14599 return SDValue();
14600
14601 EVT CondVT = Cond.getValueType();
14602 assert(CondVT.isVector() && "Vector select expects a vector selector!");
14603
14604 bool IsTAllZero = ISD::isConstantSplatVectorAllZeros(N: TVal.getNode());
14605 bool IsTAllOne = ISD::isConstantSplatVectorAllOnes(N: TVal.getNode());
14606 bool IsFAllZero = ISD::isConstantSplatVectorAllZeros(N: FVal.getNode());
14607 bool IsFAllOne = ISD::isConstantSplatVectorAllOnes(N: FVal.getNode());
14608
14609 // no vselect(cond, 0/-1, X) or vselect(cond, X, 0/-1), return
14610 if (!IsTAllZero && !IsTAllOne && !IsFAllZero && !IsFAllOne)
14611 return SDValue();
14612
14613 // select Cond, 0, 0 → 0
14614 if (IsTAllZero && IsFAllZero) {
14615 return VT.isFloatingPoint() ? DAG.getConstantFP(Val: 0.0, DL, VT)
14616 : DAG.getConstant(Val: 0, DL, VT);
14617 }
14618
14619 // check select(setgt lhs, -1), 1, -1 --> or (sra lhs, bitwidth - 1), 1
14620 APInt TValAPInt;
14621 if (Cond.getOpcode() == ISD::SETCC &&
14622 Cond.getOperand(i: 2) == DAG.getCondCode(Cond: ISD::SETGT) &&
14623 Cond.getOperand(i: 0).getValueType() == VT && VT.isSimple() &&
14624 ISD::isConstantSplatVector(N: TVal.getNode(), SplatValue&: TValAPInt) &&
14625 TValAPInt.isOne() &&
14626 ISD::isConstantSplatVectorAllOnes(N: Cond.getOperand(i: 1).getNode()) &&
14627 ISD::isConstantSplatVectorAllOnes(N: FVal.getNode()) &&
14628 !TLI.shouldAvoidTransformToShift(VT, Amount: VT.getScalarSizeInBits() - 1)) {
14629 SDValue LHS = Cond.getOperand(i: 0);
14630 SDValue ShiftC =
14631 DAG.getShiftAmountConstant(Val: VT.getScalarSizeInBits() - 1, VT, DL);
14632 SDValue Shift = DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: LHS, N2: ShiftC);
14633 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: Shift, N2: TVal);
14634 }
14635
14636 // To use the condition operand as a bitwise mask, it must have elements that
14637 // are the same size as the select elements. i.e, the condition operand must
14638 // have already been promoted from the IR select condition type <N x i1>.
14639 // Don't check if the types themselves are equal because that excludes
14640 // vector floating-point selects.
14641 if (CondVT.getScalarSizeInBits() != VT.getScalarSizeInBits())
14642 return SDValue();
14643
14644 // Cond value must be 'sign splat' to be converted to a logical op.
14645 if (DAG.ComputeNumSignBits(Op: Cond) != CondVT.getScalarSizeInBits())
14646 return SDValue();
14647
14648 // Try inverting Cond and swapping T/F if it gives all-ones/all-zeros form
14649 if (!IsTAllOne && !IsFAllZero && Cond.hasOneUse() &&
14650 Cond.getOpcode() == ISD::SETCC &&
14651 TLI.getSetCCResultType(DL: DAG.getDataLayout(), Context&: *DAG.getContext(), VT) ==
14652 CondVT) {
14653 if (IsTAllZero || IsFAllOne) {
14654 SDValue CC = Cond.getOperand(i: 2);
14655 ISD::CondCode InverseCC = ISD::getSetCCInverse(
14656 Operation: cast<CondCodeSDNode>(Val&: CC)->get(), Type: Cond.getOperand(i: 0).getValueType());
14657 Cond = DAG.getSetCC(DL, VT: CondVT, LHS: Cond.getOperand(i: 0), RHS: Cond.getOperand(i: 1),
14658 Cond: InverseCC);
14659 std::swap(a&: TVal, b&: FVal);
14660 std::swap(a&: IsTAllOne, b&: IsFAllOne);
14661 std::swap(a&: IsTAllZero, b&: IsFAllZero);
14662 }
14663 }
14664
14665 assert(DAG.ComputeNumSignBits(Cond) == CondVT.getScalarSizeInBits() &&
14666 "Select condition no longer all-sign bits");
14667
14668 // select Cond, -1, 0 → bitcast Cond
14669 if (IsTAllOne && IsFAllZero)
14670 return DAG.getBitcast(VT, V: Cond);
14671
14672 // select Cond, -1, x → or Cond, x
14673 if (IsTAllOne) {
14674 SDValue X = DAG.getBitcast(VT: CondVT, V: DAG.getFreeze(V: FVal));
14675 SDValue Or = DAG.getNode(Opcode: ISD::OR, DL, VT: CondVT, N1: Cond, N2: X);
14676 return DAG.getBitcast(VT, V: Or);
14677 }
14678
14679 // select Cond, x, 0 → and Cond, x
14680 if (IsFAllZero) {
14681 SDValue X = DAG.getBitcast(VT: CondVT, V: DAG.getFreeze(V: TVal));
14682 SDValue And = DAG.getNode(Opcode: ISD::AND, DL, VT: CondVT, N1: Cond, N2: X);
14683 return DAG.getBitcast(VT, V: And);
14684 }
14685
14686 // select Cond, 0, x -> and not(Cond), x
14687 if (IsTAllZero &&
14688 (isBitwiseNot(V: peekThroughBitcasts(V: Cond)) || TLI.hasAndNot(X: Cond))) {
14689 SDValue X = DAG.getBitcast(VT: CondVT, V: DAG.getFreeze(V: FVal));
14690 SDValue And =
14691 DAG.getNode(Opcode: ISD::AND, DL, VT: CondVT, N1: DAG.getNOT(DL, Val: Cond, VT: CondVT), N2: X);
14692 return DAG.getBitcast(VT, V: And);
14693 }
14694
14695 return SDValue();
14696}
14697
14698SDValue DAGCombiner::visitVSELECT(SDNode *N) {
14699 SDValue N0 = N->getOperand(Num: 0);
14700 SDValue N1 = N->getOperand(Num: 1);
14701 SDValue N2 = N->getOperand(Num: 2);
14702 EVT VT = N->getValueType(ResNo: 0);
14703 SDLoc DL(N);
14704
14705 if (SDValue V = DAG.simplifySelect(Cond: N0, TVal: N1, FVal: N2))
14706 return V;
14707
14708 if (SDValue V = foldBoolSelectToLogic(N, DL, DAG))
14709 return V;
14710
14711 // vselect (not Cond), N1, N2 -> vselect Cond, N2, N1
14712 if (!TLI.isTargetCanonicalSelect(N))
14713 if (SDValue F = extractBooleanFlip(V: N0, DAG, TLI, Force: false))
14714 return DAG.getSelect(DL, VT, Cond: F, LHS: N2, RHS: N1, Flags: N->getFlags());
14715
14716 // select (sext m), (add X, C), X --> (add X, (and C, (sext m))))
14717 if (N1.getOpcode() == ISD::ADD && N1.getOperand(i: 0) == N2 && N1->hasOneUse() &&
14718 DAG.isConstantIntBuildVectorOrConstantInt(N: N1.getOperand(i: 1)) &&
14719 N0.getScalarValueSizeInBits() == N1.getScalarValueSizeInBits() &&
14720 TLI.getBooleanContents(Type: N0.getValueType()) ==
14721 TargetLowering::ZeroOrNegativeOneBooleanContent) {
14722 return DAG.getNode(
14723 Opcode: ISD::ADD, DL, VT: N1.getValueType(), N1: N2,
14724 N2: DAG.getNode(Opcode: ISD::AND, DL, VT: N0.getValueType(), N1: N1.getOperand(i: 1), N2: N0));
14725 }
14726
14727 // Canonicalize integer abs.
14728 // vselect (setg[te] X, 0), X, -X ->
14729 // vselect (setgt X, -1), X, -X ->
14730 // vselect (setl[te] X, 0), -X, X ->
14731 // Y = sra (X, size(X)-1); xor (add (X, Y), Y)
14732 if (N0.getOpcode() == ISD::SETCC) {
14733 SDValue LHS = N0.getOperand(i: 0), RHS = N0.getOperand(i: 1);
14734 ISD::CondCode CC = cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get();
14735 bool isAbs = false;
14736 bool RHSIsAllZeros = ISD::isBuildVectorAllZeros(N: RHS.getNode());
14737
14738 if (((RHSIsAllZeros && (CC == ISD::SETGT || CC == ISD::SETGE)) ||
14739 (ISD::isBuildVectorAllOnes(N: RHS.getNode()) && CC == ISD::SETGT)) &&
14740 N1 == LHS && N2.getOpcode() == ISD::SUB && N1 == N2.getOperand(i: 1))
14741 isAbs = ISD::isBuildVectorAllZeros(N: N2.getOperand(i: 0).getNode());
14742 else if ((RHSIsAllZeros && (CC == ISD::SETLT || CC == ISD::SETLE)) &&
14743 N2 == LHS && N1.getOpcode() == ISD::SUB && N2 == N1.getOperand(i: 1))
14744 isAbs = ISD::isBuildVectorAllZeros(N: N1.getOperand(i: 0).getNode());
14745
14746 if (isAbs) {
14747 if (TLI.isOperationLegalOrCustom(Op: ISD::ABS, VT))
14748 return DAG.getNode(Opcode: ISD::ABS, DL, VT, Operand: LHS);
14749
14750 SDValue Shift = DAG.getNode(
14751 Opcode: ISD::SRA, DL, VT, N1: LHS,
14752 N2: DAG.getShiftAmountConstant(Val: VT.getScalarSizeInBits() - 1, VT, DL));
14753 SDValue Add = DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: LHS, N2: Shift);
14754 AddToWorklist(N: Shift.getNode());
14755 AddToWorklist(N: Add.getNode());
14756 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: Add, N2: Shift);
14757 }
14758
14759 // vselect x, y (fcmp lt x, y) -> fminnum x, y
14760 // vselect x, y (fcmp gt x, y) -> fmaxnum x, y
14761 //
14762 // This is OK if we don't care about what happens if either operand is a
14763 // NaN.
14764 //
14765 if (N0.hasOneUse() &&
14766 isLegalToCombineMinNumMaxNum(DAG, LHS, RHS, SelectFlags: N->getFlags(),
14767 CmpFlags: N0->getFlags(), TLI)) {
14768 if (SDValue FMinMax = combineMinNumMaxNum(DL, VT, LHS, RHS, True: N1, False: N2, CC))
14769 return FMinMax;
14770 }
14771
14772 if (SDValue S = PerformMinMaxFpToSatCombine(N0: LHS, N1: RHS, N2: N1, N3: N2, CC, DAG))
14773 return S;
14774 if (SDValue S = PerformUMinFpToSatCombine(N0: LHS, N1: RHS, N2: N1, N3: N2, CC, DAG))
14775 return S;
14776 if (SDValue S = performNanGuardFpToSatCombine(N, DAG))
14777 return S;
14778
14779 // If this select has a condition (setcc) with narrower operands than the
14780 // select, try to widen the compare to match the select width.
14781 // TODO: This should be extended to handle any constant.
14782 // TODO: This could be extended to handle non-loading patterns, but that
14783 // requires thorough testing to avoid regressions.
14784 if (isNullOrNullSplat(V: RHS)) {
14785 EVT NarrowVT = LHS.getValueType();
14786 EVT WideVT = N1.getValueType().changeVectorElementTypeToInteger();
14787 EVT SetCCVT = getSetCCResultType(VT: LHS.getValueType());
14788 unsigned SetCCWidth = SetCCVT.getScalarSizeInBits();
14789 unsigned WideWidth = WideVT.getScalarSizeInBits();
14790 bool IsSigned = isSignedIntSetCC(Code: CC);
14791 auto LoadExtOpcode = IsSigned ? ISD::SEXTLOAD : ISD::ZEXTLOAD;
14792 if (LHS.getOpcode() == ISD::LOAD && LHS.hasOneUse() && SetCCWidth != 1 &&
14793 SetCCWidth < WideWidth &&
14794 TLI.isOperationLegalOrCustom(Op: ISD::SETCC, VT: WideVT)) {
14795 LoadSDNode *Ld = cast<LoadSDNode>(Val&: LHS);
14796
14797 if (TLI.isLoadLegalOrCustom(ValVT: WideVT, MemVT: NarrowVT, Alignment: Ld->getAlign(),
14798 AddrSpace: Ld->getAddressSpace(), ExtType: LoadExtOpcode,
14799 Atomic: false)) {
14800 // Both compare operands can be widened for free. The LHS can use an
14801 // extended load, and the RHS is a constant:
14802 // vselect (ext (setcc load(X), C)), N1, N2 -->
14803 // vselect (setcc extload(X), C'), N1, N2
14804 auto ExtOpcode = IsSigned ? ISD::SIGN_EXTEND : ISD::ZERO_EXTEND;
14805 SDValue WideLHS = DAG.getNode(Opcode: ExtOpcode, DL, VT: WideVT, Operand: LHS);
14806 SDValue WideRHS = DAG.getNode(Opcode: ExtOpcode, DL, VT: WideVT, Operand: RHS);
14807 EVT WideSetCCVT = getSetCCResultType(VT: WideVT);
14808 SDValue WideSetCC =
14809 DAG.getSetCC(DL, VT: WideSetCCVT, LHS: WideLHS, RHS: WideRHS, Cond: CC);
14810 return DAG.getSelect(DL, VT: N1.getValueType(), Cond: WideSetCC, LHS: N1, RHS: N2);
14811 }
14812 }
14813 }
14814
14815 if (SDValue ABD = foldSelectToABD(LHS, RHS, True: N1, False: N2, CC, DL))
14816 return ABD;
14817
14818 // Match VSELECTs into add with unsigned saturation.
14819 if (hasOperation(Opcode: ISD::UADDSAT, VT)) {
14820 // Check if one of the arms of the VSELECT is vector with all bits set.
14821 // If it's on the left side invert the predicate to simplify logic below.
14822 SDValue Other;
14823 ISD::CondCode SatCC = CC;
14824 if (ISD::isConstantSplatVectorAllOnes(N: N1.getNode())) {
14825 Other = N2;
14826 SatCC = ISD::getSetCCInverse(Operation: SatCC, Type: VT.getScalarType());
14827 } else if (ISD::isConstantSplatVectorAllOnes(N: N2.getNode())) {
14828 Other = N1;
14829 }
14830
14831 if (Other && Other.getOpcode() == ISD::ADD) {
14832 SDValue CondLHS = LHS, CondRHS = RHS;
14833 SDValue OpLHS = Other.getOperand(i: 0), OpRHS = Other.getOperand(i: 1);
14834
14835 // Canonicalize condition operands.
14836 if (SatCC == ISD::SETUGE) {
14837 std::swap(a&: CondLHS, b&: CondRHS);
14838 SatCC = ISD::SETULE;
14839 }
14840
14841 // We can test against either of the addition operands.
14842 // x <= x+y ? x+y : ~0 --> uaddsat x, y
14843 // x+y >= x ? x+y : ~0 --> uaddsat x, y
14844 if (SatCC == ISD::SETULE && Other == CondRHS &&
14845 (OpLHS == CondLHS || OpRHS == CondLHS))
14846 return DAG.getNode(Opcode: ISD::UADDSAT, DL, VT, N1: OpLHS, N2: OpRHS);
14847
14848 if (OpRHS.getOpcode() == CondRHS.getOpcode() &&
14849 (OpRHS.getOpcode() == ISD::BUILD_VECTOR ||
14850 OpRHS.getOpcode() == ISD::SPLAT_VECTOR) &&
14851 CondLHS == OpLHS) {
14852 // If the RHS is a constant we have to reverse the const
14853 // canonicalization.
14854 // x >= ~C ? x+C : ~0 --> uaddsat x, C
14855 auto MatchUADDSAT = [](ConstantSDNode *Op, ConstantSDNode *Cond) {
14856 return Cond->getAPIntValue() == ~Op->getAPIntValue();
14857 };
14858 if (SatCC == ISD::SETULE &&
14859 ISD::matchBinaryPredicate(LHS: OpRHS, RHS: CondRHS, Match: MatchUADDSAT))
14860 return DAG.getNode(Opcode: ISD::UADDSAT, DL, VT, N1: OpLHS, N2: OpRHS);
14861 }
14862 }
14863 }
14864
14865 // Match VSELECTs into sub with unsigned saturation.
14866 if (hasOperation(Opcode: ISD::USUBSAT, VT)) {
14867 // Check if one of the arms of the VSELECT is a zero vector. If it's on
14868 // the left side invert the predicate to simplify logic below.
14869 SDValue Other;
14870 ISD::CondCode SatCC = CC;
14871 if (ISD::isConstantSplatVectorAllZeros(N: N1.getNode())) {
14872 Other = N2;
14873 SatCC = ISD::getSetCCInverse(Operation: SatCC, Type: VT.getScalarType());
14874 } else if (ISD::isConstantSplatVectorAllZeros(N: N2.getNode())) {
14875 Other = N1;
14876 }
14877
14878 // zext(x) >= y ? trunc(zext(x) - y) : 0
14879 // --> usubsat(trunc(zext(x)),trunc(umin(y,SatLimit)))
14880 // zext(x) > y ? trunc(zext(x) - y) : 0
14881 // --> usubsat(trunc(zext(x)),trunc(umin(y,SatLimit)))
14882 if (Other && Other.getOpcode() == ISD::TRUNCATE &&
14883 Other.getOperand(i: 0).getOpcode() == ISD::SUB &&
14884 (SatCC == ISD::SETUGE || SatCC == ISD::SETUGT)) {
14885 SDValue OpLHS = Other.getOperand(i: 0).getOperand(i: 0);
14886 SDValue OpRHS = Other.getOperand(i: 0).getOperand(i: 1);
14887 if (LHS == OpLHS && RHS == OpRHS && LHS.getOpcode() == ISD::ZERO_EXTEND)
14888 if (SDValue R = getTruncatedUSUBSAT(DstVT: VT, SrcVT: LHS.getValueType(), LHS, RHS,
14889 DAG, DL))
14890 return R;
14891 }
14892
14893 if (Other && Other.getNumOperands() == 2) {
14894 SDValue CondRHS = RHS;
14895 SDValue OpLHS = Other.getOperand(i: 0), OpRHS = Other.getOperand(i: 1);
14896
14897 if (OpLHS == LHS) {
14898 // Look for a general sub with unsigned saturation first.
14899 // x >= y ? x-y : 0 --> usubsat x, y
14900 // x > y ? x-y : 0 --> usubsat x, y
14901 if ((SatCC == ISD::SETUGE || SatCC == ISD::SETUGT) &&
14902 Other.getOpcode() == ISD::SUB && OpRHS == CondRHS)
14903 return DAG.getNode(Opcode: ISD::USUBSAT, DL, VT, N1: OpLHS, N2: OpRHS);
14904
14905 if (OpRHS.getOpcode() == ISD::BUILD_VECTOR ||
14906 OpRHS.getOpcode() == ISD::SPLAT_VECTOR) {
14907 if (CondRHS.getOpcode() == ISD::BUILD_VECTOR ||
14908 CondRHS.getOpcode() == ISD::SPLAT_VECTOR) {
14909 // If the RHS is a constant we have to reverse the const
14910 // canonicalization.
14911 // x > C-1 ? x+-C : 0 --> usubsat x, C
14912 auto MatchUSUBSAT = [](ConstantSDNode *Op, ConstantSDNode *Cond) {
14913 return (!Op && !Cond) ||
14914 (Op && Cond &&
14915 Cond->getAPIntValue() == (-Op->getAPIntValue() - 1));
14916 };
14917 if (SatCC == ISD::SETUGT && Other.getOpcode() == ISD::ADD &&
14918 ISD::matchBinaryPredicate(LHS: OpRHS, RHS: CondRHS, Match: MatchUSUBSAT,
14919 /*AllowUndefs*/ true)) {
14920 OpRHS = DAG.getNegative(Val: OpRHS, DL, VT);
14921 return DAG.getNode(Opcode: ISD::USUBSAT, DL, VT, N1: OpLHS, N2: OpRHS);
14922 }
14923
14924 // Another special case: If C was a sign bit, the sub has been
14925 // canonicalized into a xor.
14926 // FIXME: Would it be better to use computeKnownBits to
14927 // determine whether it's safe to decanonicalize the xor?
14928 // x s< 0 ? x^C : 0 --> usubsat x, C
14929 APInt SplatValue;
14930 if (SatCC == ISD::SETLT && Other.getOpcode() == ISD::XOR &&
14931 ISD::isConstantSplatVector(N: OpRHS.getNode(), SplatValue) &&
14932 ISD::isConstantSplatVectorAllZeros(N: CondRHS.getNode()) &&
14933 SplatValue.isSignMask()) {
14934 // Note that we have to rebuild the RHS constant here to
14935 // ensure we don't rely on particular values of undef lanes.
14936 OpRHS = DAG.getConstant(Val: SplatValue, DL, VT);
14937 return DAG.getNode(Opcode: ISD::USUBSAT, DL, VT, N1: OpLHS, N2: OpRHS);
14938 }
14939 }
14940 }
14941 }
14942 }
14943 }
14944
14945 // (vselect (ugt x, C), (add x, ~C), x) -> (umin (add x, ~C), x)
14946 // (vselect (ult x, C), x, (add x, -C)) -> (umin x, (add x, -C))
14947 if (SDValue UMin = foldSelectToUMin(LHS, RHS, True: N1, False: N2, CC, DL))
14948 return UMin;
14949 }
14950
14951 if (SimplifySelectOps(SELECT: N, LHS: N1, RHS: N2))
14952 return SDValue(N, 0); // Don't revisit N.
14953
14954 // Fold (vselect all_ones, N1, N2) -> N1
14955 if (ISD::isConstantSplatVectorAllOnes(N: N0.getNode()))
14956 return N1;
14957 // Fold (vselect all_zeros, N1, N2) -> N2
14958 if (ISD::isConstantSplatVectorAllZeros(N: N0.getNode()))
14959 return N2;
14960
14961 // The ConvertSelectToConcatVector function is assuming both the above
14962 // checks for (vselect (build_vector all{ones,zeros) ...) have been made
14963 // and addressed.
14964 if (N1.getOpcode() == ISD::CONCAT_VECTORS &&
14965 N2.getOpcode() == ISD::CONCAT_VECTORS &&
14966 ISD::isBuildVectorOfConstantSDNodes(N: N0.getNode())) {
14967 if (SDValue CV = ConvertSelectToConcatVector(N, DAG))
14968 return CV;
14969 }
14970
14971 if (SDValue V = foldVSelectOfConstants(N))
14972 return V;
14973
14974 if (hasOperation(Opcode: ISD::SRA, VT))
14975 if (SDValue V = foldVSelectToSignBitSplatMask(N, DAG))
14976 return V;
14977
14978 if (SimplifyDemandedVectorElts(Op: SDValue(N, 0)))
14979 return SDValue(N, 0);
14980
14981 if (SDValue V = combineVSelectWithAllOnesOrZeros(Cond: N0, TVal: N1, FVal: N2, TLI, DAG, DL))
14982 return V;
14983
14984 if (SDValue R = combineSelectToPseudoMinMax(DAG, N))
14985 return R;
14986
14987 return SDValue();
14988}
14989
14990SDValue DAGCombiner::visitSELECT_CC(SDNode *N) {
14991 SDValue N0 = N->getOperand(Num: 0);
14992 SDValue N1 = N->getOperand(Num: 1);
14993 SDValue N2 = N->getOperand(Num: 2);
14994 SDValue N3 = N->getOperand(Num: 3);
14995 SDValue N4 = N->getOperand(Num: 4);
14996 ISD::CondCode CC = cast<CondCodeSDNode>(Val&: N4)->get();
14997 SDLoc DL(N);
14998
14999 // fold select_cc lhs, rhs, x, x, cc -> x
15000 if (N2 == N3)
15001 return N2;
15002
15003 if (SDValue R = foldSelectCCOfSelect(N, DAG))
15004 return R;
15005
15006 // select_cc bool, 0, x, y, seteq -> select bool, y, x
15007 if (CC == ISD::SETEQ && !LegalTypes && N0.getValueType() == MVT::i1 &&
15008 isNullConstant(V: N1))
15009 return DAG.getSelect(DL, VT: N2.getValueType(), Cond: N0, LHS: N3, RHS: N2);
15010
15011 // Determine if the condition we're dealing with is constant
15012 if (SDValue SCC = SimplifySetCC(VT: getSetCCResultType(VT: N0.getValueType()), N0, N1,
15013 Cond: CC, DL, foldBooleans: false)) {
15014 AddToWorklist(N: SCC.getNode());
15015
15016 // cond always true -> true val
15017 // cond always false -> false val
15018 if (auto *SCCC = dyn_cast<ConstantSDNode>(Val: SCC.getNode()))
15019 return SCCC->isZero() ? N3 : N2;
15020
15021 // When the condition is UNDEF, just return the first operand. This is
15022 // coherent the DAG creation, no setcc node is created in this case
15023 if (SCC->isUndef())
15024 return N2;
15025
15026 // Fold to a simpler select_cc
15027 if (SCC.getOpcode() == ISD::SETCC) {
15028 return DAG.getNode(Opcode: ISD::SELECT_CC, DL, VT: N2.getValueType(),
15029 N1: SCC.getOperand(i: 0), N2: SCC.getOperand(i: 1), N3: N2, N4: N3,
15030 N5: SCC.getOperand(i: 2), Flags: SCC->getFlags());
15031 }
15032 }
15033
15034 // If we can fold this based on the true/false value, do so.
15035 if (SimplifySelectOps(SELECT: N, LHS: N2, RHS: N3))
15036 return SDValue(N, 0); // Don't revisit N.
15037
15038 auto [Opcode, NewLHS, NewRHS] = combineSelectCCToPseudoMinMax(
15039 DAG, DL, CC, Op0: N0, Op1: N1, LHS: N2, RHS: N3, Flags: N->getFlags(), /*IsStrict=*/false);
15040 if (Opcode)
15041 return DAG.getNode(Opcode, DL, VT: N->getValueType(ResNo: 0), N1: NewLHS, N2: NewRHS,
15042 Flags: N->getFlags());
15043
15044 // fold select_cc into other things, such as min/max/abs
15045 return SimplifySelectCC(DL, N0, N1, N2, N3, CC);
15046}
15047
15048SDValue DAGCombiner::visitSETCC(SDNode *N) {
15049 // setcc is very commonly used as an argument to brcond or cond_loop. This
15050 // pattern also lend itself to numerous combines and, as a result, it is
15051 // desired we keep the argument to a brcond as a setcc as much as possible.
15052 bool PreferSetCC =
15053 N->hasOneUse() && (N->user_begin()->getOpcode() == ISD::BRCOND ||
15054 N->user_begin()->getOpcode() == ISD::COND_LOOP);
15055
15056 ISD::CondCode Cond = cast<CondCodeSDNode>(Val: N->getOperand(Num: 2))->get();
15057 EVT VT = N->getValueType(ResNo: 0);
15058 SDValue N0 = N->getOperand(Num: 0), N1 = N->getOperand(Num: 1);
15059 SDLoc DL(N);
15060
15061 if (SDValue Combined = SimplifySetCC(VT, N0, N1, Cond, DL, foldBooleans: !PreferSetCC)) {
15062 // If we prefer to have a setcc, and we don't, we'll try our best to
15063 // recreate one using rebuildSetCC.
15064 if (PreferSetCC && Combined.getOpcode() != ISD::SETCC) {
15065 SDValue NewSetCC = rebuildSetCC(N: Combined);
15066
15067 // We don't have anything interesting to combine to.
15068 if (NewSetCC.getNode() == N)
15069 return SDValue();
15070
15071 if (NewSetCC)
15072 return NewSetCC;
15073 }
15074 return Combined;
15075 }
15076
15077 // Optimize
15078 // 1) (icmp eq/ne (and X, C0), (shift X, C1))
15079 // or
15080 // 2) (icmp eq/ne X, (rotate X, C1))
15081 // If C0 is a mask or shifted mask and the shift amt (C1) isolates the
15082 // remaining bits (i.e something like `(x64 & UINT32_MAX) == (x64 >> 32)`)
15083 // Then:
15084 // If C1 divides the bit width, then the rotate and shift+and versions are
15085 // equivalent, so we can interchange them depending on target preference.
15086 // Otherwise, if we have the shift+and version we can interchange srl/shl
15087 // which inturn affects the constant C0. We can use this to get better
15088 // constants again determined by target preference.
15089 if (Cond == ISD::SETNE || Cond == ISD::SETEQ) {
15090 auto IsAndWithShift = [](SDValue A, SDValue B) {
15091 return A.getOpcode() == ISD::AND &&
15092 (B.getOpcode() == ISD::SRL || B.getOpcode() == ISD::SHL) &&
15093 A.getOperand(i: 0) == B.getOperand(i: 0);
15094 };
15095 auto IsRotateWithOp = [](SDValue A, SDValue B) {
15096 return (B.getOpcode() == ISD::ROTL || B.getOpcode() == ISD::ROTR) &&
15097 B.getOperand(i: 0) == A;
15098 };
15099 SDValue AndOrOp = SDValue(), ShiftOrRotate = SDValue();
15100 bool IsRotate = false;
15101
15102 // Find either shift+and or rotate pattern.
15103 if (IsAndWithShift(N0, N1)) {
15104 AndOrOp = N0;
15105 ShiftOrRotate = N1;
15106 } else if (IsAndWithShift(N1, N0)) {
15107 AndOrOp = N1;
15108 ShiftOrRotate = N0;
15109 } else if (IsRotateWithOp(N0, N1)) {
15110 IsRotate = true;
15111 AndOrOp = N0;
15112 ShiftOrRotate = N1;
15113 } else if (IsRotateWithOp(N1, N0)) {
15114 IsRotate = true;
15115 AndOrOp = N1;
15116 ShiftOrRotate = N0;
15117 }
15118
15119 if (AndOrOp && ShiftOrRotate && ShiftOrRotate.hasOneUse() &&
15120 (IsRotate || AndOrOp.hasOneUse())) {
15121 EVT OpVT = N0.getValueType();
15122 // Get constant shift/rotate amount and possibly mask (if its shift+and
15123 // variant).
15124 auto GetAPIntValue = [](SDValue Op) -> std::optional<APInt> {
15125 ConstantSDNode *CNode = isConstOrConstSplat(N: Op, /*AllowUndefs*/ false,
15126 /*AllowTrunc*/ AllowTruncation: false);
15127 if (CNode == nullptr)
15128 return std::nullopt;
15129 return CNode->getAPIntValue();
15130 };
15131 std::optional<APInt> AndCMask =
15132 IsRotate ? std::nullopt : GetAPIntValue(AndOrOp.getOperand(i: 1));
15133 std::optional<APInt> ShiftCAmt =
15134 GetAPIntValue(ShiftOrRotate.getOperand(i: 1));
15135 unsigned NumBits = OpVT.getScalarSizeInBits();
15136
15137 // We found constants.
15138 if (ShiftCAmt && (IsRotate || AndCMask) && ShiftCAmt->ult(RHS: NumBits)) {
15139 unsigned ShiftOpc = ShiftOrRotate.getOpcode();
15140 // Check that the constants meet the constraints.
15141 bool CanTransform = IsRotate;
15142 if (!CanTransform) {
15143 // Check that mask and shift compliment eachother
15144 CanTransform = *ShiftCAmt == (~*AndCMask).popcount();
15145 // Check that we are comparing all bits
15146 CanTransform &= (*ShiftCAmt + AndCMask->popcount()) == NumBits;
15147 // Check that the and mask is correct for the shift
15148 CanTransform &=
15149 ShiftOpc == ISD::SHL ? (~*AndCMask).isMask() : AndCMask->isMask();
15150 }
15151
15152 // The rotate and shift+and forms are only equivalent if the shift
15153 // amount divides the bit width.
15154 bool MayTransformRotate =
15155 !ShiftCAmt->isZero() && NumBits % ShiftCAmt->getZExtValue() == 0;
15156 // See if target prefers another shift/rotate opcode.
15157 unsigned NewShiftOpc = TLI.preferedOpcodeForCmpEqPiecesOfOperand(
15158 VT: OpVT, ShiftOpc, MayTransformRotate, ShiftOrRotateAmt: *ShiftCAmt, AndMask: AndCMask);
15159 // Transform is valid and we have a new preference.
15160 if (CanTransform && NewShiftOpc != ShiftOpc) {
15161 SDValue NewShiftOrRotate =
15162 DAG.getNode(Opcode: NewShiftOpc, DL, VT: OpVT, N1: ShiftOrRotate.getOperand(i: 0),
15163 N2: ShiftOrRotate.getOperand(i: 1));
15164 SDValue NewAndOrOp = SDValue();
15165
15166 if (NewShiftOpc == ISD::SHL || NewShiftOpc == ISD::SRL) {
15167 APInt NewMask =
15168 NewShiftOpc == ISD::SHL
15169 ? APInt::getHighBitsSet(numBits: NumBits,
15170 hiBitsSet: NumBits - ShiftCAmt->getZExtValue())
15171 : APInt::getLowBitsSet(numBits: NumBits,
15172 loBitsSet: NumBits - ShiftCAmt->getZExtValue());
15173 NewAndOrOp =
15174 DAG.getNode(Opcode: ISD::AND, DL, VT: OpVT, N1: ShiftOrRotate.getOperand(i: 0),
15175 N2: DAG.getConstant(Val: NewMask, DL, VT: OpVT));
15176 } else {
15177 NewAndOrOp = ShiftOrRotate.getOperand(i: 0);
15178 }
15179
15180 return DAG.getSetCC(DL, VT, LHS: NewAndOrOp, RHS: NewShiftOrRotate, Cond);
15181 }
15182 }
15183 }
15184 }
15185 return SDValue();
15186}
15187
15188SDValue DAGCombiner::visitSETCCCARRY(SDNode *N) {
15189 SDValue LHS = N->getOperand(Num: 0);
15190 SDValue RHS = N->getOperand(Num: 1);
15191 SDValue Carry = N->getOperand(Num: 2);
15192 SDValue Cond = N->getOperand(Num: 3);
15193
15194 // If Carry is false, fold to a regular SETCC.
15195 if (isNullConstant(V: Carry))
15196 return DAG.getNode(Opcode: ISD::SETCC, DL: SDLoc(N), VTList: N->getVTList(), N1: LHS, N2: RHS, N3: Cond);
15197
15198 return SDValue();
15199}
15200
15201/// Check if N satisfies:
15202/// N is used once.
15203/// N is a Load.
15204/// The load is compatible with ExtOpcode. It means
15205/// If load has explicit zero/sign extension, ExpOpcode must have the same
15206/// extension.
15207/// Otherwise returns true.
15208static bool isCompatibleLoad(SDValue N, unsigned ExtOpcode) {
15209 if (!N.hasOneUse())
15210 return false;
15211
15212 if (!isa<LoadSDNode>(Val: N))
15213 return false;
15214
15215 LoadSDNode *Load = cast<LoadSDNode>(Val&: N);
15216 ISD::LoadExtType LoadExt = Load->getExtensionType();
15217 if (LoadExt == ISD::NON_EXTLOAD || LoadExt == ISD::EXTLOAD)
15218 return true;
15219
15220 // Now LoadExt is either SEXTLOAD or ZEXTLOAD, ExtOpcode must have the same
15221 // extension.
15222 if ((LoadExt == ISD::SEXTLOAD && ExtOpcode != ISD::SIGN_EXTEND) ||
15223 (LoadExt == ISD::ZEXTLOAD && ExtOpcode != ISD::ZERO_EXTEND))
15224 return false;
15225
15226 return true;
15227}
15228
15229/// Fold
15230/// (sext (select c, load x, load y)) -> (select c, sextload x, sextload y)
15231/// (zext (select c, load x, load y)) -> (select c, zextload x, zextload y)
15232/// (aext (select c, load x, load y)) -> (select c, extload x, extload y)
15233/// This function is called by the DAGCombiner when visiting sext/zext/aext
15234/// dag nodes (see for example method DAGCombiner::visitSIGN_EXTEND).
15235static SDValue tryToFoldExtendSelectLoad(SDNode *N, const TargetLowering &TLI,
15236 SelectionDAG &DAG, const SDLoc &DL,
15237 CombineLevel Level) {
15238 unsigned Opcode = N->getOpcode();
15239 SDValue N0 = N->getOperand(Num: 0);
15240 EVT VT = N->getValueType(ResNo: 0);
15241 assert((Opcode == ISD::SIGN_EXTEND || Opcode == ISD::ZERO_EXTEND ||
15242 Opcode == ISD::ANY_EXTEND) &&
15243 "Expected EXTEND dag node in input!");
15244
15245 SDValue Cond, Op1, Op2;
15246 if (!sd_match(N: N0, P: m_OneUse(P: m_SelectLike(Cond: m_Value(N&: Cond), T: m_Value(N&: Op1),
15247 F: m_Value(N&: Op2)))))
15248 return SDValue();
15249
15250 if (!isCompatibleLoad(N: Op1, ExtOpcode: Opcode) || !isCompatibleLoad(N: Op2, ExtOpcode: Opcode))
15251 return SDValue();
15252
15253 auto ExtLoadOpcode = ISD::EXTLOAD;
15254 if (Opcode == ISD::SIGN_EXTEND)
15255 ExtLoadOpcode = ISD::SEXTLOAD;
15256 else if (Opcode == ISD::ZERO_EXTEND)
15257 ExtLoadOpcode = ISD::ZEXTLOAD;
15258
15259 // Illegal VSELECT may ISel fail if happen after legalization (DAG
15260 // Combine2), so we should conservatively check the OperationAction.
15261 LoadSDNode *Load1 = cast<LoadSDNode>(Val&: Op1);
15262 LoadSDNode *Load2 = cast<LoadSDNode>(Val&: Op2);
15263 if (!TLI.isLoadLegal(ValVT: VT, MemVT: Load1->getMemoryVT(), Alignment: Load1->getAlign(),
15264 AddrSpace: Load1->getAddressSpace(), ExtType: ExtLoadOpcode, Atomic: false) ||
15265 !TLI.isLoadLegal(ValVT: VT, MemVT: Load2->getMemoryVT(), Alignment: Load2->getAlign(),
15266 AddrSpace: Load2->getAddressSpace(), ExtType: ExtLoadOpcode, Atomic: false) ||
15267 (N0->getOpcode() == ISD::VSELECT && Level >= AfterLegalizeTypes &&
15268 TLI.getOperationAction(Op: ISD::VSELECT, VT) != TargetLowering::Legal))
15269 return SDValue();
15270
15271 SDValue Ext1 = DAG.getNode(Opcode, DL, VT, Operand: Op1);
15272 SDValue Ext2 = DAG.getNode(Opcode, DL, VT, Operand: Op2);
15273 return DAG.getSelect(DL, VT, Cond, LHS: Ext1, RHS: Ext2);
15274}
15275
15276/// Try to fold a sext/zext/aext dag node into a ConstantSDNode or
15277/// a build_vector of constants.
15278/// This function is called by the DAGCombiner when visiting sext/zext/aext
15279/// dag nodes (see for example method DAGCombiner::visitSIGN_EXTEND).
15280/// Vector extends are not folded if operations are legal; this is to
15281/// avoid introducing illegal build_vector dag nodes.
15282static SDValue tryToFoldExtendOfConstant(SDNode *N, const SDLoc &DL,
15283 const TargetLowering &TLI,
15284 SelectionDAG &DAG, bool LegalTypes) {
15285 unsigned Opcode = N->getOpcode();
15286 SDValue N0 = N->getOperand(Num: 0);
15287 EVT VT = N->getValueType(ResNo: 0);
15288
15289 assert((ISD::isExtOpcode(Opcode) || ISD::isExtVecInRegOpcode(Opcode)) &&
15290 "Expected EXTEND dag node in input!");
15291
15292 // fold (sext c1) -> c1
15293 // fold (zext c1) -> c1
15294 // fold (aext c1) -> c1
15295 if (isa<ConstantSDNode>(Val: N0))
15296 return DAG.getNode(Opcode, DL, VT, Operand: N0);
15297
15298 // fold (sext (select cond, c1, c2)) -> (select cond, sext c1, sext c2)
15299 // fold (zext (select cond, c1, c2)) -> (select cond, zext c1, zext c2)
15300 // fold (aext (select cond, c1, c2)) -> (select cond, sext c1, sext c2)
15301 if (N0->getOpcode() == ISD::SELECT) {
15302 SDValue Op1 = N0->getOperand(Num: 1);
15303 SDValue Op2 = N0->getOperand(Num: 2);
15304 if (isa<ConstantSDNode>(Val: Op1) && isa<ConstantSDNode>(Val: Op2) &&
15305 (Opcode != ISD::ZERO_EXTEND || !TLI.isZExtFree(FromTy: N0.getValueType(), ToTy: VT))) {
15306 // For any_extend, choose sign extension of the constants to allow a
15307 // possible further transform to sign_extend_inreg.i.e.
15308 //
15309 // t1: i8 = select t0, Constant:i8<-1>, Constant:i8<0>
15310 // t2: i64 = any_extend t1
15311 // -->
15312 // t3: i64 = select t0, Constant:i64<-1>, Constant:i64<0>
15313 // -->
15314 // t4: i64 = sign_extend_inreg t3
15315 unsigned FoldOpc = Opcode;
15316 if (FoldOpc == ISD::ANY_EXTEND)
15317 FoldOpc = ISD::SIGN_EXTEND;
15318 return DAG.getSelect(DL, VT, Cond: N0->getOperand(Num: 0),
15319 LHS: DAG.getNode(Opcode: FoldOpc, DL, VT, Operand: Op1),
15320 RHS: DAG.getNode(Opcode: FoldOpc, DL, VT, Operand: Op2));
15321 }
15322 }
15323
15324 // fold (sext (build_vector AllConstants) -> (build_vector AllConstants)
15325 // fold (zext (build_vector AllConstants) -> (build_vector AllConstants)
15326 // fold (aext (build_vector AllConstants) -> (build_vector AllConstants)
15327 EVT SVT = VT.getScalarType();
15328 if (!(VT.isVector() && (!LegalTypes || TLI.isTypeLegal(VT: SVT)) &&
15329 ISD::isBuildVectorOfConstantSDNodes(N: N0.getNode())))
15330 return SDValue();
15331
15332 // We can fold this node into a build_vector.
15333 unsigned VTBits = SVT.getSizeInBits();
15334 unsigned EVTBits = N0->getValueType(ResNo: 0).getScalarSizeInBits();
15335 SmallVector<SDValue, 8> Elts;
15336 unsigned NumElts = VT.getVectorNumElements();
15337
15338 for (unsigned i = 0; i != NumElts; ++i) {
15339 SDValue Op = N0.getOperand(i);
15340 if (Op.isUndef()) {
15341 if (Opcode == ISD::ANY_EXTEND || Opcode == ISD::ANY_EXTEND_VECTOR_INREG)
15342 Elts.push_back(Elt: DAG.getUNDEF(VT: SVT));
15343 else
15344 Elts.push_back(Elt: DAG.getConstant(Val: 0, DL, VT: SVT));
15345 continue;
15346 }
15347
15348 SDLoc DL(Op);
15349 // Get the constant value and if needed trunc it to the size of the type.
15350 // Nodes like build_vector might have constants wider than the scalar type.
15351 APInt C = Op->getAsAPIntVal().zextOrTrunc(width: EVTBits);
15352 if (Opcode == ISD::SIGN_EXTEND || Opcode == ISD::SIGN_EXTEND_VECTOR_INREG)
15353 Elts.push_back(Elt: DAG.getConstant(Val: C.sext(width: VTBits), DL, VT: SVT));
15354 else
15355 Elts.push_back(Elt: DAG.getConstant(Val: C.zext(width: VTBits), DL, VT: SVT));
15356 }
15357
15358 return DAG.getBuildVector(VT, DL, Ops: Elts);
15359}
15360
15361// ExtendUsesToFormExtLoad - Trying to extend uses of a load to enable this:
15362// "fold ({s|z|a}ext (load x)) -> ({s|z|a}ext (truncate ({s|z|a}extload x)))"
15363// transformation. Returns true if extension are possible and the above
15364// mentioned transformation is profitable.
15365static bool ExtendUsesToFormExtLoad(EVT VT, SDNode *N, SDValue N0,
15366 unsigned ExtOpc,
15367 SmallVectorImpl<SDNode *> &ExtendNodes,
15368 const TargetLowering &TLI) {
15369 bool HasCopyToRegUses = false;
15370 bool isTruncFree = TLI.isTruncateFree(FromVT: VT, ToVT: N0.getValueType());
15371 for (SDUse &Use : N0->uses()) {
15372 SDNode *User = Use.getUser();
15373 if (User == N)
15374 continue;
15375 if (Use.getResNo() != N0.getResNo())
15376 continue;
15377 // FIXME: Only extend SETCC N, N and SETCC N, c for now.
15378 if (ExtOpc != ISD::ANY_EXTEND && User->getOpcode() == ISD::SETCC) {
15379 ISD::CondCode CC = cast<CondCodeSDNode>(Val: User->getOperand(Num: 2))->get();
15380 if (ExtOpc == ISD::ZERO_EXTEND && ISD::isSignedIntSetCC(Code: CC))
15381 // Sign bits will be lost after a zext.
15382 return false;
15383 bool Add = false;
15384 for (unsigned i = 0; i != 2; ++i) {
15385 SDValue UseOp = User->getOperand(Num: i);
15386 if (UseOp == N0)
15387 continue;
15388 if (!isa<ConstantSDNode>(Val: UseOp))
15389 return false;
15390 Add = true;
15391 }
15392 if (Add)
15393 ExtendNodes.push_back(Elt: User);
15394 continue;
15395 }
15396 // If truncates aren't free and there are users we can't
15397 // extend, it isn't worthwhile.
15398 if (!isTruncFree)
15399 return false;
15400 // Remember if this value is live-out.
15401 if (User->getOpcode() == ISD::CopyToReg)
15402 HasCopyToRegUses = true;
15403 }
15404
15405 if (HasCopyToRegUses) {
15406 bool BothLiveOut = false;
15407 for (SDUse &Use : N->uses()) {
15408 if (Use.getResNo() == 0 && Use.getUser()->getOpcode() == ISD::CopyToReg) {
15409 BothLiveOut = true;
15410 break;
15411 }
15412 }
15413 if (BothLiveOut)
15414 // Both unextended and extended values are live out. There had better be
15415 // a good reason for the transformation.
15416 return !ExtendNodes.empty();
15417 }
15418 return true;
15419}
15420
15421void DAGCombiner::ExtendSetCCUses(const SmallVectorImpl<SDNode *> &SetCCs,
15422 SDValue OrigLoad, SDValue ExtLoad,
15423 ISD::NodeType ExtType) {
15424 // Extend SetCC uses if necessary.
15425 SDLoc DL(ExtLoad);
15426 for (SDNode *SetCC : SetCCs) {
15427 SmallVector<SDValue, 4> Ops;
15428
15429 for (unsigned j = 0; j != 2; ++j) {
15430 SDValue SOp = SetCC->getOperand(Num: j);
15431 if (SOp == OrigLoad)
15432 Ops.push_back(Elt: ExtLoad);
15433 else
15434 Ops.push_back(Elt: DAG.getNode(Opcode: ExtType, DL, VT: ExtLoad->getValueType(ResNo: 0), Operand: SOp));
15435 }
15436
15437 Ops.push_back(Elt: SetCC->getOperand(Num: 2));
15438 CombineTo(N: SetCC, Res: DAG.getNode(Opcode: ISD::SETCC, DL, VT: SetCC->getValueType(ResNo: 0), Ops));
15439 }
15440}
15441
15442// FIXME: Bring more similar combines here, common to sext/zext (maybe aext?).
15443SDValue DAGCombiner::CombineExtLoad(SDNode *N) {
15444 SDValue N0 = N->getOperand(Num: 0);
15445 EVT DstVT = N->getValueType(ResNo: 0);
15446 EVT SrcVT = N0.getValueType();
15447
15448 assert((N->getOpcode() == ISD::SIGN_EXTEND ||
15449 N->getOpcode() == ISD::ZERO_EXTEND) &&
15450 "Unexpected node type (not an extend)!");
15451
15452 // fold (sext (load x)) to multiple smaller sextloads; same for zext.
15453 // For example, on a target with legal v4i32, but illegal v8i32, turn:
15454 // (v8i32 (sext (v8i16 (load x))))
15455 // into:
15456 // (v8i32 (concat_vectors (v4i32 (sextload x)),
15457 // (v4i32 (sextload (x + 16)))))
15458 // Where uses of the original load, i.e.:
15459 // (v8i16 (load x))
15460 // are replaced with:
15461 // (v8i16 (truncate
15462 // (v8i32 (concat_vectors (v4i32 (sextload x)),
15463 // (v4i32 (sextload (x + 16)))))))
15464 //
15465 // This combine is only applicable to illegal, but splittable, vectors.
15466 // All legal types, and illegal non-vector types, are handled elsewhere.
15467 // This combine is controlled by TargetLowering::isVectorLoadExtDesirable.
15468 //
15469 if (N0->getOpcode() != ISD::LOAD)
15470 return SDValue();
15471
15472 LoadSDNode *LN0 = cast<LoadSDNode>(Val&: N0);
15473
15474 if (!ISD::isNON_EXTLoad(N: LN0) || !ISD::isUNINDEXEDLoad(N: LN0) ||
15475 !N0.hasOneUse() || !LN0->isSimple() ||
15476 !DstVT.isVector() || !DstVT.isPow2VectorType() ||
15477 !TLI.isVectorLoadExtDesirable(ExtVal: SDValue(N, 0)))
15478 return SDValue();
15479
15480 SmallVector<SDNode *, 4> SetCCs;
15481 if (!ExtendUsesToFormExtLoad(VT: DstVT, N, N0, ExtOpc: N->getOpcode(), ExtendNodes&: SetCCs, TLI))
15482 return SDValue();
15483
15484 ISD::LoadExtType ExtType =
15485 N->getOpcode() == ISD::SIGN_EXTEND ? ISD::SEXTLOAD : ISD::ZEXTLOAD;
15486
15487 // Try to split the vector types to get down to legal types.
15488 EVT SplitSrcVT = SrcVT;
15489 EVT SplitDstVT = DstVT;
15490 while (!TLI.isLoadLegalOrCustom(ValVT: SplitDstVT, MemVT: SplitSrcVT, Alignment: LN0->getAlign(),
15491 AddrSpace: LN0->getAddressSpace(), ExtType, Atomic: false) &&
15492 SplitSrcVT.getVectorNumElements() > 1) {
15493 SplitDstVT = DAG.GetSplitDestVTs(VT: SplitDstVT).first;
15494 SplitSrcVT = DAG.GetSplitDestVTs(VT: SplitSrcVT).first;
15495 }
15496
15497 if (!TLI.isLoadLegalOrCustom(ValVT: SplitDstVT, MemVT: SplitSrcVT, Alignment: LN0->getAlign(),
15498 AddrSpace: LN0->getAddressSpace(), ExtType, Atomic: false))
15499 return SDValue();
15500
15501 assert(!DstVT.isScalableVector() && "Unexpected scalable vector type");
15502
15503 SDLoc DL(N);
15504 const unsigned NumSplits =
15505 DstVT.getVectorNumElements() / SplitDstVT.getVectorNumElements();
15506 const unsigned Stride = SplitSrcVT.getStoreSize();
15507 SmallVector<SDValue, 4> Loads;
15508 SmallVector<SDValue, 4> Chains;
15509
15510 SDValue BasePtr = LN0->getBasePtr();
15511 for (unsigned Idx = 0; Idx < NumSplits; Idx++) {
15512 const unsigned Offset = Idx * Stride;
15513
15514 SDValue SplitLoad =
15515 DAG.getExtLoad(ExtType, dl: SDLoc(LN0), VT: SplitDstVT, Chain: LN0->getChain(),
15516 Ptr: BasePtr, PtrInfo: LN0->getPointerInfo().getWithOffset(O: Offset),
15517 MemVT: SplitSrcVT, Alignment: LN0->getBaseAlign(),
15518 MMOFlags: LN0->getMemOperand()->getFlags(), Metadata: LN0->getAAInfo());
15519
15520 BasePtr = DAG.getMemBasePlusOffset(Base: BasePtr, Offset: TypeSize::getFixed(ExactSize: Stride), DL);
15521
15522 Loads.push_back(Elt: SplitLoad.getValue(R: 0));
15523 Chains.push_back(Elt: SplitLoad.getValue(R: 1));
15524 }
15525
15526 SDValue NewChain = DAG.getNode(Opcode: ISD::TokenFactor, DL, VT: MVT::Other, Ops: Chains);
15527 SDValue NewValue = DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT: DstVT, Ops: Loads);
15528
15529 // Simplify TF.
15530 AddToWorklist(N: NewChain.getNode());
15531
15532 CombineTo(N, Res: NewValue);
15533
15534 // Replace uses of the original load (before extension)
15535 // with a truncate of the concatenated sextloaded vectors.
15536 SDValue Trunc =
15537 DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(N0), VT: N0.getValueType(), Operand: NewValue);
15538 ExtendSetCCUses(SetCCs, OrigLoad: N0, ExtLoad: NewValue, ExtType: (ISD::NodeType)N->getOpcode());
15539 CombineTo(N: N0.getNode(), Res0: Trunc, Res1: NewChain);
15540 return SDValue(N, 0); // Return N so it doesn't get rechecked!
15541}
15542
15543// fold (zext (and/or/xor (shl/shr (load x), cst), cst)) ->
15544// (and/or/xor (shl/shr (zextload x), (zext cst)), (zext cst))
15545SDValue DAGCombiner::CombineZExtLogicopShiftLoad(SDNode *N) {
15546 assert(N->getOpcode() == ISD::ZERO_EXTEND);
15547 EVT VT = N->getValueType(ResNo: 0);
15548 EVT OrigVT = N->getOperand(Num: 0).getValueType();
15549 if (TLI.isZExtFree(FromTy: OrigVT, ToTy: VT))
15550 return SDValue();
15551
15552 // and/or/xor
15553 SDValue N0 = N->getOperand(Num: 0);
15554 if (!ISD::isBitwiseLogicOp(Opcode: N0.getOpcode()) ||
15555 N0.getOperand(i: 1).getOpcode() != ISD::Constant ||
15556 (LegalOperations && !TLI.isOperationLegal(Op: N0.getOpcode(), VT)))
15557 return SDValue();
15558
15559 // shl/shr
15560 SDValue N1 = N0->getOperand(Num: 0);
15561 if (!(N1.getOpcode() == ISD::SHL || N1.getOpcode() == ISD::SRL) ||
15562 N1.getOperand(i: 1).getOpcode() != ISD::Constant ||
15563 (LegalOperations && !TLI.isOperationLegal(Op: N1.getOpcode(), VT)))
15564 return SDValue();
15565
15566 // load
15567 if (!isa<LoadSDNode>(Val: N1.getOperand(i: 0)))
15568 return SDValue();
15569 LoadSDNode *Load = cast<LoadSDNode>(Val: N1.getOperand(i: 0));
15570 EVT MemVT = Load->getMemoryVT();
15571 if (!TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: Load->getAlign(), AddrSpace: Load->getAddressSpace(),
15572 ExtType: ISD::ZEXTLOAD, Atomic: false) ||
15573 Load->getExtensionType() == ISD::SEXTLOAD || Load->isIndexed())
15574 return SDValue();
15575
15576
15577 // If the shift op is SHL, the logic op must be AND, otherwise the result
15578 // will be wrong.
15579 if (N1.getOpcode() == ISD::SHL && N0.getOpcode() != ISD::AND)
15580 return SDValue();
15581
15582 if (!N0.hasOneUse() || !N1.hasOneUse())
15583 return SDValue();
15584
15585 SmallVector<SDNode*, 4> SetCCs;
15586 if (!ExtendUsesToFormExtLoad(VT, N: N1.getNode(), N0: N1.getOperand(i: 0),
15587 ExtOpc: ISD::ZERO_EXTEND, ExtendNodes&: SetCCs, TLI))
15588 return SDValue();
15589
15590 // Actually do the transformation.
15591 SDValue ExtLoad = DAG.getExtLoad(ExtType: ISD::ZEXTLOAD, dl: SDLoc(Load), VT,
15592 Chain: Load->getChain(), Ptr: Load->getBasePtr(),
15593 MemVT: Load->getMemoryVT(), MMO: Load->getMemOperand());
15594
15595 SDLoc DL1(N1);
15596 SDValue Shift = DAG.getNode(Opcode: N1.getOpcode(), DL: DL1, VT, N1: ExtLoad,
15597 N2: N1.getOperand(i: 1));
15598
15599 APInt Mask = N0.getConstantOperandAPInt(i: 1).zext(width: VT.getSizeInBits());
15600 SDLoc DL0(N0);
15601 SDValue And = DAG.getNode(Opcode: N0.getOpcode(), DL: DL0, VT, N1: Shift,
15602 N2: DAG.getConstant(Val: Mask, DL: DL0, VT));
15603
15604 ExtendSetCCUses(SetCCs, OrigLoad: N1.getOperand(i: 0), ExtLoad, ExtType: ISD::ZERO_EXTEND);
15605 CombineTo(N, Res: And);
15606 if (SDValue(Load, 0).hasOneUse()) {
15607 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Load, 1), To: ExtLoad.getValue(R: 1));
15608 } else {
15609 SDValue Trunc = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(Load),
15610 VT: Load->getValueType(ResNo: 0), Operand: ExtLoad);
15611 CombineTo(N: Load, Res0: Trunc, Res1: ExtLoad.getValue(R: 1));
15612 }
15613
15614 // N0 is dead at this point.
15615 recursivelyDeleteUnusedNodes(N: N0.getNode());
15616
15617 return SDValue(N,0); // Return N so it doesn't get rechecked!
15618}
15619
15620/// If we're narrowing or widening the result of a vector select and the final
15621/// size is the same size as a setcc (compare) feeding the select, then try to
15622/// apply the cast operation to the select's operands because matching vector
15623/// sizes for a select condition and other operands should be more efficient.
15624SDValue DAGCombiner::matchVSelectOpSizesWithSetCC(SDNode *Cast) {
15625 unsigned CastOpcode = Cast->getOpcode();
15626 assert((CastOpcode == ISD::SIGN_EXTEND || CastOpcode == ISD::ZERO_EXTEND ||
15627 CastOpcode == ISD::TRUNCATE || CastOpcode == ISD::FP_EXTEND ||
15628 CastOpcode == ISD::FP_ROUND) &&
15629 "Unexpected opcode for vector select narrowing/widening");
15630
15631 // We only do this transform before legal ops because the pattern may be
15632 // obfuscated by target-specific operations after legalization. Do not create
15633 // an illegal select op, however, because that may be difficult to lower.
15634 EVT VT = Cast->getValueType(ResNo: 0);
15635 if (LegalOperations || !TLI.isOperationLegalOrCustom(Op: ISD::VSELECT, VT))
15636 return SDValue();
15637
15638 SDValue VSel = Cast->getOperand(Num: 0);
15639 if (VSel.getOpcode() != ISD::VSELECT || !VSel.hasOneUse() ||
15640 VSel.getOperand(i: 0).getOpcode() != ISD::SETCC)
15641 return SDValue();
15642
15643 // Does the setcc have the same vector size as the casted select?
15644 SDValue SetCC = VSel.getOperand(i: 0);
15645 EVT SetCCVT = getSetCCResultType(VT: SetCC.getOperand(i: 0).getValueType());
15646 if (SetCCVT.getSizeInBits() != VT.getSizeInBits())
15647 return SDValue();
15648
15649 // cast (vsel (setcc X), A, B) --> vsel (setcc X), (cast A), (cast B)
15650 SDValue A = VSel.getOperand(i: 1);
15651 SDValue B = VSel.getOperand(i: 2);
15652 SDValue CastA, CastB;
15653 SDLoc DL(Cast);
15654 if (CastOpcode == ISD::FP_ROUND) {
15655 // FP_ROUND (fptrunc) has an extra flag operand to pass along.
15656 CastA = DAG.getNode(Opcode: CastOpcode, DL, VT, N1: A, N2: Cast->getOperand(Num: 1));
15657 CastB = DAG.getNode(Opcode: CastOpcode, DL, VT, N1: B, N2: Cast->getOperand(Num: 1));
15658 } else {
15659 CastA = DAG.getNode(Opcode: CastOpcode, DL, VT, Operand: A);
15660 CastB = DAG.getNode(Opcode: CastOpcode, DL, VT, Operand: B);
15661 }
15662 return DAG.getNode(Opcode: ISD::VSELECT, DL, VT, N1: SetCC, N2: CastA, N3: CastB);
15663}
15664
15665// fold ([s|z]ext ([s|z]extload x)) -> ([s|z]ext (truncate ([s|z]extload x)))
15666// fold ([s|z]ext ( extload x)) -> ([s|z]ext (truncate ([s|z]extload x)))
15667static SDValue tryToFoldExtOfExtload(SelectionDAG &DAG, DAGCombiner &Combiner,
15668 const TargetLowering &TLI, EVT VT,
15669 bool LegalOperations, SDNode *N,
15670 SDValue N0, ISD::LoadExtType ExtLoadType) {
15671 bool Frozen = N0.getOpcode() == ISD::FREEZE;
15672 auto *OldExtLoad = dyn_cast<LoadSDNode>(Val: Frozen ? N0.getOperand(i: 0) : N0);
15673 if (!OldExtLoad)
15674 return SDValue();
15675
15676 bool isAExtLoad = (ExtLoadType == ISD::SEXTLOAD)
15677 ? ISD::isSEXTLoad(N: OldExtLoad)
15678 : ISD::isZEXTLoad(N: OldExtLoad);
15679 if ((!isAExtLoad && !ISD::isEXTLoad(N: OldExtLoad)) ||
15680 !ISD::isUNINDEXEDLoad(N: OldExtLoad) || !OldExtLoad->hasNUsesOfValue(NUses: 1, Value: 0))
15681 return SDValue();
15682
15683 EVT MemVT = OldExtLoad->getMemoryVT();
15684 if ((LegalOperations || !OldExtLoad->isSimple() || VT.isVector()) &&
15685 !TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: OldExtLoad->getAlign(),
15686 AddrSpace: OldExtLoad->getAddressSpace(), ExtType: ExtLoadType, Atomic: false))
15687 return SDValue();
15688
15689 SDLoc DL(OldExtLoad);
15690 SDValue ExtLoad = DAG.getExtLoad(ExtType: ExtLoadType, dl: DL, VT, Chain: OldExtLoad->getChain(),
15691 Ptr: OldExtLoad->getBasePtr(), MemVT,
15692 MMO: OldExtLoad->getMemOperand());
15693 SDValue Res = ExtLoad;
15694 if (Frozen) {
15695 Res = DAG.getFreeze(V: ExtLoad);
15696 Res = DAG.getNode(
15697 Opcode: ExtLoadType == ISD::SEXTLOAD ? ISD::AssertSext : ISD::AssertZext, DL,
15698 VT: Res.getValueType(), N1: Res,
15699 N2: DAG.getValueType(OldExtLoad->getValueType(ResNo: 0).getScalarType()));
15700 }
15701 Combiner.CombineTo(N, Res);
15702 DAG.ReplaceAllUsesOfValueWith(From: SDValue(OldExtLoad, 1), To: ExtLoad.getValue(R: 1));
15703 if (N0->use_empty())
15704 Combiner.recursivelyDeleteUnusedNodes(N: N0.getNode());
15705 return SDValue(N, 0); // Return N so it doesn't get rechecked!
15706}
15707
15708// fold ([s|z]ext (load x)) -> ([s|z]ext (truncate ([s|z]extload x)))
15709// Only generate vector extloads when 1) they're legal, and 2) they are
15710// deemed desirable by the target. NonNegZExt can be set to true if a zero
15711// extend has the nonneg flag to allow use of sextload if profitable.
15712static SDValue tryToFoldExtOfLoad(SelectionDAG &DAG, DAGCombiner &Combiner,
15713 const TargetLowering &TLI, EVT VT,
15714 bool LegalOperations, SDNode *N, SDValue N0,
15715 ISD::LoadExtType ExtLoadType,
15716 ISD::NodeType ExtOpc,
15717 bool NonNegZExt = false) {
15718
15719 bool Frozen = N0.getOpcode() == ISD::FREEZE;
15720 SDValue Freeze = Frozen ? N0 : SDValue();
15721 auto *Load = dyn_cast<LoadSDNode>(Val: Frozen ? N0.getOperand(i: 0) : N0);
15722 // TODO: Support multiple uses of the load when frozen.
15723 if (!Load || !ISD::isNON_EXTLoad(N: Load) || !ISD::isUNINDEXEDLoad(N: Load) ||
15724 (Frozen && !Load->hasNUsesOfValue(NUses: 1, Value: 0)))
15725 return {};
15726
15727 // If this is zext nneg, see if it would make sense to treat it as a sext.
15728 if (NonNegZExt) {
15729 assert(ExtLoadType == ISD::ZEXTLOAD && ExtOpc == ISD::ZERO_EXTEND &&
15730 "Unexpected load type or opcode");
15731 for (SDNode *User : Load->users()) {
15732 if (User->getOpcode() == ISD::SETCC) {
15733 ISD::CondCode CC = cast<CondCodeSDNode>(Val: User->getOperand(Num: 2))->get();
15734 if (ISD::isSignedIntSetCC(Code: CC)) {
15735 ExtLoadType = ISD::SEXTLOAD;
15736 ExtOpc = ISD::SIGN_EXTEND;
15737 break;
15738 }
15739 }
15740 }
15741 }
15742
15743 // TODO: isFixedLengthVector() should be removed and any negative effects on
15744 // code generation being the result of that target's implementation of
15745 // isVectorLoadExtDesirable().
15746 if ((LegalOperations || VT.isFixedLengthVector() || !Load->isSimple()) &&
15747 !TLI.isLoadLegal(ValVT: VT, MemVT: Load->getValueType(ResNo: 0), Alignment: Load->getAlign(),
15748 AddrSpace: Load->getAddressSpace(), ExtType: ExtLoadType, Atomic: false))
15749 return {};
15750
15751 bool DoXform = true;
15752 SmallVector<SDNode *, 4> SetCCs;
15753 if (!N0->hasOneUse())
15754 DoXform = ExtendUsesToFormExtLoad(VT, N, N0: Frozen ? Freeze : SDValue(Load, 0),
15755 ExtOpc, ExtendNodes&: SetCCs, TLI);
15756 if (VT.isVector())
15757 DoXform &= TLI.isVectorLoadExtDesirable(ExtVal: SDValue(N, 0));
15758 if (!DoXform)
15759 return {};
15760
15761 SDLoc DL(Load);
15762
15763 auto SalvageDbgValue = [&](SDDbgValue *Dbg, SDValue Old, SDValue New,
15764 unsigned OldBits, unsigned NewBits,
15765 bool IsSigned) {
15766 SmallVector<SDDbgOperand> Locs = Dbg->copyLocationOps();
15767 bool Changed = false;
15768
15769 bool IsVariadic = Dbg->isVariadic();
15770 SmallVector<unsigned, 2> AffectedArgs;
15771
15772 for (unsigned I = 0, E = Locs.size(); I != E; ++I) {
15773 SDDbgOperand &Op = Locs[I];
15774 if (Op.getKind() != SDDbgOperand::SDNODE)
15775 continue;
15776
15777 if (Op.getSDNode() == Old.getNode() && Op.getResNo() == Old.getResNo()) {
15778 Op = SDDbgOperand::fromNode(Node: New.getNode(), ResNo: New.getResNo());
15779 Changed = true;
15780
15781 if (IsVariadic)
15782 AffectedArgs.push_back(Elt: I);
15783 }
15784 }
15785
15786 if (!Changed)
15787 return;
15788
15789 const DIExpression *OldExpr = Dbg->getExpression();
15790 const DIExpression *NewExpr = nullptr;
15791
15792 if (!IsVariadic) {
15793 // Do not introduce DW_OP_LLVM_arg into ordinary single-location
15794 // DBG_VALUEs.
15795 NewExpr = DIExpression::appendExt(Expr: OldExpr, FromSize: NewBits, ToSize: OldBits, Signed: IsSigned);
15796 } else {
15797 auto ExtOps = DIExpression::getExtOps(FromSize: NewBits, ToSize: OldBits, Signed: IsSigned);
15798
15799 NewExpr = DIExpression::convertToVariadicExpression(Expr: OldExpr);
15800
15801 for (unsigned ArgNo : AffectedArgs)
15802 NewExpr = DIExpression::appendOpsToArg(Expr: NewExpr, Ops: ExtOps, ArgNo,
15803 /*StackValue=*/false);
15804 }
15805
15806 SDDbgValue *NewDV = DAG.getDbgValueList(
15807 Var: Dbg->getVariable(), Expr: const_cast<DIExpression *>(NewExpr), Locs,
15808 Dependencies: Dbg->getAdditionalDependencies(), IsIndirect: Dbg->isIndirect(), DL: Dbg->getDebugLoc(),
15809 O: Dbg->getOrder(), IsVariadic: Dbg->isVariadic());
15810
15811 Dbg->setIsInvalidated();
15812 Dbg->setIsEmitted();
15813 DAG.AddDbgValue(DB: NewDV, /*isParameter=*/false);
15814 };
15815
15816 // Because we are replacing a load and a s|z ext with a load-s|z ext
15817 // instruction, the dbg_value attached to the load will be of a smaller bit
15818 // width, and we have to add a DW_OP_LLVM_convert expression to get the
15819 // correct size.
15820 auto SalvageToOldLoadSize = [&](SDValue Old, SDValue New, bool IsSigned) {
15821 SmallVector<SDDbgValue *, 4> DbgVals(
15822 DAG.GetDbgValues(SD: Old.getNode()).begin(),
15823 DAG.GetDbgValues(SD: Old.getNode()).end());
15824
15825 unsigned VarBitsOld = Old.getValueSizeInBits();
15826 unsigned VarBitsNew = New.getValueSizeInBits();
15827
15828 for (SDDbgValue *Dbg : DbgVals) {
15829 if (Dbg->isInvalidated())
15830 continue;
15831
15832 SalvageDbgValue(Dbg, Old, New, VarBitsOld, VarBitsNew, IsSigned);
15833 }
15834 };
15835
15836 SDValue ExtLoad =
15837 DAG.getExtLoad(ExtType: ExtLoadType, dl: DL, VT, Chain: Load->getChain(), Ptr: Load->getBasePtr(),
15838 MemVT: Load->getValueType(ResNo: 0), MMO: Load->getMemOperand());
15839 SDValue Res = ExtLoad;
15840 if (Frozen) {
15841 Res = DAG.getFreeze(V: ExtLoad);
15842 Res = DAG.getNode(Opcode: ExtLoadType == ISD::SEXTLOAD ? ISD::AssertSext
15843 : ISD::AssertZext,
15844 DL, VT: Res.getValueType(), N1: Res,
15845 N2: DAG.getValueType(Load->getValueType(ResNo: 0).getScalarType()));
15846 }
15847 Combiner.ExtendSetCCUses(SetCCs, OrigLoad: N0, ExtLoad: Res, ExtType: ExtOpc);
15848 // If the load value is used only by N, replace it via CombineTo N.
15849 bool NoReplaceTrunc = N0.hasOneUse();
15850 if (N->getHasDebugValue()) {
15851 SDValue OldExtValue(N, 0);
15852 DAG.transferDbgValues(From: OldExtValue, To: ExtLoad);
15853 }
15854 if (NoReplaceTrunc) {
15855 bool IsSigned = N->getOpcode() == ISD::SIGN_EXTEND;
15856 if (Load->getHasDebugValue()) {
15857 SDValue OldLoadVal(Load, 0);
15858 SalvageToOldLoadSize(OldLoadVal, ExtLoad, IsSigned);
15859 }
15860 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Load, 1), To: ExtLoad.getValue(R: 1));
15861 Combiner.CombineTo(N, Res);
15862 Combiner.recursivelyDeleteUnusedNodes(N: N0.getNode());
15863 } else {
15864 Combiner.CombineTo(N, Res);
15865 SDValue Trunc = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: Load->getValueType(ResNo: 0), Operand: Res);
15866 if (Frozen) {
15867 Combiner.CombineTo(N: Freeze.getNode(), Res: Trunc);
15868 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Load, 1), To: ExtLoad.getValue(R: 1));
15869 } else {
15870 Combiner.CombineTo(N: Load, Res0: Trunc, Res1: ExtLoad.getValue(R: 1));
15871 }
15872 }
15873 return SDValue(N, 0); // Return N so it doesn't get rechecked!
15874}
15875
15876static SDValue
15877tryToFoldExtOfMaskedLoad(SelectionDAG &DAG, const TargetLowering &TLI, EVT VT,
15878 bool LegalOperations, SDNode *N, SDValue N0,
15879 ISD::LoadExtType ExtLoadType, ISD::NodeType ExtOpc) {
15880 if (!N0.hasOneUse())
15881 return SDValue();
15882
15883 MaskedLoadSDNode *Ld = dyn_cast<MaskedLoadSDNode>(Val&: N0);
15884 if (!Ld || Ld->getExtensionType() != ISD::NON_EXTLOAD)
15885 return SDValue();
15886
15887 if ((LegalOperations || !cast<MaskedLoadSDNode>(Val&: N0)->isSimple()) &&
15888 !TLI.isLoadLegalOrCustom(ValVT: VT, MemVT: Ld->getValueType(ResNo: 0), Alignment: Ld->getAlign(),
15889 AddrSpace: Ld->getAddressSpace(), ExtType: ExtLoadType, Atomic: false))
15890 return SDValue();
15891
15892 if (!TLI.isVectorLoadExtDesirable(ExtVal: SDValue(N, 0)))
15893 return SDValue();
15894
15895 SDLoc dl(Ld);
15896 SDValue PassThru = DAG.getNode(Opcode: ExtOpc, DL: dl, VT, Operand: Ld->getPassThru());
15897 SDValue NewLoad = DAG.getMaskedLoad(
15898 VT, dl, Chain: Ld->getChain(), Base: Ld->getBasePtr(), Offset: Ld->getOffset(), Mask: Ld->getMask(),
15899 Src0: PassThru, MemVT: Ld->getMemoryVT(), MMO: Ld->getMemOperand(), AM: Ld->getAddressingMode(),
15900 ExtLoadType, IsExpanding: Ld->isExpandingLoad());
15901 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Ld, 1), To: SDValue(NewLoad.getNode(), 1));
15902 return NewLoad;
15903}
15904
15905// fold ([s|z]ext (atomic_load)) -> ([s|z]ext (truncate ([s|z]ext atomic_load)))
15906static SDValue tryToFoldExtOfAtomicLoad(SelectionDAG &DAG,
15907 const TargetLowering &TLI, EVT VT,
15908 SDValue N0,
15909 ISD::LoadExtType ExtLoadType) {
15910 auto *ALoad = dyn_cast<AtomicSDNode>(Val&: N0);
15911 if (!ALoad || ALoad->getOpcode() != ISD::ATOMIC_LOAD)
15912 return {};
15913 EVT MemoryVT = ALoad->getMemoryVT();
15914 if (!TLI.isLoadLegal(ValVT: VT, MemVT: MemoryVT, Alignment: ALoad->getAlign(),
15915 AddrSpace: ALoad->getAddressSpace(), ExtType: ExtLoadType, Atomic: true))
15916 return {};
15917 // Can't fold into ALoad if it is already extending differently.
15918 ISD::LoadExtType ALoadExtTy = ALoad->getExtensionType();
15919 if ((ALoadExtTy == ISD::ZEXTLOAD && ExtLoadType == ISD::SEXTLOAD) ||
15920 (ALoadExtTy == ISD::SEXTLOAD && ExtLoadType == ISD::ZEXTLOAD))
15921 return {};
15922
15923 EVT OrigVT = ALoad->getValueType(ResNo: 0);
15924 assert(OrigVT.getSizeInBits() < VT.getSizeInBits() && "VT should be wider.");
15925 auto *NewALoad = cast<AtomicSDNode>(Val: DAG.getAtomicLoad(
15926 ExtType: ExtLoadType, dl: SDLoc(ALoad), MemVT: MemoryVT, VT, Chain: ALoad->getChain(),
15927 Ptr: ALoad->getBasePtr(), MMO: ALoad->getMemOperand()));
15928 DAG.ReplaceAllUsesOfValueWith(
15929 From: SDValue(ALoad, 0),
15930 To: DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(ALoad), VT: OrigVT, Operand: SDValue(NewALoad, 0)));
15931 // Update the chain uses.
15932 DAG.ReplaceAllUsesOfValueWith(From: SDValue(ALoad, 1), To: SDValue(NewALoad, 1));
15933 return SDValue(NewALoad, 0);
15934}
15935
15936static SDValue foldExtendedSignBitTest(SDNode *N, SelectionDAG &DAG,
15937 bool LegalOperations) {
15938 assert((N->getOpcode() == ISD::SIGN_EXTEND ||
15939 N->getOpcode() == ISD::ZERO_EXTEND) && "Expected sext or zext");
15940
15941 SDValue SetCC = N->getOperand(Num: 0);
15942 if (LegalOperations || SetCC.getOpcode() != ISD::SETCC ||
15943 !SetCC.hasOneUse() || SetCC.getValueType() != MVT::i1)
15944 return SDValue();
15945
15946 SDValue X = SetCC.getOperand(i: 0);
15947 SDValue Ones = SetCC.getOperand(i: 1);
15948 ISD::CondCode CC = cast<CondCodeSDNode>(Val: SetCC.getOperand(i: 2))->get();
15949 EVT VT = N->getValueType(ResNo: 0);
15950 EVT XVT = X.getValueType();
15951 // setge X, C is canonicalized to setgt, so we do not need to match that
15952 // pattern. The setlt sibling is folded in SimplifySelectCC() because it does
15953 // not require the 'not' op.
15954 if (CC == ISD::SETGT && isAllOnesConstant(V: Ones) && VT == XVT) {
15955 // Invert and smear/shift the sign bit:
15956 // sext i1 (setgt iN X, -1) --> sra (not X), (N - 1)
15957 // zext i1 (setgt iN X, -1) --> srl (not X), (N - 1)
15958 SDLoc DL(N);
15959 unsigned ShCt = VT.getSizeInBits() - 1;
15960 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
15961 if (!TLI.shouldAvoidTransformToShift(VT, Amount: ShCt)) {
15962 SDValue NotX = DAG.getNOT(DL, Val: X, VT);
15963 SDValue ShiftAmount = DAG.getConstant(Val: ShCt, DL, VT);
15964 auto ShiftOpcode =
15965 N->getOpcode() == ISD::SIGN_EXTEND ? ISD::SRA : ISD::SRL;
15966 return DAG.getNode(Opcode: ShiftOpcode, DL, VT, N1: NotX, N2: ShiftAmount);
15967 }
15968 }
15969 return SDValue();
15970}
15971
15972SDValue DAGCombiner::foldSextSetcc(SDNode *N) {
15973 SDValue N0 = N->getOperand(Num: 0);
15974 if (N0.getOpcode() != ISD::SETCC)
15975 return SDValue();
15976
15977 SDValue N00 = N0.getOperand(i: 0);
15978 SDValue N01 = N0.getOperand(i: 1);
15979 ISD::CondCode CC = cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get();
15980 EVT VT = N->getValueType(ResNo: 0);
15981 EVT N00VT = N00.getValueType();
15982 SDLoc DL(N);
15983
15984 // Propagate fast-math-flags.
15985 SDNodeFlags Flags = N0->getFlags();
15986
15987 // On some architectures (such as SSE/NEON/etc) the SETCC result type is
15988 // the same size as the compared operands. Try to optimize sext(setcc())
15989 // if this is the case.
15990 if (VT.isVector() && !LegalOperations &&
15991 TLI.getBooleanContents(Type: N00VT) ==
15992 TargetLowering::ZeroOrNegativeOneBooleanContent) {
15993 EVT SVT = getSetCCResultType(VT: N00VT);
15994
15995 // If we already have the desired type, don't change it.
15996 if (SVT != N0.getValueType()) {
15997 // We know that the # elements of the results is the same as the
15998 // # elements of the compare (and the # elements of the compare result
15999 // for that matter). Check to see that they are the same size. If so,
16000 // we know that the element size of the sext'd result matches the
16001 // element size of the compare operands.
16002 if (VT.getSizeInBits() == SVT.getSizeInBits())
16003 return DAG.getSetCC(DL, VT, LHS: N00, RHS: N01, Cond: CC, /*Chain=*/{},
16004 /*Signaling=*/IsSignaling: false, Flags);
16005
16006 // If the desired elements are smaller or larger than the source
16007 // elements, we can use a matching integer vector type and then
16008 // truncate/sign extend.
16009 EVT MatchingVecType = N00VT.changeVectorElementTypeToInteger();
16010 if (SVT == MatchingVecType) {
16011 SDValue VsetCC = DAG.getSetCC(DL, VT: MatchingVecType, LHS: N00, RHS: N01, Cond: CC,
16012 /*Chain=*/{}, /*Signaling=*/IsSignaling: false, Flags);
16013 return DAG.getSExtOrTrunc(Op: VsetCC, DL, VT);
16014 }
16015 }
16016
16017 // Try to eliminate the sext of a setcc by zexting the compare operands.
16018 if (N0.hasOneUse() && TLI.isOperationLegalOrCustom(Op: ISD::SETCC, VT) &&
16019 !TLI.isOperationLegalOrCustom(Op: ISD::SETCC, VT: SVT)) {
16020 bool IsSignedCmp = ISD::isSignedIntSetCC(Code: CC);
16021 unsigned LoadOpcode = IsSignedCmp ? ISD::SEXTLOAD : ISD::ZEXTLOAD;
16022 unsigned ExtOpcode = IsSignedCmp ? ISD::SIGN_EXTEND : ISD::ZERO_EXTEND;
16023
16024 // We have an unsupported narrow vector compare op that would be legal
16025 // if extended to the destination type. See if the compare operands
16026 // can be freely extended to the destination type.
16027 auto IsFreeToExtend = [&](SDValue V) {
16028 if (isConstantOrConstantVector(N: V, /*NoOpaques*/ true))
16029 return true;
16030 // Match a simple, non-extended load that can be converted to a
16031 // legal {z/s}ext-load.
16032 // TODO: Allow widening of an existing {z/s}ext-load?
16033 if (!(ISD::isNON_EXTLoad(N: V.getNode()) &&
16034 ISD::isUNINDEXEDLoad(N: V.getNode())))
16035 return false;
16036
16037 LoadSDNode *Ld = cast<LoadSDNode>(Val: V.getNode());
16038
16039 if (!Ld->isSimple() ||
16040 !TLI.isLoadLegal(ValVT: VT, MemVT: V.getValueType(), Alignment: Ld->getAlign(),
16041 AddrSpace: Ld->getAddressSpace(), ExtType: LoadOpcode, Atomic: false))
16042 return false;
16043
16044 // Non-chain users of this value must either be the setcc in this
16045 // sequence or extends that can be folded into the new {z/s}ext-load.
16046 for (SDUse &Use : V->uses()) {
16047 // Skip uses of the chain and the setcc.
16048 SDNode *User = Use.getUser();
16049 if (Use.getResNo() != 0 || User == N0.getNode())
16050 continue;
16051 // Extra users must have exactly the same cast we are about to create.
16052 // TODO: This restriction could be eased if ExtendUsesToFormExtLoad()
16053 // is enhanced similarly.
16054 if (User->getOpcode() != ExtOpcode || User->getValueType(ResNo: 0) != VT)
16055 return false;
16056 }
16057 return true;
16058 };
16059
16060 if (IsFreeToExtend(N00) && IsFreeToExtend(N01)) {
16061 SDValue Ext0 = DAG.getNode(Opcode: ExtOpcode, DL, VT, Operand: N00);
16062 SDValue Ext1 = DAG.getNode(Opcode: ExtOpcode, DL, VT, Operand: N01);
16063 return DAG.getSetCC(DL, VT, LHS: Ext0, RHS: Ext1, Cond: CC, /*Chain=*/{},
16064 /*Signaling=*/IsSignaling: false, Flags);
16065 }
16066 }
16067 }
16068
16069 // sext(setcc x, y, cc) -> (select (setcc x, y, cc), T, 0)
16070 // Here, T can be 1 or -1, depending on the type of the setcc and
16071 // getBooleanContents().
16072 unsigned SetCCWidth = N0.getScalarValueSizeInBits();
16073
16074 // To determine the "true" side of the select, we need to know the high bit
16075 // of the value returned by the setcc if it evaluates to true.
16076 // If the type of the setcc is i1, then the true case of the select is just
16077 // sext(i1 1), that is, -1.
16078 // If the type of the setcc is larger (say, i8) then the value of the high
16079 // bit depends on getBooleanContents(), so ask TLI for a real "true" value
16080 // of the appropriate width.
16081 SDValue ExtTrueVal = (SetCCWidth == 1)
16082 ? DAG.getAllOnesConstant(DL, VT)
16083 : DAG.getBoolConstant(V: true, DL, VT, OpVT: N00VT);
16084 SDValue Zero = DAG.getConstant(Val: 0, DL, VT);
16085 if (SDValue SCC = SimplifySelectCC(DL, N0: N00, N1: N01, N2: ExtTrueVal, N3: Zero, CC, NotExtCompare: true))
16086 return SCC;
16087
16088 if (!VT.isVector() && !shouldConvertSelectOfConstantsToMath(Cond: N0, VT, TLI)) {
16089 EVT SetCCVT = getSetCCResultType(VT: N00VT);
16090 // Don't do this transform for i1 because there's a select transform
16091 // that would reverse it.
16092 // TODO: We should not do this transform at all without a target hook
16093 // because a sext is likely cheaper than a select?
16094 if (SetCCVT.getScalarSizeInBits() != 1 &&
16095 (!LegalOperations || TLI.isOperationLegal(Op: ISD::SETCC, VT: N00VT))) {
16096 SDValue SetCC = DAG.getSetCC(DL, VT: SetCCVT, LHS: N00, RHS: N01, Cond: CC, /*Chain=*/{},
16097 /*Signaling=*/IsSignaling: false, Flags);
16098 return DAG.getSelect(DL, VT, Cond: SetCC, LHS: ExtTrueVal, RHS: Zero, Flags);
16099 }
16100 }
16101
16102 return SDValue();
16103}
16104
16105SDValue DAGCombiner::visitSIGN_EXTEND(SDNode *N) {
16106 SDValue N0 = N->getOperand(Num: 0);
16107 EVT VT = N->getValueType(ResNo: 0);
16108 SDLoc DL(N);
16109
16110 if (VT.isVector())
16111 if (SDValue FoldedVOp = SimplifyVCastOp(N, DL))
16112 return FoldedVOp;
16113
16114 // sext(undef) = 0 because the top bit will all be the same.
16115 if (N0.isUndef())
16116 return DAG.getConstant(Val: 0, DL, VT);
16117
16118 if (SDValue Res = tryToFoldExtendOfConstant(N, DL, TLI, DAG, LegalTypes))
16119 return Res;
16120
16121 // fold (sext (sext x)) -> (sext x)
16122 // fold (sext (aext x)) -> (sext x)
16123 if (N0.getOpcode() == ISD::SIGN_EXTEND || N0.getOpcode() == ISD::ANY_EXTEND)
16124 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: N0.getOperand(i: 0));
16125
16126 // fold (sext (aext_extend_vector_inreg x)) -> (sext_extend_vector_inreg x)
16127 // fold (sext (sext_extend_vector_inreg x)) -> (sext_extend_vector_inreg x)
16128 if (N0.getOpcode() == ISD::ANY_EXTEND_VECTOR_INREG ||
16129 N0.getOpcode() == ISD::SIGN_EXTEND_VECTOR_INREG)
16130 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_VECTOR_INREG, DL: SDLoc(N), VT,
16131 Operand: N0.getOperand(i: 0));
16132
16133 if (N0.getOpcode() == ISD::SIGN_EXTEND_INREG) {
16134 SDValue N00 = N0.getOperand(i: 0);
16135 EVT ExtVT = cast<VTSDNode>(Val: N0->getOperand(Num: 1))->getVT();
16136 if (N00.getOpcode() == ISD::TRUNCATE || TLI.isTruncateFree(Val: N00, VT2: ExtVT)) {
16137 // fold (sext (sext_inreg x)) -> (sext (trunc x))
16138 if ((!LegalTypes || TLI.isTypeLegal(VT: ExtVT))) {
16139 SDValue T = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: ExtVT, Operand: N00);
16140 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: T);
16141 }
16142
16143 // If the trunc wasn't legal, try to fold to (sext_inreg (anyext x))
16144 if (!LegalTypes || TLI.isTypeLegal(VT)) {
16145 SDValue ExtSrc = DAG.getAnyExtOrTrunc(Op: N00, DL, VT);
16146 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT, N1: ExtSrc,
16147 N2: N0->getOperand(Num: 1));
16148 }
16149 }
16150 }
16151
16152 if (N0.getOpcode() == ISD::TRUNCATE) {
16153 // fold (sext (truncate (load x))) -> (sext (smaller load x))
16154 // fold (sext (truncate (srl (load x), c))) -> (sext (smaller load (x+c/n)))
16155 if (SDValue NarrowLoad = reduceLoadWidth(N: N0.getNode())) {
16156 SDNode *oye = N0.getOperand(i: 0).getNode();
16157 if (NarrowLoad.getNode() != N0.getNode()) {
16158 CombineTo(N: N0.getNode(), Res: NarrowLoad);
16159 // CombineTo deleted the truncate, if needed, but not what's under it.
16160 AddToWorklist(N: oye);
16161 }
16162 return SDValue(N, 0); // Return N so it doesn't get rechecked!
16163 }
16164
16165 // See if the value being truncated is already sign extended. If so, just
16166 // eliminate the trunc/sext pair.
16167 SDValue Op = N0.getOperand(i: 0);
16168 unsigned OpBits = Op.getScalarValueSizeInBits();
16169 unsigned MidBits = N0.getScalarValueSizeInBits();
16170 unsigned DestBits = VT.getScalarSizeInBits();
16171
16172 if (N0->getFlags().hasNoSignedWrap() ||
16173 DAG.ComputeNumSignBits(Op) > OpBits - MidBits) {
16174 if (OpBits == DestBits) {
16175 // Op is i32, Mid is i8, and Dest is i32. If Op has more than 24 sign
16176 // bits, it is already ready.
16177 return Op;
16178 }
16179
16180 if (OpBits < DestBits) {
16181 // Op is i32, Mid is i8, and Dest is i64. If Op has more than 24 sign
16182 // bits, just sext from i32.
16183 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: Op);
16184 }
16185
16186 // Op is i64, Mid is i8, and Dest is i32. If Op has more than 56 sign
16187 // bits, just truncate to i32.
16188 SDNodeFlags Flags;
16189 Flags.setNoSignedWrap(true);
16190 Flags.setNoUnsignedWrap(N0->getFlags().hasNoUnsignedWrap());
16191 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: Op, Flags);
16192 }
16193
16194 // fold (sext (truncate x)) -> (sextinreg x).
16195 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::SIGN_EXTEND_INREG,
16196 VT: N0.getValueType())) {
16197 if (OpBits < DestBits)
16198 Op = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL: SDLoc(N0), VT, Operand: Op);
16199 else if (OpBits > DestBits)
16200 Op = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(N0), VT, Operand: Op);
16201 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT, N1: Op,
16202 N2: DAG.getValueType(N0.getValueType()));
16203 }
16204 }
16205
16206 // Try to simplify (sext (load x)).
16207 if (SDValue foldedExt =
16208 tryToFoldExtOfLoad(DAG, Combiner&: *this, TLI, VT, LegalOperations, N, N0,
16209 ExtLoadType: ISD::SEXTLOAD, ExtOpc: ISD::SIGN_EXTEND))
16210 return foldedExt;
16211
16212 if (SDValue foldedExt =
16213 tryToFoldExtOfMaskedLoad(DAG, TLI, VT, LegalOperations, N, N0,
16214 ExtLoadType: ISD::SEXTLOAD, ExtOpc: ISD::SIGN_EXTEND))
16215 return foldedExt;
16216
16217 // fold (sext (load x)) to multiple smaller sextloads.
16218 // Only on illegal but splittable vectors.
16219 if (SDValue ExtLoad = CombineExtLoad(N))
16220 return ExtLoad;
16221
16222 // Try to simplify (sext (sextload x)).
16223 if (SDValue foldedExt = tryToFoldExtOfExtload(
16224 DAG, Combiner&: *this, TLI, VT, LegalOperations, N, N0, ExtLoadType: ISD::SEXTLOAD))
16225 return foldedExt;
16226
16227 // Try to simplify (sext (atomic_load x)).
16228 if (SDValue foldedExt =
16229 tryToFoldExtOfAtomicLoad(DAG, TLI, VT, N0, ExtLoadType: ISD::SEXTLOAD))
16230 return foldedExt;
16231
16232 // fold (sext (and/or/xor (load x), cst)) ->
16233 // (and/or/xor (sextload x), (sext cst))
16234 if (ISD::isBitwiseLogicOp(Opcode: N0.getOpcode()) &&
16235 isa<LoadSDNode>(Val: N0.getOperand(i: 0)) &&
16236 N0.getOperand(i: 1).getOpcode() == ISD::Constant &&
16237 (!LegalOperations && TLI.isOperationLegal(Op: N0.getOpcode(), VT))) {
16238 LoadSDNode *LN00 = cast<LoadSDNode>(Val: N0.getOperand(i: 0));
16239 EVT MemVT = LN00->getMemoryVT();
16240 if (TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: LN00->getAlign(), AddrSpace: LN00->getAddressSpace(),
16241 ExtType: ISD::SEXTLOAD, Atomic: false) &&
16242 LN00->getExtensionType() != ISD::ZEXTLOAD && LN00->isUnindexed()) {
16243 SmallVector<SDNode*, 4> SetCCs;
16244 bool DoXform = ExtendUsesToFormExtLoad(VT, N: N0.getNode(), N0: N0.getOperand(i: 0),
16245 ExtOpc: ISD::SIGN_EXTEND, ExtendNodes&: SetCCs, TLI);
16246 if (DoXform) {
16247 SDValue ExtLoad = DAG.getExtLoad(ExtType: ISD::SEXTLOAD, dl: SDLoc(LN00), VT,
16248 Chain: LN00->getChain(), Ptr: LN00->getBasePtr(),
16249 MemVT: LN00->getMemoryVT(),
16250 MMO: LN00->getMemOperand());
16251 APInt Mask = N0.getConstantOperandAPInt(i: 1).sext(width: VT.getSizeInBits());
16252 SDValue And = DAG.getNode(Opcode: N0.getOpcode(), DL, VT,
16253 N1: ExtLoad, N2: DAG.getConstant(Val: Mask, DL, VT));
16254 ExtendSetCCUses(SetCCs, OrigLoad: N0.getOperand(i: 0), ExtLoad, ExtType: ISD::SIGN_EXTEND);
16255 bool NoReplaceTruncAnd = !N0.hasOneUse();
16256 bool NoReplaceTrunc = SDValue(LN00, 0).hasOneUse();
16257 CombineTo(N, Res: And);
16258 // If N0 has multiple uses, change other uses as well.
16259 if (NoReplaceTruncAnd) {
16260 SDValue TruncAnd =
16261 DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: N0.getValueType(), Operand: And);
16262 CombineTo(N: N0.getNode(), Res: TruncAnd);
16263 }
16264 if (NoReplaceTrunc) {
16265 DAG.ReplaceAllUsesOfValueWith(From: SDValue(LN00, 1), To: ExtLoad.getValue(R: 1));
16266 } else {
16267 SDValue Trunc = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(LN00),
16268 VT: LN00->getValueType(ResNo: 0), Operand: ExtLoad);
16269 CombineTo(N: LN00, Res0: Trunc, Res1: ExtLoad.getValue(R: 1));
16270 }
16271 return SDValue(N,0); // Return N so it doesn't get rechecked!
16272 }
16273 }
16274 }
16275
16276 if (SDValue V = foldExtendedSignBitTest(N, DAG, LegalOperations))
16277 return V;
16278
16279 if (SDValue V = foldSextSetcc(N))
16280 return V;
16281
16282 // fold (sext x) -> (zext x) if the sign bit is known zero.
16283 if (!TLI.isSExtCheaperThanZExt(FromTy: N0.getValueType(), ToTy: VT) &&
16284 (!LegalOperations || TLI.isOperationLegal(Op: ISD::ZERO_EXTEND, VT)) &&
16285 DAG.SignBitIsZero(Op: N0))
16286 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: N0, Flags: SDNodeFlags::NonNeg);
16287
16288 if (SDValue NewVSel = matchVSelectOpSizesWithSetCC(Cast: N))
16289 return NewVSel;
16290
16291 // Eliminate this sign extend by doing a negation in the destination type:
16292 // sext i32 (0 - (zext i8 X to i32)) to i64 --> 0 - (zext i8 X to i64)
16293 if (N0.getOpcode() == ISD::SUB && N0.hasOneUse() &&
16294 isNullOrNullSplat(V: N0.getOperand(i: 0)) &&
16295 N0.getOperand(i: 1).getOpcode() == ISD::ZERO_EXTEND &&
16296 TLI.isOperationLegalOrCustom(Op: ISD::SUB, VT)) {
16297 SDValue Zext = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 1).getOperand(i: 0), DL, VT);
16298 return DAG.getNegative(Val: Zext, DL, VT);
16299 }
16300 // Eliminate this sign extend by doing a decrement in the destination type:
16301 // sext i32 ((zext i8 X to i32) + (-1)) to i64 --> (zext i8 X to i64) + (-1)
16302 if (N0.getOpcode() == ISD::ADD && N0.hasOneUse() &&
16303 isAllOnesOrAllOnesSplat(V: N0.getOperand(i: 1)) &&
16304 N0.getOperand(i: 0).getOpcode() == ISD::ZERO_EXTEND &&
16305 TLI.isOperationLegalOrCustom(Op: ISD::ADD, VT)) {
16306 SDValue Zext = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 0).getOperand(i: 0), DL, VT);
16307 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Zext, N2: DAG.getAllOnesConstant(DL, VT));
16308 }
16309
16310 // fold sext (not i1 X) -> add (zext i1 X), -1
16311 // TODO: This could be extended to handle bool vectors.
16312 if (N0.getValueType() == MVT::i1 && isBitwiseNot(V: N0) && N0.hasOneUse() &&
16313 (!LegalOperations || (TLI.isOperationLegal(Op: ISD::ZERO_EXTEND, VT) &&
16314 TLI.isOperationLegal(Op: ISD::ADD, VT)))) {
16315 // If we can eliminate the 'not', the sext form should be better
16316 if (SDValue NewXor = visitXOR(N: N0.getNode())) {
16317 // Returning N0 is a form of in-visit replacement that may have
16318 // invalidated N0.
16319 if (NewXor.getNode() == N0.getNode()) {
16320 // Return SDValue here as the xor should have already been replaced in
16321 // this sext.
16322 return SDValue();
16323 }
16324
16325 // Return a new sext with the new xor.
16326 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: NewXor);
16327 }
16328
16329 SDValue Zext = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: N0.getOperand(i: 0));
16330 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Zext, N2: DAG.getAllOnesConstant(DL, VT));
16331 }
16332
16333 if (SDValue Res = tryToFoldExtendSelectLoad(N, TLI, DAG, DL, Level))
16334 return Res;
16335
16336 return SDValue();
16337}
16338
16339/// Given an extending node with a pop-count operand, if the target does not
16340/// support a pop-count in the narrow source type but does support it in the
16341/// destination type, widen the pop-count to the destination type.
16342static SDValue widenCtPop(SDNode *Extend, SelectionDAG &DAG, const SDLoc &DL) {
16343 assert((Extend->getOpcode() == ISD::ZERO_EXTEND ||
16344 Extend->getOpcode() == ISD::ANY_EXTEND) &&
16345 "Expected extend op");
16346
16347 SDValue CtPop = Extend->getOperand(Num: 0);
16348 if (CtPop.getOpcode() != ISD::CTPOP || !CtPop.hasOneUse())
16349 return SDValue();
16350
16351 EVT VT = Extend->getValueType(ResNo: 0);
16352 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
16353 if (TLI.isOperationLegalOrCustom(Op: ISD::CTPOP, VT: CtPop.getValueType()) ||
16354 !TLI.isOperationLegalOrCustom(Op: ISD::CTPOP, VT))
16355 return SDValue();
16356
16357 // zext (ctpop X) --> ctpop (zext X)
16358 SDValue NewZext = DAG.getZExtOrTrunc(Op: CtPop.getOperand(i: 0), DL, VT);
16359 return DAG.getNode(Opcode: ISD::CTPOP, DL, VT, Operand: NewZext);
16360}
16361
16362// If we have (zext (abs X)) where X is a type that will be promoted by type
16363// legalization, convert to (abs_min_poison (sext X)). But do not extend
16364// past a legal type.
16365static SDValue widenAbs(SDNode *Extend, SelectionDAG &DAG) {
16366 assert(Extend->getOpcode() == ISD::ZERO_EXTEND && "Expected zero extend.");
16367
16368 EVT VT = Extend->getValueType(ResNo: 0);
16369 if (VT.isVector())
16370 return SDValue();
16371
16372 SDValue Abs = Extend->getOperand(Num: 0);
16373 if (!ISD::isAbsOpcode(Opcode: Abs.getOpcode()) || !Abs.hasOneUse())
16374 return SDValue();
16375
16376 EVT AbsVT = Abs.getValueType();
16377 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
16378 if (TLI.getTypeAction(Context&: *DAG.getContext(), VT: AbsVT) !=
16379 TargetLowering::TypePromoteInteger)
16380 return SDValue();
16381
16382 EVT LegalVT = TLI.getTypeToTransformTo(Context&: *DAG.getContext(), VT: AbsVT);
16383
16384 SDValue SExt =
16385 DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL: SDLoc(Abs), VT: LegalVT, Operand: Abs.getOperand(i: 0));
16386 SDValue NewAbs = DAG.getNode(Opcode: ISD::ABS_MIN_POISON, DL: SDLoc(Abs), VT: LegalVT, Operand: SExt);
16387 return DAG.getZExtOrTrunc(Op: NewAbs, DL: SDLoc(Extend), VT);
16388}
16389
16390SDValue DAGCombiner::visitZERO_EXTEND(SDNode *N) {
16391 SDValue N0 = N->getOperand(Num: 0);
16392 EVT VT = N->getValueType(ResNo: 0);
16393 SDLoc DL(N);
16394
16395 if (VT.isVector())
16396 if (SDValue FoldedVOp = SimplifyVCastOp(N, DL))
16397 return FoldedVOp;
16398
16399 // zext(undef) = 0
16400 if (N0.isUndef())
16401 return DAG.getConstant(Val: 0, DL, VT);
16402
16403 if (SDValue Res = tryToFoldExtendOfConstant(N, DL, TLI, DAG, LegalTypes))
16404 return Res;
16405
16406 // fold (zext (zext x)) -> (zext x)
16407 // fold (zext (aext x)) -> (zext x)
16408 if (N0.getOpcode() == ISD::ZERO_EXTEND || N0.getOpcode() == ISD::ANY_EXTEND) {
16409 SDNodeFlags Flags;
16410 if (N0.getOpcode() == ISD::ZERO_EXTEND)
16411 Flags.setNonNeg(N0->getFlags().hasNonNeg());
16412 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: N0.getOperand(i: 0), Flags);
16413 }
16414
16415 // fold (zext (aext_extend_vector_inreg x)) -> (zext_extend_vector_inreg x)
16416 // fold (zext (zext_extend_vector_inreg x)) -> (zext_extend_vector_inreg x)
16417 if (N0.getOpcode() == ISD::ANY_EXTEND_VECTOR_INREG ||
16418 N0.getOpcode() == ISD::ZERO_EXTEND_VECTOR_INREG)
16419 return DAG.getNode(Opcode: ISD::ZERO_EXTEND_VECTOR_INREG, DL, VT, Operand: N0.getOperand(i: 0));
16420
16421 // fold (zext (truncate x)) -> (zext x) or
16422 // (zext (truncate x)) -> (truncate x)
16423 // This is valid when the truncated bits of x are already zero.
16424 SDValue Op;
16425 KnownBits Known;
16426 if (isTruncateOf(DAG, N: N0, Op, Known)) {
16427 APInt TruncatedBits =
16428 (Op.getScalarValueSizeInBits() == N0.getScalarValueSizeInBits()) ?
16429 APInt(Op.getScalarValueSizeInBits(), 0) :
16430 APInt::getBitsSet(numBits: Op.getScalarValueSizeInBits(),
16431 loBit: N0.getScalarValueSizeInBits(),
16432 hiBit: std::min(a: Op.getScalarValueSizeInBits(),
16433 b: VT.getScalarSizeInBits()));
16434 if (TruncatedBits.isSubsetOf(RHS: Known.Zero)) {
16435 SDValue ZExtOrTrunc = DAG.getZExtOrTrunc(Op, DL, VT);
16436 DAG.salvageDebugInfo(N&: *N0.getNode());
16437
16438 return ZExtOrTrunc;
16439 }
16440 }
16441
16442 // fold (zext (truncate x)) -> (and x, mask)
16443 if (N0.getOpcode() == ISD::TRUNCATE) {
16444 // fold (zext (truncate (load x))) -> (zext (smaller load x))
16445 // fold (zext (truncate (srl (load x), c))) -> (zext (smaller load (x+c/n)))
16446 if (SDValue NarrowLoad = reduceLoadWidth(N: N0.getNode())) {
16447 SDNode *oye = N0.getOperand(i: 0).getNode();
16448 if (NarrowLoad.getNode() != N0.getNode()) {
16449 CombineTo(N: N0.getNode(), Res: NarrowLoad);
16450 // CombineTo deleted the truncate, if needed, but not what's under it.
16451 AddToWorklist(N: oye);
16452 }
16453 return SDValue(N, 0); // Return N so it doesn't get rechecked!
16454 }
16455
16456 EVT SrcVT = N0.getOperand(i: 0).getValueType();
16457 EVT MinVT = N0.getValueType();
16458
16459 if (N->getFlags().hasNonNeg()) {
16460 SDValue Op = N0.getOperand(i: 0);
16461 unsigned OpBits = SrcVT.getScalarSizeInBits();
16462 unsigned MidBits = MinVT.getScalarSizeInBits();
16463 unsigned DestBits = VT.getScalarSizeInBits();
16464
16465 if (N0->getFlags().hasNoSignedWrap() ||
16466 DAG.ComputeNumSignBits(Op) > OpBits - MidBits) {
16467 if (OpBits == DestBits) {
16468 // Op is i32, Mid is i8, and Dest is i32. If Op has more than 24 sign
16469 // bits, it is already ready.
16470 return Op;
16471 }
16472
16473 if (OpBits < DestBits) {
16474 // Op is i32, Mid is i8, and Dest is i64. If Op has more than 24 sign
16475 // bits, just sext from i32.
16476 // FIXME: This can probably be ZERO_EXTEND nneg?
16477 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: Op);
16478 }
16479
16480 // Op is i64, Mid is i8, and Dest is i32. If Op has more than 56 sign
16481 // bits, just truncate to i32.
16482 SDNodeFlags Flags;
16483 Flags.setNoSignedWrap(true);
16484 Flags.setNoUnsignedWrap(true);
16485 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: Op, Flags);
16486 }
16487 }
16488
16489 // Try to mask before the extension to avoid having to generate a larger mask,
16490 // possibly over several sub-vectors.
16491 if (SrcVT.bitsLT(VT) && VT.isVector()) {
16492 if (!LegalOperations || (TLI.isOperationLegal(Op: ISD::AND, VT: SrcVT) &&
16493 TLI.isOperationLegal(Op: ISD::ZERO_EXTEND, VT))) {
16494 SDValue Op = N0.getOperand(i: 0);
16495 Op = DAG.getZeroExtendInReg(Op, DL, VT: MinVT);
16496 AddToWorklist(N: Op.getNode());
16497 SDValue ZExtOrTrunc = DAG.getZExtOrTrunc(Op, DL, VT);
16498 // Transfer the debug info; the new node is equivalent to N0.
16499 DAG.transferDbgValues(From: N0, To: ZExtOrTrunc);
16500 return ZExtOrTrunc;
16501 }
16502 }
16503
16504 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::AND, VT)) {
16505 SDValue Op = DAG.getAnyExtOrTrunc(Op: N0.getOperand(i: 0), DL, VT);
16506 AddToWorklist(N: Op.getNode());
16507 SDValue And = DAG.getZeroExtendInReg(Op, DL, VT: MinVT);
16508 // We may safely transfer the debug info describing the truncate node over
16509 // to the equivalent and operation.
16510 DAG.transferDbgValues(From: N0, To: And);
16511 return And;
16512 }
16513 }
16514
16515 // Fold (zext (and (trunc x), cst)) -> (and x, cst),
16516 // if either of the casts is not free.
16517 // Also handles (zext (and (bitcast (extract_subvector vNi1, 0)) cst))
16518 // by treating the bitcast+extract as equivalent to a truncate of the
16519 // wider bitcast, e.g. on AVX512DQ where v8i1 extract replaces truncate.
16520 if (N0.getOpcode() == ISD::AND &&
16521 N0.getOperand(i: 1).getOpcode() == ISD::Constant) {
16522 SDValue AndSrc = N0.getOperand(i: 0);
16523 SDValue X;
16524 if (AndSrc.getOpcode() == ISD::TRUNCATE) {
16525 X = AndSrc.getOperand(i: 0);
16526 } else if (AndSrc.getOpcode() == ISD::BITCAST &&
16527 AndSrc.getOperand(i: 0).getOpcode() == ISD::EXTRACT_SUBVECTOR &&
16528 AndSrc.getOperand(i: 0).getConstantOperandVal(i: 1) == 0) {
16529 // (bitcast (extract_subvector vNi1, 0) -> iK) is equivalent to
16530 // (truncate (bitcast vNi1 -> iN) -> iK); use the wider vNi1 as X.
16531 SDValue Src = AndSrc.getOperand(i: 0).getOperand(i: 0);
16532 EVT SrcVT = Src.getValueType();
16533 if (SrcVT.isFixedLengthVectorOf(EltVT: MVT::i1)) {
16534 EVT WideIntVT =
16535 EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: SrcVT.getSizeInBits());
16536 if (TLI.isTypeLegal(VT: WideIntVT))
16537 X = DAG.getBitcast(VT: WideIntVT, V: Src);
16538 }
16539 }
16540 if (X && (!TLI.isTruncateFree(Val: X, VT2: N0.getValueType()) ||
16541 !TLI.isZExtFree(FromTy: N0.getValueType(), ToTy: VT))) {
16542 X = DAG.getAnyExtOrTrunc(Op: X, DL: SDLoc(X), VT);
16543 APInt Mask = N0.getConstantOperandAPInt(i: 1).zext(width: VT.getSizeInBits());
16544 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: X, N2: DAG.getConstant(Val: Mask, DL, VT));
16545 }
16546 }
16547
16548 // Try to simplify (zext (load x)).
16549 if (SDValue foldedExt = tryToFoldExtOfLoad(
16550 DAG, Combiner&: *this, TLI, VT, LegalOperations, N, N0, ExtLoadType: ISD::ZEXTLOAD,
16551 ExtOpc: ISD::ZERO_EXTEND, NonNegZExt: N->getFlags().hasNonNeg()))
16552 return foldedExt;
16553
16554 if (SDValue foldedExt =
16555 tryToFoldExtOfMaskedLoad(DAG, TLI, VT, LegalOperations, N, N0,
16556 ExtLoadType: ISD::ZEXTLOAD, ExtOpc: ISD::ZERO_EXTEND))
16557 return foldedExt;
16558
16559 // fold (zext (load x)) to multiple smaller zextloads.
16560 // Only on illegal but splittable vectors.
16561 if (SDValue ExtLoad = CombineExtLoad(N))
16562 return ExtLoad;
16563
16564 // Try to simplify (zext (atomic_load x)).
16565 if (SDValue foldedExt =
16566 tryToFoldExtOfAtomicLoad(DAG, TLI, VT, N0, ExtLoadType: ISD::ZEXTLOAD))
16567 return foldedExt;
16568
16569 // fold (zext (and/or/xor (load x), cst)) ->
16570 // (and/or/xor (zextload x), (zext cst))
16571 // Unless (and (load x) cst) will match as a zextload already and has
16572 // additional users, or the zext is already free.
16573 if (ISD::isBitwiseLogicOp(Opcode: N0.getOpcode()) && !TLI.isZExtFree(Val: N0, VT2: VT) &&
16574 isa<LoadSDNode>(Val: N0.getOperand(i: 0)) &&
16575 N0.getOperand(i: 1).getOpcode() == ISD::Constant &&
16576 (!LegalOperations && TLI.isOperationLegal(Op: N0.getOpcode(), VT))) {
16577 LoadSDNode *LN00 = cast<LoadSDNode>(Val: N0.getOperand(i: 0));
16578 EVT MemVT = LN00->getMemoryVT();
16579 if (TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: LN00->getAlign(), AddrSpace: LN00->getAddressSpace(),
16580 ExtType: ISD::ZEXTLOAD, Atomic: false) &&
16581 LN00->getExtensionType() != ISD::SEXTLOAD && LN00->isUnindexed()) {
16582 bool DoXform = true;
16583 SmallVector<SDNode*, 4> SetCCs;
16584 if (!N0.hasOneUse()) {
16585 if (N0.getOpcode() == ISD::AND) {
16586 auto *AndC = cast<ConstantSDNode>(Val: N0.getOperand(i: 1));
16587 EVT LoadResultTy = AndC->getValueType(ResNo: 0);
16588 EVT ExtVT;
16589 if (isAndLoadExtLoad(AndC, LoadN: LN00, LoadResultTy, ExtVT))
16590 DoXform = false;
16591 }
16592 }
16593 if (DoXform)
16594 DoXform = ExtendUsesToFormExtLoad(VT, N: N0.getNode(), N0: N0.getOperand(i: 0),
16595 ExtOpc: ISD::ZERO_EXTEND, ExtendNodes&: SetCCs, TLI);
16596 if (DoXform) {
16597 SDValue ExtLoad = DAG.getExtLoad(ExtType: ISD::ZEXTLOAD, dl: SDLoc(LN00), VT,
16598 Chain: LN00->getChain(), Ptr: LN00->getBasePtr(),
16599 MemVT: LN00->getMemoryVT(),
16600 MMO: LN00->getMemOperand());
16601 APInt Mask = N0.getConstantOperandAPInt(i: 1).zext(width: VT.getSizeInBits());
16602 SDValue And = DAG.getNode(Opcode: N0.getOpcode(), DL, VT,
16603 N1: ExtLoad, N2: DAG.getConstant(Val: Mask, DL, VT));
16604 ExtendSetCCUses(SetCCs, OrigLoad: N0.getOperand(i: 0), ExtLoad, ExtType: ISD::ZERO_EXTEND);
16605 bool NoReplaceTruncAnd = !N0.hasOneUse();
16606 bool NoReplaceTrunc = SDValue(LN00, 0).hasOneUse();
16607 CombineTo(N, Res: And);
16608 // If N0 has multiple uses, change other uses as well.
16609 if (NoReplaceTruncAnd) {
16610 SDValue TruncAnd =
16611 DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: N0.getValueType(), Operand: And);
16612 CombineTo(N: N0.getNode(), Res: TruncAnd);
16613 }
16614 if (NoReplaceTrunc) {
16615 DAG.ReplaceAllUsesOfValueWith(From: SDValue(LN00, 1), To: ExtLoad.getValue(R: 1));
16616 } else {
16617 SDValue Trunc = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(LN00),
16618 VT: LN00->getValueType(ResNo: 0), Operand: ExtLoad);
16619 CombineTo(N: LN00, Res0: Trunc, Res1: ExtLoad.getValue(R: 1));
16620 }
16621 return SDValue(N,0); // Return N so it doesn't get rechecked!
16622 }
16623 }
16624 }
16625
16626 // fold (zext (and/or/xor (shl/shr (load x), cst), cst)) ->
16627 // (and/or/xor (shl/shr (zextload x), (zext cst)), (zext cst))
16628 if (SDValue ZExtLoad = CombineZExtLogicopShiftLoad(N))
16629 return ZExtLoad;
16630
16631 // Try to simplify (zext (zextload x)).
16632 if (SDValue foldedExt = tryToFoldExtOfExtload(
16633 DAG, Combiner&: *this, TLI, VT, LegalOperations, N, N0, ExtLoadType: ISD::ZEXTLOAD))
16634 return foldedExt;
16635
16636 if (SDValue V = foldExtendedSignBitTest(N, DAG, LegalOperations))
16637 return V;
16638
16639 if (N0.getOpcode() == ISD::SETCC) {
16640 // Propagate fast-math-flags.
16641 SelectionDAG::FlagInserter FlagsInserter(DAG, N0->getFlags());
16642
16643 // Only do this before legalize for now.
16644 if (!LegalOperations && VT.isVector() &&
16645 N0.getValueType().getVectorElementType() == MVT::i1) {
16646 EVT N00VT = N0.getOperand(i: 0).getValueType();
16647 if (getSetCCResultType(VT: N00VT) == N0.getValueType())
16648 return SDValue();
16649
16650 // We know that the # elements of the results is the same as the #
16651 // elements of the compare (and the # elements of the compare result for
16652 // that matter). Check to see that they are the same size. If so, we know
16653 // that the element size of the sext'd result matches the element size of
16654 // the compare operands.
16655 if (VT.getSizeInBits() == N00VT.getSizeInBits()) {
16656 // zext(setcc) -> zext_in_reg(vsetcc) for vectors.
16657 SDValue VSetCC = DAG.getNode(Opcode: ISD::SETCC, DL, VT, N1: N0.getOperand(i: 0),
16658 N2: N0.getOperand(i: 1), N3: N0.getOperand(i: 2));
16659 return DAG.getZeroExtendInReg(Op: VSetCC, DL, VT: N0.getValueType());
16660 }
16661
16662 // If the desired elements are smaller or larger than the source
16663 // elements we can use a matching integer vector type and then
16664 // truncate/any extend followed by zext_in_reg.
16665 EVT MatchingVectorType = N00VT.changeVectorElementTypeToInteger();
16666 SDValue VsetCC =
16667 DAG.getNode(Opcode: ISD::SETCC, DL, VT: MatchingVectorType, N1: N0.getOperand(i: 0),
16668 N2: N0.getOperand(i: 1), N3: N0.getOperand(i: 2));
16669 return DAG.getZeroExtendInReg(Op: DAG.getAnyExtOrTrunc(Op: VsetCC, DL, VT), DL,
16670 VT: N0.getValueType());
16671 }
16672
16673 // zext(setcc x,y,cc) -> zext(select x, y, true, false, cc)
16674 EVT N0VT = N0.getValueType();
16675 EVT N00VT = N0.getOperand(i: 0).getValueType();
16676 if (SDValue SCC = SimplifySelectCC(
16677 DL, N0: N0.getOperand(i: 0), N1: N0.getOperand(i: 1),
16678 N2: DAG.getBoolConstant(V: true, DL, VT: N0VT, OpVT: N00VT),
16679 N3: DAG.getBoolConstant(V: false, DL, VT: N0VT, OpVT: N00VT),
16680 CC: cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get(), NotExtCompare: true))
16681 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: SCC);
16682 }
16683
16684 // (zext (shl (zext x), cst)) -> (shl (zext x), cst)
16685 if ((N0.getOpcode() == ISD::SHL || N0.getOpcode() == ISD::SRL) &&
16686 !TLI.isZExtFree(Val: N0, VT2: VT)) {
16687 SDValue ShVal = N0.getOperand(i: 0);
16688 SDValue ShAmt = N0.getOperand(i: 1);
16689 if (auto *ShAmtC = dyn_cast<ConstantSDNode>(Val&: ShAmt)) {
16690 if (ShVal.getOpcode() == ISD::ZERO_EXTEND && N0.hasOneUse()) {
16691 if (N0.getOpcode() == ISD::SHL) {
16692 // If the original shl may be shifting out bits, do not perform this
16693 // transformation.
16694 unsigned KnownZeroBits = ShVal.getValueSizeInBits() -
16695 ShVal.getOperand(i: 0).getValueSizeInBits();
16696 if (ShAmtC->getAPIntValue().ugt(RHS: KnownZeroBits)) {
16697 // If the shift is too large, then see if we can deduce that the
16698 // shift is safe anyway.
16699
16700 // Check if the bits being shifted out are known to be zero.
16701 KnownBits KnownShVal = DAG.computeKnownBits(Op: ShVal);
16702 if (ShAmtC->getAPIntValue().ugt(RHS: KnownShVal.countMinLeadingZeros()))
16703 return SDValue();
16704 }
16705 }
16706
16707 // Ensure that the shift amount is wide enough for the shifted value.
16708 if (Log2_32_Ceil(Value: VT.getSizeInBits()) > ShAmt.getValueSizeInBits())
16709 ShAmt = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: MVT::i32, Operand: ShAmt);
16710
16711 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT,
16712 N1: DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: ShVal), N2: ShAmt);
16713 }
16714 }
16715 }
16716
16717 if (SDValue NewVSel = matchVSelectOpSizesWithSetCC(Cast: N))
16718 return NewVSel;
16719
16720 if (SDValue NewCtPop = widenCtPop(Extend: N, DAG, DL))
16721 return NewCtPop;
16722
16723 if (SDValue V = widenAbs(Extend: N, DAG))
16724 return V;
16725
16726 if (SDValue Res = tryToFoldExtendSelectLoad(N, TLI, DAG, DL, Level))
16727 return Res;
16728
16729 // CSE zext nneg with sext if the zext is not free.
16730 if (N->getFlags().hasNonNeg() && !TLI.isZExtFree(FromTy: N0.getValueType(), ToTy: VT)) {
16731 SDNode *CSENode = DAG.getNodeIfExists(Opcode: ISD::SIGN_EXTEND, VTList: N->getVTList(), Ops: N0);
16732 if (CSENode)
16733 return SDValue(CSENode, 0);
16734 }
16735
16736 return SDValue();
16737}
16738
16739/// For a unary operation, propagate poison or undef from the operand \p N0 to
16740/// the result type \p VT. Returns an empty SDValue if \p N0 is neither.
16741static SDValue propagateUnaryUndef(SelectionDAG &DAG, SDValue N0, EVT VT) {
16742 if (N0.getOpcode() == ISD::POISON)
16743 return DAG.getPOISON(VT);
16744 if (N0.getOpcode() == ISD::UNDEF)
16745 return DAG.getUNDEF(VT);
16746 return SDValue();
16747}
16748
16749SDValue DAGCombiner::visitANY_EXTEND(SDNode *N) {
16750 SDValue N0 = N->getOperand(Num: 0);
16751 EVT VT = N->getValueType(ResNo: 0);
16752 SDLoc DL(N);
16753
16754 // aext(undef) = undef, aext(poison) = poison
16755 if (SDValue R = propagateUnaryUndef(DAG, N0, VT))
16756 return R;
16757
16758 if (SDValue Res = tryToFoldExtendOfConstant(N, DL, TLI, DAG, LegalTypes))
16759 return Res;
16760
16761 // fold (aext (aext x)) -> (aext x)
16762 // fold (aext (zext x)) -> (zext x)
16763 // fold (aext (sext x)) -> (sext x)
16764 if (N0.getOpcode() == ISD::ANY_EXTEND || N0.getOpcode() == ISD::ZERO_EXTEND ||
16765 N0.getOpcode() == ISD::SIGN_EXTEND) {
16766 SDNodeFlags Flags;
16767 if (N0.getOpcode() == ISD::ZERO_EXTEND)
16768 Flags.setNonNeg(N0->getFlags().hasNonNeg());
16769 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, Operand: N0.getOperand(i: 0), Flags);
16770 }
16771
16772 // fold (aext (aext_extend_vector_inreg x)) -> (aext_extend_vector_inreg x)
16773 // fold (aext (zext_extend_vector_inreg x)) -> (zext_extend_vector_inreg x)
16774 // fold (aext (sext_extend_vector_inreg x)) -> (sext_extend_vector_inreg x)
16775 if (N0.getOpcode() == ISD::ANY_EXTEND_VECTOR_INREG ||
16776 N0.getOpcode() == ISD::ZERO_EXTEND_VECTOR_INREG ||
16777 N0.getOpcode() == ISD::SIGN_EXTEND_VECTOR_INREG)
16778 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, Operand: N0.getOperand(i: 0));
16779
16780 // fold (aext (truncate (load x))) -> (aext (smaller load x))
16781 // fold (aext (truncate (srl (load x), c))) -> (aext (small load (x+c/n)))
16782 if (N0.getOpcode() == ISD::TRUNCATE) {
16783 if (SDValue NarrowLoad = reduceLoadWidth(N: N0.getNode())) {
16784 SDNode *oye = N0.getOperand(i: 0).getNode();
16785 if (NarrowLoad.getNode() != N0.getNode()) {
16786 CombineTo(N: N0.getNode(), Res: NarrowLoad);
16787 // CombineTo deleted the truncate, if needed, but not what's under it.
16788 AddToWorklist(N: oye);
16789 }
16790 return SDValue(N, 0); // Return N so it doesn't get rechecked!
16791 }
16792 }
16793
16794 // fold (aext (truncate x))
16795 if (N0.getOpcode() == ISD::TRUNCATE)
16796 return DAG.getAnyExtOrTrunc(Op: N0.getOperand(i: 0), DL, VT);
16797
16798 // Fold (aext (and (trunc x), cst)) -> (and x, cst)
16799 // if either of the casts is not free, and sign-extending the narrow type is
16800 // not cheaper than zero-extending it (which would indicate the target prefers
16801 // to keep operations at the narrower width).
16802 // Also handles (aext (and (bitcast (extract_subvector vNi1, 0)) cst))
16803 // which arises on AVX512DQ where v8i1 extract replaces truncate.
16804 if (N0.getOpcode() == ISD::AND &&
16805 N0.getOperand(i: 1).getOpcode() == ISD::Constant) {
16806 SDValue AndSrc = N0.getOperand(i: 0);
16807 SDValue X;
16808 if (AndSrc.getOpcode() == ISD::TRUNCATE) {
16809 X = AndSrc.getOperand(i: 0);
16810 } else if (AndSrc.getOpcode() == ISD::BITCAST &&
16811 AndSrc.getOperand(i: 0).getOpcode() == ISD::EXTRACT_SUBVECTOR &&
16812 AndSrc.getOperand(i: 0).getConstantOperandVal(i: 1) == 0) {
16813 SDValue Src = AndSrc.getOperand(i: 0).getOperand(i: 0);
16814 EVT SrcVT = Src.getValueType();
16815 if (SrcVT.isFixedLengthVectorOf(EltVT: MVT::i1)) {
16816 EVT WideIntVT =
16817 EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: SrcVT.getSizeInBits());
16818 if (TLI.isTypeLegal(VT: WideIntVT))
16819 X = DAG.getBitcast(VT: WideIntVT, V: Src);
16820 }
16821 }
16822 if (X && (!TLI.isTruncateFree(Val: X, VT2: N0.getValueType()) ||
16823 (!TLI.isZExtFree(FromTy: N0.getValueType(), ToTy: VT) &&
16824 !TLI.isSExtCheaperThanZExt(FromTy: N0.getValueType(), ToTy: VT)))) {
16825 X = DAG.getAnyExtOrTrunc(Op: X, DL, VT);
16826 APInt Mask = N0.getConstantOperandAPInt(i: 1).zext(width: VT.getSizeInBits());
16827 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: X, N2: DAG.getConstant(Val: Mask, DL, VT));
16828 }
16829 }
16830
16831 // fold (aext (load x)) -> (aext (truncate (extload x)))
16832 // None of the supported targets knows how to perform load and any_ext
16833 // on vectors in one instruction, so attempt to fold to zext instead.
16834 if (VT.isVector()) {
16835 // Try to simplify (zext (load x)).
16836 if (SDValue foldedExt =
16837 tryToFoldExtOfLoad(DAG, Combiner&: *this, TLI, VT, LegalOperations, N, N0,
16838 ExtLoadType: ISD::ZEXTLOAD, ExtOpc: ISD::ZERO_EXTEND))
16839 return foldedExt;
16840 } else if (ISD::isNON_EXTLoad(N: N0.getNode()) &&
16841 ISD::isUNINDEXEDLoad(N: N0.getNode())) {
16842 LoadSDNode *LN0 = cast<LoadSDNode>(Val&: N0);
16843 if (TLI.isLoadLegalOrCustom(ValVT: VT, MemVT: N0.getValueType(), Alignment: LN0->getAlign(),
16844 AddrSpace: LN0->getAddressSpace(), ExtType: ISD::EXTLOAD, Atomic: false)) {
16845 bool DoXform = true;
16846 SmallVector<SDNode *, 4> SetCCs;
16847 if (!N0.hasOneUse())
16848 DoXform =
16849 ExtendUsesToFormExtLoad(VT, N, N0, ExtOpc: ISD::ANY_EXTEND, ExtendNodes&: SetCCs, TLI);
16850 if (DoXform) {
16851 SDValue ExtLoad = DAG.getExtLoad(ExtType: ISD::EXTLOAD, dl: DL, VT, Chain: LN0->getChain(),
16852 Ptr: LN0->getBasePtr(), MemVT: N0.getValueType(),
16853 MMO: LN0->getMemOperand());
16854 ExtendSetCCUses(SetCCs, OrigLoad: N0, ExtLoad, ExtType: ISD::ANY_EXTEND);
16855 // If the load value is used only by N, replace it via CombineTo N.
16856 bool NoReplaceTrunc = N0.hasOneUse();
16857 CombineTo(N, Res: ExtLoad);
16858 if (NoReplaceTrunc) {
16859 DAG.ReplaceAllUsesOfValueWith(From: SDValue(LN0, 1), To: ExtLoad.getValue(R: 1));
16860 recursivelyDeleteUnusedNodes(N: LN0);
16861 } else {
16862 SDValue Trunc =
16863 DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(N0), VT: N0.getValueType(), Operand: ExtLoad);
16864 CombineTo(N: LN0, Res0: Trunc, Res1: ExtLoad.getValue(R: 1));
16865 }
16866 return SDValue(N, 0); // Return N so it doesn't get rechecked!
16867 }
16868 }
16869 }
16870
16871 // fold (aext (zextload x)) -> (aext (truncate (zextload x)))
16872 // fold (aext (sextload x)) -> (aext (truncate (sextload x)))
16873 // fold (aext ( extload x)) -> (aext (truncate (extload x)))
16874 if (N0.getOpcode() == ISD::LOAD && !ISD::isNON_EXTLoad(N: N0.getNode()) &&
16875 ISD::isUNINDEXEDLoad(N: N0.getNode()) && N0.hasOneUse()) {
16876 LoadSDNode *LN0 = cast<LoadSDNode>(Val&: N0);
16877 ISD::LoadExtType ExtType = LN0->getExtensionType();
16878 EVT MemVT = LN0->getMemoryVT();
16879 if (!LegalOperations ||
16880 TLI.isLoadLegal(ValVT: VT, MemVT, Alignment: LN0->getAlign(), AddrSpace: LN0->getAddressSpace(),
16881 ExtType, Atomic: false)) {
16882 SDValue ExtLoad =
16883 DAG.getExtLoad(ExtType, dl: DL, VT, Chain: LN0->getChain(), Ptr: LN0->getBasePtr(),
16884 MemVT, MMO: LN0->getMemOperand());
16885 CombineTo(N, Res: ExtLoad);
16886 DAG.ReplaceAllUsesOfValueWith(From: SDValue(LN0, 1), To: ExtLoad.getValue(R: 1));
16887 recursivelyDeleteUnusedNodes(N: LN0);
16888 return SDValue(N, 0); // Return N so it doesn't get rechecked!
16889 }
16890 }
16891
16892 if (N0.getOpcode() == ISD::SETCC) {
16893 // Propagate fast-math-flags.
16894 SDNodeFlags Flags = N0->getFlags();
16895 SelectionDAG::FlagInserter FlagsInserter(DAG, Flags);
16896
16897 // For vectors:
16898 // aext(setcc) -> vsetcc
16899 // aext(setcc) -> truncate(vsetcc)
16900 // aext(setcc) -> aext(vsetcc)
16901 // Only do this before legalize for now.
16902 if (VT.isVector() && !LegalOperations) {
16903 EVT N00VT = N0.getOperand(i: 0).getValueType();
16904 if (getSetCCResultType(VT: N00VT) == N0.getValueType())
16905 return SDValue();
16906
16907 // We know that the # elements of the results is the same as the
16908 // # elements of the compare (and the # elements of the compare result
16909 // for that matter). Check to see that they are the same size. If so,
16910 // we know that the element size of the sext'd result matches the
16911 // element size of the compare operands.
16912 if (VT.getSizeInBits() == N00VT.getSizeInBits())
16913 return DAG.getSetCC(DL, VT, LHS: N0.getOperand(i: 0), RHS: N0.getOperand(i: 1),
16914 Cond: cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get(),
16915 /*Chain=*/{}, /*Signaling=*/IsSignaling: false, Flags);
16916
16917 // If the desired elements are smaller or larger than the source
16918 // elements we can use a matching integer vector type and then
16919 // truncate/any extend
16920 EVT MatchingVectorType = N00VT.changeVectorElementTypeToInteger();
16921 SDValue VsetCC = DAG.getSetCC(
16922 DL, VT: MatchingVectorType, LHS: N0.getOperand(i: 0), RHS: N0.getOperand(i: 1),
16923 Cond: cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get(), /*Chain=*/{},
16924 /*Signaling=*/IsSignaling: false, Flags);
16925 return DAG.getAnyExtOrTrunc(Op: VsetCC, DL, VT);
16926 }
16927
16928 // aext(setcc x,y,cc) -> select_cc x, y, 1, 0, cc
16929 if (SDValue SCC = SimplifySelectCC(
16930 DL, N0: N0.getOperand(i: 0), N1: N0.getOperand(i: 1), N2: DAG.getConstant(Val: 1, DL, VT),
16931 N3: DAG.getConstant(Val: 0, DL, VT),
16932 CC: cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get(), NotExtCompare: true))
16933 return SCC;
16934 }
16935
16936 if (SDValue NewCtPop = widenCtPop(Extend: N, DAG, DL))
16937 return NewCtPop;
16938
16939 if (SDValue Res = tryToFoldExtendSelectLoad(N, TLI, DAG, DL, Level))
16940 return Res;
16941
16942 return SDValue();
16943}
16944
16945SDValue DAGCombiner::visitAssertExt(SDNode *N) {
16946 unsigned Opcode = N->getOpcode();
16947 SDValue N0 = N->getOperand(Num: 0);
16948 SDValue N1 = N->getOperand(Num: 1);
16949 EVT AssertVT = cast<VTSDNode>(Val&: N1)->getVT();
16950
16951 // fold (assert?ext (assert?ext x, vt), vt) -> (assert?ext x, vt)
16952 if (N0.getOpcode() == Opcode &&
16953 AssertVT == cast<VTSDNode>(Val: N0.getOperand(i: 1))->getVT())
16954 return N0;
16955
16956 // fold (assert?ext c, vt) -> c
16957 if (isa<ConstantSDNode>(Val: N0))
16958 return N0;
16959
16960 if (N0.getOpcode() == ISD::TRUNCATE && N0.hasOneUse() &&
16961 N0.getOperand(i: 0).getOpcode() == Opcode) {
16962 // We have an assert, truncate, assert sandwich. Make one stronger assert
16963 // by asserting on the smallest asserted type to the larger source type.
16964 // This eliminates the later assert:
16965 // assert (trunc (assert X, i8) to iN), i1 --> trunc (assert X, i1) to iN
16966 // assert (trunc (assert X, i1) to iN), i8 --> trunc (assert X, i1) to iN
16967 SDLoc DL(N);
16968 SDValue BigA = N0.getOperand(i: 0);
16969 EVT BigA_AssertVT = cast<VTSDNode>(Val: BigA.getOperand(i: 1))->getVT();
16970 EVT MinAssertVT = AssertVT.bitsLT(VT: BigA_AssertVT) ? AssertVT : BigA_AssertVT;
16971 SDValue MinAssertVTVal = DAG.getValueType(MinAssertVT);
16972 SDValue NewAssert = DAG.getNode(Opcode, DL, VT: BigA.getValueType(),
16973 N1: BigA.getOperand(i: 0), N2: MinAssertVTVal);
16974 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: N->getValueType(ResNo: 0), Operand: NewAssert);
16975 }
16976
16977 // If we have (AssertZext (truncate (AssertSext X, iX)), iY) and Y is smaller
16978 // than X. Just move the AssertZext in front of the truncate and drop the
16979 // AssertSExt.
16980 if (N0.getOpcode() == ISD::TRUNCATE && N0.hasOneUse() &&
16981 N0.getOperand(i: 0).getOpcode() == ISD::AssertSext &&
16982 Opcode == ISD::AssertZext) {
16983 SDValue BigA = N0.getOperand(i: 0);
16984 EVT BigA_AssertVT = cast<VTSDNode>(Val: BigA.getOperand(i: 1))->getVT();
16985 if (AssertVT.bitsLT(VT: BigA_AssertVT)) {
16986 SDLoc DL(N);
16987 SDValue NewAssert = DAG.getNode(Opcode, DL, VT: BigA.getValueType(),
16988 N1: BigA.getOperand(i: 0), N2: N1);
16989 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: N->getValueType(ResNo: 0), Operand: NewAssert);
16990 }
16991 }
16992
16993 if (Opcode == ISD::AssertZext && N0.getOpcode() == ISD::AND &&
16994 isa<ConstantSDNode>(Val: N0.getOperand(i: 1))) {
16995 const APInt &Mask = N0.getConstantOperandAPInt(i: 1);
16996
16997 // If we have (AssertZext (and (AssertSext X, iX), M), iY) and Y is smaller
16998 // than X, and the And doesn't change the lower iX bits, we can move the
16999 // AssertZext in front of the And and drop the AssertSext.
17000 if (N0.getOperand(i: 0).getOpcode() == ISD::AssertSext && N0.hasOneUse()) {
17001 SDValue BigA = N0.getOperand(i: 0);
17002 EVT BigA_AssertVT = cast<VTSDNode>(Val: BigA.getOperand(i: 1))->getVT();
17003 if (AssertVT.bitsLT(VT: BigA_AssertVT) &&
17004 Mask.countr_one() >= BigA_AssertVT.getScalarSizeInBits()) {
17005 SDLoc DL(N);
17006 SDValue NewAssert =
17007 DAG.getNode(Opcode, DL, VT: N->getValueType(ResNo: 0), N1: BigA.getOperand(i: 0), N2: N1);
17008 return DAG.getNode(Opcode: ISD::AND, DL, VT: N->getValueType(ResNo: 0), N1: NewAssert,
17009 N2: N0.getOperand(i: 1));
17010 }
17011 }
17012
17013 // Remove AssertZext entirely if the mask guarantees the assertion cannot
17014 // fail.
17015 // TODO: Use KB countMinLeadingZeros to handle non-constant masks?
17016 if (Mask.isIntN(N: AssertVT.getScalarSizeInBits()))
17017 return N0;
17018 }
17019
17020 return SDValue();
17021}
17022
17023SDValue DAGCombiner::visitAssertAlign(SDNode *N) {
17024 SDLoc DL(N);
17025
17026 Align AL = cast<AssertAlignSDNode>(Val: N)->getAlign();
17027 SDValue N0 = N->getOperand(Num: 0);
17028
17029 // Fold (assertalign (assertalign x, AL0), AL1) ->
17030 // (assertalign x, max(AL0, AL1))
17031 if (auto *AAN = dyn_cast<AssertAlignSDNode>(Val&: N0))
17032 return DAG.getAssertAlign(DL, V: N0.getOperand(i: 0),
17033 A: std::max(a: AL, b: AAN->getAlign()));
17034
17035 // In rare cases, there are trivial arithmetic ops in source operands. Sink
17036 // this assert down to source operands so that those arithmetic ops could be
17037 // exposed to the DAG combining.
17038 switch (N0.getOpcode()) {
17039 default:
17040 break;
17041 case ISD::ADD:
17042 case ISD::PTRADD:
17043 case ISD::SUB: {
17044 unsigned AlignShift = Log2(A: AL);
17045 SDValue LHS = N0.getOperand(i: 0);
17046 SDValue RHS = N0.getOperand(i: 1);
17047 unsigned LHSAlignShift = DAG.computeKnownBits(Op: LHS).countMinTrailingZeros();
17048 unsigned RHSAlignShift = DAG.computeKnownBits(Op: RHS).countMinTrailingZeros();
17049 if (LHSAlignShift >= AlignShift || RHSAlignShift >= AlignShift) {
17050 if (LHSAlignShift < AlignShift)
17051 LHS = DAG.getAssertAlign(DL, V: LHS, A: AL);
17052 if (RHSAlignShift < AlignShift)
17053 RHS = DAG.getAssertAlign(DL, V: RHS, A: AL);
17054 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT: N0.getValueType(), N1: LHS, N2: RHS);
17055 }
17056 break;
17057 }
17058 }
17059
17060 return SDValue();
17061}
17062
17063SDValue DAGCombiner::visitIS_FPCLASS(SDNode *N) {
17064 SDValue Src = N->getOperand(Num: 0);
17065 FPClassTest Mask = static_cast<FPClassTest>(N->getConstantOperandVal(Num: 1));
17066 EVT VT = N->getValueType(ResNo: 0);
17067 SDLoc DL(N);
17068
17069 // is.fpclass(poison, mask) -> poison
17070 if (Src.getOpcode() == ISD::POISON)
17071 return DAG.getPOISON(VT);
17072
17073 KnownFPClass Known = DAG.computeKnownFPClass(Op: Src, InterestedClasses: Mask);
17074
17075 // All possible classes are within the mask: result is always true.
17076 if ((~Mask & Known.getKnownFPClasses()) == fcNone)
17077 return DAG.getBoolConstant(V: true, DL, VT, OpVT: Src.getValueType());
17078
17079 // Clear test bits we know must be false from the source value.
17080 // fp_class (nnan x), qnan|snan|other -> fp_class (nnan x), other
17081 // fp_class (ninf x), ninf|pinf|other -> fp_class (ninf x), other
17082 if ((Mask & Known.getKnownFPClasses()) != Mask) {
17083 return DAG.getNode(
17084 Opcode: ISD::IS_FPCLASS, DL, VT, N1: Src,
17085 N2: DAG.getTargetConstant(Val: Mask & Known.getKnownFPClasses(), DL, VT: MVT::i32),
17086 Flags: N->getFlags());
17087 }
17088
17089 return SDValue();
17090}
17091
17092/// If the result of a load is shifted/masked/truncated to an effectively
17093/// narrower type, try to transform the load to a narrower type and/or
17094/// use an extending load.
17095SDValue DAGCombiner::reduceLoadWidth(SDNode *N) {
17096 unsigned Opc = N->getOpcode();
17097
17098 ISD::LoadExtType ExtType = ISD::NON_EXTLOAD;
17099 SDValue N0 = N->getOperand(Num: 0);
17100 EVT VT = N->getValueType(ResNo: 0);
17101 EVT ExtVT = VT;
17102
17103 // This transformation isn't valid for vector loads.
17104 if (VT.isVector())
17105 return SDValue();
17106
17107 // The ShAmt variable is used to indicate that we've consumed a right
17108 // shift. I.e. we want to narrow the width of the load by skipping to load the
17109 // ShAmt least significant bits.
17110 unsigned ShAmt = 0;
17111 // A special case is when the least significant bits from the load are masked
17112 // away, but using an AND rather than a right shift. HasShiftedOffset is used
17113 // to indicate that the narrowed load should be left-shifted ShAmt bits to get
17114 // the result.
17115 unsigned ShiftedOffset = 0;
17116 // Special case: SIGN_EXTEND_INREG is basically truncating to ExtVT then
17117 // extended to VT.
17118 if (Opc == ISD::SIGN_EXTEND_INREG) {
17119 ExtType = ISD::SEXTLOAD;
17120 ExtVT = cast<VTSDNode>(Val: N->getOperand(Num: 1))->getVT();
17121 } else if (Opc == ISD::SRL || Opc == ISD::SRA) {
17122 // Another special-case: SRL/SRA is basically zero/sign-extending a narrower
17123 // value, or it may be shifting a higher subword, half or byte into the
17124 // lowest bits.
17125
17126 // Only handle shift with constant shift amount, and the shiftee must be a
17127 // load.
17128 auto *LN = dyn_cast<LoadSDNode>(Val&: N0);
17129 auto *N1C = dyn_cast<ConstantSDNode>(Val: N->getOperand(Num: 1));
17130 if (!N1C || !LN)
17131 return SDValue();
17132 // If the shift amount is larger than the memory type then we're not
17133 // accessing any of the loaded bytes.
17134 ShAmt = N1C->getZExtValue();
17135 uint64_t MemoryWidth = LN->getMemoryVT().getScalarSizeInBits();
17136 if (MemoryWidth <= ShAmt)
17137 return SDValue();
17138 // Attempt to fold away the SRL by using ZEXTLOAD and SRA by using SEXTLOAD.
17139 ExtType = Opc == ISD::SRL ? ISD::ZEXTLOAD : ISD::SEXTLOAD;
17140 ExtVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: MemoryWidth - ShAmt);
17141 // If original load is a SEXTLOAD then we can't simply replace it by a
17142 // ZEXTLOAD (we could potentially replace it by a more narrow SEXTLOAD
17143 // followed by a ZEXT, but that is not handled at the moment). Similarly if
17144 // the original load is a ZEXTLOAD and we want to use a SEXTLOAD.
17145 if ((LN->getExtensionType() == ISD::SEXTLOAD ||
17146 LN->getExtensionType() == ISD::ZEXTLOAD) &&
17147 LN->getExtensionType() != ExtType)
17148 return SDValue();
17149 } else if (Opc == ISD::AND) {
17150 // An AND with a constant mask is the same as a truncate + zero-extend.
17151 auto AndC = dyn_cast<ConstantSDNode>(Val: N->getOperand(Num: 1));
17152 if (!AndC)
17153 return SDValue();
17154
17155 const APInt &Mask = AndC->getAPIntValue();
17156 unsigned ActiveBits = 0;
17157 if (Mask.isMask()) {
17158 ActiveBits = Mask.countr_one();
17159 } else if (Mask.isShiftedMask(MaskIdx&: ShAmt, MaskLen&: ActiveBits)) {
17160 ShiftedOffset = ShAmt;
17161 } else {
17162 return SDValue();
17163 }
17164
17165 ExtType = ISD::ZEXTLOAD;
17166 ExtVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: ActiveBits);
17167 }
17168
17169 // In case Opc==SRL we've already prepared ExtVT/ExtType/ShAmt based on doing
17170 // a right shift. Here we redo some of those checks, to possibly adjust the
17171 // ExtVT even further based on "a masking AND". We could also end up here for
17172 // other reasons (e.g. based on Opc==TRUNCATE) and that is why some checks
17173 // need to be done here as well.
17174 if (Opc == ISD::SRL || N0.getOpcode() == ISD::SRL) {
17175 SDValue SRL = Opc == ISD::SRL ? SDValue(N, 0) : N0;
17176 // Bail out when the SRL has more than one use. This is done for historical
17177 // (undocumented) reasons. Maybe intent was to guard the AND-masking below
17178 // check below? And maybe it could be non-profitable to do the transform in
17179 // case the SRL has multiple uses and we get here with Opc!=ISD::SRL?
17180 // FIXME: Can't we just skip this check for the Opc==ISD::SRL case.
17181 if (!SRL.hasOneUse())
17182 return SDValue();
17183
17184 // Only handle shift with constant shift amount, and the shiftee must be a
17185 // load.
17186 auto *LN = dyn_cast<LoadSDNode>(Val: SRL.getOperand(i: 0));
17187 auto *SRL1C = dyn_cast<ConstantSDNode>(Val: SRL.getOperand(i: 1));
17188 if (!SRL1C || !LN)
17189 return SDValue();
17190
17191 // If the shift amount is larger than the input type then we're not
17192 // accessing any of the loaded bytes. If the load was a zextload/extload
17193 // then the result of the shift+trunc is zero/undef (handled elsewhere).
17194 ShAmt = SRL1C->getZExtValue();
17195 uint64_t MemoryWidth = LN->getMemoryVT().getSizeInBits();
17196 if (ShAmt >= MemoryWidth)
17197 return SDValue();
17198
17199 // Because a SRL must be assumed to *need* to zero-extend the high bits
17200 // (as opposed to anyext the high bits), we can't combine the zextload
17201 // lowering of SRL and an sextload.
17202 if (LN->getExtensionType() == ISD::SEXTLOAD)
17203 return SDValue();
17204
17205 // Avoid reading outside the memory accessed by the original load (could
17206 // happened if we only adjust the load base pointer by ShAmt). Instead we
17207 // try to narrow the load even further. The typical scenario here is:
17208 // (i64 (truncate (i96 (srl (load x), 64)))) ->
17209 // (i64 (truncate (i96 (zextload (load i32 + offset) from i32))))
17210 if (ExtVT.getScalarSizeInBits() > MemoryWidth - ShAmt) {
17211 // Don't replace sextload by zextload.
17212 if (ExtType == ISD::SEXTLOAD)
17213 return SDValue();
17214 // Narrow the load.
17215 ExtType = ISD::ZEXTLOAD;
17216 ExtVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: MemoryWidth - ShAmt);
17217 }
17218
17219 // If the SRL is only used by a masking AND, we may be able to adjust
17220 // the ExtVT to make the AND redundant.
17221 SDNode *Mask = *(SRL->user_begin());
17222 if (SRL.hasOneUse() && Mask->getOpcode() == ISD::AND &&
17223 isa<ConstantSDNode>(Val: Mask->getOperand(Num: 1))) {
17224 unsigned Offset, ActiveBits;
17225 const APInt& ShiftMask = Mask->getConstantOperandAPInt(Num: 1);
17226 if (ShiftMask.isMask()) {
17227 EVT MaskedVT =
17228 EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: ShiftMask.countr_one());
17229 // If the mask is smaller, recompute the type.
17230 if ((ExtVT.getScalarSizeInBits() > MaskedVT.getScalarSizeInBits()) &&
17231 TLI.isLoadLegal(ValVT: SRL.getValueType(), MemVT: MaskedVT, Alignment: LN->getAlign(),
17232 AddrSpace: LN->getAddressSpace(), ExtType, Atomic: false))
17233 ExtVT = MaskedVT;
17234 } else if (ExtType == ISD::ZEXTLOAD &&
17235 ShiftMask.isShiftedMask(MaskIdx&: Offset, MaskLen&: ActiveBits) &&
17236 (Offset + ShAmt) < VT.getScalarSizeInBits()) {
17237 EVT MaskedVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: ActiveBits);
17238 // If the mask is shifted we can use a narrower load and a shl to insert
17239 // the trailing zeros.
17240 if (((Offset + ActiveBits) <= ExtVT.getScalarSizeInBits()) &&
17241 TLI.isLoadLegal(ValVT: SRL.getValueType(), MemVT: MaskedVT, Alignment: LN->getAlign(),
17242 AddrSpace: LN->getAddressSpace(), ExtType, Atomic: false)) {
17243 ExtVT = MaskedVT;
17244 ShAmt = Offset + ShAmt;
17245 ShiftedOffset = Offset;
17246 }
17247 }
17248 }
17249
17250 N0 = SRL.getOperand(i: 0);
17251 }
17252
17253 // If the load is shifted left (and the result isn't shifted back right), we
17254 // can fold a truncate through the shift. The typical scenario is that N
17255 // points at a TRUNCATE here so the attempted fold is:
17256 // (truncate (shl (load x), c))) -> (shl (narrow load x), c)
17257 // ShLeftAmt will indicate how much a narrowed load should be shifted left.
17258 unsigned ShLeftAmt = 0;
17259 if (ShAmt == 0 && N0.getOpcode() == ISD::SHL && N0.hasOneUse() &&
17260 ExtVT == VT && TLI.isNarrowingProfitable(N, SrcVT: N0.getValueType(), DestVT: VT)) {
17261 if (ConstantSDNode *N01 = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1))) {
17262 ShLeftAmt = N01->getZExtValue();
17263 N0 = N0.getOperand(i: 0);
17264 }
17265 }
17266
17267 // Look through a freeze if present between the operation and the load.
17268 // The freeze will be preserved on the narrowed result.
17269 SDValue FreezeNode;
17270 if (N0.getOpcode() == ISD::FREEZE) {
17271 FreezeNode = N0;
17272 N0 = N0.getOperand(i: 0);
17273 }
17274
17275 // If we haven't found a load, we can't narrow it.
17276 if (!isa<LoadSDNode>(Val: N0))
17277 return SDValue();
17278
17279 LoadSDNode *LN0 = cast<LoadSDNode>(Val&: N0);
17280 // Reducing the width of a volatile load is illegal. For atomics, we may be
17281 // able to reduce the width provided we never widen again. (see D66309)
17282 if (!LN0->isSimple() ||
17283 !isLegalNarrowLdSt(LDST: LN0, ExtType, MemVT&: ExtVT, ShAmt))
17284 return SDValue();
17285
17286 // Bail early when looking through a multi-use freeze, since other users of
17287 // the freeze can depend on the full load value. But its still safe to change
17288 // the extension type from anyext to zext.
17289 if (FreezeNode && !FreezeNode.hasOneUse() &&
17290 (LN0->getMemoryVT().bitsGT(VT: ExtVT) || ExtType != ISD::ZEXTLOAD ||
17291 (LN0->getExtensionType() != ISD::EXTLOAD &&
17292 LN0->getExtensionType() != ISD::ZEXTLOAD)))
17293 return SDValue();
17294
17295 auto AdjustBigEndianShift = [&](unsigned ShAmt) {
17296 unsigned LVTStoreBits =
17297 LN0->getMemoryVT().getStoreSizeInBits().getFixedValue();
17298 unsigned EVTStoreBits = ExtVT.getStoreSizeInBits().getFixedValue();
17299 return LVTStoreBits - EVTStoreBits - ShAmt;
17300 };
17301
17302 // We need to adjust the pointer to the load by ShAmt bits in order to load
17303 // the correct bytes.
17304 unsigned PtrAdjustmentInBits =
17305 DAG.getDataLayout().isBigEndian() ? AdjustBigEndianShift(ShAmt) : ShAmt;
17306
17307 uint64_t PtrOff = PtrAdjustmentInBits / 8;
17308 SDLoc DL(LN0);
17309 // The original load itself didn't wrap, so an offset within it doesn't.
17310 SDValue NewPtr =
17311 DAG.getMemBasePlusOffset(Base: LN0->getBasePtr(), Offset: TypeSize::getFixed(ExactSize: PtrOff),
17312 DL, Flags: SDNodeFlags::NoUnsignedWrap);
17313 AddToWorklist(N: NewPtr.getNode());
17314
17315 SDValue Load;
17316 if (ExtType == ISD::NON_EXTLOAD) {
17317 const MDNode *OldRanges = LN0->getRanges();
17318 const MDNode *NewRanges = nullptr;
17319 // If LSBs are loaded and the truncated ConstantRange for the OldRanges
17320 // metadata is not the full-set for the new width then create a NewRanges
17321 // metadata for the truncated load
17322 if (ShAmt == 0 && OldRanges) {
17323 ConstantRange CR = getConstantRangeFromMetadata(RangeMD: *OldRanges);
17324 unsigned BitSize = VT.getScalarSizeInBits();
17325
17326 // It is possible for an 8-bit extending load with 8-bit range
17327 // metadata to be narrowed to an 8-bit load. This guard is necessary to
17328 // ensure that truncation is strictly smaller.
17329 if (CR.getBitWidth() > BitSize) {
17330 ConstantRange TruncatedCR = CR.truncate(BitWidth: BitSize);
17331 if (!TruncatedCR.isFullSet()) {
17332 Metadata *Bounds[2] = {
17333 ConstantAsMetadata::get(
17334 C: ConstantInt::get(Context&: *DAG.getContext(), V: TruncatedCR.getLower())),
17335 ConstantAsMetadata::get(
17336 C: ConstantInt::get(Context&: *DAG.getContext(), V: TruncatedCR.getUpper()))};
17337 NewRanges = MDNode::get(Context&: *DAG.getContext(), MDs: Bounds);
17338 }
17339 } else if (CR.getBitWidth() == BitSize)
17340 NewRanges = OldRanges;
17341 }
17342 Load = DAG.getLoad(VT, dl: DL, Chain: LN0->getChain(), Ptr: NewPtr,
17343 PtrInfo: LN0->getPointerInfo().getWithOffset(O: PtrOff),
17344 Alignment: LN0->getBaseAlign(), MMOFlags: LN0->getMemOperand()->getFlags(),
17345 Metadata: MMOMetadata(LN0->getAAInfo(), NewRanges));
17346 } else
17347 Load = DAG.getExtLoad(ExtType, dl: DL, VT, Chain: LN0->getChain(), Ptr: NewPtr,
17348 PtrInfo: LN0->getPointerInfo().getWithOffset(O: PtrOff), MemVT: ExtVT,
17349 Alignment: LN0->getBaseAlign(), MMOFlags: LN0->getMemOperand()->getFlags(),
17350 Metadata: LN0->getAAInfo());
17351
17352 // Replace the old load's chain with the new load's chain.
17353 WorklistRemover DeadNodes(*this);
17354 DAG.ReplaceAllUsesOfValueWith(From: N0.getValue(R: 1), To: Load.getValue(R: 1));
17355
17356 // Replace old load value for multi-use freeze so all users benefit.
17357 if (FreezeNode && !FreezeNode.hasOneUse())
17358 DAG.ReplaceAllUsesOfValueWith(From: N0.getValue(R: 0), To: Load.getValue(R: 0));
17359
17360 // If we looked through a freeze, rewrap the narrowed result and add an
17361 // Assert node so downstream analyses can see the range.
17362 SDValue Result = Load;
17363 if (FreezeNode) {
17364 Result = DAG.getNode(Opcode: ISD::FREEZE, DL, VT, Operand: Result);
17365 if (ExtType == ISD::ZEXTLOAD)
17366 Result =
17367 DAG.getNode(Opcode: ISD::AssertZext, DL, VT, N1: Result, N2: DAG.getValueType(ExtVT));
17368 else if (ExtType == ISD::SEXTLOAD)
17369 Result =
17370 DAG.getNode(Opcode: ISD::AssertSext, DL, VT, N1: Result, N2: DAG.getValueType(ExtVT));
17371 }
17372
17373 // Shift the result left, if we've swallowed a left shift.
17374 if (ShLeftAmt != 0) {
17375 // If the shift amount is as large as the result size (but, presumably,
17376 // no larger than the source) then the useful bits of the result are
17377 // zero; we can't simply return the shortened shift, because the result
17378 // of that operation is undefined.
17379 if (ShLeftAmt >= VT.getScalarSizeInBits())
17380 Result = DAG.getConstant(Val: 0, DL, VT);
17381 else
17382 Result = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Result,
17383 N2: DAG.getShiftAmountConstant(Val: ShLeftAmt, VT, DL));
17384 }
17385
17386 if (ShiftedOffset != 0) {
17387 // We're using a shifted mask, so the load now has an offset. This means
17388 // that data has been loaded into the lower bytes than it would have been
17389 // before, so we need to shl the loaded data into the correct position in the
17390 // register.
17391 SDValue ShiftC = DAG.getConstant(Val: ShiftedOffset, DL, VT);
17392 Result = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Result, N2: ShiftC);
17393 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 0), To: Result);
17394 }
17395
17396 // Return the new loaded value.
17397 return Result;
17398}
17399
17400SDValue DAGCombiner::visitSIGN_EXTEND_INREG(SDNode *N) {
17401 SDValue N0 = N->getOperand(Num: 0);
17402 SDValue N1 = N->getOperand(Num: 1);
17403 EVT VT = N->getValueType(ResNo: 0);
17404 EVT ExtVT = cast<VTSDNode>(Val&: N1)->getVT();
17405 unsigned VTBits = VT.getScalarSizeInBits();
17406 unsigned ExtVTBits = ExtVT.getScalarSizeInBits();
17407 SDLoc DL(N);
17408
17409 // sext_vector_inreg(undef) = 0 because the top bit will all be the same.
17410 if (N0.isUndef())
17411 return DAG.getConstant(Val: 0, DL, VT);
17412
17413 // fold (sext_in_reg c1) -> c1
17414 if (SDValue C =
17415 DAG.FoldConstantArithmetic(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT, Ops: {N0, N1}))
17416 return C;
17417
17418 // If the input is already sign extended, just drop the extension.
17419 if (ExtVTBits >= DAG.ComputeMaxSignificantBits(Op: N0))
17420 return N0;
17421
17422 // fold (sext_in_reg (sext_in_reg x, VT2), VT1) -> (sext_in_reg x, minVT) pt2
17423 if (N0.getOpcode() == ISD::SIGN_EXTEND_INREG &&
17424 ExtVT.bitsLT(VT: cast<VTSDNode>(Val: N0.getOperand(i: 1))->getVT()))
17425 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT, N1: N0.getOperand(i: 0), N2: N1);
17426
17427 // fold (sext_in_reg (sext x)) -> (sext x)
17428 // fold (sext_in_reg (aext x)) -> (sext x)
17429 // if x is small enough or if we know that x has more than 1 sign bit and the
17430 // sign_extend_inreg is extending from one of them.
17431 if (N0.getOpcode() == ISD::SIGN_EXTEND || N0.getOpcode() == ISD::ANY_EXTEND) {
17432 SDValue N00 = N0.getOperand(i: 0);
17433 unsigned N00Bits = N00.getScalarValueSizeInBits();
17434 if ((N00Bits <= ExtVTBits ||
17435 DAG.ComputeMaxSignificantBits(Op: N00) <= ExtVTBits) &&
17436 (!LegalOperations || TLI.isOperationLegal(Op: ISD::SIGN_EXTEND, VT)))
17437 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: N00);
17438 }
17439
17440 // fold (sext_in_reg (*_extend_vector_inreg x)) -> (sext_vector_inreg x)
17441 // if x is small enough or if we know that x has more than 1 sign bit and the
17442 // sign_extend_inreg is extending from one of them.
17443 if (ISD::isExtVecInRegOpcode(Opcode: N0.getOpcode())) {
17444 SDValue N00 = N0.getOperand(i: 0);
17445 unsigned N00Bits = N00.getScalarValueSizeInBits();
17446 bool IsZext = N0.getOpcode() == ISD::ZERO_EXTEND_VECTOR_INREG;
17447 if ((N00Bits == ExtVTBits ||
17448 (!IsZext && (N00Bits < ExtVTBits ||
17449 DAG.ComputeMaxSignificantBits(Op: N00) <= ExtVTBits))) &&
17450 (!LegalOperations ||
17451 TLI.isOperationLegal(Op: ISD::SIGN_EXTEND_VECTOR_INREG, VT)))
17452 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_VECTOR_INREG, DL, VT, Operand: N00);
17453 }
17454
17455 // fold (sext_in_reg (zext x)) -> (sext x)
17456 // iff we are extending the source sign bit.
17457 if (N0.getOpcode() == ISD::ZERO_EXTEND) {
17458 SDValue N00 = N0.getOperand(i: 0);
17459 if (N00.getScalarValueSizeInBits() == ExtVTBits &&
17460 (!LegalOperations || TLI.isOperationLegal(Op: ISD::SIGN_EXTEND, VT)))
17461 return DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT, Operand: N00);
17462 }
17463
17464 // fold (sext_in_reg x) -> (zext_in_reg x) if the sign bit is known zero.
17465 if (DAG.MaskedValueIsZero(Op: N0, Mask: APInt::getOneBitSet(numBits: VTBits, BitNo: ExtVTBits - 1)))
17466 return DAG.getZeroExtendInReg(Op: N0, DL, VT: ExtVT);
17467
17468 // fold operands of sext_in_reg based on knowledge that the top bits are not
17469 // demanded.
17470 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
17471 return SDValue(N, 0);
17472
17473 // fold (sext_in_reg (load x)) -> (smaller sextload x)
17474 // fold (sext_in_reg (srl (load x), c)) -> (smaller sextload (x+c/evtbits))
17475 if (SDValue NarrowLoad = reduceLoadWidth(N))
17476 return NarrowLoad;
17477
17478 // fold (sext_in_reg (srl X, 24), i8) -> (sra X, 24)
17479 // fold (sext_in_reg (srl X, 23), i8) -> (sra X, 23) iff possible.
17480 // We already fold "(sext_in_reg (srl X, 25), i8) -> srl X, 25" above.
17481 if (N0.getOpcode() == ISD::SRL) {
17482 if (auto *ShAmt = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 1)))
17483 if (ShAmt->getAPIntValue().ule(RHS: VTBits - ExtVTBits)) {
17484 // We can turn this into an SRA iff the input to the SRL is already sign
17485 // extended enough.
17486 unsigned InSignBits = DAG.ComputeNumSignBits(Op: N0.getOperand(i: 0));
17487 if (((VTBits - ExtVTBits) - ShAmt->getZExtValue()) < InSignBits)
17488 return DAG.getNode(Opcode: ISD::SRA, DL, VT, N1: N0.getOperand(i: 0),
17489 N2: N0.getOperand(i: 1));
17490 }
17491 }
17492
17493 // fold (sext_inreg (extload x)) -> (sextload x)
17494 // If sextload is not supported by target, we can only do the combine when
17495 // load has one use. Doing otherwise can block folding the extload with other
17496 // extends that the target does support.
17497 if (ISD::isEXTLoad(N: N0.getNode()) && ISD::isUNINDEXEDLoad(N: N0.getNode())) {
17498 auto *LN0 = cast<LoadSDNode>(Val&: N0);
17499 if (ExtVT == LN0->getMemoryVT() &&
17500 ((!LegalOperations && LN0->isSimple() && N0.hasOneUse()) ||
17501 TLI.isLoadLegal(ValVT: VT, MemVT: ExtVT, Alignment: LN0->getAlign(), AddrSpace: LN0->getAddressSpace(),
17502 ExtType: ISD::SEXTLOAD, Atomic: false))) {
17503 SDValue ExtLoad =
17504 DAG.getExtLoad(ExtType: ISD::SEXTLOAD, dl: DL, VT, Chain: LN0->getChain(),
17505 Ptr: LN0->getBasePtr(), MemVT: ExtVT, MMO: LN0->getMemOperand());
17506 CombineTo(N, Res: ExtLoad);
17507 CombineTo(N: N0.getNode(), Res0: ExtLoad, Res1: ExtLoad.getValue(R: 1));
17508 AddToWorklist(N: ExtLoad.getNode());
17509 return SDValue(N, 0); // Return N so it doesn't get rechecked!
17510 }
17511 }
17512
17513 // fold (sext_inreg (zextload x)) -> (sextload x) iff load has one use
17514 if (ISD::isZEXTLoad(N: N0.getNode()) && ISD::isUNINDEXEDLoad(N: N0.getNode())) {
17515 auto *LN0 = cast<LoadSDNode>(Val&: N0);
17516
17517 if (N0.hasOneUse() && ExtVT == LN0->getMemoryVT() &&
17518 ((!LegalOperations && LN0->isSimple()) &&
17519 TLI.isLoadLegal(ValVT: VT, MemVT: ExtVT, Alignment: LN0->getAlign(), AddrSpace: LN0->getAddressSpace(),
17520 ExtType: ISD::SEXTLOAD, Atomic: false))) {
17521 SDValue ExtLoad =
17522 DAG.getExtLoad(ExtType: ISD::SEXTLOAD, dl: DL, VT, Chain: LN0->getChain(),
17523 Ptr: LN0->getBasePtr(), MemVT: ExtVT, MMO: LN0->getMemOperand());
17524 CombineTo(N, Res: ExtLoad);
17525 CombineTo(N: N0.getNode(), Res0: ExtLoad, Res1: ExtLoad.getValue(R: 1));
17526 return SDValue(N, 0); // Return N so it doesn't get rechecked!
17527 }
17528 }
17529
17530 // fold (sext_inreg (masked_load x)) -> (sext_masked_load x)
17531 // ignore it if the masked load is already sign extended
17532 bool Frozen = N0.getOpcode() == ISD::FREEZE && N0.hasOneUse();
17533 if (auto *Ld = dyn_cast<MaskedLoadSDNode>(Val: Frozen ? N0.getOperand(i: 0) : N0)) {
17534 if (ExtVT == Ld->getMemoryVT() && Ld->hasNUsesOfValue(NUses: 1, Value: 0) &&
17535 Ld->getExtensionType() != ISD::LoadExtType::NON_EXTLOAD &&
17536 TLI.isLoadLegal(ValVT: VT, MemVT: ExtVT, Alignment: Ld->getAlign(), AddrSpace: Ld->getAddressSpace(),
17537 ExtType: ISD::SEXTLOAD, Atomic: false)) {
17538 SDValue ExtMaskedLoad = DAG.getMaskedLoad(
17539 VT, dl: DL, Chain: Ld->getChain(), Base: Ld->getBasePtr(), Offset: Ld->getOffset(),
17540 Mask: Ld->getMask(), Src0: Ld->getPassThru(), MemVT: ExtVT, MMO: Ld->getMemOperand(),
17541 AM: Ld->getAddressingMode(), ISD::SEXTLOAD, IsExpanding: Ld->isExpandingLoad());
17542 CombineTo(N, Res: Frozen ? N0 : ExtMaskedLoad);
17543 CombineTo(N: Ld, Res0: ExtMaskedLoad, Res1: ExtMaskedLoad.getValue(R: 1));
17544 return SDValue(N, 0); // Return N so it doesn't get rechecked!
17545 }
17546 }
17547
17548 // fold (sext_inreg (masked_gather x)) -> (sext_masked_gather x)
17549 if (auto *GN0 = dyn_cast<MaskedGatherSDNode>(Val&: N0)) {
17550 if (SDValue(GN0, 0).hasOneUse() && ExtVT == GN0->getMemoryVT() &&
17551 TLI.isVectorLoadExtDesirable(ExtVal: SDValue(N, 0))) {
17552 SDValue Ops[] = {GN0->getChain(), GN0->getPassThru(), GN0->getMask(),
17553 GN0->getBasePtr(), GN0->getIndex(), GN0->getScale()};
17554
17555 SDValue ExtLoad = DAG.getMaskedGather(
17556 VTs: DAG.getVTList(VT1: VT, VT2: MVT::Other), MemVT: ExtVT, dl: DL, Ops, MMO: GN0->getMemOperand(),
17557 IndexType: GN0->getIndexType(), ExtTy: ISD::SEXTLOAD);
17558
17559 CombineTo(N, Res: ExtLoad);
17560 CombineTo(N: N0.getNode(), Res0: ExtLoad, Res1: ExtLoad.getValue(R: 1));
17561 AddToWorklist(N: ExtLoad.getNode());
17562 return SDValue(N, 0); // Return N so it doesn't get rechecked!
17563 }
17564 }
17565
17566 // Form (sext_inreg (bswap >> 16)) or (sext_inreg (rotl (bswap) 16))
17567 if (ExtVTBits <= 16 && N0.getOpcode() == ISD::OR) {
17568 if (SDValue BSwap = MatchBSwapHWordLow(N: N0.getNode(), N0: N0.getOperand(i: 0),
17569 N1: N0.getOperand(i: 1), DemandHighBits: false))
17570 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT, N1: BSwap, N2: N1);
17571 }
17572
17573 // Fold (iM_signext_inreg
17574 // (extract_subvector (zext|anyext|sext iN_v to _) _)
17575 // from iN)
17576 // -> (extract_subvector (signext iN_v to iM))
17577 if (N0.getOpcode() == ISD::EXTRACT_SUBVECTOR && N0.hasOneUse() &&
17578 ISD::isExtOpcode(Opcode: N0.getOperand(i: 0).getOpcode())) {
17579 SDValue InnerExt = N0.getOperand(i: 0);
17580 EVT InnerExtVT = InnerExt->getValueType(ResNo: 0);
17581 SDValue Extendee = InnerExt->getOperand(Num: 0);
17582
17583 if (ExtVTBits == Extendee.getValueType().getScalarSizeInBits() &&
17584 (!LegalOperations ||
17585 TLI.isOperationLegal(Op: ISD::SIGN_EXTEND, VT: InnerExtVT))) {
17586 SDValue SignExtExtendee =
17587 DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL, VT: InnerExtVT, Operand: Extendee);
17588 return DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT, N1: SignExtExtendee,
17589 N2: N0.getOperand(i: 1));
17590 }
17591 }
17592
17593 return SDValue();
17594}
17595
17596static SDValue foldExtendVectorInregToExtendOfSubvector(
17597 SDNode *N, const SDLoc &DL, const TargetLowering &TLI, SelectionDAG &DAG,
17598 bool LegalOperations) {
17599 unsigned InregOpcode = N->getOpcode();
17600 unsigned Opcode = DAG.getOpcode_EXTEND(Opcode: InregOpcode);
17601
17602 SDValue Src = N->getOperand(Num: 0);
17603 EVT VT = N->getValueType(ResNo: 0);
17604 EVT SrcVT = VT.changeVectorElementType(
17605 Context&: *DAG.getContext(), EltVT: Src.getValueType().getVectorElementType());
17606
17607 assert(ISD::isExtVecInRegOpcode(InregOpcode) &&
17608 "Expected EXTEND_VECTOR_INREG dag node in input!");
17609
17610 // Profitability check: our operand must be an one-use CONCAT_VECTORS.
17611 // FIXME: one-use check may be overly restrictive
17612 if (!Src.hasOneUse() || Src.getOpcode() != ISD::CONCAT_VECTORS)
17613 return SDValue();
17614
17615 // Profitability check: we must be extending exactly one of it's operands.
17616 // FIXME: this is probably overly restrictive.
17617 Src = Src.getOperand(i: 0);
17618 if (Src.getValueType() != SrcVT)
17619 return SDValue();
17620
17621 if (LegalOperations && !TLI.isOperationLegal(Op: Opcode, VT))
17622 return SDValue();
17623
17624 return DAG.getNode(Opcode, DL, VT, Operand: Src);
17625}
17626
17627SDValue DAGCombiner::visitEXTEND_VECTOR_INREG(SDNode *N) {
17628 SDValue N0 = N->getOperand(Num: 0);
17629 EVT VT = N->getValueType(ResNo: 0);
17630 SDLoc DL(N);
17631
17632 if (N0.isUndef()) {
17633 // aext_vector_inreg(undef) = undef because the top bits are undefined.
17634 // {s/z}ext_vector_inreg(undef) = 0 because the top bits must be the same.
17635 return N->getOpcode() == ISD::ANY_EXTEND_VECTOR_INREG
17636 ? DAG.getUNDEF(VT)
17637 : DAG.getConstant(Val: 0, DL, VT);
17638 }
17639
17640 if (SDValue Res = tryToFoldExtendOfConstant(N, DL, TLI, DAG, LegalTypes))
17641 return Res;
17642
17643 if (SimplifyDemandedVectorElts(Op: SDValue(N, 0)))
17644 return SDValue(N, 0);
17645
17646 if (SDValue R = foldExtendVectorInregToExtendOfSubvector(N, DL, TLI, DAG,
17647 LegalOperations))
17648 return R;
17649
17650 return SDValue();
17651}
17652
17653SDValue DAGCombiner::visitTRUNCATE_USAT_U(SDNode *N) {
17654 EVT VT = N->getValueType(ResNo: 0);
17655 SDValue N0 = N->getOperand(Num: 0);
17656
17657 SDValue FPVal;
17658 if (sd_match(N: N0, P: m_FPToUI(Op: m_Value(N&: FPVal))) &&
17659 DAG.getTargetLoweringInfo().shouldConvertFpToSat(
17660 Op: ISD::FP_TO_UINT_SAT, FPVT: FPVal.getValueType(), VT))
17661 return DAG.getNode(Opcode: ISD::FP_TO_UINT_SAT, DL: SDLoc(N0), VT, N1: FPVal,
17662 N2: DAG.getValueType(VT.getScalarType()));
17663
17664 return SDValue();
17665}
17666
17667/// Detect patterns of truncation with unsigned saturation:
17668///
17669/// (truncate (umin (x, unsigned_max_of_dest_type)) to dest_type).
17670/// Return the source value x to be truncated or SDValue() if the pattern was
17671/// not matched.
17672///
17673static SDValue detectUSatUPattern(SDValue In, EVT VT) {
17674 unsigned NumDstBits = VT.getScalarSizeInBits();
17675 unsigned NumSrcBits = In.getScalarValueSizeInBits();
17676 // Saturation with truncation. We truncate from InVT to VT.
17677 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
17678
17679 SDValue Min;
17680 APInt UnsignedMax = APInt::getMaxValue(numBits: NumDstBits).zext(width: NumSrcBits);
17681 if (sd_match(N: In, P: m_UMin(L: m_Value(N&: Min), R: m_SpecificInt(V: UnsignedMax))))
17682 return Min;
17683
17684 return SDValue();
17685}
17686
17687/// Detect patterns of truncation with signed saturation:
17688/// (truncate (smin (smax (x, signed_min_of_dest_type),
17689/// signed_max_of_dest_type)) to dest_type)
17690/// or:
17691/// (truncate (smax (smin (x, signed_max_of_dest_type),
17692/// signed_min_of_dest_type)) to dest_type).
17693///
17694/// Return the source value to be truncated or SDValue() if the pattern was not
17695/// matched.
17696static SDValue detectSSatSPattern(SDValue In, EVT VT) {
17697 unsigned NumDstBits = VT.getScalarSizeInBits();
17698 unsigned NumSrcBits = In.getScalarValueSizeInBits();
17699 // Saturation with truncation. We truncate from InVT to VT.
17700 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
17701
17702 SDValue Val;
17703 APInt SignedMax = APInt::getSignedMaxValue(numBits: NumDstBits).sext(width: NumSrcBits);
17704 APInt SignedMin = APInt::getSignedMinValue(numBits: NumDstBits).sext(width: NumSrcBits);
17705
17706 if (sd_match(N: In, P: m_SMin(L: m_SMax(L: m_Value(N&: Val), R: m_SpecificInt(V: SignedMin)),
17707 R: m_SpecificInt(V: SignedMax))))
17708 return Val;
17709
17710 if (sd_match(N: In, P: m_SMax(L: m_SMin(L: m_Value(N&: Val), R: m_SpecificInt(V: SignedMax)),
17711 R: m_SpecificInt(V: SignedMin))))
17712 return Val;
17713
17714 return SDValue();
17715}
17716
17717/// Detect patterns of truncation with unsigned saturation:
17718static SDValue detectSSatUPattern(SDValue In, EVT VT, SelectionDAG &DAG,
17719 const SDLoc &DL) {
17720 unsigned NumDstBits = VT.getScalarSizeInBits();
17721 unsigned NumSrcBits = In.getScalarValueSizeInBits();
17722 // Saturation with truncation. We truncate from InVT to VT.
17723 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
17724
17725 SDValue Val;
17726 APInt UnsignedMax = APInt::getMaxValue(numBits: NumDstBits).zext(width: NumSrcBits);
17727 // Min == 0, Max is unsigned max of destination type.
17728 if (sd_match(N: In, P: m_SMax(L: m_SMin(L: m_Value(N&: Val), R: m_SpecificInt(V: UnsignedMax)),
17729 R: m_Zero())))
17730 return Val;
17731
17732 if (sd_match(N: In, P: m_SMin(L: m_SMax(L: m_Value(N&: Val), R: m_Zero()),
17733 R: m_SpecificInt(V: UnsignedMax))))
17734 return Val;
17735
17736 if (sd_match(N: In, P: m_UMin(L: m_SMax(L: m_Value(N&: Val), R: m_Zero()),
17737 R: m_SpecificInt(V: UnsignedMax))))
17738 return Val;
17739
17740 return SDValue();
17741}
17742
17743static SDValue foldToSaturated(SDNode *N, EVT &VT, SDValue &Src, EVT &SrcVT,
17744 SDLoc &DL, const TargetLowering &TLI,
17745 SelectionDAG &DAG) {
17746 auto AllowedTruncateSat = [&](unsigned Opc, EVT SrcVT, EVT VT) -> bool {
17747 return (TLI.isOperationLegalOrCustom(Op: Opc, VT: SrcVT) &&
17748 TLI.isTypeDesirableForOp(Opc, VT));
17749 };
17750
17751 if (Src.getOpcode() == ISD::SMIN || Src.getOpcode() == ISD::SMAX) {
17752 if (AllowedTruncateSat(ISD::TRUNCATE_SSAT_S, SrcVT, VT))
17753 if (SDValue SSatVal = detectSSatSPattern(In: Src, VT))
17754 return DAG.getNode(Opcode: ISD::TRUNCATE_SSAT_S, DL, VT, Operand: SSatVal);
17755 if (AllowedTruncateSat(ISD::TRUNCATE_SSAT_U, SrcVT, VT))
17756 if (SDValue SSatVal = detectSSatUPattern(In: Src, VT, DAG, DL))
17757 return DAG.getNode(Opcode: ISD::TRUNCATE_SSAT_U, DL, VT, Operand: SSatVal);
17758 } else if (Src.getOpcode() == ISD::UMIN) {
17759 if (AllowedTruncateSat(ISD::TRUNCATE_SSAT_U, SrcVT, VT))
17760 if (SDValue SSatVal = detectSSatUPattern(In: Src, VT, DAG, DL))
17761 return DAG.getNode(Opcode: ISD::TRUNCATE_SSAT_U, DL, VT, Operand: SSatVal);
17762 if (AllowedTruncateSat(ISD::TRUNCATE_USAT_U, SrcVT, VT))
17763 if (SDValue USatVal = detectUSatUPattern(In: Src, VT))
17764 return DAG.getNode(Opcode: ISD::TRUNCATE_USAT_U, DL, VT, Operand: USatVal);
17765 }
17766
17767 return SDValue();
17768}
17769
17770SDValue DAGCombiner::visitTRUNCATE(SDNode *N) {
17771 SDValue N0 = N->getOperand(Num: 0);
17772 EVT VT = N->getValueType(ResNo: 0);
17773 EVT SrcVT = N0.getValueType();
17774 bool isLE = DAG.getDataLayout().isLittleEndian();
17775 SDLoc DL(N);
17776
17777 // trunc(undef) = undef, trunc(poison) = poison
17778 if (SDValue R = propagateUnaryUndef(DAG, N0, VT))
17779 return R;
17780
17781 // fold (truncate (truncate x)) -> (truncate x)
17782 if (N0.getOpcode() == ISD::TRUNCATE)
17783 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N0.getOperand(i: 0));
17784
17785 // fold saturated truncate
17786 if (SDValue SaturatedTR = foldToSaturated(N, VT, Src&: N0, SrcVT, DL, TLI, DAG))
17787 return SaturatedTR;
17788
17789 // fold (truncate c1) -> c1
17790 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::TRUNCATE, DL, VT, Ops: {N0}))
17791 return C;
17792
17793 // fold (truncate (ext x)) -> (ext x) or (truncate x) or x
17794 if (N0.getOpcode() == ISD::ZERO_EXTEND ||
17795 N0.getOpcode() == ISD::SIGN_EXTEND ||
17796 N0.getOpcode() == ISD::ANY_EXTEND) {
17797 // if the source is smaller than the dest, we still need an extend.
17798 if (N0.getOperand(i: 0).getValueType().bitsLT(VT)) {
17799 SDNodeFlags Flags;
17800 if (N0.getOpcode() == ISD::ZERO_EXTEND)
17801 Flags.setNonNeg(N0->getFlags().hasNonNeg());
17802 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, Operand: N0.getOperand(i: 0), Flags);
17803 }
17804 // if the source is larger than the dest, than we just need the truncate.
17805 if (N0.getOperand(i: 0).getValueType().bitsGT(VT))
17806 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N0.getOperand(i: 0));
17807 // if the source and dest are the same type, we can drop both the extend
17808 // and the truncate.
17809 return N0.getOperand(i: 0);
17810 }
17811
17812 // Try to narrow a truncate-of-sext_in_reg to the destination type:
17813 // trunc (sign_ext_inreg X, iM) to iN --> sign_ext_inreg (trunc X to iN), iM
17814 if (!LegalTypes && N0.getOpcode() == ISD::SIGN_EXTEND_INREG &&
17815 N0.hasOneUse()) {
17816 SDValue X = N0.getOperand(i: 0);
17817 SDValue ExtVal = N0.getOperand(i: 1);
17818 EVT ExtVT = cast<VTSDNode>(Val&: ExtVal)->getVT();
17819 if (ExtVT.bitsLT(VT) && TLI.preferSextInRegOfTruncate(TruncVT: VT, VT: SrcVT, ExtVT)) {
17820 SDValue TrX = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: X);
17821 return DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT, N1: TrX, N2: ExtVal);
17822 }
17823 }
17824
17825 // If this is anyext(trunc), don't fold it, allow ourselves to be folded.
17826 if (N->hasOneUse() && (N->user_begin()->getOpcode() == ISD::ANY_EXTEND))
17827 return SDValue();
17828
17829 // Fold extract-and-trunc into a narrow extract. For example:
17830 // i64 x = EXTRACT_VECTOR_ELT(v2i64 val, i32 1)
17831 // i32 y = TRUNCATE(i64 x)
17832 // -- becomes --
17833 // v16i8 b = BITCAST (v2i64 val)
17834 // i8 x = EXTRACT_VECTOR_ELT(v16i8 b, i32 8)
17835 //
17836 // Note: We only run this optimization after type legalization (which often
17837 // creates this pattern) and before operation legalization after which
17838 // we need to be more careful about the vector instructions that we generate.
17839 if (LegalTypes && !LegalOperations && VT.isScalarInteger() && VT != MVT::i1 &&
17840 N0->hasOneUse()) {
17841 EVT TrTy = N->getValueType(ResNo: 0);
17842 SDValue Src = N0;
17843
17844 // Check for cases where we shift down an upper element before truncation.
17845 int EltOffset = 0;
17846 if (Src.getOpcode() == ISD::SRL && Src.getOperand(i: 0)->hasOneUse()) {
17847 if (auto ShAmt = DAG.getValidShiftAmount(V: Src)) {
17848 if ((*ShAmt % TrTy.getSizeInBits()) == 0) {
17849 Src = Src.getOperand(i: 0);
17850 EltOffset = *ShAmt / TrTy.getSizeInBits();
17851 }
17852 }
17853 }
17854
17855 if (Src.getOpcode() == ISD::EXTRACT_VECTOR_ELT) {
17856 EVT VecTy = Src.getOperand(i: 0).getValueType();
17857 EVT ExTy = Src.getValueType();
17858
17859 auto EltCnt = VecTy.getVectorElementCount();
17860 unsigned SizeRatio = ExTy.getSizeInBits() / TrTy.getSizeInBits();
17861 auto NewEltCnt = EltCnt * SizeRatio;
17862
17863 EVT NVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: TrTy, EC: NewEltCnt);
17864 assert(NVT.getSizeInBits() == VecTy.getSizeInBits() && "Invalid Size");
17865
17866 SDValue EltNo = Src->getOperand(Num: 1);
17867 if (isa<ConstantSDNode>(Val: EltNo) && isTypeLegal(VT: NVT)) {
17868 int Elt = EltNo->getAsZExtVal();
17869 int Index = isLE ? (Elt * SizeRatio + EltOffset)
17870 : (Elt * SizeRatio + (SizeRatio - 1) - EltOffset);
17871 return DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: TrTy,
17872 N1: DAG.getBitcast(VT: NVT, V: Src.getOperand(i: 0)),
17873 N2: DAG.getVectorIdxConstant(Val: Index, DL));
17874 }
17875 }
17876 }
17877
17878 // trunc (select c, a, b) -> select c, (trunc a), (trunc b)
17879 if (N0.getOpcode() == ISD::SELECT && N0.hasOneUse() &&
17880 TLI.isTruncateFree(FromVT: SrcVT, ToVT: VT)) {
17881 if (!LegalOperations ||
17882 (TLI.isOperationLegal(Op: ISD::SELECT, VT: SrcVT) &&
17883 TLI.isNarrowingProfitable(N: N0.getNode(), SrcVT, DestVT: VT))) {
17884 SDLoc SL(N0);
17885 SDValue Cond = N0.getOperand(i: 0);
17886 SDValue TruncOp0 = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SL, VT, Operand: N0.getOperand(i: 1));
17887 SDValue TruncOp1 = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SL, VT, Operand: N0.getOperand(i: 2));
17888 return DAG.getNode(Opcode: ISD::SELECT, DL, VT, N1: Cond, N2: TruncOp0, N3: TruncOp1);
17889 }
17890 }
17891
17892 // trunc (shl x, K) -> shl (trunc x), K => K < VT.getScalarSizeInBits()
17893 if (N0.getOpcode() == ISD::SHL && N0.hasOneUse() &&
17894 (!LegalOperations || TLI.isOperationLegal(Op: ISD::SHL, VT)) &&
17895 TLI.isTypeDesirableForOp(ISD::SHL, VT)) {
17896 SDValue Amt = N0.getOperand(i: 1);
17897 KnownBits Known = DAG.computeKnownBits(Op: Amt);
17898 unsigned Size = VT.getScalarSizeInBits();
17899 if (Known.countMaxActiveBits() <= Log2_32(Value: Size)) {
17900 EVT AmtVT = TLI.getShiftAmountTy(LHSTy: VT, DL: DAG.getDataLayout());
17901 SDValue Trunc = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N0.getOperand(i: 0));
17902 if (AmtVT != Amt.getValueType()) {
17903 Amt = DAG.getZExtOrTrunc(Op: Amt, DL, VT: AmtVT);
17904 AddToWorklist(N: Amt.getNode());
17905 }
17906 return DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: Trunc, N2: Amt);
17907 }
17908 }
17909
17910 if (SDValue V = foldSubToUSubSat(DstVT: VT, N: N0.getNode(), DL))
17911 return V;
17912
17913 if (SDValue ABD = foldABSToABD(N, DL))
17914 return ABD;
17915
17916 // Attempt to pre-truncate BUILD_VECTOR sources.
17917 if (N0.getOpcode() == ISD::BUILD_VECTOR && !LegalOperations &&
17918 N0.hasOneUse() &&
17919 // Avoid creating illegal types if running after type legalizer.
17920 (!LegalTypes || TLI.isTypeLegal(VT: VT.getScalarType()))) {
17921 if (TLI.isTruncateFree(FromVT: SrcVT.getScalarType(), ToVT: VT.getScalarType()))
17922 return DAG.UnrollVectorOp(N);
17923
17924 // trunc(build_vector(ext(x), ext(x)) -> build_vector(x,x)
17925 if (SDValue SplatVal = DAG.getSplatValue(V: N0)) {
17926 if (ISD::isExtOpcode(Opcode: SplatVal.getOpcode()) &&
17927 SrcVT.getScalarType() == SplatVal.getValueType())
17928 return DAG.UnrollVectorOp(N);
17929 }
17930 }
17931
17932 // trunc (splat_vector x) -> splat_vector (trunc x)
17933 if (N0.getOpcode() == ISD::SPLAT_VECTOR &&
17934 (!LegalTypes || TLI.isTypeLegal(VT: VT.getScalarType())) &&
17935 (!LegalOperations || TLI.isOperationLegal(Op: ISD::SPLAT_VECTOR, VT))) {
17936 EVT SVT = VT.getScalarType();
17937 return DAG.getSplatVector(
17938 VT, DL, Op: DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: SVT, Operand: N0->getOperand(Num: 0)));
17939 }
17940
17941 // Fold a series of buildvector, bitcast, and truncate if possible.
17942 // For example fold
17943 // (2xi32 trunc (bitcast ((4xi32)buildvector x, x, y, y) 2xi64)) to
17944 // (2xi32 (buildvector x, y)).
17945 if (Level == AfterLegalizeVectorOps && VT.isVector() &&
17946 N0.getOpcode() == ISD::BITCAST && N0.hasOneUse() &&
17947 N0.getOperand(i: 0).getOpcode() == ISD::BUILD_VECTOR &&
17948 N0.getOperand(i: 0).hasOneUse()) {
17949 SDValue BuildVect = N0.getOperand(i: 0);
17950 EVT BuildVectEltTy = BuildVect.getValueType().getVectorElementType();
17951 EVT TruncVecEltTy = VT.getVectorElementType();
17952
17953 // Check that the element types match.
17954 if (BuildVectEltTy == TruncVecEltTy) {
17955 // Now we only need to compute the offset of the truncated elements.
17956 unsigned BuildVecNumElts = BuildVect.getNumOperands();
17957 unsigned TruncVecNumElts = VT.getVectorNumElements();
17958 unsigned TruncEltOffset = BuildVecNumElts / TruncVecNumElts;
17959 unsigned FirstElt = isLE ? 0 : (TruncEltOffset - 1);
17960
17961 assert((BuildVecNumElts % TruncVecNumElts) == 0 &&
17962 "Invalid number of elements");
17963
17964 SmallVector<SDValue, 8> Opnds;
17965 for (unsigned i = FirstElt, e = BuildVecNumElts; i < e;
17966 i += TruncEltOffset)
17967 Opnds.push_back(Elt: BuildVect.getOperand(i));
17968
17969 return DAG.getBuildVector(VT, DL, Ops: Opnds);
17970 }
17971 }
17972
17973 // fold (truncate (load x)) -> (smaller load x)
17974 // fold (truncate (srl (load x), c)) -> (smaller load (x+c/evtbits))
17975 if (!LegalTypes || TLI.isTypeDesirableForOp(N: N0.getNode(), VT)) {
17976 if (SDValue Reduced = reduceLoadWidth(N))
17977 return Reduced;
17978
17979 // Handle the case where the truncated result is at least as wide as the
17980 // loaded type.
17981 if (N0.hasOneUse() && ISD::isUNINDEXEDLoad(N: N0.getNode())) {
17982 auto *LN0 = cast<LoadSDNode>(Val&: N0);
17983 if (LN0->isSimple() && LN0->getMemoryVT().bitsLE(VT)) {
17984 SDValue NewLoad = DAG.getExtLoad(
17985 ExtType: LN0->getExtensionType(), dl: SDLoc(LN0), VT, Chain: LN0->getChain(),
17986 Ptr: LN0->getBasePtr(), MemVT: LN0->getMemoryVT(), MMO: LN0->getMemOperand());
17987 DAG.ReplaceAllUsesOfValueWith(From: N0.getValue(R: 1), To: NewLoad.getValue(R: 1));
17988 return NewLoad;
17989 }
17990 }
17991 }
17992
17993 // fold (trunc (concat ... x ...)) -> (concat ..., (trunc x), ...)),
17994 // where ... are all 'undef'.
17995 if (N0.getOpcode() == ISD::CONCAT_VECTORS && !LegalTypes) {
17996 SmallVector<EVT, 8> VTs;
17997 SDValue V;
17998 unsigned Idx = 0;
17999 unsigned NumDefs = 0;
18000
18001 for (unsigned i = 0, e = N0.getNumOperands(); i != e; ++i) {
18002 SDValue X = N0.getOperand(i);
18003 if (!X.isUndef()) {
18004 V = X;
18005 Idx = i;
18006 NumDefs++;
18007 }
18008 // Stop if more than one members are non-undef.
18009 if (NumDefs > 1)
18010 break;
18011
18012 VTs.push_back(Elt: EVT::getVectorVT(Context&: *DAG.getContext(),
18013 VT: VT.getVectorElementType(),
18014 EC: X.getValueType().getVectorElementCount()));
18015 }
18016
18017 if (NumDefs == 0)
18018 return DAG.getUNDEF(VT);
18019
18020 if (NumDefs == 1) {
18021 assert(V.getNode() && "The single defined operand is empty!");
18022 SmallVector<SDValue, 8> Opnds;
18023 for (unsigned i = 0, e = VTs.size(); i != e; ++i) {
18024 if (i != Idx) {
18025 Opnds.push_back(Elt: DAG.getUNDEF(VT: VTs[i]));
18026 continue;
18027 }
18028 SDValue NV = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(V), VT: VTs[i], Operand: V);
18029 AddToWorklist(N: NV.getNode());
18030 Opnds.push_back(Elt: NV);
18031 }
18032 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT, Ops: Opnds);
18033 }
18034 }
18035
18036 // Fold truncate of a bitcast of a vector to an extract of the low vector
18037 // element.
18038 //
18039 // e.g. trunc (i64 (bitcast v2i32:x)) -> extract_vector_elt v2i32:x, idx
18040 if (N0.getOpcode() == ISD::BITCAST && !VT.isVector()) {
18041 SDValue VecSrc = N0.getOperand(i: 0);
18042 EVT VecSrcVT = VecSrc.getValueType();
18043 if (VecSrcVT.isVectorOf(EltVT: VT) &&
18044 (!LegalOperations ||
18045 TLI.isOperationLegal(Op: ISD::EXTRACT_VECTOR_ELT, VT: VecSrcVT))) {
18046 unsigned Idx = isLE ? 0 : VecSrcVT.getVectorNumElements() - 1;
18047 return DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT, N1: VecSrc,
18048 N2: DAG.getVectorIdxConstant(Val: Idx, DL));
18049 }
18050 }
18051
18052 // Simplify the operands using demanded-bits information.
18053 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
18054 return SDValue(N, 0);
18055
18056 // fold (truncate (extract_subvector(ext x))) ->
18057 // (extract_subvector x)
18058 // TODO: This can be generalized to cover cases where the truncate and extract
18059 // do not fully cancel each other out.
18060 if (!LegalTypes && N0.getOpcode() == ISD::EXTRACT_SUBVECTOR) {
18061 SDValue N00 = N0.getOperand(i: 0);
18062 if (N00.getOpcode() == ISD::SIGN_EXTEND ||
18063 N00.getOpcode() == ISD::ZERO_EXTEND ||
18064 N00.getOpcode() == ISD::ANY_EXTEND) {
18065 if (N00.getOperand(i: 0)->getValueType(ResNo: 0).getVectorElementType() ==
18066 VT.getVectorElementType())
18067 return DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL: SDLoc(N0->getOperand(Num: 0)), VT,
18068 N1: N00.getOperand(i: 0), N2: N0.getOperand(i: 1));
18069 }
18070 }
18071
18072 if (SDValue NewVSel = matchVSelectOpSizesWithSetCC(Cast: N))
18073 return NewVSel;
18074
18075 // Narrow a suitable binary operation with a non-opaque constant operand by
18076 // moving it ahead of the truncate. This is limited to pre-legalization
18077 // because targets may prefer a wider type during later combines and invert
18078 // this transform.
18079 switch (N0.getOpcode()) {
18080 case ISD::ADD:
18081 case ISD::SUB:
18082 case ISD::MUL:
18083 case ISD::AND:
18084 case ISD::OR:
18085 case ISD::XOR:
18086 if (!LegalOperations && N0.hasOneUse() &&
18087 (N0.getOperand(i: 0) == N0.getOperand(i: 1) ||
18088 isConstantOrConstantVector(N: N0.getOperand(i: 0), NoOpaques: true) ||
18089 isConstantOrConstantVector(N: N0.getOperand(i: 1), NoOpaques: true))) {
18090 // TODO: We already restricted this to pre-legalization, but for vectors
18091 // we are extra cautious to not create an unsupported operation.
18092 // Target-specific changes are likely needed to avoid regressions here.
18093 if (VT.isScalarInteger() || TLI.isOperationLegal(Op: N0.getOpcode(), VT)) {
18094 SDValue NarrowL = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N0.getOperand(i: 0));
18095 SDValue NarrowR = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N0.getOperand(i: 1));
18096 SDNodeFlags Flags;
18097 // Propagate nuw for sub.
18098 if (N0->getOpcode() == ISD::SUB && N0->getFlags().hasNoUnsignedWrap() &&
18099 DAG.MaskedValueIsZero(
18100 Op: N0->getOperand(Num: 0),
18101 Mask: APInt::getBitsSetFrom(numBits: SrcVT.getScalarSizeInBits(),
18102 loBit: VT.getScalarSizeInBits())))
18103 Flags.setNoUnsignedWrap(true);
18104 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT, N1: NarrowL, N2: NarrowR, Flags);
18105 }
18106 }
18107 break;
18108 case ISD::ADDE:
18109 case ISD::UADDO_CARRY:
18110 // (trunc adde(X, Y, Carry)) -> (adde trunc(X), trunc(Y), Carry)
18111 // (trunc uaddo_carry(X, Y, Carry)) ->
18112 // (uaddo_carry trunc(X), trunc(Y), Carry)
18113 // When the adde's carry is not used.
18114 // We only do for uaddo_carry before legalize operation
18115 if (((!LegalOperations && N0.getOpcode() == ISD::UADDO_CARRY) ||
18116 TLI.isOperationLegal(Op: N0.getOpcode(), VT)) &&
18117 N0.hasOneUse() && !N0->hasAnyUseOfValue(Value: 1)) {
18118 SDValue X = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N0.getOperand(i: 0));
18119 SDValue Y = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: N0.getOperand(i: 1));
18120 SDVTList VTs = DAG.getVTList(VT1: VT, VT2: N0->getValueType(ResNo: 1));
18121 return DAG.getNode(Opcode: N0.getOpcode(), DL, VTList: VTs, N1: X, N2: Y, N3: N0.getOperand(i: 2));
18122 }
18123 break;
18124 case ISD::USUBSAT:
18125 // Truncate the USUBSAT only if LHS is a known zero-extension, its not
18126 // enough to know that the upper bits are zero we must ensure that we don't
18127 // introduce an extra truncate.
18128 if (!LegalOperations && N0.hasOneUse() &&
18129 N0.getOperand(i: 0).getOpcode() == ISD::ZERO_EXTEND &&
18130 N0.getOperand(i: 0).getOperand(i: 0).getScalarValueSizeInBits() <=
18131 VT.getScalarSizeInBits() &&
18132 hasOperation(Opcode: N0.getOpcode(), VT)) {
18133 return getTruncatedUSUBSAT(DstVT: VT, SrcVT, LHS: N0.getOperand(i: 0), RHS: N0.getOperand(i: 1),
18134 DAG, DL);
18135 }
18136 break;
18137 case ISD::AVGCEILS:
18138 case ISD::AVGCEILU:
18139 // trunc (avgceilu (sext (x), sext (y))) -> avgceils(x, y)
18140 // trunc (avgceils (zext (x), zext (y))) -> avgceilu(x, y)
18141 if (N0.hasOneUse()) {
18142 SDValue Op0 = N0.getOperand(i: 0);
18143 SDValue Op1 = N0.getOperand(i: 1);
18144 if (N0.getOpcode() == ISD::AVGCEILU) {
18145 if (TLI.isOperationLegalOrCustom(Op: ISD::AVGCEILS, VT) &&
18146 Op0.getOpcode() == ISD::SIGN_EXTEND &&
18147 Op1.getOpcode() == ISD::SIGN_EXTEND &&
18148 Op0.getOperand(i: 0).getValueType() == VT &&
18149 Op1.getOperand(i: 0).getValueType() == VT)
18150 return DAG.getNode(Opcode: ISD::AVGCEILS, DL, VT, N1: Op0.getOperand(i: 0),
18151 N2: Op1.getOperand(i: 0));
18152 } else {
18153 if (TLI.isOperationLegalOrCustom(Op: ISD::AVGCEILU, VT) &&
18154 Op0.getOpcode() == ISD::ZERO_EXTEND &&
18155 Op1.getOpcode() == ISD::ZERO_EXTEND &&
18156 Op0.getOperand(i: 0).getValueType() == VT &&
18157 Op1.getOperand(i: 0).getValueType() == VT)
18158 return DAG.getNode(Opcode: ISD::AVGCEILU, DL, VT, N1: Op0.getOperand(i: 0),
18159 N2: Op1.getOperand(i: 0));
18160 }
18161 }
18162 [[fallthrough]];
18163 case ISD::AVGFLOORS:
18164 case ISD::AVGFLOORU:
18165 case ISD::ABDS:
18166 case ISD::ABDU:
18167 // (trunc (avg a, b)) -> (avg (trunc a), (trunc b))
18168 // (trunc (abdu/abds a, b)) -> (abdu/abds (trunc a), (trunc b))
18169 if (!LegalOperations && N0.hasOneUse() &&
18170 TLI.isOperationLegal(Op: N0.getOpcode(), VT)) {
18171 EVT TruncVT = VT;
18172 unsigned SrcBits = SrcVT.getScalarSizeInBits();
18173 unsigned TruncBits = TruncVT.getScalarSizeInBits();
18174
18175 SDValue A = N0.getOperand(i: 0);
18176 SDValue B = N0.getOperand(i: 1);
18177 bool CanFold = false;
18178
18179 if (N0.getOpcode() == ISD::AVGFLOORU || N0.getOpcode() == ISD::AVGCEILU ||
18180 N0.getOpcode() == ISD::ABDU) {
18181 APInt UpperBits = APInt::getBitsSetFrom(numBits: SrcBits, loBit: TruncBits);
18182 CanFold = DAG.MaskedValueIsZero(Op: B, Mask: UpperBits) &&
18183 DAG.MaskedValueIsZero(Op: A, Mask: UpperBits);
18184 } else {
18185 unsigned NeededBits = SrcBits - TruncBits;
18186 CanFold = DAG.ComputeNumSignBits(Op: B) > NeededBits &&
18187 DAG.ComputeNumSignBits(Op: A) > NeededBits;
18188 }
18189
18190 if (CanFold) {
18191 SDValue NewA = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: TruncVT, Operand: A);
18192 SDValue NewB = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: TruncVT, Operand: B);
18193 return DAG.getNode(Opcode: N0.getOpcode(), DL, VT: TruncVT, N1: NewA, N2: NewB);
18194 }
18195 }
18196 break;
18197 }
18198
18199 return SDValue();
18200}
18201
18202static SDNode *getBuildPairElt(SDNode *N, unsigned i) {
18203 SDValue Elt = N->getOperand(Num: i);
18204 if (Elt.getOpcode() != ISD::MERGE_VALUES)
18205 return Elt.getNode();
18206 return Elt.getOperand(i: Elt.getResNo()).getNode();
18207}
18208
18209/// build_pair (load, load) -> load
18210/// if load locations are consecutive.
18211SDValue DAGCombiner::CombineConsecutiveLoads(SDNode *N, EVT VT) {
18212 assert(N->getOpcode() == ISD::BUILD_PAIR);
18213
18214 auto *LD1 = dyn_cast<LoadSDNode>(Val: getBuildPairElt(N, i: 0));
18215 auto *LD2 = dyn_cast<LoadSDNode>(Val: getBuildPairElt(N, i: 1));
18216
18217 // A BUILD_PAIR is always having the least significant part in elt 0 and the
18218 // most significant part in elt 1. So when combining into one large load, we
18219 // need to consider the endianness.
18220 if (DAG.getDataLayout().isBigEndian())
18221 std::swap(a&: LD1, b&: LD2);
18222
18223 if (!LD1 || !LD2 || !ISD::isNON_EXTLoad(N: LD1) || !ISD::isNON_EXTLoad(N: LD2) ||
18224 !LD1->hasOneUse() || !LD2->hasOneUse() ||
18225 LD1->getAddressSpace() != LD2->getAddressSpace())
18226 return SDValue();
18227
18228 unsigned LD1Fast = 0;
18229 EVT LD1VT = LD1->getValueType(ResNo: 0);
18230 unsigned LD1Bytes = LD1VT.getStoreSize();
18231 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::LOAD, VT)) &&
18232 DAG.areNonVolatileConsecutiveLoads(LD: LD2, Base: LD1, Bytes: LD1Bytes, Dist: 1) &&
18233 TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT,
18234 MMO: *LD1->getMemOperand(), Fast: &LD1Fast) && LD1Fast)
18235 return DAG.getLoad(VT, dl: SDLoc(N), Chain: LD1->getChain(), Ptr: LD1->getBasePtr(),
18236 PtrInfo: LD1->getPointerInfo(), Alignment: LD1->getAlign());
18237
18238 return SDValue();
18239}
18240
18241static unsigned getPPCf128HiElementSelector(const SelectionDAG &DAG) {
18242 // On little-endian machines, bitcasting from ppcf128 to i128 does swap the Hi
18243 // and Lo parts; on big-endian machines it doesn't.
18244 return DAG.getDataLayout().isBigEndian() ? 1 : 0;
18245}
18246
18247SDValue DAGCombiner::foldBitcastedFPLogic(SDNode *N, SelectionDAG &DAG,
18248 const TargetLowering &TLI) {
18249 // If this is not a bitcast to an FP type or if the target doesn't have
18250 // IEEE754-compliant FP logic, we're done.
18251 EVT VT = N->getValueType(ResNo: 0);
18252 SDValue N0 = N->getOperand(Num: 0);
18253 EVT SourceVT = N0.getValueType();
18254
18255 if (!VT.isFloatingPoint())
18256 return SDValue();
18257
18258 // TODO: Handle cases where the integer constant is a different scalar
18259 // bitwidth to the FP.
18260 if (VT.getScalarSizeInBits() != SourceVT.getScalarSizeInBits())
18261 return SDValue();
18262
18263 unsigned FPOpcode;
18264 APInt SignMask;
18265 switch (N0.getOpcode()) {
18266 case ISD::AND:
18267 FPOpcode = ISD::FABS;
18268 SignMask = ~APInt::getSignMask(BitWidth: SourceVT.getScalarSizeInBits());
18269 break;
18270 case ISD::XOR:
18271 FPOpcode = ISD::FNEG;
18272 SignMask = APInt::getSignMask(BitWidth: SourceVT.getScalarSizeInBits());
18273 break;
18274 case ISD::OR:
18275 FPOpcode = ISD::FABS;
18276 SignMask = APInt::getSignMask(BitWidth: SourceVT.getScalarSizeInBits());
18277 break;
18278 default:
18279 return SDValue();
18280 }
18281
18282 if (LegalOperations && !TLI.isOperationLegal(Op: FPOpcode, VT))
18283 return SDValue();
18284
18285 // This needs to be the inverse of logic in foldSignChangeInBitcast.
18286 // FIXME: I don't think looking for bitcast intrinsically makes sense, but
18287 // removing this would require more changes.
18288 auto IsBitCastOrFree = [&TLI, FPOpcode](SDValue Op, EVT VT) {
18289 if (sd_match(N: Op, P: m_BitCast(Op: m_SpecificVT(RefVT: VT))))
18290 return true;
18291
18292 return FPOpcode == ISD::FABS ? TLI.isFAbsFree(VT) : TLI.isFNegFree(VT);
18293 };
18294
18295 // Fold (bitcast int (and (bitcast fp X to int), 0x7fff...) to fp) -> fabs X
18296 // Fold (bitcast int (xor (bitcast fp X to int), 0x8000...) to fp) -> fneg X
18297 // Fold (bitcast int (or (bitcast fp X to int), 0x8000...) to fp) ->
18298 // fneg (fabs X)
18299 SDValue LogicOp0 = N0.getOperand(i: 0);
18300 ConstantSDNode *LogicOp1 = isConstOrConstSplat(N: N0.getOperand(i: 1), AllowUndefs: true);
18301 if (LogicOp1 && LogicOp1->getAPIntValue() == SignMask &&
18302 IsBitCastOrFree(LogicOp0, VT)) {
18303 SDValue CastOp0 = DAG.getNode(Opcode: ISD::BITCAST, DL: SDLoc(N), VT, Operand: LogicOp0);
18304 SDValue FPOp = DAG.getNode(Opcode: FPOpcode, DL: SDLoc(N), VT, Operand: CastOp0);
18305 NumFPLogicOpsConv++;
18306 if (N0.getOpcode() == ISD::OR)
18307 return DAG.getNode(Opcode: ISD::FNEG, DL: SDLoc(N), VT, Operand: FPOp);
18308 return FPOp;
18309 }
18310
18311 return SDValue();
18312}
18313
18314SDValue DAGCombiner::visitBITCAST(SDNode *N) {
18315 SDValue N0 = N->getOperand(Num: 0);
18316 EVT VT = N->getValueType(ResNo: 0);
18317
18318 if (SDValue R = propagateUnaryUndef(DAG, N0, VT))
18319 return R;
18320
18321 // If the input is a BUILD_VECTOR with all constant elements, fold this now.
18322 // Only do this before legalize types, unless both types are integer and the
18323 // scalar type is legal. Only do this before legalize ops, since the target
18324 // maybe depending on the bitcast.
18325 // First check to see if this is all constant.
18326 // TODO: Support FP bitcasts after legalize types.
18327 if (VT.isVector() &&
18328 (!LegalTypes ||
18329 (!LegalOperations && VT.isInteger() && N0.getValueType().isInteger() &&
18330 TLI.isTypeLegal(VT: VT.getVectorElementType()))) &&
18331 N0.getOpcode() == ISD::BUILD_VECTOR && N0->hasOneUse() &&
18332 cast<BuildVectorSDNode>(Val&: N0)->isConstant())
18333 return DAG.FoldConstantBuildVector(BV: cast<BuildVectorSDNode>(Val&: N0), DL: SDLoc(N),
18334 DstEltVT: VT.getVectorElementType());
18335
18336 // If the input is a constant, let getNode fold it.
18337 if (isIntOrFPConstant(V: N0)) {
18338 // If we can't allow illegal operations, we need to check that this is just
18339 // a fp -> int or int -> conversion and that the resulting operation will
18340 // be legal.
18341 if (!LegalOperations ||
18342 (isa<ConstantSDNode>(Val: N0) && VT.isFloatingPoint() && !VT.isVector() &&
18343 TLI.isOperationLegal(Op: ISD::ConstantFP, VT)) ||
18344 (isa<ConstantFPSDNode>(Val: N0) && VT.isInteger() && !VT.isVector() &&
18345 TLI.isOperationLegal(Op: ISD::Constant, VT))) {
18346 SDValue C = DAG.getBitcast(VT, V: N0);
18347 if (C.getNode() != N)
18348 return C;
18349 }
18350 }
18351
18352 // (conv (conv x, t1), t2) -> (conv x, t2)
18353 if (N0.getOpcode() == ISD::BITCAST)
18354 return DAG.getBitcast(VT, V: N0.getOperand(i: 0));
18355
18356 // fold (conv (logicop (conv x), (c))) -> (logicop x, (conv c))
18357 // iff the current bitwise logicop type isn't legal
18358 if (ISD::isBitwiseLogicOp(Opcode: N0.getOpcode()) && VT.isInteger() &&
18359 !TLI.isTypeLegal(VT: N0.getOperand(i: 0).getValueType())) {
18360 auto IsFreeBitcast = [VT](SDValue V) {
18361 return (V.getOpcode() == ISD::BITCAST &&
18362 V.getOperand(i: 0).getValueType() == VT) ||
18363 (ISD::isBuildVectorOfConstantSDNodes(N: V.getNode()) &&
18364 V->hasOneUse());
18365 };
18366 if (IsFreeBitcast(N0.getOperand(i: 0)) && IsFreeBitcast(N0.getOperand(i: 1)))
18367 return DAG.getNode(Opcode: N0.getOpcode(), DL: SDLoc(N), VT,
18368 N1: DAG.getBitcast(VT, V: N0.getOperand(i: 0)),
18369 N2: DAG.getBitcast(VT, V: N0.getOperand(i: 1)));
18370 }
18371
18372 // fold (conv (load x)) -> (load (conv*)x)
18373 // fold (conv (freeze (load x))) -> (freeze (load (conv*)x))
18374 // If the resultant load doesn't need a higher alignment than the original!
18375 auto CastLoad = [this, &VT](SDValue N0, const SDLoc &DL) {
18376 // Peek through scalar_to_vector if the scalar is same size as VT - often a
18377 // leftover from legalization.
18378 if (N0.getOpcode() == ISD::SCALAR_TO_VECTOR && N0.hasOneUse() &&
18379 N0.getOperand(i: 0).getValueSizeInBits() == VT.getSizeInBits())
18380 N0 = N0.getOperand(i: 0);
18381 if (N0.getOpcode() == ISD::AssertNoFPClass)
18382 N0 = N0.getOperand(i: 0);
18383 if (!ISD::isNormalLoad(N: N0.getNode()) || !N0.hasOneUse())
18384 return SDValue();
18385
18386 // Do not remove the cast if the types differ in endian layout.
18387 if (TLI.hasBigEndianPartOrdering(VT: N0.getValueType(), DL: DAG.getDataLayout()) !=
18388 TLI.hasBigEndianPartOrdering(VT, DL: DAG.getDataLayout()))
18389 return SDValue();
18390
18391 // If the load is volatile, we only want to change the load type if the
18392 // resulting load is legal. Otherwise we might increase the number of
18393 // memory accesses. We don't care if the original type was legal or not
18394 // as we assume software couldn't rely on the number of accesses of an
18395 // illegal type.
18396 auto *LN0 = cast<LoadSDNode>(Val&: N0);
18397 if ((LegalOperations || !LN0->isSimple()) &&
18398 !TLI.isOperationLegal(Op: ISD::LOAD, VT))
18399 return SDValue();
18400
18401 if (!TLI.isLoadBitCastBeneficial(LoadVT: N0.getValueType(), BitcastVT: VT, DAG,
18402 MMO: *LN0->getMemOperand()))
18403 return SDValue();
18404
18405 // If the range metadata type does not match the new memory
18406 // operation type, remove the range metadata.
18407 if (const MDNode *MD = LN0->getRanges()) {
18408 ConstantInt *Lower = mdconst::extract<ConstantInt>(MD: MD->getOperand(I: 0));
18409 if (Lower->getBitWidth() != VT.getScalarSizeInBits() || !VT.isInteger()) {
18410 LN0->getMemOperand()->clearRanges();
18411 }
18412 }
18413 SDValue Load = DAG.getLoad(VT, dl: DL, Chain: LN0->getChain(), Ptr: LN0->getBasePtr(),
18414 MMO: LN0->getMemOperand());
18415 DAG.ReplaceAllUsesOfValueWith(From: N0.getValue(R: 1), To: Load.getValue(R: 1));
18416 return Load;
18417 };
18418
18419 if (SDValue NewLd = CastLoad(N0, SDLoc(N)))
18420 return NewLd;
18421
18422 if (N0.getOpcode() == ISD::FREEZE && N0.hasOneUse())
18423 if (SDValue NewLd = CastLoad(N0.getOperand(i: 0), SDLoc(N)))
18424 return DAG.getFreeze(V: NewLd);
18425
18426 if (SDValue V = foldBitcastedFPLogic(N, DAG, TLI))
18427 return V;
18428
18429 // fold (bitconvert (fneg x)) -> (xor (bitconvert x), signbit)
18430 // fold (bitconvert (fabs x)) -> (and (bitconvert x), (not signbit))
18431 //
18432 // For ppc_fp128:
18433 // fold (bitcast (fneg x)) ->
18434 // flipbit = signbit
18435 // (xor (bitcast x) (build_pair flipbit, flipbit))
18436 //
18437 // fold (bitcast (fabs x)) ->
18438 // flipbit = (and (extract_element (bitcast x), 0), signbit)
18439 // (xor (bitcast x) (build_pair flipbit, flipbit))
18440 // This often reduces constant pool loads.
18441 if (((N0.getOpcode() == ISD::FNEG && !TLI.isFNegFree(VT: N0.getValueType())) ||
18442 (N0.getOpcode() == ISD::FABS && !TLI.isFAbsFree(VT: N0.getValueType()))) &&
18443 N0->hasOneUse() && VT.isInteger() && !VT.isVector() &&
18444 !N0.getValueType().isVector()) {
18445 SDValue NewConv = DAG.getBitcast(VT, V: N0.getOperand(i: 0));
18446 AddToWorklist(N: NewConv.getNode());
18447
18448 SDLoc DL(N);
18449 if (N0.getValueType() == MVT::ppcf128 && !LegalTypes) {
18450 assert(VT.getSizeInBits() == 128);
18451 SDValue SignBit = DAG.getConstant(
18452 Val: APInt::getSignMask(BitWidth: VT.getSizeInBits() / 2), DL: SDLoc(N0), VT: MVT::i64);
18453 SDValue FlipBit;
18454 if (N0.getOpcode() == ISD::FNEG) {
18455 FlipBit = SignBit;
18456 AddToWorklist(N: FlipBit.getNode());
18457 } else {
18458 assert(N0.getOpcode() == ISD::FABS);
18459 SDValue Hi =
18460 DAG.getNode(Opcode: ISD::EXTRACT_ELEMENT, DL: SDLoc(NewConv), VT: MVT::i64, N1: NewConv,
18461 N2: DAG.getIntPtrConstant(Val: getPPCf128HiElementSelector(DAG),
18462 DL: SDLoc(NewConv)));
18463 AddToWorklist(N: Hi.getNode());
18464 FlipBit = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(N0), VT: MVT::i64, N1: Hi, N2: SignBit);
18465 AddToWorklist(N: FlipBit.getNode());
18466 }
18467 SDValue FlipBits =
18468 DAG.getNode(Opcode: ISD::BUILD_PAIR, DL: SDLoc(N0), VT, N1: FlipBit, N2: FlipBit);
18469 AddToWorklist(N: FlipBits.getNode());
18470 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: NewConv, N2: FlipBits);
18471 }
18472 APInt SignBit = APInt::getSignMask(BitWidth: VT.getSizeInBits());
18473 if (N0.getOpcode() == ISD::FNEG)
18474 return DAG.getNode(Opcode: ISD::XOR, DL, VT,
18475 N1: NewConv, N2: DAG.getConstant(Val: SignBit, DL, VT));
18476 assert(N0.getOpcode() == ISD::FABS);
18477 return DAG.getNode(Opcode: ISD::AND, DL, VT,
18478 N1: NewConv, N2: DAG.getConstant(Val: ~SignBit, DL, VT));
18479 }
18480
18481 // fold (bitconvert (fcopysign cst, x)) ->
18482 // (or (and (bitconvert x), sign), (and cst, (not sign)))
18483 // Note that we don't handle (copysign x, cst) because this can always be
18484 // folded to an fneg or fabs.
18485 //
18486 // For ppc_fp128:
18487 // fold (bitcast (fcopysign cst, x)) ->
18488 // flipbit = (and (extract_element
18489 // (xor (bitcast cst), (bitcast x)), 0),
18490 // signbit)
18491 // (xor (bitcast cst) (build_pair flipbit, flipbit))
18492 if (N0.getOpcode() == ISD::FCOPYSIGN && N0->hasOneUse() &&
18493 isa<ConstantFPSDNode>(Val: N0.getOperand(i: 0)) && VT.isInteger() &&
18494 !VT.isVector()) {
18495 unsigned OrigXWidth = N0.getOperand(i: 1).getValueSizeInBits();
18496 EVT IntXVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: OrigXWidth);
18497 if (isTypeLegal(VT: IntXVT)) {
18498 SDValue X = DAG.getBitcast(VT: IntXVT, V: N0.getOperand(i: 1));
18499 AddToWorklist(N: X.getNode());
18500
18501 // If X has a different width than the result/lhs, sext it or truncate it.
18502 unsigned VTWidth = VT.getSizeInBits();
18503 if (OrigXWidth < VTWidth) {
18504 X = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL: SDLoc(N), VT, Operand: X);
18505 AddToWorklist(N: X.getNode());
18506 } else if (OrigXWidth > VTWidth) {
18507 // To get the sign bit in the right place, we have to shift it right
18508 // before truncating.
18509 SDLoc DL(X);
18510 X = DAG.getNode(Opcode: ISD::SRL, DL,
18511 VT: X.getValueType(), N1: X,
18512 N2: DAG.getConstant(Val: OrigXWidth-VTWidth, DL,
18513 VT: X.getValueType()));
18514 AddToWorklist(N: X.getNode());
18515 X = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(X), VT, Operand: X);
18516 AddToWorklist(N: X.getNode());
18517 }
18518
18519 if (N0.getValueType() == MVT::ppcf128 && !LegalTypes) {
18520 APInt SignBit = APInt::getSignMask(BitWidth: VT.getSizeInBits() / 2);
18521 SDValue Cst = DAG.getBitcast(VT, V: N0.getOperand(i: 0));
18522 AddToWorklist(N: Cst.getNode());
18523 SDValue X = DAG.getBitcast(VT, V: N0.getOperand(i: 1));
18524 AddToWorklist(N: X.getNode());
18525 SDValue XorResult = DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N0), VT, N1: Cst, N2: X);
18526 AddToWorklist(N: XorResult.getNode());
18527 SDValue XorResult64 = DAG.getNode(
18528 Opcode: ISD::EXTRACT_ELEMENT, DL: SDLoc(XorResult), VT: MVT::i64, N1: XorResult,
18529 N2: DAG.getIntPtrConstant(Val: getPPCf128HiElementSelector(DAG),
18530 DL: SDLoc(XorResult)));
18531 AddToWorklist(N: XorResult64.getNode());
18532 SDValue FlipBit =
18533 DAG.getNode(Opcode: ISD::AND, DL: SDLoc(XorResult64), VT: MVT::i64, N1: XorResult64,
18534 N2: DAG.getConstant(Val: SignBit, DL: SDLoc(XorResult64), VT: MVT::i64));
18535 AddToWorklist(N: FlipBit.getNode());
18536 SDValue FlipBits =
18537 DAG.getNode(Opcode: ISD::BUILD_PAIR, DL: SDLoc(N0), VT, N1: FlipBit, N2: FlipBit);
18538 AddToWorklist(N: FlipBits.getNode());
18539 return DAG.getNode(Opcode: ISD::XOR, DL: SDLoc(N), VT, N1: Cst, N2: FlipBits);
18540 }
18541 APInt SignBit = APInt::getSignMask(BitWidth: VT.getSizeInBits());
18542 X = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(X), VT,
18543 N1: X, N2: DAG.getConstant(Val: SignBit, DL: SDLoc(X), VT));
18544 AddToWorklist(N: X.getNode());
18545
18546 SDValue Cst = DAG.getBitcast(VT, V: N0.getOperand(i: 0));
18547 Cst = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(Cst), VT,
18548 N1: Cst, N2: DAG.getConstant(Val: ~SignBit, DL: SDLoc(Cst), VT));
18549 AddToWorklist(N: Cst.getNode());
18550
18551 return DAG.getNode(Opcode: ISD::OR, DL: SDLoc(N), VT, N1: X, N2: Cst);
18552 }
18553 }
18554
18555 // bitconvert(build_pair(ld, ld)) -> ld iff load locations are consecutive.
18556 if (N0.getOpcode() == ISD::BUILD_PAIR)
18557 if (SDValue CombineLD = CombineConsecutiveLoads(N: N0.getNode(), VT))
18558 return CombineLD;
18559
18560 // int_vt (bitcast (vec_vt (scalar_to_vector elt_vt:x)))
18561 // => int_vt (any_extend elt_vt:x)
18562 if (DAG.getDataLayout().isLittleEndian() &&
18563 N0.getOpcode() == ISD::SCALAR_TO_VECTOR && VT.isScalarInteger()) {
18564 SDValue SrcScalar = N0.getOperand(i: 0);
18565 EVT SrcVT = SrcScalar.getValueType();
18566 if (SrcVT.isScalarInteger() && VT.bitsGT(VT: SrcVT))
18567 return DAG.getNode(Opcode: ISD::ANY_EXTEND, DL: SDLoc(N), VT, Operand: SrcScalar);
18568 }
18569
18570 // vt (bitcast (scalar_to_vector vt:x)) -> x
18571 if (N0.getOpcode() == ISD::SCALAR_TO_VECTOR &&
18572 N0.getOperand(i: 0).getValueType() == VT)
18573 return N0.getOperand(i: 0);
18574
18575 // Remove double bitcasts from shuffles - this is often a legacy of
18576 // XformToShuffleWithZero being used to combine bitmaskings (of
18577 // float vectors bitcast to integer vectors) into shuffles.
18578 // bitcast(shuffle(bitcast(s0),bitcast(s1))) -> shuffle(s0,s1)
18579 if (Level < AfterLegalizeDAG && TLI.isTypeLegal(VT) && VT.isVector() &&
18580 N0->getOpcode() == ISD::VECTOR_SHUFFLE && N0.hasOneUse() &&
18581 VT.getVectorNumElements() >= N0.getValueType().getVectorNumElements() &&
18582 !(VT.getVectorNumElements() % N0.getValueType().getVectorNumElements())) {
18583 ShuffleVectorSDNode *SVN = cast<ShuffleVectorSDNode>(Val&: N0);
18584
18585 // If operands are a bitcast, peek through if it casts the original VT.
18586 // If operands are a constant, just bitcast back to original VT.
18587 auto PeekThroughBitcast = [&](SDValue Op) {
18588 if (Op.getOpcode() == ISD::BITCAST &&
18589 Op.getOperand(i: 0).getValueType() == VT)
18590 return SDValue(Op.getOperand(i: 0));
18591 if (Op.isUndef() || isAnyConstantBuildVector(V: Op))
18592 return DAG.getBitcast(VT, V: Op);
18593 return SDValue();
18594 };
18595
18596 // FIXME: If either input vector is bitcast, try to convert the shuffle to
18597 // the result type of this bitcast. This would eliminate at least one
18598 // bitcast. See the transform in InstCombine.
18599 SDValue SV0 = PeekThroughBitcast(N0->getOperand(Num: 0));
18600 SDValue SV1 = PeekThroughBitcast(N0->getOperand(Num: 1));
18601 if (!(SV0 && SV1))
18602 return SDValue();
18603
18604 int MaskScale =
18605 VT.getVectorNumElements() / N0.getValueType().getVectorNumElements();
18606 SmallVector<int, 8> NewMask;
18607 for (int M : SVN->getMask())
18608 for (int i = 0; i != MaskScale; ++i)
18609 NewMask.push_back(Elt: M < 0 ? -1 : M * MaskScale + i);
18610
18611 SDValue LegalShuffle =
18612 TLI.buildLegalVectorShuffle(VT, DL: SDLoc(N), N0: SV0, N1: SV1, Mask: NewMask, DAG);
18613 if (LegalShuffle)
18614 return LegalShuffle;
18615 }
18616
18617 return SDValue();
18618}
18619
18620SDValue DAGCombiner::visitBUILD_PAIR(SDNode *N) {
18621 EVT VT = N->getValueType(ResNo: 0);
18622 return CombineConsecutiveLoads(N, VT);
18623}
18624
18625SDValue DAGCombiner::visitFREEZE(SDNode *N) {
18626 SDValue N0 = N->getOperand(Num: 0);
18627
18628 if (DAG.isGuaranteedNotToBeUndefOrPoison(Op: N0, Kind: UndefPoisonKind::UndefOrPoison))
18629 return N0;
18630
18631 // If we have frozen and unfrozen users of N0, update so everything uses N.
18632 if (!N0.isUndef() && !N0.hasOneUse()) {
18633 SDValue FrozenN0(N, 0);
18634 // Unfreeze all (possibly nested) uses of N to avoid double deleting N from
18635 // the CSE map.
18636 while (!N->use_empty())
18637 DAG.ReplaceAllUsesOfValueWith(From: FrozenN0, To: N0);
18638 DAG.ReplaceAllUsesOfValueWith(From: N0, To: FrozenN0);
18639 // ReplaceAllUsesOfValueWith will have also updated the use in N, thus
18640 // creating a cycle in a DAG. Let's undo that by mutating the freeze.
18641 assert(N->getOperand(0) == FrozenN0 && "Expected cycle in DAG");
18642 DAG.UpdateNodeOperands(N, Op: N0);
18643 // Revisit the node.
18644 AddToWorklist(N);
18645 return FrozenN0;
18646 }
18647
18648 // We currently avoid folding freeze over SRL, due to the problems seen
18649 // with (freeze (assert ext)) blocking simplifications of SRL. See for
18650 // example https://reviews.llvm.org/D136529#4120959.
18651 if (N0.getOpcode() == ISD::SRL)
18652 return SDValue();
18653
18654 // Fold freeze(op(x, ...)) -> op(freeze(x), ...).
18655 // Try to push freeze through instructions that propagate but don't produce
18656 // poison as far as possible. If an operand of freeze follows three
18657 // conditions 1) one-use, 2) does not produce poison, and 3) has all but one
18658 // guaranteed-non-poison operands (or is a BUILD_VECTOR or similar) then push
18659 // the freeze through to the operands that are not guaranteed non-poison.
18660 // NOTE: we will strip poison-generating flags, so ignore them here.
18661 if (DAG.canCreateUndefOrPoison(Op: N0, Kind: UndefPoisonKind::UndefOrPoison,
18662 /*ConsiderFlags*/ false) ||
18663 N0->getNumValues() != 1 || !N0->hasOneUse())
18664 return SDValue();
18665
18666 // TOOD: we should always allow multiple operands, however this increases the
18667 // likelihood of infinite loops due to the ReplaceAllUsesOfValueWith call
18668 // below causing later nodes that share frozen operands to fold again and no
18669 // longer being able to confirm other operands are not poison due to recursion
18670 // depth limits on isGuaranteedNotToBeUndefOrPoison.
18671 bool AllowMultipleMaybePoisonOperands =
18672 N0.getOpcode() == ISD::SELECT_CC || N0.getOpcode() == ISD::SETCC ||
18673 N0.getOpcode() == ISD::BUILD_VECTOR ||
18674 N0.getOpcode() == ISD::INSERT_SUBVECTOR ||
18675 N0.getOpcode() == ISD::BUILD_PAIR ||
18676 N0.getOpcode() == ISD::VECTOR_SHUFFLE ||
18677 N0.getOpcode() == ISD::CONCAT_VECTORS || N0.getOpcode() == ISD::FMUL;
18678
18679 // Avoid turning a BUILD_VECTOR that can be recognized as "all zeros", "all
18680 // ones" or "constant" into something that depends on FrozenUndef. We can
18681 // instead pick undef values to keep those properties, while at the same time
18682 // folding away the freeze.
18683 // If we implement a more general solution for folding away freeze(undef) in
18684 // the future, then this special handling can be removed.
18685 if (N0.getOpcode() == ISD::BUILD_VECTOR) {
18686 SDLoc DL(N0);
18687 EVT VT = N0.getValueType();
18688 if (llvm::ISD::isBuildVectorAllOnes(N: N0.getNode()) && VT.isInteger())
18689 return DAG.getAllOnesConstant(DL, VT);
18690 if (llvm::ISD::isBuildVectorOfConstantSDNodes(N: N0.getNode())) {
18691 SmallVector<SDValue, 8> NewVecC;
18692 for (const SDValue &Op : N0->op_values())
18693 NewVecC.push_back(
18694 Elt: Op.isUndef() ? DAG.getConstant(Val: 0, DL, VT: Op.getValueType()) : Op);
18695 return DAG.getBuildVector(VT, DL, Ops: NewVecC);
18696 }
18697 }
18698
18699 SmallSet<SDValue, 8> MaybePoisonOperands;
18700 SmallVector<unsigned, 8> MaybePoisonOperandNumbers;
18701 for (auto [OpNo, Op] : enumerate(First: N0->ops())) {
18702 if (DAG.isGuaranteedNotToBeUndefOrPoison(Op,
18703 Kind: UndefPoisonKind::UndefOrPoison))
18704 continue;
18705 bool HadMaybePoisonOperands = !MaybePoisonOperands.empty();
18706 bool IsNewMaybePoisonOperand = MaybePoisonOperands.insert(V: Op).second;
18707 if (IsNewMaybePoisonOperand)
18708 MaybePoisonOperandNumbers.push_back(Elt: OpNo);
18709 if (!HadMaybePoisonOperands)
18710 continue;
18711 if (IsNewMaybePoisonOperand && !AllowMultipleMaybePoisonOperands) {
18712 // Multiple maybe-poison ops when not allowed - bail out.
18713 return SDValue();
18714 }
18715 }
18716 // NOTE: the whole op may be not guaranteed to not be undef or poison because
18717 // it could create undef or poison due to it's poison-generating flags.
18718 // So not finding any maybe-poison operands is fine.
18719
18720 for (unsigned OpNo : MaybePoisonOperandNumbers) {
18721 // N0 can mutate during iteration, so make sure to refetch the maybe poison
18722 // operands via the operand numbers. The typical scenario is that we have
18723 // something like this
18724 // t262: i32 = freeze t181
18725 // t150: i32 = ctlz_zero_poison t262
18726 // t184: i32 = ctlz_zero_poison t181
18727 // t268: i32 = select_cc t181, Constant:i32<0>, t184, t186, setne:ch
18728 // When freezing the t181 operand we get t262 back, and then the
18729 // ReplaceAllUsesOfValueWith call will not only replace t181 by t262, but
18730 // also recursively replace t184 by t150.
18731 SDValue MaybePoisonOperand = N->getOperand(Num: 0).getOperand(i: OpNo);
18732 // Don't replace every single UNDEF everywhere with frozen UNDEF, though.
18733 if (MaybePoisonOperand.isUndef())
18734 continue;
18735 // First, freeze each offending operand.
18736 SDValue FrozenMaybePoisonOperand = DAG.getFreeze(V: MaybePoisonOperand);
18737 // Then, change all other uses of unfrozen operand to use frozen operand.
18738 DAG.ReplaceAllUsesOfValueWith(From: MaybePoisonOperand, To: FrozenMaybePoisonOperand);
18739 if (FrozenMaybePoisonOperand.getOpcode() == ISD::FREEZE &&
18740 FrozenMaybePoisonOperand.getOperand(i: 0) == FrozenMaybePoisonOperand) {
18741 // But, that also updated the use in the freeze we just created, thus
18742 // creating a cycle in a DAG. Let's undo that by mutating the freeze.
18743 DAG.UpdateNodeOperands(N: FrozenMaybePoisonOperand.getNode(),
18744 Op: MaybePoisonOperand);
18745 }
18746
18747 // This node has been merged with another.
18748 if (N->getOpcode() == ISD::DELETED_NODE)
18749 return SDValue(N, 0);
18750 }
18751
18752 assert(N->getOpcode() != ISD::DELETED_NODE && "Node was deleted!");
18753
18754 // The whole node may have been updated, so the value we were holding
18755 // may no longer be valid. Re-fetch the operand we're `freeze`ing.
18756 N0 = N->getOperand(Num: 0);
18757
18758 // Finally, recreate the node, it's operands were updated to use
18759 // frozen operands, so we just need to use it's "original" operands.
18760 SmallVector<SDValue> Ops(N0->ops());
18761 // TODO: ISD::UNDEF and ISD::POISON should get separate handling, but best
18762 // leave for a future patch.
18763 for (SDValue &Op : Ops) {
18764 if (Op.isUndef())
18765 Op = DAG.getFreeze(V: Op);
18766 }
18767
18768 SDLoc DL(N0);
18769
18770 // Special case handling for ShuffleVectorSDNode nodes.
18771 if (auto *SVN = dyn_cast<ShuffleVectorSDNode>(Val&: N0))
18772 return DAG.getVectorShuffle(VT: N0.getValueType(), dl: DL, N1: Ops[0], N2: Ops[1],
18773 Mask: SVN->getMask());
18774
18775 // NOTE: this strips poison generating flags.
18776 // Folding freeze(op(x, ...)) -> op(freeze(x), ...) does not require nnan,
18777 // ninf, nsz, or fast.
18778 // However, contract, reassoc, afn, and arcp should be preserved,
18779 // as these fast-math flags do not introduce poison values.
18780 SDNodeFlags SrcFlags = N0->getFlags();
18781 SDNodeFlags SafeFlags;
18782 SafeFlags.setAllowContract(SrcFlags.hasAllowContract());
18783 SafeFlags.setAllowReassociation(SrcFlags.hasAllowReassociation());
18784 SafeFlags.setApproximateFuncs(SrcFlags.hasApproximateFuncs());
18785 SafeFlags.setAllowReciprocal(SrcFlags.hasAllowReciprocal());
18786 return DAG.getNode(Opcode: N0.getOpcode(), DL, VTList: N0->getVTList(), Ops, Flags: SafeFlags);
18787}
18788
18789// Returns true if floating point contraction is allowed on the FMUL-SDValue
18790// `N`
18791static bool isContractableFMUL(SDValue N) {
18792 assert(N.getOpcode() == ISD::FMUL);
18793
18794 return N->getFlags().hasAllowContract();
18795}
18796
18797/// Try to perform FMA combining on a given FADD node.
18798SDValue DAGCombiner::visitFADDForFMACombine(SDNode *N) {
18799 SDValue N0 = N->getOperand(Num: 0);
18800 SDValue N1 = N->getOperand(Num: 1);
18801 EVT VT = N->getValueType(ResNo: 0);
18802 SDLoc SL(N);
18803
18804 // Floating-point multiply-add with intermediate rounding.
18805 bool HasFMAD = (LegalOperations && TLI.isFMADLegal(DAG, N));
18806
18807 // Floating-point multiply-add without intermediate rounding.
18808 bool HasFMA =
18809 (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::FMA, VT)) &&
18810 TLI.isFMAFasterThanFMulAndFAdd(MF: DAG.getMachineFunction(), VT);
18811
18812 // No valid opcode, do not combine.
18813 if (!HasFMAD && !HasFMA)
18814 return SDValue();
18815
18816 // FMAD (with intermediate rounding) is always safe to form; FMA requires the
18817 // contract fast-math flag.
18818 bool AllowFusionGlobally = HasFMAD;
18819 // If the addition is not contractable, do not combine.
18820 if (!AllowFusionGlobally && !N->getFlags().hasAllowContract())
18821 return SDValue();
18822
18823 // Folding fadd (fmul x, y), (fmul x, y) -> fma x, y, (fmul x, y) is never
18824 // beneficial. It does not reduce latency. It increases register pressure. It
18825 // replaces an fadd with an fma which is a more complex instruction, so is
18826 // likely to have a larger encoding, use more functional units, etc.
18827 if (N0 == N1)
18828 return SDValue();
18829
18830 if (TLI.generateFMAsInMachineCombiner(VT, OptLevel))
18831 return SDValue();
18832
18833 // Always prefer FMAD to FMA for precision.
18834 unsigned PreferredFusedOpcode = HasFMAD ? ISD::FMAD : ISD::FMA;
18835 bool Aggressive = TLI.enableAggressiveFMAFusion(VT);
18836
18837 auto isFusedOp = [&](SDValue N) {
18838 unsigned Opcode = N.getOpcode();
18839 return Opcode == ISD::FMA || Opcode == ISD::FMAD;
18840 };
18841
18842 // Is the node an FMUL and contractable either due to global flags or
18843 // SDNodeFlags.
18844 auto isContractableFMUL = [AllowFusionGlobally](SDValue N) {
18845 if (N.getOpcode() != ISD::FMUL)
18846 return false;
18847 return AllowFusionGlobally || N->getFlags().hasAllowContract();
18848 };
18849 // If we have two choices trying to fold (fadd (fmul u, v), (fmul x, y)),
18850 // prefer to fold the multiply with fewer uses.
18851 if (Aggressive && isContractableFMUL(N0) && isContractableFMUL(N1)) {
18852 if (N0->use_size() > N1->use_size())
18853 std::swap(a&: N0, b&: N1);
18854 }
18855
18856 // fold (fadd (fmul x, y), z) -> (fma x, y, z)
18857 if (isContractableFMUL(N0) && (Aggressive || N0->hasOneUse())) {
18858 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: N0.getOperand(i: 0),
18859 N2: N0.getOperand(i: 1), N3: N1);
18860 }
18861
18862 // fold (fadd x, (fmul y, z)) -> (fma y, z, x)
18863 // Note: Commutes FADD operands.
18864 if (isContractableFMUL(N1) && (Aggressive || N1->hasOneUse())) {
18865 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: N1.getOperand(i: 0),
18866 N2: N1.getOperand(i: 1), N3: N0);
18867 }
18868
18869 // fadd (fma A, B, (fmul C, D)), E --> fma A, B, (fma C, D, E)
18870 // fadd E, (fma A, B, (fmul C, D)) --> fma A, B, (fma C, D, E)
18871 // This also works with nested fma instructions:
18872 // fadd (fma A, B, (fma (C, D, (fmul (E, F))))), G -->
18873 // fma A, B, (fma C, D, fma (E, F, G))
18874 // fadd (G, (fma A, B, (fma (C, D, (fmul (E, F)))))) -->
18875 // fma A, B, (fma C, D, fma (E, F, G)).
18876 // This requires reassociation because it changes the order of operations.
18877 bool CanReassociate = N->getFlags().hasAllowReassociation();
18878 if (CanReassociate) {
18879 SDValue FMA, E;
18880 if (isFusedOp(N0) && N0.hasOneUse()) {
18881 FMA = N0;
18882 E = N1;
18883 } else if (isFusedOp(N1) && N1.hasOneUse()) {
18884 FMA = N1;
18885 E = N0;
18886 }
18887
18888 SDValue TmpFMA = FMA;
18889 while (E && isFusedOp(TmpFMA) && TmpFMA.hasOneUse()) {
18890 SDValue FMul = TmpFMA->getOperand(Num: 2);
18891 if (FMul.getOpcode() == ISD::FMUL && FMul.hasOneUse()) {
18892 SDValue C = FMul.getOperand(i: 0);
18893 SDValue D = FMul.getOperand(i: 1);
18894 SDValue CDE = DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: C, N2: D, N3: E);
18895 DAG.ReplaceAllUsesOfValueWith(From: FMul, To: CDE);
18896 // Replacing the inner FMul could cause the outer FMA to be simplified
18897 // away.
18898 return FMA.getOpcode() == ISD::DELETED_NODE ? SDValue(N, 0) : FMA;
18899 }
18900
18901 TmpFMA = TmpFMA->getOperand(Num: 2);
18902 }
18903 }
18904
18905 // Look through FP_EXTEND nodes to do more combining.
18906
18907 // fold (fadd (fpext (fmul x, y)), z) -> (fma (fpext x), (fpext y), z)
18908 if (N0.getOpcode() == ISD::FP_EXTEND) {
18909 SDValue N00 = N0.getOperand(i: 0);
18910 if (isContractableFMUL(N00) &&
18911 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
18912 SrcVT: N00.getValueType())) {
18913 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
18914 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N00.getOperand(i: 0)),
18915 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N00.getOperand(i: 1)),
18916 N3: N1);
18917 }
18918 }
18919
18920 // fold (fadd x, (fpext (fmul y, z))) -> (fma (fpext y), (fpext z), x)
18921 // Note: Commutes FADD operands.
18922 if (N1.getOpcode() == ISD::FP_EXTEND) {
18923 SDValue N10 = N1.getOperand(i: 0);
18924 if (isContractableFMUL(N10) &&
18925 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
18926 SrcVT: N10.getValueType())) {
18927 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
18928 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N10.getOperand(i: 0)),
18929 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N10.getOperand(i: 1)),
18930 N3: N0);
18931 }
18932 }
18933
18934 // More folding opportunities when target permits.
18935 if (Aggressive) {
18936 // fold (fadd (fma x, y, (fpext (fmul u, v))), z)
18937 // -> (fma x, y, (fma (fpext u), (fpext v), z))
18938 auto FoldFAddFMAFPExtFMul = [&](SDValue X, SDValue Y, SDValue U, SDValue V,
18939 SDValue Z) {
18940 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: X, N2: Y,
18941 N3: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
18942 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: U),
18943 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: V),
18944 N3: Z));
18945 };
18946 if (isFusedOp(N0)) {
18947 SDValue N02 = N0.getOperand(i: 2);
18948 if (N02.getOpcode() == ISD::FP_EXTEND) {
18949 SDValue N020 = N02.getOperand(i: 0);
18950 if (isContractableFMUL(N020) &&
18951 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
18952 SrcVT: N020.getValueType())) {
18953 return FoldFAddFMAFPExtFMul(N0.getOperand(i: 0), N0.getOperand(i: 1),
18954 N020.getOperand(i: 0), N020.getOperand(i: 1),
18955 N1);
18956 }
18957 }
18958 }
18959
18960 // fold (fadd (fpext (fma x, y, (fmul u, v))), z)
18961 // -> (fma (fpext x), (fpext y), (fma (fpext u), (fpext v), z))
18962 // FIXME: This turns two single-precision and one double-precision
18963 // operation into two double-precision operations, which might not be
18964 // interesting for all targets, especially GPUs.
18965 auto FoldFAddFPExtFMAFMul = [&](SDValue X, SDValue Y, SDValue U, SDValue V,
18966 SDValue Z) {
18967 return DAG.getNode(
18968 Opcode: PreferredFusedOpcode, DL: SL, VT, N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: X),
18969 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: Y),
18970 N3: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
18971 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: U),
18972 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: V), N3: Z));
18973 };
18974 if (N0.getOpcode() == ISD::FP_EXTEND) {
18975 SDValue N00 = N0.getOperand(i: 0);
18976 if (isFusedOp(N00)) {
18977 SDValue N002 = N00.getOperand(i: 2);
18978 if (isContractableFMUL(N002) &&
18979 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
18980 SrcVT: N00.getValueType())) {
18981 return FoldFAddFPExtFMAFMul(N00.getOperand(i: 0), N00.getOperand(i: 1),
18982 N002.getOperand(i: 0), N002.getOperand(i: 1),
18983 N1);
18984 }
18985 }
18986 }
18987
18988 // fold (fadd x, (fma y, z, (fpext (fmul u, v)))
18989 // -> (fma y, z, (fma (fpext u), (fpext v), x))
18990 if (isFusedOp(N1)) {
18991 SDValue N12 = N1.getOperand(i: 2);
18992 if (N12.getOpcode() == ISD::FP_EXTEND) {
18993 SDValue N120 = N12.getOperand(i: 0);
18994 if (isContractableFMUL(N120) &&
18995 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
18996 SrcVT: N120.getValueType())) {
18997 return FoldFAddFMAFPExtFMul(N1.getOperand(i: 0), N1.getOperand(i: 1),
18998 N120.getOperand(i: 0), N120.getOperand(i: 1),
18999 N0);
19000 }
19001 }
19002 }
19003
19004 // fold (fadd x, (fpext (fma y, z, (fmul u, v)))
19005 // -> (fma (fpext y), (fpext z), (fma (fpext u), (fpext v), x))
19006 // FIXME: This turns two single-precision and one double-precision
19007 // operation into two double-precision operations, which might not be
19008 // interesting for all targets, especially GPUs.
19009 if (N1.getOpcode() == ISD::FP_EXTEND) {
19010 SDValue N10 = N1.getOperand(i: 0);
19011 if (isFusedOp(N10)) {
19012 SDValue N102 = N10.getOperand(i: 2);
19013 if (isContractableFMUL(N102) &&
19014 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19015 SrcVT: N10.getValueType())) {
19016 return FoldFAddFPExtFMAFMul(N10.getOperand(i: 0), N10.getOperand(i: 1),
19017 N102.getOperand(i: 0), N102.getOperand(i: 1),
19018 N0);
19019 }
19020 }
19021 }
19022 }
19023
19024 return SDValue();
19025}
19026
19027/// Try to perform FMA combining on a given FSUB node.
19028SDValue DAGCombiner::visitFSUBForFMACombine(SDNode *N) {
19029 SDValue N0 = N->getOperand(Num: 0);
19030 SDValue N1 = N->getOperand(Num: 1);
19031 EVT VT = N->getValueType(ResNo: 0);
19032 SDLoc SL(N);
19033
19034 // Floating-point multiply-add with intermediate rounding.
19035 bool HasFMAD = (LegalOperations && TLI.isFMADLegal(DAG, N));
19036
19037 // Floating-point multiply-add without intermediate rounding.
19038 bool HasFMA =
19039 (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::FMA, VT)) &&
19040 TLI.isFMAFasterThanFMulAndFAdd(MF: DAG.getMachineFunction(), VT);
19041
19042 // No valid opcode, do not combine.
19043 if (!HasFMAD && !HasFMA)
19044 return SDValue();
19045
19046 const SDNodeFlags Flags = N->getFlags();
19047 // FMAD (with intermediate rounding) is always safe to form; FMA requires the
19048 // contract fast-math flag.
19049 bool AllowFusionGlobally = HasFMAD;
19050
19051 // If the subtraction is not contractable, do not combine.
19052 if (!AllowFusionGlobally && !N->getFlags().hasAllowContract())
19053 return SDValue();
19054
19055 if (TLI.generateFMAsInMachineCombiner(VT, OptLevel))
19056 return SDValue();
19057
19058 // Always prefer FMAD to FMA for precision.
19059 unsigned PreferredFusedOpcode = HasFMAD ? ISD::FMAD : ISD::FMA;
19060 bool Aggressive = TLI.enableAggressiveFMAFusion(VT);
19061 bool NoSignedZero = Flags.hasNoSignedZeros();
19062
19063 // Is the node an FMUL and contractable either due to global flags or
19064 // SDNodeFlags.
19065 auto isContractableFMUL = [AllowFusionGlobally](SDValue N) {
19066 if (N.getOpcode() != ISD::FMUL)
19067 return false;
19068 return AllowFusionGlobally || N->getFlags().hasAllowContract();
19069 };
19070
19071 // fold (fsub (fmul x, y), z) -> (fma x, y, (fneg z))
19072 auto tryToFoldXYSubZ = [&](SDValue XY, SDValue Z) {
19073 if (isContractableFMUL(XY) && (Aggressive || XY->hasOneUse())) {
19074 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: XY.getOperand(i: 0),
19075 N2: XY.getOperand(i: 1), N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: Z));
19076 }
19077 return SDValue();
19078 };
19079
19080 // fold (fsub x, (fmul y, z)) -> (fma (fneg y), z, x)
19081 // Note: Commutes FSUB operands.
19082 auto tryToFoldXSubYZ = [&](SDValue X, SDValue YZ) {
19083 if (isContractableFMUL(YZ) && (Aggressive || YZ->hasOneUse())) {
19084 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19085 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: YZ.getOperand(i: 0)),
19086 N2: YZ.getOperand(i: 1), N3: X);
19087 }
19088 return SDValue();
19089 };
19090
19091 // If we have two choices trying to fold (fsub (fmul u, v), (fmul x, y)),
19092 // prefer to fold the multiply with fewer uses.
19093 if (isContractableFMUL(N0) && isContractableFMUL(N1) &&
19094 (N0->use_size() > N1->use_size())) {
19095 // fold (fsub (fmul a, b), (fmul c, d)) -> (fma (fneg c), d, (fmul a, b))
19096 if (SDValue V = tryToFoldXSubYZ(N0, N1))
19097 return V;
19098 // fold (fsub (fmul a, b), (fmul c, d)) -> (fma a, b, (fneg (fmul c, d)))
19099 if (SDValue V = tryToFoldXYSubZ(N0, N1))
19100 return V;
19101 } else {
19102 // fold (fsub (fmul x, y), z) -> (fma x, y, (fneg z))
19103 if (SDValue V = tryToFoldXYSubZ(N0, N1))
19104 return V;
19105 // fold (fsub x, (fmul y, z)) -> (fma (fneg y), z, x)
19106 if (SDValue V = tryToFoldXSubYZ(N0, N1))
19107 return V;
19108 }
19109
19110 // fold (fsub (fneg (fmul, x, y)), z) -> (fma (fneg x), y, (fneg z))
19111 if (N0.getOpcode() == ISD::FNEG && isContractableFMUL(N0.getOperand(i: 0)) &&
19112 (Aggressive || (N0->hasOneUse() && N0.getOperand(i: 0).hasOneUse()))) {
19113 SDValue N00 = N0.getOperand(i: 0).getOperand(i: 0);
19114 SDValue N01 = N0.getOperand(i: 0).getOperand(i: 1);
19115 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19116 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N00), N2: N01,
19117 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N1));
19118 }
19119
19120 // Look through FP_EXTEND nodes to do more combining.
19121
19122 // fold (fsub (fpext (fmul x, y)), z)
19123 // -> (fma (fpext x), (fpext y), (fneg z))
19124 if (N0.getOpcode() == ISD::FP_EXTEND) {
19125 SDValue N00 = N0.getOperand(i: 0);
19126 if (isContractableFMUL(N00) &&
19127 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19128 SrcVT: N00.getValueType())) {
19129 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19130 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N00.getOperand(i: 0)),
19131 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N00.getOperand(i: 1)),
19132 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N1));
19133 }
19134 }
19135
19136 // fold (fsub x, (fpext (fmul y, z)))
19137 // -> (fma (fneg (fpext y)), (fpext z), x)
19138 // Note: Commutes FSUB operands.
19139 if (N1.getOpcode() == ISD::FP_EXTEND) {
19140 SDValue N10 = N1.getOperand(i: 0);
19141 if (isContractableFMUL(N10) &&
19142 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19143 SrcVT: N10.getValueType())) {
19144 return DAG.getNode(
19145 Opcode: PreferredFusedOpcode, DL: SL, VT,
19146 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT,
19147 Operand: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N10.getOperand(i: 0))),
19148 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N10.getOperand(i: 1)), N3: N0);
19149 }
19150 }
19151
19152 // fold (fsub (fpext (fneg (fmul, x, y))), z)
19153 // -> (fneg (fma (fpext x), (fpext y), z))
19154 // Note: This could be removed with appropriate canonicalization of the
19155 // input expression into (fneg (fadd (fpext (fmul, x, y)), z)). However, the
19156 // command line flag -fp-contract=fast and fast-math flag contract prevent
19157 // from implementing the canonicalization in visitFSUB.
19158 if (N0.getOpcode() == ISD::FP_EXTEND) {
19159 SDValue N00 = N0.getOperand(i: 0);
19160 if (N00.getOpcode() == ISD::FNEG) {
19161 SDValue N000 = N00.getOperand(i: 0);
19162 if (isContractableFMUL(N000) &&
19163 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19164 SrcVT: N00.getValueType())) {
19165 return DAG.getNode(
19166 Opcode: ISD::FNEG, DL: SL, VT,
19167 Operand: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19168 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N000.getOperand(i: 0)),
19169 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N000.getOperand(i: 1)),
19170 N3: N1));
19171 }
19172 }
19173 }
19174
19175 // fold (fsub (fneg (fpext (fmul, x, y))), z)
19176 // -> (fneg (fma (fpext x)), (fpext y), z)
19177 // Note: This could be removed with appropriate canonicalization of the
19178 // input expression into (fneg (fadd (fpext (fmul, x, y)), z). However, the
19179 // command line flag -fp-contract=fast and fast-math flag contract prevent
19180 // from implementing the canonicalization in visitFSUB.
19181 if (N0.getOpcode() == ISD::FNEG) {
19182 SDValue N00 = N0.getOperand(i: 0);
19183 if (N00.getOpcode() == ISD::FP_EXTEND) {
19184 SDValue N000 = N00.getOperand(i: 0);
19185 if (isContractableFMUL(N000) &&
19186 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19187 SrcVT: N000.getValueType())) {
19188 return DAG.getNode(
19189 Opcode: ISD::FNEG, DL: SL, VT,
19190 Operand: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19191 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N000.getOperand(i: 0)),
19192 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N000.getOperand(i: 1)),
19193 N3: N1));
19194 }
19195 }
19196 }
19197
19198 auto isContractableAndReassociableFMUL = [&isContractableFMUL](SDValue N) {
19199 return isContractableFMUL(N) && N->getFlags().hasAllowReassociation();
19200 };
19201
19202 auto isFusedOp = [&](SDValue N) {
19203 unsigned Opcode = N.getOpcode();
19204 return Opcode == ISD::FMA || Opcode == ISD::FMAD;
19205 };
19206
19207 // More folding opportunities when target permits.
19208 if (Aggressive && N->getFlags().hasAllowReassociation()) {
19209 bool CanFuse = N->getFlags().hasAllowContract();
19210 // fold (fsub (fma x, y, (fmul u, v)), z)
19211 // -> (fma x, y (fma u, v, (fneg z)))
19212 if (CanFuse && isFusedOp(N0) &&
19213 isContractableAndReassociableFMUL(N0.getOperand(i: 2)) &&
19214 N0->hasOneUse() && N0.getOperand(i: 2)->hasOneUse()) {
19215 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: N0.getOperand(i: 0),
19216 N2: N0.getOperand(i: 1),
19217 N3: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19218 N1: N0.getOperand(i: 2).getOperand(i: 0),
19219 N2: N0.getOperand(i: 2).getOperand(i: 1),
19220 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N1)));
19221 }
19222
19223 // fold (fsub x, (fma y, z, (fmul u, v)))
19224 // -> (fma (fneg y), z, (fma (fneg u), v, x))
19225 if (CanFuse && isFusedOp(N1) &&
19226 isContractableAndReassociableFMUL(N1.getOperand(i: 2)) &&
19227 N1->hasOneUse() && NoSignedZero) {
19228 SDValue N20 = N1.getOperand(i: 2).getOperand(i: 0);
19229 SDValue N21 = N1.getOperand(i: 2).getOperand(i: 1);
19230 return DAG.getNode(
19231 Opcode: PreferredFusedOpcode, DL: SL, VT,
19232 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N1.getOperand(i: 0)), N2: N1.getOperand(i: 1),
19233 N3: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19234 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N20), N2: N21, N3: N0));
19235 }
19236
19237 // fold (fsub (fma x, y, (fpext (fmul u, v))), z)
19238 // -> (fma x, y (fma (fpext u), (fpext v), (fneg z)))
19239 if (isFusedOp(N0) && N0->hasOneUse()) {
19240 SDValue N02 = N0.getOperand(i: 2);
19241 if (N02.getOpcode() == ISD::FP_EXTEND) {
19242 SDValue N020 = N02.getOperand(i: 0);
19243 if (isContractableAndReassociableFMUL(N020) &&
19244 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19245 SrcVT: N020.getValueType())) {
19246 return DAG.getNode(
19247 Opcode: PreferredFusedOpcode, DL: SL, VT, N1: N0.getOperand(i: 0), N2: N0.getOperand(i: 1),
19248 N3: DAG.getNode(
19249 Opcode: PreferredFusedOpcode, DL: SL, VT,
19250 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N020.getOperand(i: 0)),
19251 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N020.getOperand(i: 1)),
19252 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N1)));
19253 }
19254 }
19255 }
19256
19257 // fold (fsub (fpext (fma x, y, (fmul u, v))), z)
19258 // -> (fma (fpext x), (fpext y),
19259 // (fma (fpext u), (fpext v), (fneg z)))
19260 // FIXME: This turns two single-precision and one double-precision
19261 // operation into two double-precision operations, which might not be
19262 // interesting for all targets, especially GPUs.
19263 if (N0.getOpcode() == ISD::FP_EXTEND) {
19264 SDValue N00 = N0.getOperand(i: 0);
19265 if (isFusedOp(N00)) {
19266 SDValue N002 = N00.getOperand(i: 2);
19267 if (isContractableAndReassociableFMUL(N002) &&
19268 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19269 SrcVT: N00.getValueType())) {
19270 return DAG.getNode(
19271 Opcode: PreferredFusedOpcode, DL: SL, VT,
19272 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N00.getOperand(i: 0)),
19273 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N00.getOperand(i: 1)),
19274 N3: DAG.getNode(
19275 Opcode: PreferredFusedOpcode, DL: SL, VT,
19276 N1: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N002.getOperand(i: 0)),
19277 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N002.getOperand(i: 1)),
19278 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N1)));
19279 }
19280 }
19281 }
19282
19283 // fold (fsub x, (fma y, z, (fpext (fmul u, v))))
19284 // -> (fma (fneg y), z, (fma (fneg (fpext u)), (fpext v), x))
19285 if (isFusedOp(N1) && N1.getOperand(i: 2).getOpcode() == ISD::FP_EXTEND &&
19286 N1->hasOneUse()) {
19287 SDValue N120 = N1.getOperand(i: 2).getOperand(i: 0);
19288 if (isContractableAndReassociableFMUL(N120) &&
19289 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19290 SrcVT: N120.getValueType())) {
19291 SDValue N1200 = N120.getOperand(i: 0);
19292 SDValue N1201 = N120.getOperand(i: 1);
19293 return DAG.getNode(
19294 Opcode: PreferredFusedOpcode, DL: SL, VT,
19295 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: N1.getOperand(i: 0)), N2: N1.getOperand(i: 1),
19296 N3: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19297 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT,
19298 Operand: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N1200)),
19299 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N1201), N3: N0));
19300 }
19301 }
19302
19303 // fold (fsub x, (fpext (fma y, z, (fmul u, v))))
19304 // -> (fma (fneg (fpext y)), (fpext z),
19305 // (fma (fneg (fpext u)), (fpext v), x))
19306 // FIXME: This turns two single-precision and one double-precision
19307 // operation into two double-precision operations, which might not be
19308 // interesting for all targets, especially GPUs.
19309 if (N1.getOpcode() == ISD::FP_EXTEND && isFusedOp(N1.getOperand(i: 0))) {
19310 SDValue CvtSrc = N1.getOperand(i: 0);
19311 SDValue N100 = CvtSrc.getOperand(i: 0);
19312 SDValue N101 = CvtSrc.getOperand(i: 1);
19313 SDValue N102 = CvtSrc.getOperand(i: 2);
19314 if (isContractableAndReassociableFMUL(N102) &&
19315 TLI.isFPExtFoldable(DAG, Opcode: PreferredFusedOpcode, DestVT: VT,
19316 SrcVT: CvtSrc.getValueType())) {
19317 SDValue N1020 = N102.getOperand(i: 0);
19318 SDValue N1021 = N102.getOperand(i: 1);
19319 return DAG.getNode(
19320 Opcode: PreferredFusedOpcode, DL: SL, VT,
19321 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT,
19322 Operand: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N100)),
19323 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N101),
19324 N3: DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19325 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT,
19326 Operand: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N1020)),
19327 N2: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SL, VT, Operand: N1021), N3: N0));
19328 }
19329 }
19330 }
19331
19332 return SDValue();
19333}
19334
19335/// Try to perform FMA combining on a given FMUL node based on the distributive
19336/// law x * (y + 1) = x * y + x and variants thereof (commuted versions,
19337/// subtraction instead of addition).
19338SDValue DAGCombiner::visitFMULForFMADistributiveCombine(SDNode *N) {
19339 SDValue N0 = N->getOperand(Num: 0);
19340 SDValue N1 = N->getOperand(Num: 1);
19341 EVT VT = N->getValueType(ResNo: 0);
19342 SDLoc SL(N);
19343
19344 assert(N->getOpcode() == ISD::FMUL && "Expected FMUL Operation");
19345
19346 // The transforms below are incorrect when x == 0 and y == inf, because the
19347 // intermediate multiplication produces a nan.
19348 SDValue FAdd = N0.getOpcode() == ISD::FADD ? N0 : N1;
19349 if (!FAdd->getFlags().hasNoInfs())
19350 return SDValue();
19351
19352 // Floating-point multiply-add without intermediate rounding.
19353 bool HasFMA =
19354 isContractableFMUL(N: SDValue(N, 0)) &&
19355 (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::FMA, VT)) &&
19356 TLI.isFMAFasterThanFMulAndFAdd(MF: DAG.getMachineFunction(), VT);
19357
19358 // Floating-point multiply-add with intermediate rounding. This can result
19359 // in a less precise result due to the changed rounding order.
19360 bool HasFMAD = LegalOperations && TLI.isFMADLegal(DAG, N);
19361
19362 // No valid opcode, do not combine.
19363 if (!HasFMAD && !HasFMA)
19364 return SDValue();
19365
19366 // Always prefer FMAD to FMA for precision.
19367 unsigned PreferredFusedOpcode = HasFMAD ? ISD::FMAD : ISD::FMA;
19368 bool Aggressive = TLI.enableAggressiveFMAFusion(VT);
19369
19370 // fold (fmul (fadd x0, +1.0), y) -> (fma x0, y, y)
19371 // fold (fmul (fadd x0, -1.0), y) -> (fma x0, y, (fneg y))
19372 auto FuseFADD = [&](SDValue X, SDValue Y) {
19373 if (X.getOpcode() == ISD::FADD && (Aggressive || X->hasOneUse())) {
19374 if (auto *C = isConstOrConstSplatFP(N: X.getOperand(i: 1), AllowUndefs: true)) {
19375 if (C->isOne())
19376 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: X.getOperand(i: 0), N2: Y,
19377 N3: Y);
19378 if (C->isMinusOne())
19379 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: X.getOperand(i: 0), N2: Y,
19380 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: Y));
19381 }
19382 }
19383 return SDValue();
19384 };
19385
19386 if (SDValue FMA = FuseFADD(N0, N1))
19387 return FMA;
19388 if (SDValue FMA = FuseFADD(N1, N0))
19389 return FMA;
19390
19391 // fold (fmul (fsub +1.0, x1), y) -> (fma (fneg x1), y, y)
19392 // fold (fmul (fsub -1.0, x1), y) -> (fma (fneg x1), y, (fneg y))
19393 // fold (fmul (fsub x0, +1.0), y) -> (fma x0, y, (fneg y))
19394 // fold (fmul (fsub x0, -1.0), y) -> (fma x0, y, y)
19395 auto FuseFSUB = [&](SDValue X, SDValue Y) {
19396 if (X.getOpcode() == ISD::FSUB && (Aggressive || X->hasOneUse())) {
19397 if (auto *C0 = isConstOrConstSplatFP(N: X.getOperand(i: 0), AllowUndefs: true)) {
19398 if (C0->isOne())
19399 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19400 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: X.getOperand(i: 1)), N2: Y,
19401 N3: Y);
19402 if (C0->isMinusOne())
19403 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT,
19404 N1: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: X.getOperand(i: 1)), N2: Y,
19405 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: Y));
19406 }
19407 if (auto *C1 = isConstOrConstSplatFP(N: X.getOperand(i: 1), AllowUndefs: true)) {
19408 if (C1->isOne())
19409 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: X.getOperand(i: 0), N2: Y,
19410 N3: DAG.getNode(Opcode: ISD::FNEG, DL: SL, VT, Operand: Y));
19411 if (C1->isMinusOne())
19412 return DAG.getNode(Opcode: PreferredFusedOpcode, DL: SL, VT, N1: X.getOperand(i: 0), N2: Y,
19413 N3: Y);
19414 }
19415 }
19416 return SDValue();
19417 };
19418
19419 if (SDValue FMA = FuseFSUB(N0, N1))
19420 return FMA;
19421 if (SDValue FMA = FuseFSUB(N1, N0))
19422 return FMA;
19423
19424 return SDValue();
19425}
19426
19427SDValue DAGCombiner::visitFADD(SDNode *N) {
19428 SDValue N0 = N->getOperand(Num: 0);
19429 SDValue N1 = N->getOperand(Num: 1);
19430 bool N0CFP = DAG.isConstantFPBuildVectorOrConstantFP(N: N0);
19431 bool N1CFP = DAG.isConstantFPBuildVectorOrConstantFP(N: N1);
19432 EVT VT = N->getValueType(ResNo: 0);
19433 SDLoc DL(N);
19434 SDNodeFlags Flags = N->getFlags();
19435 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
19436
19437 if (SDValue R = DAG.simplifyFPBinop(Opcode: N->getOpcode(), X: N0, Y: N1, Flags))
19438 return R;
19439
19440 // fold (fadd c1, c2) -> c1 + c2
19441 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FADD, DL, VT, Ops: {N0, N1}))
19442 return C;
19443
19444 // canonicalize constant to RHS
19445 if (N0CFP && !N1CFP)
19446 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1, N2: N0);
19447
19448 // fold vector ops
19449 if (VT.isVector())
19450 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
19451 return FoldedVOp;
19452
19453 // N0 + -0.0 --> N0 (also allowed with +0.0 and fast-math)
19454 ConstantFPSDNode *N1C = isConstOrConstSplatFP(N: N1, AllowUndefs: true);
19455 if (N1C && N1C->isZero())
19456 if (N1C->isNegative() || DAG.canIgnoreSignBitOfZero(Op: SDValue(N, 0)))
19457 return N0;
19458
19459 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
19460 return NewSel;
19461
19462 // fold (fadd A, (fneg B)) -> (fsub A, B)
19463 if (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::FSUB, VT))
19464 if (SDValue NegN1 = TLI.getCheaperNegatedExpression(
19465 Op: N1, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize))
19466 return DAG.getNode(Opcode: ISD::FSUB, DL, VT, N1: N0, N2: NegN1);
19467
19468 // fold (fadd (fneg A), B) -> (fsub B, A)
19469 if (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::FSUB, VT))
19470 if (SDValue NegN0 = TLI.getCheaperNegatedExpression(
19471 Op: N0, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize))
19472 return DAG.getNode(Opcode: ISD::FSUB, DL, VT, N1, N2: NegN0);
19473
19474 auto isFMulNegTwo = [](SDValue FMul) {
19475 if (!FMul.hasOneUse() || FMul.getOpcode() != ISD::FMUL)
19476 return false;
19477 auto *C = isConstOrConstSplatFP(N: FMul.getOperand(i: 1), AllowUndefs: true);
19478 return C && C->isExactlyValue(V: -2.0);
19479 };
19480
19481 // fadd (fmul B, -2.0), A --> fsub A, (fadd B, B)
19482 if (isFMulNegTwo(N0)) {
19483 SDValue B = N0.getOperand(i: 0);
19484 SDValue Add = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: B, N2: B);
19485 return DAG.getNode(Opcode: ISD::FSUB, DL, VT, N1, N2: Add);
19486 }
19487 // fadd A, (fmul B, -2.0) --> fsub A, (fadd B, B)
19488 if (isFMulNegTwo(N1)) {
19489 SDValue B = N1.getOperand(i: 0);
19490 SDValue Add = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: B, N2: B);
19491 return DAG.getNode(Opcode: ISD::FSUB, DL, VT, N1: N0, N2: Add);
19492 }
19493
19494 // No FP constant should be created after legalization as Instruction
19495 // Selection pass has a hard time dealing with FP constants.
19496 bool AllowNewConst = (Level < AfterLegalizeDAG);
19497
19498 // If nnan is enabled, fold lots of things.
19499 if (Flags.hasNoNaNs() && AllowNewConst) {
19500 // If allowed, fold (fadd (fneg x), x) -> 0.0
19501 if (N0.getOpcode() == ISD::FNEG && N0.getOperand(i: 0) == N1)
19502 return DAG.getConstantFP(Val: 0.0, DL, VT);
19503
19504 // If allowed, fold (fadd x, (fneg x)) -> 0.0
19505 if (N1.getOpcode() == ISD::FNEG && N1.getOperand(i: 0) == N0)
19506 return DAG.getConstantFP(Val: 0.0, DL, VT);
19507 }
19508
19509 // If reassoc and nsz, fold lots of things.
19510 // TODO: break out portions of the transformations below for which Unsafe is
19511 // considered and which do not require both nsz and reassoc
19512 if (Flags.hasAllowReassociation() && Flags.hasNoSignedZeros() &&
19513 AllowNewConst) {
19514 // fadd (fadd x, c1), c2 -> fadd x, c1 + c2
19515 if (N1CFP && N0.getOpcode() == ISD::FADD &&
19516 DAG.isConstantFPBuildVectorOrConstantFP(N: N0.getOperand(i: 1))) {
19517 SDValue NewC = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0.getOperand(i: 1), N2: N1);
19518 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0.getOperand(i: 0), N2: NewC);
19519 }
19520
19521 // We can fold chains of FADD's of the same value into multiplications.
19522 // This transform is not safe in general because we are reducing the number
19523 // of rounding steps.
19524 if ((!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::FMUL, VT)) &&
19525 !N0CFP && !N1CFP) {
19526 if (N0.getOpcode() == ISD::FMUL) {
19527 bool CFP00 = DAG.isConstantFPBuildVectorOrConstantFP(N: N0.getOperand(i: 0));
19528 bool CFP01 = DAG.isConstantFPBuildVectorOrConstantFP(N: N0.getOperand(i: 1));
19529
19530 // (fadd (fmul x, c), x) -> (fmul x, c+1)
19531 if (CFP01 && !CFP00 && N0.getOperand(i: 0) == N1) {
19532 SDValue NewCFP = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0.getOperand(i: 1),
19533 N2: DAG.getConstantFP(Val: 1.0, DL, VT));
19534 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1, N2: NewCFP);
19535 }
19536
19537 // (fadd (fmul x, c), (fadd x, x)) -> (fmul x, c+2)
19538 if (CFP01 && !CFP00 && N1.getOpcode() == ISD::FADD &&
19539 N1.getOperand(i: 0) == N1.getOperand(i: 1) &&
19540 N0.getOperand(i: 0) == N1.getOperand(i: 0)) {
19541 SDValue NewCFP = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0.getOperand(i: 1),
19542 N2: DAG.getConstantFP(Val: 2.0, DL, VT));
19543 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0.getOperand(i: 0), N2: NewCFP);
19544 }
19545 }
19546
19547 if (N1.getOpcode() == ISD::FMUL) {
19548 bool CFP10 = DAG.isConstantFPBuildVectorOrConstantFP(N: N1.getOperand(i: 0));
19549 bool CFP11 = DAG.isConstantFPBuildVectorOrConstantFP(N: N1.getOperand(i: 1));
19550
19551 // (fadd x, (fmul x, c)) -> (fmul x, c+1)
19552 if (CFP11 && !CFP10 && N1.getOperand(i: 0) == N0) {
19553 SDValue NewCFP = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N1.getOperand(i: 1),
19554 N2: DAG.getConstantFP(Val: 1.0, DL, VT));
19555 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0, N2: NewCFP);
19556 }
19557
19558 // (fadd (fadd x, x), (fmul x, c)) -> (fmul x, c+2)
19559 if (CFP11 && !CFP10 && N0.getOpcode() == ISD::FADD &&
19560 N0.getOperand(i: 0) == N0.getOperand(i: 1) &&
19561 N1.getOperand(i: 0) == N0.getOperand(i: 0)) {
19562 SDValue NewCFP = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N1.getOperand(i: 1),
19563 N2: DAG.getConstantFP(Val: 2.0, DL, VT));
19564 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N1.getOperand(i: 0), N2: NewCFP);
19565 }
19566 }
19567
19568 if (N0.getOpcode() == ISD::FADD) {
19569 bool CFP00 = DAG.isConstantFPBuildVectorOrConstantFP(N: N0.getOperand(i: 0));
19570 // (fadd (fadd x, x), x) -> (fmul x, 3.0)
19571 if (!CFP00 && N0.getOperand(i: 0) == N0.getOperand(i: 1) &&
19572 (N0.getOperand(i: 0) == N1)) {
19573 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1,
19574 N2: DAG.getConstantFP(Val: 3.0, DL, VT));
19575 }
19576 }
19577
19578 if (N1.getOpcode() == ISD::FADD) {
19579 bool CFP10 = DAG.isConstantFPBuildVectorOrConstantFP(N: N1.getOperand(i: 0));
19580 // (fadd x, (fadd x, x)) -> (fmul x, 3.0)
19581 if (!CFP10 && N1.getOperand(i: 0) == N1.getOperand(i: 1) &&
19582 N1.getOperand(i: 0) == N0) {
19583 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0,
19584 N2: DAG.getConstantFP(Val: 3.0, DL, VT));
19585 }
19586 }
19587
19588 // (fadd (fadd x, x), (fadd x, x)) -> (fmul x, 4.0)
19589 if (N0.getOpcode() == ISD::FADD && N1.getOpcode() == ISD::FADD &&
19590 N0.getOperand(i: 0) == N0.getOperand(i: 1) &&
19591 N1.getOperand(i: 0) == N1.getOperand(i: 1) &&
19592 N0.getOperand(i: 0) == N1.getOperand(i: 0)) {
19593 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0.getOperand(i: 0),
19594 N2: DAG.getConstantFP(Val: 4.0, DL, VT));
19595 }
19596 }
19597 } // reassoc && nsz && AllowNewConst
19598
19599 if (Flags.hasAllowReassociation() && Flags.hasNoSignedZeros()) {
19600 // Fold fadd(vecreduce(x), vecreduce(y)) -> vecreduce(fadd(x, y))
19601 if (SDValue SD = reassociateReduction(RedOpc: ISD::VECREDUCE_FADD, Opc: ISD::FADD, DL,
19602 VT, N0, N1, Flags))
19603 return SD;
19604 }
19605
19606 // FADD -> FMA combines:
19607 if (SDValue Fused = visitFADDForFMACombine(N)) {
19608 if (Fused.getOpcode() != ISD::DELETED_NODE)
19609 AddToWorklist(N: Fused.getNode());
19610 return Fused;
19611 }
19612 return SDValue();
19613}
19614
19615SDValue DAGCombiner::visitSTRICT_FADD(SDNode *N) {
19616 SDValue Chain = N->getOperand(Num: 0);
19617 SDValue N0 = N->getOperand(Num: 1);
19618 SDValue N1 = N->getOperand(Num: 2);
19619 EVT VT = N->getValueType(ResNo: 0);
19620 EVT ChainVT = N->getValueType(ResNo: 1);
19621 SDLoc DL(N);
19622 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
19623
19624 // fold (strict_fadd A, (fneg B)) -> (strict_fsub A, B)
19625 if (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::STRICT_FSUB, VT))
19626 if (SDValue NegN1 = TLI.getCheaperNegatedExpression(
19627 Op: N1, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize)) {
19628 return DAG.getNode(Opcode: ISD::STRICT_FSUB, DL, VTList: DAG.getVTList(VT1: VT, VT2: ChainVT),
19629 Ops: {Chain, N0, NegN1});
19630 }
19631
19632 // fold (strict_fadd (fneg A), B) -> (strict_fsub B, A)
19633 if (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::STRICT_FSUB, VT))
19634 if (SDValue NegN0 = TLI.getCheaperNegatedExpression(
19635 Op: N0, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize)) {
19636 return DAG.getNode(Opcode: ISD::STRICT_FSUB, DL, VTList: DAG.getVTList(VT1: VT, VT2: ChainVT),
19637 Ops: {Chain, N1, NegN0});
19638 }
19639 return SDValue();
19640}
19641
19642SDValue DAGCombiner::visitFSUB(SDNode *N) {
19643 SDValue N0 = N->getOperand(Num: 0);
19644 SDValue N1 = N->getOperand(Num: 1);
19645 ConstantFPSDNode *N0CFP = isConstOrConstSplatFP(N: N0, AllowUndefs: true);
19646 ConstantFPSDNode *N1CFP = isConstOrConstSplatFP(N: N1, AllowUndefs: true);
19647 EVT VT = N->getValueType(ResNo: 0);
19648 SDLoc DL(N);
19649 const SDNodeFlags Flags = N->getFlags();
19650 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
19651
19652 if (SDValue R = DAG.simplifyFPBinop(Opcode: N->getOpcode(), X: N0, Y: N1, Flags))
19653 return R;
19654
19655 // fold (fsub c1, c2) -> c1-c2
19656 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FSUB, DL, VT, Ops: {N0, N1}))
19657 return C;
19658
19659 // fold vector ops
19660 if (VT.isVector())
19661 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
19662 return FoldedVOp;
19663
19664 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
19665 return NewSel;
19666
19667 // (fsub A, 0) -> A
19668 if (N1CFP && N1CFP->isZero()) {
19669 if (!N1CFP->isNegative() || DAG.canIgnoreSignBitOfZero(Op: SDValue(N, 0))) {
19670 return N0;
19671 }
19672 }
19673
19674 if (N0 == N1) {
19675 // (fsub x, x) -> 0.0
19676 if (Flags.hasNoNaNs())
19677 return DAG.getConstantFP(Val: 0.0f, DL, VT);
19678 }
19679
19680 // (fsub -0.0, N1) -> -N1
19681 if (N0CFP && N0CFP->isZero()) {
19682 if (N0CFP->isNegative() || DAG.canIgnoreSignBitOfZero(Op: SDValue(N, 0))) {
19683 // We cannot replace an FSUB(+-0.0,X) with FNEG(X) when denormals are
19684 // flushed to zero, unless all users treat denorms as zero (DAZ).
19685 // FIXME: This transform will change the sign of a NaN and the behavior
19686 // of a signaling NaN. It is only valid when a NoNaN flag is present.
19687 DenormalMode DenormMode = DAG.getDenormalMode(VT);
19688 if (DenormMode == DenormalMode::getIEEE()) {
19689 if (SDValue NegN1 =
19690 TLI.getNegatedExpression(Op: N1, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize))
19691 return NegN1;
19692 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::FNEG, VT))
19693 return DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: N1);
19694 }
19695 }
19696 }
19697
19698 if (Flags.hasAllowReassociation() && Flags.hasNoSignedZeros() &&
19699 N1.getOpcode() == ISD::FADD) {
19700 // X - (X + Y) -> -Y
19701 if (N0 == N1->getOperand(Num: 0))
19702 return DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: N1->getOperand(Num: 1));
19703 // X - (Y + X) -> -Y
19704 if (N0 == N1->getOperand(Num: 1))
19705 return DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: N1->getOperand(Num: 0));
19706 }
19707
19708 // fold (fsub A, (fneg B)) -> (fadd A, B)
19709 if (SDValue NegN1 =
19710 TLI.getNegatedExpression(Op: N1, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize))
19711 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0, N2: NegN1);
19712
19713 // FSUB -> FMA combines:
19714 if (SDValue Fused = visitFSUBForFMACombine(N)) {
19715 AddToWorklist(N: Fused.getNode());
19716 return Fused;
19717 }
19718
19719 return SDValue();
19720}
19721
19722// Transform IEEE Floats:
19723// (fmul C, (uitofp Pow2))
19724// -> (bitcast_to_FP (add (bitcast_to_INT C), Log2(Pow2) << mantissa))
19725// (fdiv C, (uitofp Pow2))
19726// -> (bitcast_to_FP (sub (bitcast_to_INT C), Log2(Pow2) << mantissa))
19727//
19728// The rationale is fmul/fdiv by a power of 2 is just change the exponent, so
19729// there is no need for more than an add/sub.
19730//
19731// This is valid under the following circumstances:
19732// 1) We are dealing with IEEE floats
19733// 2) C is normal
19734// 3) The fmul/fdiv add/sub will not go outside of min/max exponent bounds.
19735// TODO: Much of this could also be used for generating `ldexp` on targets the
19736// prefer it.
19737SDValue DAGCombiner::combineFMulOrFDivWithIntPow2(SDNode *N) {
19738 EVT VT = N->getValueType(ResNo: 0);
19739 if (!APFloat::isIEEELikeFP(VT.getFltSemantics()))
19740 return SDValue();
19741
19742 SDValue ConstOp, Pow2Op;
19743
19744 std::optional<int> Mantissa;
19745 auto GetConstAndPow2Ops = [&](unsigned ConstOpIdx) {
19746 if (ConstOpIdx == 1 && N->getOpcode() == ISD::FDIV)
19747 return false;
19748
19749 ConstOp = peekThroughBitcasts(V: N->getOperand(Num: ConstOpIdx));
19750 Pow2Op = N->getOperand(Num: 1 - ConstOpIdx);
19751 unsigned Pow2Opc = Pow2Op.getOpcode();
19752 if (Pow2Opc != ISD::UINT_TO_FP && Pow2Opc != ISD::SINT_TO_FP)
19753 return false;
19754
19755 Pow2Op = Pow2Op.getOperand(i: 0);
19756
19757 KnownBits Pow2OpKnownBits = DAG.computeKnownBits(Op: Pow2Op);
19758 if (Pow2Opc == ISD::SINT_TO_FP && !Pow2OpKnownBits.isNonNegative())
19759 return false;
19760
19761 int MaxExpChange = Pow2OpKnownBits.countMaxActiveBits();
19762
19763 auto IsFPConstValid = [N, MaxExpChange, &Mantissa](ConstantFPSDNode *CFP) {
19764 if (CFP == nullptr)
19765 return false;
19766
19767 const APFloat &APF = CFP->getValueAPF();
19768
19769 // Make sure we have normal constant.
19770 if (!APF.isNormal())
19771 return false;
19772
19773 // Make sure the floats exponent is within the bounds that this transform
19774 // produces bitwise equals value.
19775 int CurExp = ilogb(Arg: APF);
19776 // FMul by pow2 will only increase exponent.
19777 int MinExp =
19778 N->getOpcode() == ISD::FMUL ? CurExp : (CurExp - MaxExpChange);
19779 // FDiv by pow2 will only decrease exponent.
19780 int MaxExp =
19781 N->getOpcode() == ISD::FDIV ? CurExp : (CurExp + MaxExpChange);
19782 if (MinExp <= APFloat::semanticsMinExponent(APF.getSemantics()) ||
19783 MaxExp >= APFloat::semanticsMaxExponent(APF.getSemantics()))
19784 return false;
19785
19786 // Finally make sure we actually know the mantissa for the float type.
19787 int ThisMantissa = APFloat::semanticsPrecision(APF.getSemantics()) - 1;
19788 if (!Mantissa)
19789 Mantissa = ThisMantissa;
19790
19791 return *Mantissa == ThisMantissa && ThisMantissa > 0;
19792 };
19793
19794 // TODO: We may be able to include undefs.
19795 return ISD::matchUnaryFpPredicate(Op: ConstOp, Match: IsFPConstValid);
19796 };
19797
19798 if (!GetConstAndPow2Ops(0) && !GetConstAndPow2Ops(1))
19799 return SDValue();
19800
19801 if (!TLI.optimizeFMulOrFDivAsShiftAddBitcast(N, FPConst: ConstOp, IntPow2: Pow2Op))
19802 return SDValue();
19803
19804 // Get log2 after all other checks have taken place. This is because
19805 // BuildLogBase2 may create a new node.
19806 SDLoc DL(N);
19807 // Get Log2 type with same bitwidth as the float type (VT).
19808 EVT NewIntVT = VT.changeElementType(
19809 Context&: *DAG.getContext(),
19810 EltVT: EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: VT.getScalarSizeInBits()));
19811
19812 SDValue Log2 = BuildLogBase2(V: Pow2Op, DL, KnownNeverZero: DAG.isKnownNeverZero(Op: Pow2Op),
19813 /*InexpensiveOnly*/ true, OutVT: NewIntVT);
19814 if (!Log2)
19815 return SDValue();
19816
19817 // Perform actual transform.
19818 SDValue MantissaShiftCnt =
19819 DAG.getShiftAmountConstant(Val: *Mantissa, VT: NewIntVT, DL);
19820 // TODO: Sometimes Log2 is of form `(X + C)`. `(X + C) << C1` should fold to
19821 // `(X << C1) + (C << C1)`, but that isn't always the case because of the
19822 // cast. We could implement that by handle here to handle the casts.
19823 SDValue Shift = DAG.getNode(Opcode: ISD::SHL, DL, VT: NewIntVT, N1: Log2, N2: MantissaShiftCnt);
19824 SDValue ResAsInt =
19825 DAG.getNode(Opcode: N->getOpcode() == ISD::FMUL ? ISD::ADD : ISD::SUB, DL,
19826 VT: NewIntVT, N1: DAG.getBitcast(VT: NewIntVT, V: ConstOp), N2: Shift);
19827 SDValue ResAsFP = DAG.getBitcast(VT, V: ResAsInt);
19828 return ResAsFP;
19829}
19830
19831SDValue DAGCombiner::visitFMUL(SDNode *N) {
19832 SDValue N0 = N->getOperand(Num: 0);
19833 SDValue N1 = N->getOperand(Num: 1);
19834 ConstantFPSDNode *N1CFP = isConstOrConstSplatFP(N: N1, AllowUndefs: true);
19835 EVT VT = N->getValueType(ResNo: 0);
19836 SDLoc DL(N);
19837 const SDNodeFlags Flags = N->getFlags();
19838 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
19839
19840 if (SDValue R = DAG.simplifyFPBinop(Opcode: N->getOpcode(), X: N0, Y: N1, Flags))
19841 return R;
19842
19843 // fold (fmul c1, c2) -> c1*c2
19844 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FMUL, DL, VT, Ops: {N0, N1}))
19845 return C;
19846
19847 // canonicalize constant to RHS
19848 if (DAG.isConstantFPBuildVectorOrConstantFP(N: N0) &&
19849 !DAG.isConstantFPBuildVectorOrConstantFP(N: N1))
19850 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1, N2: N0);
19851
19852 // fold vector ops
19853 if (VT.isVector())
19854 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
19855 return FoldedVOp;
19856
19857 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
19858 return NewSel;
19859
19860 if (Flags.hasAllowReassociation()) {
19861 // fmul (fmul X, C1), C2 -> fmul X, C1 * C2
19862 if (DAG.isConstantFPBuildVectorOrConstantFP(N: N1) &&
19863 N0.getOpcode() == ISD::FMUL) {
19864 SDValue N00 = N0.getOperand(i: 0);
19865 SDValue N01 = N0.getOperand(i: 1);
19866 // Avoid an infinite loop by making sure that N00 is not a constant
19867 // (the inner multiply has not been constant folded yet).
19868 if (DAG.isConstantFPBuildVectorOrConstantFP(N: N01) &&
19869 !DAG.isConstantFPBuildVectorOrConstantFP(N: N00)) {
19870 SDValue MulConsts = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N01, N2: N1);
19871 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N00, N2: MulConsts);
19872 }
19873 }
19874
19875 // Match a special-case: we convert X * 2.0 into fadd.
19876 // fmul (fadd X, X), C -> fmul X, 2.0 * C
19877 if (N0.getOpcode() == ISD::FADD && N0.hasOneUse() &&
19878 N0.getOperand(i: 0) == N0.getOperand(i: 1)) {
19879 const SDValue Two = DAG.getConstantFP(Val: 2.0, DL, VT);
19880 SDValue MulConsts = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Two, N2: N1);
19881 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0.getOperand(i: 0), N2: MulConsts);
19882 }
19883
19884 // Fold fmul(vecreduce(x), vecreduce(y)) -> vecreduce(fmul(x, y))
19885 if (SDValue SD = reassociateReduction(RedOpc: ISD::VECREDUCE_FMUL, Opc: ISD::FMUL, DL,
19886 VT, N0, N1, Flags))
19887 return SD;
19888 }
19889
19890 // fold (fmul X, 2.0) -> (fadd X, X)
19891 if (N1CFP && N1CFP->isExactlyValue(V: +2.0))
19892 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0, N2: N0);
19893
19894 // fold (fmul X, -1.0) -> (fsub -0.0, X)
19895 if (N1CFP && N1CFP->isMinusOne()) {
19896 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::FSUB, VT)) {
19897 return DAG.getNode(Opcode: ISD::FSUB, DL, VT,
19898 N1: DAG.getConstantFP(Val: -0.0, DL, VT), N2: N0, Flags);
19899 }
19900 }
19901
19902 // -N0 * -N1 --> N0 * N1
19903 TargetLowering::NegatibleCost CostN0 =
19904 TargetLowering::NegatibleCost::Expensive;
19905 TargetLowering::NegatibleCost CostN1 =
19906 TargetLowering::NegatibleCost::Expensive;
19907 SDValue NegN0 =
19908 TLI.getNegatedExpression(Op: N0, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize, Cost&: CostN0);
19909 if (NegN0) {
19910 HandleSDNode NegN0Handle(NegN0);
19911 SDValue NegN1 =
19912 TLI.getNegatedExpression(Op: N1, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize, Cost&: CostN1);
19913 if (NegN1 && (CostN0 == TargetLowering::NegatibleCost::Cheaper ||
19914 CostN1 == TargetLowering::NegatibleCost::Cheaper))
19915 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: NegN0, N2: NegN1);
19916 }
19917
19918 // fold (fmul X, (select (fcmp X > 0.0), -1.0, 1.0)) -> (fneg (fabs X))
19919 // fold (fmul X, (select (fcmp X > 0.0), 1.0, -1.0)) -> (fabs X)
19920 if (Flags.hasNoNaNs() && Flags.hasNoSignedZeros() &&
19921 (N0.getOpcode() == ISD::SELECT || N1.getOpcode() == ISD::SELECT) &&
19922 TLI.isOperationLegal(Op: ISD::FABS, VT)) {
19923 SDValue Select = N0, X = N1;
19924 if (Select.getOpcode() != ISD::SELECT)
19925 std::swap(a&: Select, b&: X);
19926
19927 SDValue Cond = Select.getOperand(i: 0);
19928 auto TrueOpnd = dyn_cast<ConstantFPSDNode>(Val: Select.getOperand(i: 1));
19929 auto FalseOpnd = dyn_cast<ConstantFPSDNode>(Val: Select.getOperand(i: 2));
19930
19931 if (TrueOpnd && FalseOpnd && Cond.getOpcode() == ISD::SETCC &&
19932 Cond.getOperand(i: 0) == X && isa<ConstantFPSDNode>(Val: Cond.getOperand(i: 1)) &&
19933 cast<ConstantFPSDNode>(Val: Cond.getOperand(i: 1))->isPosZero()) {
19934 ISD::CondCode CC = cast<CondCodeSDNode>(Val: Cond.getOperand(i: 2))->get();
19935 switch (CC) {
19936 default: break;
19937 case ISD::SETOLT:
19938 case ISD::SETULT:
19939 case ISD::SETOLE:
19940 case ISD::SETULE:
19941 case ISD::SETLT:
19942 case ISD::SETLE:
19943 std::swap(a&: TrueOpnd, b&: FalseOpnd);
19944 [[fallthrough]];
19945 case ISD::SETOGT:
19946 case ISD::SETUGT:
19947 case ISD::SETOGE:
19948 case ISD::SETUGE:
19949 case ISD::SETGT:
19950 case ISD::SETGE:
19951 if (TrueOpnd->isMinusOne() && FalseOpnd->isOne() &&
19952 TLI.isOperationLegal(Op: ISD::FNEG, VT))
19953 return DAG.getNode(Opcode: ISD::FNEG, DL, VT,
19954 Operand: DAG.getNode(Opcode: ISD::FABS, DL, VT, Operand: X));
19955 if (TrueOpnd->isOne() && FalseOpnd->isMinusOne())
19956 return DAG.getNode(Opcode: ISD::FABS, DL, VT, Operand: X);
19957
19958 break;
19959 }
19960 }
19961 }
19962
19963 // FMUL -> FMA combines:
19964 if (SDValue Fused = visitFMULForFMADistributiveCombine(N)) {
19965 AddToWorklist(N: Fused.getNode());
19966 return Fused;
19967 }
19968
19969 // Don't do `combineFMulOrFDivWithIntPow2` until after FMUL -> FMA has been
19970 // able to run.
19971 if (SDValue R = combineFMulOrFDivWithIntPow2(N))
19972 return R;
19973
19974 return SDValue();
19975}
19976
19977SDValue DAGCombiner::visitFMA(SDNode *N) {
19978 SDValue N0 = N->getOperand(Num: 0);
19979 SDValue N1 = N->getOperand(Num: 1);
19980 SDValue N2 = N->getOperand(Num: 2);
19981 ConstantFPSDNode *N0CFP = dyn_cast<ConstantFPSDNode>(Val&: N0);
19982 ConstantFPSDNode *N1CFP = dyn_cast<ConstantFPSDNode>(Val&: N1);
19983 ConstantFPSDNode *N2CFP = dyn_cast<ConstantFPSDNode>(Val&: N2);
19984 EVT VT = N->getValueType(ResNo: 0);
19985 SDLoc DL(N);
19986 // FMA nodes have flags that propagate to the created nodes.
19987 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
19988
19989 // Constant fold FMA.
19990 if (SDValue C =
19991 DAG.FoldConstantArithmetic(Opcode: N->getOpcode(), DL, VT, Ops: {N0, N1, N2}))
19992 return C;
19993
19994 // (-N0 * -N1) + N2 --> (N0 * N1) + N2
19995 TargetLowering::NegatibleCost CostN0 =
19996 TargetLowering::NegatibleCost::Expensive;
19997 TargetLowering::NegatibleCost CostN1 =
19998 TargetLowering::NegatibleCost::Expensive;
19999 SDValue NegN0 =
20000 TLI.getNegatedExpression(Op: N0, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize, Cost&: CostN0);
20001 if (NegN0) {
20002 HandleSDNode NegN0Handle(NegN0);
20003 SDValue NegN1 =
20004 TLI.getNegatedExpression(Op: N1, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize, Cost&: CostN1);
20005 if (NegN1 && (CostN0 == TargetLowering::NegatibleCost::Cheaper ||
20006 CostN1 == TargetLowering::NegatibleCost::Cheaper))
20007 return DAG.getNode(Opcode: ISD::FMA, DL, VT, N1: NegN0, N2: NegN1, N3: N2);
20008 }
20009
20010 if (N->getFlags().hasNoNaNs() && N->getFlags().hasNoInfs()) {
20011 if (N->getFlags().hasNoSignedZeros() || (N2CFP && !N2CFP->isNegZero())) {
20012 if (N0CFP && N0CFP->isZero())
20013 return N2;
20014 if (N1CFP && N1CFP->isZero())
20015 return N2;
20016 }
20017 }
20018
20019 if (N0CFP && N0CFP->isOne())
20020 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1, N2);
20021 if (N1CFP && N1CFP->isOne())
20022 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0, N2);
20023
20024 // Canonicalize (fma c, x, y) -> (fma x, c, y)
20025 if (DAG.isConstantFPBuildVectorOrConstantFP(N: N0) &&
20026 !DAG.isConstantFPBuildVectorOrConstantFP(N: N1))
20027 return DAG.getNode(Opcode: ISD::FMA, DL, VT, N1, N2: N0, N3: N2);
20028
20029 bool CanReassociate = N->getFlags().hasAllowReassociation();
20030 if (CanReassociate) {
20031 // (fma x, c1, (fmul x, c2)) -> (fmul x, c1+c2)
20032 if (N2.getOpcode() == ISD::FMUL && N0 == N2.getOperand(i: 0) &&
20033 DAG.isConstantFPBuildVectorOrConstantFP(N: N1) &&
20034 DAG.isConstantFPBuildVectorOrConstantFP(N: N2.getOperand(i: 1))) {
20035 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0,
20036 N2: DAG.getNode(Opcode: ISD::FADD, DL, VT, N1, N2: N2.getOperand(i: 1)));
20037 }
20038
20039 // (fma (fmul x, c1), c2, y) -> (fma x, c1*c2, y)
20040 if (N0.getOpcode() == ISD::FMUL &&
20041 DAG.isConstantFPBuildVectorOrConstantFP(N: N1) &&
20042 DAG.isConstantFPBuildVectorOrConstantFP(N: N0.getOperand(i: 1))) {
20043 return DAG.getNode(Opcode: ISD::FMA, DL, VT, N1: N0.getOperand(i: 0),
20044 N2: DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1, N2: N0.getOperand(i: 1)),
20045 N3: N2);
20046 }
20047 }
20048
20049 // (fma x, -1, y) -> (fadd (fneg x), y)
20050 if (N1CFP) {
20051 if (N1CFP->isOne())
20052 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N0, N2);
20053
20054 if (N1CFP->isMinusOne() &&
20055 (!LegalOperations || TLI.isOperationLegal(Op: ISD::FNEG, VT))) {
20056 SDValue RHSNeg = DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: N0);
20057 AddToWorklist(N: RHSNeg.getNode());
20058 return DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: N2, N2: RHSNeg);
20059 }
20060
20061 // fma (fneg x), K, y -> fma x -K, y
20062 if (N0.getOpcode() == ISD::FNEG &&
20063 (TLI.isOperationLegal(Op: ISD::ConstantFP, VT) ||
20064 (N1.hasOneUse() &&
20065 !TLI.isFPImmLegal(N1CFP->getValueAPF(), VT, ForCodeSize)))) {
20066 return DAG.getNode(Opcode: ISD::FMA, DL, VT, N1: N0.getOperand(i: 0),
20067 N2: DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: N1), N3: N2);
20068 }
20069 }
20070
20071 if (CanReassociate) {
20072 // (fma x, c, x) -> (fmul x, (c+1))
20073 if (N1CFP && N0 == N2) {
20074 return DAG.getNode(
20075 Opcode: ISD::FMUL, DL, VT, N1: N0,
20076 N2: DAG.getNode(Opcode: ISD::FADD, DL, VT, N1, N2: DAG.getConstantFP(Val: 1.0, DL, VT)));
20077 }
20078
20079 // (fma x, c, (fneg x)) -> (fmul x, (c-1))
20080 if (N1CFP && N2.getOpcode() == ISD::FNEG && N2.getOperand(i: 0) == N0) {
20081 return DAG.getNode(
20082 Opcode: ISD::FMUL, DL, VT, N1: N0,
20083 N2: DAG.getNode(Opcode: ISD::FADD, DL, VT, N1, N2: DAG.getConstantFP(Val: -1.0, DL, VT)));
20084 }
20085 }
20086
20087 // fold ((fma (fneg X), Y, (fneg Z)) -> fneg (fma X, Y, Z))
20088 // fold ((fma X, (fneg Y), (fneg Z)) -> fneg (fma X, Y, Z))
20089 if (!TLI.isFNegFree(VT))
20090 if (SDValue Neg = TLI.getCheaperNegatedExpression(
20091 Op: SDValue(N, 0), DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize))
20092 return DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: Neg);
20093 return SDValue();
20094}
20095
20096SDValue DAGCombiner::visitFMAD(SDNode *N) {
20097 SDValue N0 = N->getOperand(Num: 0);
20098 SDValue N1 = N->getOperand(Num: 1);
20099 SDValue N2 = N->getOperand(Num: 2);
20100 EVT VT = N->getValueType(ResNo: 0);
20101 SDLoc DL(N);
20102
20103 // Constant fold FMAD.
20104 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FMAD, DL, VT, Ops: {N0, N1, N2}))
20105 return C;
20106
20107 return SDValue();
20108}
20109
20110SDValue DAGCombiner::visitFMULADD(SDNode *N) {
20111 SDValue N0 = N->getOperand(Num: 0);
20112 SDValue N1 = N->getOperand(Num: 1);
20113 SDValue N2 = N->getOperand(Num: 2);
20114 EVT VT = N->getValueType(ResNo: 0);
20115 SDLoc DL(N);
20116
20117 // Constant fold FMULADD.
20118 if (SDValue C =
20119 DAG.FoldConstantArithmetic(Opcode: ISD::FMULADD, DL, VT, Ops: {N0, N1, N2}))
20120 return C;
20121
20122 return SDValue();
20123}
20124
20125// Combine multiple FDIVs with the same divisor into multiple FMULs by the
20126// reciprocal.
20127// E.g., (a / D; b / D;) -> (recip = 1.0 / D; a * recip; b * recip)
20128// Notice that this is not always beneficial. One reason is different targets
20129// may have different costs for FDIV and FMUL, so sometimes the cost of two
20130// FDIVs may be lower than the cost of one FDIV and two FMULs. Another reason
20131// is the critical path is increased from "one FDIV" to "one FDIV + one FMUL".
20132SDValue DAGCombiner::combineRepeatedFPDivisors(SDNode *N) {
20133 // TODO: Limit this transform based on optsize/minsize - it always creates at
20134 // least 1 extra instruction. But the perf win may be substantial enough
20135 // that only minsize should restrict this.
20136 const SDNodeFlags Flags = N->getFlags();
20137 if (LegalDAG || !Flags.hasAllowReciprocal())
20138 return SDValue();
20139
20140 // Skip if current node is a reciprocal/fneg-reciprocal.
20141 SDValue N0 = N->getOperand(Num: 0), N1 = N->getOperand(Num: 1);
20142 ConstantFPSDNode *N0CFP = isConstOrConstSplatFP(N: N0, /* AllowUndefs */ true);
20143 if (N0CFP && (N0CFP->isOne() || N0CFP->isMinusOne()))
20144 return SDValue();
20145
20146 // Exit early if the target does not want this transform or if there can't
20147 // possibly be enough uses of the divisor to make the transform worthwhile.
20148 unsigned MinUses = TLI.combineRepeatedFPDivisors();
20149
20150 // For splat vectors, scale the number of uses by the splat factor. If we can
20151 // convert the division into a scalar op, that will likely be much faster.
20152 unsigned NumElts = 1;
20153 EVT VT = N->getValueType(ResNo: 0);
20154 if (VT.isVector() && DAG.isSplatValue(V: N1))
20155 NumElts = VT.getVectorMinNumElements();
20156
20157 if (!MinUses || (N1->use_size() * NumElts) < MinUses)
20158 return SDValue();
20159
20160 // Find all FDIV users of the same divisor.
20161 // Use a set because duplicates may be present in the user list.
20162 SetVector<SDNode *> Users;
20163 for (auto *U : N1->users()) {
20164 if (U->getOpcode() == ISD::FDIV && U->getOperand(Num: 1) == N1) {
20165 // Skip X/sqrt(X) that has not been simplified to sqrt(X) yet.
20166 if (U->getOperand(Num: 1).getOpcode() == ISD::FSQRT &&
20167 U->getOperand(Num: 0) == U->getOperand(Num: 1).getOperand(i: 0) &&
20168 U->getFlags().hasAllowReassociation() &&
20169 U->getFlags().hasNoSignedZeros())
20170 continue;
20171
20172 // This division is eligible for optimization only if global unsafe math
20173 // is enabled or if this division allows reciprocal formation.
20174 if (U->getFlags().hasAllowReciprocal())
20175 Users.insert(X: U);
20176 }
20177 }
20178
20179 // Now that we have the actual number of divisor uses, make sure it meets
20180 // the minimum threshold specified by the target.
20181 if ((Users.size() * NumElts) < MinUses)
20182 return SDValue();
20183
20184 SDLoc DL(N);
20185 SDValue FPOne = DAG.getConstantFP(Val: 1.0, DL, VT);
20186 SDValue Reciprocal = DAG.getNode(Opcode: ISD::FDIV, DL, VT, N1: FPOne, N2: N1, Flags);
20187
20188 // Dividend / Divisor -> Dividend * Reciprocal
20189 for (auto *U : Users) {
20190 SDValue Dividend = U->getOperand(Num: 0);
20191 if (Dividend != FPOne) {
20192 SDValue NewNode = DAG.getNode(Opcode: ISD::FMUL, DL: SDLoc(U), VT, N1: Dividend,
20193 N2: Reciprocal, Flags);
20194 CombineTo(N: U, Res: NewNode);
20195 } else if (U != Reciprocal.getNode()) {
20196 // In the absence of fast-math-flags, this user node is always the
20197 // same node as Reciprocal, but with FMF they may be different nodes.
20198 CombineTo(N: U, Res: Reciprocal);
20199 }
20200 }
20201 return SDValue(N, 0); // N was replaced.
20202}
20203
20204SDValue DAGCombiner::visitFDIV(SDNode *N) {
20205 SDValue N0 = N->getOperand(Num: 0);
20206 SDValue N1 = N->getOperand(Num: 1);
20207 EVT VT = N->getValueType(ResNo: 0);
20208 SDLoc DL(N);
20209 SDNodeFlags Flags = N->getFlags();
20210 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
20211
20212 if (SDValue R = DAG.simplifyFPBinop(Opcode: N->getOpcode(), X: N0, Y: N1, Flags))
20213 return R;
20214
20215 // fold (fdiv c1, c2) -> c1/c2
20216 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FDIV, DL, VT, Ops: {N0, N1}))
20217 return C;
20218
20219 // fold vector ops
20220 if (VT.isVector())
20221 if (SDValue FoldedVOp = SimplifyVBinOp(N, DL))
20222 return FoldedVOp;
20223
20224 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
20225 return NewSel;
20226
20227 if (SDValue V = combineRepeatedFPDivisors(N))
20228 return V;
20229
20230 // fold (fdiv X, c2) -> (fmul X, 1/c2) if there is no loss in precision, or
20231 // the loss is acceptable with AllowReciprocal.
20232 if (auto *N1CFP = isConstOrConstSplatFP(N: N1, AllowUndefs: true)) {
20233 // Compute the reciprocal 1.0 / c2.
20234 const APFloat &N1APF = N1CFP->getValueAPF();
20235 APFloat Recip = APFloat::getOne(Sem: N1APF.getSemantics());
20236 APFloat::opStatus st = Recip.divide(RHS: N1APF, RM: APFloat::rmNearestTiesToEven);
20237 // Only do the transform if the reciprocal is a legal fp immediate that
20238 // isn't too nasty (eg NaN, denormal, ...).
20239 if (((st == APFloat::opOK && !Recip.isDenormal()) ||
20240 (st == APFloat::opInexact && Flags.hasAllowReciprocal())) &&
20241 (!LegalOperations ||
20242 // FIXME: custom lowering of ConstantFP might fail (see e.g. ARM
20243 // backend)... we should handle this gracefully after Legalize.
20244 // TLI.isOperationLegalOrCustom(ISD::ConstantFP, VT) ||
20245 TLI.isOperationLegal(Op: ISD::ConstantFP, VT) ||
20246 TLI.isFPImmLegal(Recip, VT, ForCodeSize)))
20247 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0,
20248 N2: DAG.getConstantFP(Val: Recip, DL, VT));
20249 }
20250
20251 if (Flags.hasAllowReciprocal()) {
20252 // If this FDIV is part of a reciprocal square root, it may be folded
20253 // into a target-specific square root estimate instruction.
20254 bool N1AllowReciprocal = N1->getFlags().hasAllowReciprocal();
20255 if (N1.getOpcode() == ISD::FSQRT) {
20256 if (SDValue RV = buildRsqrtEstimate(Op: N1.getOperand(i: 0), Flags: N1->getFlags()))
20257 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0, N2: RV);
20258 } else if (N1.getOpcode() == ISD::FP_EXTEND &&
20259 N1.getOperand(i: 0).getOpcode() == ISD::FSQRT &&
20260 N1AllowReciprocal) {
20261 if (SDValue RV = buildRsqrtEstimate(Op: N1.getOperand(i: 0).getOperand(i: 0),
20262 Flags: N1.getOperand(i: 0)->getFlags())) {
20263 RV = DAG.getNode(Opcode: ISD::FP_EXTEND, DL: SDLoc(N1), VT, Operand: RV);
20264 AddToWorklist(N: RV.getNode());
20265 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0, N2: RV);
20266 }
20267 } else if (N1.getOpcode() == ISD::FP_ROUND &&
20268 N1.getOperand(i: 0).getOpcode() == ISD::FSQRT) {
20269 if (SDValue RV = buildRsqrtEstimate(Op: N1.getOperand(i: 0).getOperand(i: 0),
20270 Flags: N1.getOperand(i: 0)->getFlags())) {
20271 RV = DAG.getNode(Opcode: ISD::FP_ROUND, DL: SDLoc(N1), VT, N1: RV, N2: N1.getOperand(i: 1));
20272 AddToWorklist(N: RV.getNode());
20273 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0, N2: RV);
20274 }
20275 } else if (N1.getOpcode() == ISD::FMUL) {
20276 // Look through an FMUL. Even though this won't remove the FDIV directly,
20277 // it's still worthwhile to get rid of the FSQRT if possible.
20278 SDValue Sqrt, Y;
20279 if (N1.getOperand(i: 0).getOpcode() == ISD::FSQRT) {
20280 Sqrt = N1.getOperand(i: 0);
20281 Y = N1.getOperand(i: 1);
20282 } else if (N1.getOperand(i: 1).getOpcode() == ISD::FSQRT) {
20283 Sqrt = N1.getOperand(i: 1);
20284 Y = N1.getOperand(i: 0);
20285 }
20286 if (Sqrt.getNode()) {
20287 // If the other multiply operand is known positive, pull it into the
20288 // sqrt. That will eliminate the division if we convert to an estimate.
20289 if (Flags.hasAllowReassociation() && N1.hasOneUse() &&
20290 N1->getFlags().hasAllowReassociation() && Sqrt.hasOneUse()) {
20291 SDValue A;
20292 if (Y.getOpcode() == ISD::FABS && Y.hasOneUse())
20293 A = Y.getOperand(i: 0);
20294 else if (Y == Sqrt.getOperand(i: 0))
20295 A = Y;
20296 if (A) {
20297 // X / (fabs(A) * sqrt(Z)) --> X / sqrt(A*A*Z) --> X * rsqrt(A*A*Z)
20298 // X / (A * sqrt(A)) --> X / sqrt(A*A*A) --> X * rsqrt(A*A*A)
20299 SDValue AA = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: A, N2: A);
20300 SDValue AAZ =
20301 DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: AA, N2: Sqrt.getOperand(i: 0));
20302 if (SDValue Rsqrt = buildRsqrtEstimate(Op: AAZ, Flags: Sqrt->getFlags()))
20303 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0, N2: Rsqrt);
20304
20305 // Estimate creation failed. Clean up speculatively created nodes.
20306 recursivelyDeleteUnusedNodes(N: AAZ.getNode());
20307 }
20308 }
20309
20310 // We found a FSQRT, so try to make this fold:
20311 // X / (Y * sqrt(Z)) -> X * (rsqrt(Z) / Y)
20312 if (SDValue Rsqrt =
20313 buildRsqrtEstimate(Op: Sqrt.getOperand(i: 0), Flags: Sqrt->getFlags())) {
20314 SDValue Div = DAG.getNode(Opcode: ISD::FDIV, DL: SDLoc(N1), VT, N1: Rsqrt, N2: Y);
20315 AddToWorklist(N: Div.getNode());
20316 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N0, N2: Div);
20317 }
20318 }
20319 }
20320
20321 // Fold into a reciprocal estimate and multiply instead of a real divide.
20322 if (Flags.hasNoInfs())
20323 if (SDValue RV = BuildDivEstimate(N: N0, Op: N1, Flags))
20324 return RV;
20325 }
20326
20327 // Fold X/Sqrt(X) -> Sqrt(X)
20328 if (DAG.canIgnoreSignBitOfZero(Op: SDValue(N, 0)) &&
20329 Flags.hasAllowReassociation())
20330 if (N1.getOpcode() == ISD::FSQRT && N0 == N1.getOperand(i: 0))
20331 return N1;
20332
20333 // (fdiv (fneg X), (fneg Y)) -> (fdiv X, Y)
20334 TargetLowering::NegatibleCost CostN0 =
20335 TargetLowering::NegatibleCost::Expensive;
20336 TargetLowering::NegatibleCost CostN1 =
20337 TargetLowering::NegatibleCost::Expensive;
20338 SDValue NegN0 =
20339 TLI.getNegatedExpression(Op: N0, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize, Cost&: CostN0);
20340 if (NegN0) {
20341 HandleSDNode NegN0Handle(NegN0);
20342 SDValue NegN1 =
20343 TLI.getNegatedExpression(Op: N1, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize, Cost&: CostN1);
20344 if (NegN1 && (CostN0 == TargetLowering::NegatibleCost::Cheaper ||
20345 CostN1 == TargetLowering::NegatibleCost::Cheaper))
20346 return DAG.getNode(Opcode: ISD::FDIV, DL, VT, N1: NegN0, N2: NegN1);
20347 }
20348
20349 if (SDValue R = combineFMulOrFDivWithIntPow2(N))
20350 return R;
20351
20352 return SDValue();
20353}
20354
20355SDValue DAGCombiner::visitFREM(SDNode *N) {
20356 SDValue N0 = N->getOperand(Num: 0);
20357 SDValue N1 = N->getOperand(Num: 1);
20358 EVT VT = N->getValueType(ResNo: 0);
20359 SDNodeFlags Flags = N->getFlags();
20360 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
20361 SDLoc DL(N);
20362
20363 if (SDValue R = DAG.simplifyFPBinop(Opcode: N->getOpcode(), X: N0, Y: N1, Flags))
20364 return R;
20365
20366 // fold (frem c1, c2) -> fmod(c1,c2)
20367 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FREM, DL, VT, Ops: {N0, N1}))
20368 return C;
20369
20370 if (SDValue NewSel = foldBinOpIntoSelect(BO: N))
20371 return NewSel;
20372
20373 // Lower frem N0, N1 => x - trunc(N0 / N1) * N1, providing N1 is an integer
20374 // power of 2.
20375 if (!TLI.isOperationLegal(Op: ISD::FREM, VT) &&
20376 TLI.isOperationLegalOrCustom(Op: ISD::FMUL, VT) &&
20377 TLI.isOperationLegalOrCustom(Op: ISD::FDIV, VT) &&
20378 TLI.isOperationLegalOrCustom(Op: ISD::FTRUNC, VT) &&
20379 DAG.isKnownToBeAPowerOfTwoFP(Val: N1)) {
20380 bool NeedsCopySign = !DAG.canIgnoreSignBitOfZero(Op: SDValue(N, 0)) &&
20381 !DAG.cannotBeOrderedNegativeFP(Op: N0);
20382 SDValue Div = DAG.getNode(Opcode: ISD::FDIV, DL, VT, N1: N0, N2: N1);
20383 SDValue Rnd = DAG.getNode(Opcode: ISD::FTRUNC, DL, VT, Operand: Div);
20384 SDValue MLA;
20385 if (TLI.isFMAFasterThanFMulAndFAdd(MF: DAG.getMachineFunction(), VT)) {
20386 MLA = DAG.getNode(Opcode: ISD::FMA, DL, VT, N1: DAG.getNode(Opcode: ISD::FNEG, DL, VT, Operand: Rnd),
20387 N2: N1, N3: N0);
20388 } else {
20389 SDValue Mul = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Rnd, N2: N1);
20390 MLA = DAG.getNode(Opcode: ISD::FSUB, DL, VT, N1: N0, N2: Mul);
20391 }
20392 return NeedsCopySign ? DAG.getNode(Opcode: ISD::FCOPYSIGN, DL, VT, N1: MLA, N2: N0) : MLA;
20393 }
20394
20395 return SDValue();
20396}
20397
20398SDValue DAGCombiner::visitFSQRT(SDNode *N) {
20399 SDNodeFlags Flags = N->getFlags();
20400
20401 // Require 'ninf' flag since sqrt(+Inf) = +Inf, but the estimation goes as:
20402 // sqrt(+Inf) == rsqrt(+Inf) * +Inf = 0 * +Inf = NaN
20403 if (!Flags.hasApproximateFuncs() || !Flags.hasNoInfs())
20404 return SDValue();
20405
20406 SDValue N0 = N->getOperand(Num: 0);
20407 if (TLI.isFsqrtCheap(X: N0, DAG))
20408 return SDValue();
20409
20410 // FSQRT nodes have flags that propagate to the created nodes.
20411 SelectionDAG::FlagInserter FlagInserter(DAG, Flags);
20412 // TODO: If this is N0/sqrt(N0), and we reach this node before trying to
20413 // transform the fdiv, we may produce a sub-optimal estimate sequence
20414 // because the reciprocal calculation may not have to filter out a
20415 // 0.0 input.
20416 return buildSqrtEstimate(Op: N0, Flags);
20417}
20418
20419/// copysign(x, fp_extend(y)) -> copysign(x, y)
20420/// copysign(x, fp_round(y)) -> copysign(x, y)
20421/// Operands to the functions are the type of X and Y respectively.
20422static inline bool CanCombineFCOPYSIGN_EXTEND_ROUND(EVT XTy, EVT YTy) {
20423 // Always fold no-op FP casts.
20424 if (XTy == YTy)
20425 return true;
20426
20427 // Do not optimize out type conversion of f128 type yet.
20428 // For some targets like x86_64, configuration is changed to keep one f128
20429 // value in one SSE register, but instruction selection cannot handle
20430 // FCOPYSIGN on SSE registers yet.
20431 if (YTy == MVT::f128)
20432 return false;
20433
20434 // Avoid mismatched vector operand types, for better instruction selection.
20435 return !YTy.isVector();
20436}
20437
20438static inline bool CanCombineFCOPYSIGN_EXTEND_ROUND(SDNode *N) {
20439 SDValue N1 = N->getOperand(Num: 1);
20440 if (N1.getOpcode() != ISD::FP_EXTEND &&
20441 N1.getOpcode() != ISD::FP_ROUND)
20442 return false;
20443 EVT N1VT = N1->getValueType(ResNo: 0);
20444 EVT N1Op0VT = N1->getOperand(Num: 0).getValueType();
20445 return CanCombineFCOPYSIGN_EXTEND_ROUND(XTy: N1VT, YTy: N1Op0VT);
20446}
20447
20448SDValue DAGCombiner::visitFCOPYSIGN(SDNode *N) {
20449 SDValue N0 = N->getOperand(Num: 0);
20450 SDValue N1 = N->getOperand(Num: 1);
20451 EVT VT = N->getValueType(ResNo: 0);
20452 SDLoc DL(N);
20453
20454 // fold (fcopysign c1, c2) -> fcopysign(c1,c2)
20455 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FCOPYSIGN, DL, VT, Ops: {N0, N1}))
20456 return C;
20457
20458 // copysign(x, fp_extend(y)) -> copysign(x, y)
20459 // copysign(x, fp_round(y)) -> copysign(x, y)
20460 if (CanCombineFCOPYSIGN_EXTEND_ROUND(N))
20461 return DAG.getNode(Opcode: ISD::FCOPYSIGN, DL, VT, N1: N0, N2: N1.getOperand(i: 0));
20462
20463 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
20464 return SDValue(N, 0);
20465
20466 if (VT != N1.getValueType())
20467 return SDValue();
20468
20469 // If this is equivalent to a disjoint or, replace it with one. This can
20470 // happen if the sign operand is a sign mask (i.e., x << sign_bit_position).
20471 if (DAG.SignBitIsZeroFP(Op: N0) &&
20472 DAG.computeKnownBits(Op: N1).Zero.isMaxSignedValue()) {
20473 // TODO: Just directly match the shift pattern. computeKnownBits is heavy
20474 // for a such a narrowly targeted case.
20475 EVT IntVT = VT.changeTypeToInteger();
20476 // TODO: It appears to be profitable in some situations to unconditionally
20477 // emit a fabs(n0) to perform this combine.
20478 SDValue CastSrc0 = DAG.getNode(Opcode: ISD::BITCAST, DL, VT: IntVT, Operand: N0);
20479 SDValue CastSrc1 = DAG.getNode(Opcode: ISD::BITCAST, DL, VT: IntVT, Operand: N1);
20480
20481 SDValue SignOr = DAG.getNode(Opcode: ISD::OR, DL, VT: IntVT, N1: CastSrc0, N2: CastSrc1,
20482 Flags: SDNodeFlags::Disjoint);
20483 return DAG.getNode(Opcode: ISD::BITCAST, DL, VT, Operand: SignOr);
20484 }
20485
20486 return SDValue();
20487}
20488
20489SDValue DAGCombiner::visitFPOW(SDNode *N) {
20490 ConstantFPSDNode *ExponentC = isConstOrConstSplatFP(N: N->getOperand(Num: 1));
20491 if (!ExponentC)
20492 return SDValue();
20493 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
20494
20495 // Try to convert x ** (1/3) into cube root.
20496 // TODO: Handle the various flavors of long double.
20497 // TODO: Since we're approximating, we don't need an exact 1/3 exponent.
20498 // Some range near 1/3 should be fine.
20499 EVT VT = N->getValueType(ResNo: 0);
20500 EVT ScalarVT = VT.getScalarType();
20501 if ((ScalarVT == MVT::f32 &&
20502 ExponentC->getValueAPF().isExactlyValue(V: 1.0f / 3.0f)) ||
20503 (ScalarVT == MVT::f64 &&
20504 ExponentC->getValueAPF().isExactlyValue(V: 1.0 / 3.0))) {
20505 // pow(-0.0, 1/3) = +0.0; cbrt(-0.0) = -0.0.
20506 // pow(-inf, 1/3) = +inf; cbrt(-inf) = -inf.
20507 // pow(-val, 1/3) = nan; cbrt(-val) = -num.
20508 // For regular numbers, rounding may cause the results to differ.
20509 // Therefore, we require { nsz ninf nnan afn } for this transform.
20510 // TODO: We could select out the special cases if we don't have nsz/ninf.
20511 SDNodeFlags Flags = N->getFlags();
20512 if (!Flags.hasNoSignedZeros() || !Flags.hasNoInfs() || !Flags.hasNoNaNs() ||
20513 !Flags.hasApproximateFuncs())
20514 return SDValue();
20515
20516 // Do not create a cbrt() libcall if the target does not have it, and do not
20517 // turn a pow that has lowering support into a cbrt() libcall.
20518 RTLIB::Libcall LC = RTLIB::getCBRT(VT);
20519 bool HasLibCall =
20520 DAG.getLibcalls().getLibcallImpl(Call: LC) != RTLIB::Unsupported;
20521 if (!HasLibCall ||
20522 (!DAG.getTargetLoweringInfo().isOperationExpand(Op: ISD::FPOW, VT) &&
20523 DAG.getTargetLoweringInfo().isOperationExpand(Op: ISD::FCBRT, VT)))
20524 return SDValue();
20525
20526 return DAG.getNode(Opcode: ISD::FCBRT, DL: SDLoc(N), VT, Operand: N->getOperand(Num: 0));
20527 }
20528
20529 // Try to convert x ** (1/4) and x ** (3/4) into square roots.
20530 // x ** (1/2) is canonicalized to sqrt, so we do not bother with that case.
20531 // TODO: This could be extended (using a target hook) to handle smaller
20532 // power-of-2 fractional exponents.
20533 bool ExponentIs025 = ExponentC->getValueAPF().isExactlyValue(V: 0.25);
20534 bool ExponentIs075 = ExponentC->getValueAPF().isExactlyValue(V: 0.75);
20535 if (ExponentIs025 || ExponentIs075) {
20536 // pow(-0.0, 0.25) = +0.0; sqrt(sqrt(-0.0)) = -0.0.
20537 // pow(-inf, 0.25) = +inf; sqrt(sqrt(-inf)) = NaN.
20538 // pow(-0.0, 0.75) = +0.0; sqrt(-0.0) * sqrt(sqrt(-0.0)) = +0.0.
20539 // pow(-inf, 0.75) = +inf; sqrt(-inf) * sqrt(sqrt(-inf)) = NaN.
20540 // For regular numbers, rounding may cause the results to differ.
20541 // Therefore, we require { nsz ninf afn } for this transform.
20542 // TODO: We could select out the special cases if we don't have nsz/ninf.
20543 SDNodeFlags Flags = N->getFlags();
20544
20545 // We only need no signed zeros for the 0.25 case.
20546 if ((!Flags.hasNoSignedZeros() && ExponentIs025) || !Flags.hasNoInfs() ||
20547 !Flags.hasApproximateFuncs())
20548 return SDValue();
20549
20550 // Don't double the number of libcalls. We are trying to inline fast code.
20551 if (!DAG.getTargetLoweringInfo().isOperationLegalOrCustom(Op: ISD::FSQRT, VT))
20552 return SDValue();
20553
20554 // Assume that libcalls are the smallest code.
20555 // TODO: This restriction should probably be lifted for vectors.
20556 if (ForCodeSize)
20557 return SDValue();
20558
20559 // pow(X, 0.25) --> sqrt(sqrt(X))
20560 SDLoc DL(N);
20561 SDValue Sqrt = DAG.getNode(Opcode: ISD::FSQRT, DL, VT, Operand: N->getOperand(Num: 0));
20562 SDValue SqrtSqrt = DAG.getNode(Opcode: ISD::FSQRT, DL, VT, Operand: Sqrt);
20563 if (ExponentIs025)
20564 return SqrtSqrt;
20565 // pow(X, 0.75) --> sqrt(X) * sqrt(sqrt(X))
20566 return DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Sqrt, N2: SqrtSqrt);
20567 }
20568
20569 return SDValue();
20570}
20571
20572static SDValue foldFPToIntToFP(SDNode *N, const SDLoc &DL, SelectionDAG &DAG,
20573 const TargetLowering &TLI) {
20574 // We can fold the fpto[us]i -> [us]itofp pattern into a single ftrunc.
20575 // Additionally, if there are clamps ([us]min or [us]max) around
20576 // the fpto[us]i, we can fold those into fminnum/fmaxnum around the ftrunc.
20577 // If NoSignedZerosFPMath is enabled, this is a direct replacement.
20578 // Otherwise, for strict math, we must handle edge cases:
20579 // 1. For unsigned conversions, use FABS to handle negative cases. Take -0.0
20580 // as example, it first becomes integer 0, and is converted back to +0.0.
20581 // FTRUNC on its own could produce -0.0.
20582
20583 // FIXME: We should be able to use node-level FMF here.
20584 EVT VT = N->getValueType(ResNo: 0);
20585 if (!TLI.isOperationLegalOrCustom(Op: ISD::FTRUNC, VT))
20586 return SDValue();
20587
20588 bool IsUnsigned = N->getOpcode() == ISD::UINT_TO_FP;
20589 bool IsSigned = N->getOpcode() == ISD::SINT_TO_FP;
20590 assert(IsSigned || IsUnsigned);
20591
20592 // Don't fold if the individual cast operations are already legal,
20593 // as FTRUNC may have a more expensive custom expansion.
20594 EVT IntVT = N->getOperand(Num: 0).getValueType();
20595 EVT LegalIntVT = TLI.getTypeToTransformTo(Context&: *DAG.getContext(), VT: IntVT);
20596 unsigned FPToIntOp = IsUnsigned ? ISD::FP_TO_UINT : ISD::FP_TO_SINT;
20597 unsigned IntToFPOp = N->getOpcode(); // UINT_TO_FP or SINT_TO_FP
20598 if (!TLI.isOperationLegal(Op: ISD::FTRUNC, VT) &&
20599 TLI.isOperationLegal(Op: FPToIntOp, VT: LegalIntVT) &&
20600 TLI.isOperationLegal(Op: IntToFPOp, VT))
20601 return SDValue();
20602
20603 bool IsSignedZeroSafe = DAG.canIgnoreSignBitOfZero(Op: SDValue(N, 0));
20604 // For signed conversions: The optimization changes signed zero behavior.
20605 if (IsSigned && !IsSignedZeroSafe)
20606 return SDValue();
20607 // For unsigned conversions, we need FABS to canonicalize -0.0 to +0.0
20608 // (unless outputting a signed zero is OK).
20609 if (IsUnsigned && !IsSignedZeroSafe && !TLI.isFAbsFree(VT))
20610 return SDValue();
20611
20612 // Collect potential clamp operations (outermost to innermost) and peel.
20613 struct ClampInfo {
20614 bool IsMin;
20615 SDValue Constant;
20616 };
20617 constexpr unsigned MaxClamps = 2;
20618 SmallVector<ClampInfo, MaxClamps> Clamps;
20619 unsigned MinOp = IsUnsigned ? ISD::UMIN : ISD::SMIN;
20620 unsigned MaxOp = IsUnsigned ? ISD::UMAX : ISD::SMAX;
20621 SDValue IntVal = N->getOperand(Num: 0);
20622 for (unsigned Level = 0; Level < MaxClamps; ++Level) {
20623 if (!IntVal.hasOneUse() ||
20624 (IntVal.getOpcode() != MinOp && IntVal.getOpcode() != MaxOp))
20625 break;
20626 SDValue RHS = IntVal.getOperand(i: 1);
20627 APInt IntConst;
20628 if (auto *IntConstNode = dyn_cast<ConstantSDNode>(Val&: RHS))
20629 IntConst = IntConstNode->getAPIntValue();
20630 else if (!ISD::isConstantSplatVector(N: RHS.getNode(), SplatValue&: IntConst))
20631 return SDValue();
20632 APFloat FPConst(VT.getFltSemantics());
20633 FPConst.convertFromAPInt(Input: IntConst, IsSigned, RM: APFloat::rmNearestTiesToEven);
20634 // Verify roundtrip exactness.
20635 APSInt RoundTrip(IntConst.getBitWidth(), IsUnsigned);
20636 bool IsExact;
20637 if (FPConst.convertToInteger(Result&: RoundTrip, RM: APFloat::rmTowardZero, IsExact: &IsExact) !=
20638 APFloat::opOK ||
20639 !IsExact || static_cast<const APInt &>(RoundTrip) != IntConst)
20640 return SDValue();
20641 bool IsMin = IntVal.getOpcode() == MinOp;
20642 Clamps.push_back(Elt: {.IsMin: IsMin, .Constant: DAG.getConstantFP(Val: FPConst, DL, VT)});
20643 IntVal = IntVal.getOperand(i: 0);
20644 }
20645
20646 // Check that the sequence ends with the correct kind of fpto[us]i.
20647 if (IntVal.getOpcode() != FPToIntOp ||
20648 IntVal.getOperand(i: 0).getValueType() != VT)
20649 return SDValue();
20650
20651 SDValue Result = IntVal.getOperand(i: 0);
20652 if (IsUnsigned && !IsSignedZeroSafe && TLI.isFAbsFree(VT))
20653 Result = DAG.getNode(Opcode: ISD::FABS, DL, VT, Operand: Result);
20654 Result = DAG.getNode(Opcode: ISD::FTRUNC, DL, VT, Operand: Result);
20655 // Apply clamps, if any, in reverse order (innermost first).
20656 for (const ClampInfo &Clamp : reverse(C&: Clamps)) {
20657 unsigned FPClampOp =
20658 getMinMaxOpcodeForClamp(IsMin: Clamp.IsMin, Operand1: Result, Operand2: Clamp.Constant, DAG, TLI);
20659 if (FPClampOp == ISD::DELETED_NODE)
20660 return SDValue();
20661 Result = DAG.getNode(Opcode: FPClampOp, DL, VT, N1: Result, N2: Clamp.Constant);
20662 }
20663 return Result;
20664}
20665
20666SDValue DAGCombiner::visitSINT_TO_FP(SDNode *N) {
20667 SDValue N0 = N->getOperand(Num: 0);
20668 EVT VT = N->getValueType(ResNo: 0);
20669 EVT OpVT = N0.getValueType();
20670 SDLoc DL(N);
20671
20672 // [us]itofp(undef) = 0, because the result value is bounded.
20673 if (N0.isUndef())
20674 return DAG.getConstantFP(Val: 0.0, DL, VT);
20675
20676 // fold (sint_to_fp c1) -> c1fp
20677 // ...but only if the target supports immediate floating-point values
20678 if ((!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::ConstantFP, VT)))
20679 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::SINT_TO_FP, DL, VT, Ops: {N0}))
20680 return C;
20681
20682 // If the input is a legal type, and SINT_TO_FP is not legal on this target,
20683 // but UINT_TO_FP is legal on this target, try to convert.
20684 if (!hasOperation(Opcode: ISD::SINT_TO_FP, VT: OpVT) &&
20685 hasOperation(Opcode: ISD::UINT_TO_FP, VT: OpVT)) {
20686 // If the sign bit is known to be zero, we can change this to UINT_TO_FP.
20687 if (DAG.SignBitIsZero(Op: N0))
20688 return DAG.getNode(Opcode: ISD::UINT_TO_FP, DL, VT, Operand: N0);
20689 }
20690
20691 // The next optimizations are desirable only if SELECT_CC can be lowered.
20692 // fold (sint_to_fp (setcc x, y, cc)) -> (select (setcc x, y, cc), -1.0, 0.0)
20693 if (N0.getOpcode() == ISD::SETCC && N0.getValueType() == MVT::i1 &&
20694 !VT.isVector() &&
20695 (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::ConstantFP, VT)))
20696 return DAG.getSelect(DL, VT, Cond: N0, LHS: DAG.getConstantFP(Val: -1.0, DL, VT),
20697 RHS: DAG.getConstantFP(Val: 0.0, DL, VT));
20698
20699 // fold (sint_to_fp (zext (setcc x, y, cc))) ->
20700 // (select (setcc x, y, cc), 1.0, 0.0)
20701 if (N0.getOpcode() == ISD::ZERO_EXTEND &&
20702 N0.getOperand(i: 0).getOpcode() == ISD::SETCC && !VT.isVector() &&
20703 (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::ConstantFP, VT)))
20704 return DAG.getSelect(DL, VT, Cond: N0.getOperand(i: 0),
20705 LHS: DAG.getConstantFP(Val: 1.0, DL, VT),
20706 RHS: DAG.getConstantFP(Val: 0.0, DL, VT));
20707
20708 if (SDValue FTrunc = foldFPToIntToFP(N, DL, DAG, TLI))
20709 return FTrunc;
20710
20711 // fold (sint_to_fp (trunc nsw x)) -> (sint_to_fp x)
20712 if (N0.getOpcode() == ISD::TRUNCATE && N0->getFlags().hasNoSignedWrap() &&
20713 TLI.isTypeDesirableForOp(ISD::SINT_TO_FP,
20714 VT: N0.getOperand(i: 0).getValueType()))
20715 return DAG.getNode(Opcode: ISD::SINT_TO_FP, DL, VT, Operand: N0.getOperand(i: 0));
20716
20717 return SDValue();
20718}
20719
20720SDValue DAGCombiner::visitUINT_TO_FP(SDNode *N) {
20721 SDValue N0 = N->getOperand(Num: 0);
20722 EVT VT = N->getValueType(ResNo: 0);
20723 EVT OpVT = N0.getValueType();
20724 SDLoc DL(N);
20725
20726 // [us]itofp(undef) = 0, because the result value is bounded.
20727 if (N0.isUndef())
20728 return DAG.getConstantFP(Val: 0.0, DL, VT);
20729
20730 // fold (uint_to_fp c1) -> c1fp
20731 // ...but only if the target supports immediate floating-point values
20732 if ((!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::ConstantFP, VT)))
20733 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::UINT_TO_FP, DL, VT, Ops: {N0}))
20734 return C;
20735
20736 // If the input is a legal type, and UINT_TO_FP is not legal on this target,
20737 // but SINT_TO_FP is legal on this target, try to convert.
20738 if (!hasOperation(Opcode: ISD::UINT_TO_FP, VT: OpVT) &&
20739 hasOperation(Opcode: ISD::SINT_TO_FP, VT: OpVT)) {
20740 // If the sign bit is known to be zero, we can change this to SINT_TO_FP.
20741 if (DAG.SignBitIsZero(Op: N0))
20742 return DAG.getNode(Opcode: ISD::SINT_TO_FP, DL, VT, Operand: N0);
20743 }
20744
20745 // fold (uint_to_fp (setcc x, y, cc)) -> (select (setcc x, y, cc), 1.0, 0.0)
20746 if (N0.getOpcode() == ISD::SETCC && !VT.isVector() &&
20747 (!LegalOperations || TLI.isOperationLegalOrCustom(Op: ISD::ConstantFP, VT)))
20748 return DAG.getSelect(DL, VT, Cond: N0, LHS: DAG.getConstantFP(Val: 1.0, DL, VT),
20749 RHS: DAG.getConstantFP(Val: 0.0, DL, VT));
20750
20751 if (SDValue FTrunc = foldFPToIntToFP(N, DL, DAG, TLI))
20752 return FTrunc;
20753
20754 // fold (uint_to_fp (trunc nuw x)) -> (uint_to_fp x)
20755 if (N0.getOpcode() == ISD::TRUNCATE && N0->getFlags().hasNoUnsignedWrap() &&
20756 TLI.isTypeDesirableForOp(ISD::UINT_TO_FP,
20757 VT: N0.getOperand(i: 0).getValueType()))
20758 return DAG.getNode(Opcode: ISD::UINT_TO_FP, DL, VT, Operand: N0.getOperand(i: 0));
20759
20760 return SDValue();
20761}
20762
20763// Fold (fp_to_{s/u}int ({s/u}int_to_fpx)) -> zext x, sext x, trunc x, or x
20764static SDValue FoldIntToFPToInt(SDNode *N, const SDLoc &DL, SelectionDAG &DAG) {
20765 SDValue N0 = N->getOperand(Num: 0);
20766 EVT VT = N->getValueType(ResNo: 0);
20767
20768 if (N0.getOpcode() != ISD::UINT_TO_FP && N0.getOpcode() != ISD::SINT_TO_FP)
20769 return SDValue();
20770
20771 SDValue Src = N0.getOperand(i: 0);
20772 EVT SrcVT = Src.getValueType();
20773 bool IsInputSigned = N0.getOpcode() == ISD::SINT_TO_FP;
20774 bool IsOutputSigned = N->getOpcode() == ISD::FP_TO_SINT;
20775
20776 // We can safely assume the conversion won't overflow the output range,
20777 // because (for example) (uint8_t)18293.f is undefined behavior.
20778
20779 // Since we can assume the conversion won't overflow, our decision as to
20780 // whether the input will fit in the float should depend on the minimum
20781 // of the input range and output range.
20782
20783 // This means this is also safe for a signed input and unsigned output, since
20784 // a negative input would lead to undefined behavior.
20785 unsigned InputSize = (int)SrcVT.getScalarSizeInBits() - IsInputSigned;
20786 unsigned OutputSize = (int)VT.getScalarSizeInBits();
20787 unsigned ActualSize = std::min(a: InputSize, b: OutputSize);
20788 const fltSemantics &Sem = N0.getValueType().getFltSemantics();
20789
20790 // We can only fold away the float conversion if the input range can be
20791 // represented exactly in the float range.
20792 if (APFloat::semanticsPrecision(Sem) >= ActualSize) {
20793 if (VT.getScalarSizeInBits() > SrcVT.getScalarSizeInBits()) {
20794 unsigned ExtOp =
20795 IsInputSigned && IsOutputSigned ? ISD::SIGN_EXTEND : ISD::ZERO_EXTEND;
20796 return DAG.getNode(Opcode: ExtOp, DL, VT, Operand: Src);
20797 }
20798 if (VT.getScalarSizeInBits() < SrcVT.getScalarSizeInBits())
20799 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT, Operand: Src);
20800 return DAG.getBitcast(VT, V: Src);
20801 }
20802 return SDValue();
20803}
20804
20805SDValue DAGCombiner::visitFP_TO_SINT(SDNode *N) {
20806 SDValue N0 = N->getOperand(Num: 0);
20807 EVT VT = N->getValueType(ResNo: 0);
20808 SDLoc DL(N);
20809
20810 // fold (fp_to_sint undef) -> undef
20811 if (SDValue R = propagateUnaryUndef(DAG, N0, VT))
20812 return R;
20813
20814 // fold (fp_to_sint c1fp) -> c1
20815 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FP_TO_SINT, DL, VT, Ops: {N0}))
20816 return C;
20817
20818 return FoldIntToFPToInt(N, DL, DAG);
20819}
20820
20821SDValue DAGCombiner::visitFP_TO_UINT(SDNode *N) {
20822 SDValue N0 = N->getOperand(Num: 0);
20823 EVT VT = N->getValueType(ResNo: 0);
20824 SDLoc DL(N);
20825
20826 // fold (fp_to_uint undef) -> undef
20827 if (SDValue R = propagateUnaryUndef(DAG, N0, VT))
20828 return R;
20829
20830 // fold (fp_to_uint c1fp) -> c1
20831 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FP_TO_UINT, DL, VT, Ops: {N0}))
20832 return C;
20833
20834 return FoldIntToFPToInt(N, DL, DAG);
20835}
20836
20837SDValue DAGCombiner::visitXROUND(SDNode *N) {
20838 SDValue N0 = N->getOperand(Num: 0);
20839 EVT VT = N->getValueType(ResNo: 0);
20840
20841 // fold (lrint|llrint undef) -> undef
20842 // fold (lround|llround undef) -> undef
20843 if (SDValue R = propagateUnaryUndef(DAG, N0, VT))
20844 return R;
20845
20846 // fold (lrint|llrint c1fp) -> c1
20847 // fold (lround|llround c1fp) -> c1
20848 if (SDValue C =
20849 DAG.FoldConstantArithmetic(Opcode: N->getOpcode(), DL: SDLoc(N), VT, Ops: {N0}))
20850 return C;
20851
20852 return SDValue();
20853}
20854
20855SDValue DAGCombiner::visitFP_ROUND(SDNode *N) {
20856 SDValue N0 = N->getOperand(Num: 0);
20857 SDValue N1 = N->getOperand(Num: 1);
20858 EVT VT = N->getValueType(ResNo: 0);
20859 SDLoc DL(N);
20860
20861 // fold (fp_round c1fp) -> c1fp
20862 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FP_ROUND, DL, VT, Ops: {N0, N1}))
20863 return C;
20864
20865 // fold (fp_round (fp_extend x)) -> x
20866 if (N0.getOpcode() == ISD::FP_EXTEND && VT == N0.getOperand(i: 0).getValueType())
20867 return N0.getOperand(i: 0);
20868
20869 // fold (fp_round (fp_round x)) -> (fp_round x)
20870 if (N0.getOpcode() == ISD::FP_ROUND) {
20871 const bool NIsTrunc = N->getConstantOperandVal(Num: 1) == 1;
20872 const bool N0IsTrunc = N0.getConstantOperandVal(i: 1) == 1;
20873
20874 // Avoid folding legal fp_rounds into non-legal ones.
20875 if (!hasOperation(Opcode: ISD::FP_ROUND, VT))
20876 return SDValue();
20877
20878 // Skip this folding if it results in an fp_round from f80 to f16.
20879 //
20880 // f80 to f16 always generates an expensive (and as yet, unimplemented)
20881 // libcall to __truncxfhf2 instead of selecting native f16 conversion
20882 // instructions from f32 or f64. Moreover, the first (value-preserving)
20883 // fp_round from f80 to either f32 or f64 may become a NOP in platforms like
20884 // x86.
20885 if (N0.getOperand(i: 0).getValueType() == MVT::f80 && VT == MVT::f16)
20886 return SDValue();
20887
20888 // If the first fp_round isn't a value preserving truncation, it might
20889 // introduce a tie in the second fp_round, that wouldn't occur in the
20890 // single-step fp_round we want to fold to.
20891 // In other words, double rounding isn't the same as rounding.
20892 // Also, this is a value preserving truncation iff both fp_round's are.
20893 if ((N->getFlags().hasAllowContract() &&
20894 N0->getFlags().hasAllowContract()) ||
20895 N0IsTrunc)
20896 return DAG.getNode(
20897 Opcode: ISD::FP_ROUND, DL, VT, N1: N0.getOperand(i: 0),
20898 N2: DAG.getIntPtrConstant(Val: NIsTrunc && N0IsTrunc, DL, /*isTarget=*/true));
20899 }
20900
20901 // fold (fp_round (copysign X, Y)) -> (copysign (fp_round X), Y)
20902 // Note: From a legality perspective, this is a two step transform. First,
20903 // we duplicate the fp_round to the arguments of the copysign, then we
20904 // eliminate the fp_round on Y. The second step requires an additional
20905 // predicate to match the implementation above.
20906 if (N0.getOpcode() == ISD::FCOPYSIGN && N0->hasOneUse() &&
20907 CanCombineFCOPYSIGN_EXTEND_ROUND(XTy: VT,
20908 YTy: N0.getValueType())) {
20909 SDValue Tmp = DAG.getNode(Opcode: ISD::FP_ROUND, DL: SDLoc(N0), VT,
20910 N1: N0.getOperand(i: 0), N2: N1);
20911 AddToWorklist(N: Tmp.getNode());
20912 return DAG.getNode(Opcode: ISD::FCOPYSIGN, DL, VT, N1: Tmp, N2: N0.getOperand(i: 1));
20913 }
20914
20915 if (SDValue NewVSel = matchVSelectOpSizesWithSetCC(Cast: N))
20916 return NewVSel;
20917
20918 return SDValue();
20919}
20920
20921// Eliminate a floating-point widening of a narrowed value if the fast math
20922// flags allow it.
20923static SDValue eliminateFPCastPair(SDNode *N) {
20924 SDValue N0 = N->getOperand(Num: 0);
20925 EVT VT = N->getValueType(ResNo: 0);
20926
20927 unsigned NarrowingOp;
20928 switch (N->getOpcode()) {
20929 case ISD::FP16_TO_FP:
20930 NarrowingOp = ISD::FP_TO_FP16;
20931 break;
20932 case ISD::BF16_TO_FP:
20933 NarrowingOp = ISD::FP_TO_BF16;
20934 break;
20935 case ISD::FP_EXTEND:
20936 NarrowingOp = ISD::FP_ROUND;
20937 break;
20938 default:
20939 llvm_unreachable("Expected widening FP cast");
20940 }
20941
20942 if (N0.getOpcode() == NarrowingOp && N0.getOperand(i: 0).getValueType() == VT) {
20943 const SDNodeFlags NarrowFlags = N0->getFlags();
20944 const SDNodeFlags WidenFlags = N->getFlags();
20945 // Narrowing can introduce inf and change the encoding of a nan, so the
20946 // widen must have the nnan and ninf flags to indicate that we don't need to
20947 // care about that. We are also removing a rounding step, and that requires
20948 // both the narrow and widen to allow contraction.
20949 if (WidenFlags.hasNoNaNs() && WidenFlags.hasNoInfs() &&
20950 NarrowFlags.hasAllowContract() && WidenFlags.hasAllowContract()) {
20951 return N0.getOperand(i: 0);
20952 }
20953 }
20954
20955 return SDValue();
20956}
20957
20958SDValue DAGCombiner::visitFP_EXTEND(SDNode *N) {
20959 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
20960 SDValue N0 = N->getOperand(Num: 0);
20961 EVT VT = N->getValueType(ResNo: 0);
20962 SDLoc DL(N);
20963
20964 if (VT.isVector())
20965 if (SDValue FoldedVOp = SimplifyVCastOp(N, DL))
20966 return FoldedVOp;
20967
20968 // If this is fp_round(fpextend), don't fold it, allow ourselves to be folded.
20969 if (N->hasOneUse() && N->user_begin()->getOpcode() == ISD::FP_ROUND)
20970 return SDValue();
20971
20972 // fold (fp_extend c1fp) -> c1fp
20973 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FP_EXTEND, DL, VT, Ops: {N0}))
20974 return C;
20975
20976 // fold (fp_extend (fp16_to_fp op)) -> (fp16_to_fp op)
20977 if (N0.getOpcode() == ISD::FP16_TO_FP &&
20978 TLI.getOperationAction(Op: ISD::FP16_TO_FP, VT) == TargetLowering::Legal)
20979 return DAG.getNode(Opcode: ISD::FP16_TO_FP, DL, VT, Operand: N0.getOperand(i: 0));
20980
20981 // Turn fp_extend(fp_round(X, 1)) -> x since the fp_round doesn't affect the
20982 // value of X.
20983 if (N0.getOpcode() == ISD::FP_ROUND && N0.getConstantOperandVal(i: 1) == 1) {
20984 SDValue In = N0.getOperand(i: 0);
20985 if (In.getValueType() == VT) return In;
20986 if (VT.bitsLT(VT: In.getValueType()))
20987 return DAG.getNode(Opcode: ISD::FP_ROUND, DL, VT, N1: In, N2: N0.getOperand(i: 1));
20988 return DAG.getNode(Opcode: ISD::FP_EXTEND, DL, VT, Operand: In);
20989 }
20990
20991 // fold (fpext (load x)) -> (fpext (fptrunc (extload x)))
20992 if (ISD::isNormalLoad(N: N0.getNode()) && N0.hasOneUse()) {
20993 LoadSDNode *LN0 = cast<LoadSDNode>(Val&: N0);
20994 if (TLI.isLoadLegalOrCustom(ValVT: VT, MemVT: N0.getValueType(), Alignment: LN0->getAlign(),
20995 AddrSpace: LN0->getAddressSpace(), ExtType: ISD::EXTLOAD, Atomic: false)) {
20996 SDValue ExtLoad = DAG.getExtLoad(ExtType: ISD::EXTLOAD, dl: DL, VT, Chain: LN0->getChain(),
20997 Ptr: LN0->getBasePtr(), MemVT: N0.getValueType(),
20998 MMO: LN0->getMemOperand());
20999 CombineTo(N, Res: ExtLoad);
21000 CombineTo(
21001 N: N0.getNode(),
21002 Res0: DAG.getNode(Opcode: ISD::FP_ROUND, DL: SDLoc(N0), VT: N0.getValueType(), N1: ExtLoad,
21003 N2: DAG.getIntPtrConstant(Val: 1, DL: SDLoc(N0), /*isTarget=*/true)),
21004 Res1: ExtLoad.getValue(R: 1));
21005 return SDValue(N, 0); // Return N so it doesn't get rechecked!
21006 }
21007 }
21008
21009 if (SDValue NewVSel = matchVSelectOpSizesWithSetCC(Cast: N))
21010 return NewVSel;
21011
21012 if (SDValue CastEliminated = eliminateFPCastPair(N))
21013 return CastEliminated;
21014
21015 return SDValue();
21016}
21017
21018SDValue DAGCombiner::visitFCEIL(SDNode *N) {
21019 SDValue N0 = N->getOperand(Num: 0);
21020 EVT VT = N->getValueType(ResNo: 0);
21021
21022 // fold (fceil c1) -> fceil(c1)
21023 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FCEIL, DL: SDLoc(N), VT, Ops: {N0}))
21024 return C;
21025
21026 return SDValue();
21027}
21028
21029SDValue DAGCombiner::visitFTRUNC(SDNode *N) {
21030 SDValue N0 = N->getOperand(Num: 0);
21031 EVT VT = N->getValueType(ResNo: 0);
21032
21033 // fold (ftrunc c1) -> ftrunc(c1)
21034 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FTRUNC, DL: SDLoc(N), VT, Ops: {N0}))
21035 return C;
21036
21037 // fold ftrunc (known rounded int x) -> x
21038 // ftrunc is a part of fptosi/fptoui expansion on some targets, so this is
21039 // likely to be generated to extract integer from a rounded floating value.
21040 switch (N0.getOpcode()) {
21041 default: break;
21042 case ISD::FRINT:
21043 case ISD::FTRUNC:
21044 case ISD::FNEARBYINT:
21045 case ISD::FROUND:
21046 case ISD::FROUNDEVEN:
21047 case ISD::FFLOOR:
21048 case ISD::FCEIL:
21049 return N0;
21050 }
21051
21052 return SDValue();
21053}
21054
21055SDValue DAGCombiner::visitFFREXP(SDNode *N) {
21056 SDValue N0 = N->getOperand(Num: 0);
21057
21058 // fold (ffrexp c1) -> ffrexp(c1)
21059 if (DAG.isConstantFPBuildVectorOrConstantFP(N: N0))
21060 return DAG.getNode(Opcode: ISD::FFREXP, DL: SDLoc(N), VTList: N->getVTList(), N: N0);
21061 return SDValue();
21062}
21063
21064SDValue DAGCombiner::visitFFLOOR(SDNode *N) {
21065 SDValue N0 = N->getOperand(Num: 0);
21066 EVT VT = N->getValueType(ResNo: 0);
21067
21068 // fold (ffloor c1) -> ffloor(c1)
21069 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FFLOOR, DL: SDLoc(N), VT, Ops: {N0}))
21070 return C;
21071
21072 return SDValue();
21073}
21074
21075SDValue DAGCombiner::visitFNEG(SDNode *N) {
21076 SDValue N0 = N->getOperand(Num: 0);
21077 EVT VT = N->getValueType(ResNo: 0);
21078 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
21079
21080 // Constant fold FNEG.
21081 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FNEG, DL: SDLoc(N), VT, Ops: {N0}))
21082 return C;
21083
21084 if (SDValue NegN0 =
21085 TLI.getNegatedExpression(Op: N0, DAG, LegalOps: LegalOperations, OptForSize: ForCodeSize))
21086 return NegN0;
21087
21088 // -(X-Y) -> (Y-X) is unsafe because when X==Y, -0.0 != +0.0
21089 // FIXME: This is duplicated in getNegatibleCost, but getNegatibleCost doesn't
21090 // know it was called from a context with a nsz flag if the input fsub does
21091 // not.
21092 if (N0.getOpcode() == ISD::FSUB && N->getFlags().hasNoSignedZeros() &&
21093 N0.hasOneUse()) {
21094 return DAG.getNode(Opcode: ISD::FSUB, DL: SDLoc(N), VT, N1: N0.getOperand(i: 1),
21095 N2: N0.getOperand(i: 0));
21096 }
21097
21098 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
21099 return SDValue(N, 0);
21100
21101 if (SDValue Cast = foldSignChangeInBitcast(N))
21102 return Cast;
21103
21104 return SDValue();
21105}
21106
21107SDValue DAGCombiner::visitFMinMax(SDNode *N) {
21108 SDValue N0 = N->getOperand(Num: 0);
21109 SDValue N1 = N->getOperand(Num: 1);
21110 EVT VT = N->getValueType(ResNo: 0);
21111 const SDNodeFlags Flags = N->getFlags();
21112 unsigned Opc = N->getOpcode();
21113 bool PropAllNaNsToQNaNs = Opc == ISD::FMINIMUM || Opc == ISD::FMAXIMUM;
21114 bool PropOnlySNaNsToQNaNs = Opc == ISD::FMINNUM || Opc == ISD::FMAXNUM;
21115 bool IsMin =
21116 Opc == ISD::FMINNUM || Opc == ISD::FMINIMUM || Opc == ISD::FMINIMUMNUM;
21117 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
21118
21119 // Constant fold.
21120 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: Opc, DL: SDLoc(N), VT, Ops: {N0, N1}))
21121 return C;
21122
21123 // Canonicalize to constant on RHS.
21124 if (DAG.isConstantFPBuildVectorOrConstantFP(N: N0) &&
21125 !DAG.isConstantFPBuildVectorOrConstantFP(N: N1))
21126 return DAG.getNode(Opcode: N->getOpcode(), DL: SDLoc(N), VT, N1, N2: N0);
21127
21128 if (const ConstantFPSDNode *N1CFP = isConstOrConstSplatFP(N: N1)) {
21129 const APFloat &AF = N1CFP->getValueAPF();
21130
21131 // minnum(X, qnan) -> X
21132 // maxnum(X, qnan) -> X
21133 // minnum(X, snan) -> qnan
21134 // maxnum(X, snan) -> qnan
21135 // minimum(X, nan) -> qnan
21136 // maximum(X, nan) -> qnan
21137 // minimumnum(X, nan) -> X
21138 // maximumnum(X, nan) -> X
21139 if (AF.isNaN()) {
21140 if (PropAllNaNsToQNaNs || (AF.isSignaling() && PropOnlySNaNsToQNaNs)) {
21141 if (AF.isSignaling())
21142 return DAG.getConstantFP(Val: AF.makeQuiet(), DL: SDLoc(N), VT);
21143 return N->getOperand(Num: 1);
21144 }
21145 return N->getOperand(Num: 0);
21146 }
21147
21148 // In the following folds, inf can be replaced with the largest finite
21149 // float, if the ninf flag is set.
21150 if (AF.isInfinity() || (Flags.hasNoInfs() && AF.isLargest())) {
21151 // minnum(X, -inf) -> -inf (ignoring sNaN -> qNaN propagation)
21152 // maxnum(X, +inf) -> +inf (ignoring sNaN -> qNaN propagation)
21153 // minimum(X, -inf) -> -inf if nnan
21154 // maximum(X, +inf) -> +inf if nnan
21155 // minimumnum(X, -inf) -> -inf
21156 // maximumnum(X, +inf) -> +inf
21157 if (IsMin == AF.isNegative() &&
21158 (!PropAllNaNsToQNaNs || Flags.hasNoNaNs()))
21159 return N->getOperand(Num: 1);
21160
21161 // minnum(X, +inf) -> X if nnan
21162 // maxnum(X, -inf) -> X if nnan
21163 // minimum(X, +inf) -> X (ignoring quieting of sNaNs)
21164 // maximum(X, -inf) -> X (ignoring quieting of sNaNs)
21165 // minimumnum(X, +inf) -> X if nnan
21166 // maximumnum(X, -inf) -> X if nnan
21167 if (IsMin != AF.isNegative() && (PropAllNaNsToQNaNs || Flags.hasNoNaNs()))
21168 return N->getOperand(Num: 0);
21169 }
21170 }
21171
21172 unsigned ReduceOpc;
21173 if (PropAllNaNsToQNaNs)
21174 ReduceOpc = IsMin ? ISD::VECREDUCE_FMINIMUM : ISD::VECREDUCE_FMAXIMUM;
21175 else if (PropOnlySNaNsToQNaNs)
21176 ReduceOpc = IsMin ? ISD::VECREDUCE_FMIN : ISD::VECREDUCE_FMAX;
21177 else
21178 ReduceOpc = IsMin ? ISD::VECREDUCE_FMINIMUMNUM : ISD::VECREDUCE_FMAXIMUMNUM;
21179
21180 if (SDValue SD =
21181 reassociateReduction(RedOpc: ReduceOpc, Opc, DL: SDLoc(N), VT, N0, N1, Flags))
21182 return SD;
21183
21184 return SDValue();
21185}
21186
21187SDValue DAGCombiner::visitFABS(SDNode *N) {
21188 SDValue N0 = N->getOperand(Num: 0);
21189 EVT VT = N->getValueType(ResNo: 0);
21190 SDLoc DL(N);
21191
21192 // fold (fabs c1) -> fabs(c1)
21193 if (SDValue C = DAG.FoldConstantArithmetic(Opcode: ISD::FABS, DL, VT, Ops: {N0}))
21194 return C;
21195
21196 if (SimplifyDemandedBits(Op: SDValue(N, 0)))
21197 return SDValue(N, 0);
21198
21199 if (SDValue Cast = foldSignChangeInBitcast(N))
21200 return Cast;
21201
21202 return SDValue();
21203}
21204
21205SDValue DAGCombiner::visitBRCOND(SDNode *N) {
21206 SDValue Chain = N->getOperand(Num: 0);
21207 SDValue N1 = N->getOperand(Num: 1);
21208 SDValue N2 = N->getOperand(Num: 2);
21209
21210 // BRCOND(FREEZE(cond)) is equivalent to BRCOND(cond) (both are
21211 // nondeterministic jumps).
21212 if (N1->getOpcode() == ISD::FREEZE && N1.hasOneUse()) {
21213 return DAG.getNode(Opcode: ISD::BRCOND, DL: SDLoc(N), VT: MVT::Other, N1: Chain,
21214 N2: N1->getOperand(Num: 0), N3: N2, Flags: N->getFlags());
21215 }
21216
21217 // Variant of the previous fold where there is a SETCC in between:
21218 // BRCOND(SETCC(FREEZE(X), CONST, Cond))
21219 // =>
21220 // BRCOND(FREEZE(SETCC(X, CONST, Cond)))
21221 // =>
21222 // BRCOND(SETCC(X, CONST, Cond))
21223 // This is correct if FREEZE(X) has one use and SETCC(FREEZE(X), CONST, Cond)
21224 // isn't equivalent to true or false.
21225 // For example, SETCC(FREEZE(X), -128, SETULT) cannot be folded to
21226 // FREEZE(SETCC(X, -128, SETULT)) because X can be poison.
21227 if (N1->getOpcode() == ISD::SETCC && N1.hasOneUse()) {
21228 SDValue S0 = N1->getOperand(Num: 0), S1 = N1->getOperand(Num: 1);
21229 ISD::CondCode Cond = cast<CondCodeSDNode>(Val: N1->getOperand(Num: 2))->get();
21230 ConstantSDNode *S0C = dyn_cast<ConstantSDNode>(Val&: S0);
21231 ConstantSDNode *S1C = dyn_cast<ConstantSDNode>(Val&: S1);
21232 bool Updated = false;
21233
21234 // Is 'X Cond C' always true or false?
21235 auto IsAlwaysTrueOrFalse = [](ISD::CondCode Cond, ConstantSDNode *C) {
21236 bool False = (Cond == ISD::SETULT && C->isZero()) ||
21237 (Cond == ISD::SETLT && C->isMinSignedValue()) ||
21238 (Cond == ISD::SETUGT && C->isAllOnes()) ||
21239 (Cond == ISD::SETGT && C->isMaxSignedValue());
21240 bool True = (Cond == ISD::SETULE && C->isAllOnes()) ||
21241 (Cond == ISD::SETLE && C->isMaxSignedValue()) ||
21242 (Cond == ISD::SETUGE && C->isZero()) ||
21243 (Cond == ISD::SETGE && C->isMinSignedValue());
21244 return True || False;
21245 };
21246
21247 if (S0->getOpcode() == ISD::FREEZE && S0.hasOneUse() && S1C) {
21248 if (!IsAlwaysTrueOrFalse(Cond, S1C)) {
21249 S0 = S0->getOperand(Num: 0);
21250 Updated = true;
21251 }
21252 }
21253 if (S1->getOpcode() == ISD::FREEZE && S1.hasOneUse() && S0C) {
21254 if (!IsAlwaysTrueOrFalse(ISD::getSetCCSwappedOperands(Operation: Cond), S0C)) {
21255 S1 = S1->getOperand(Num: 0);
21256 Updated = true;
21257 }
21258 }
21259
21260 if (Updated)
21261 return DAG.getNode(
21262 Opcode: ISD::BRCOND, DL: SDLoc(N), VT: MVT::Other, N1: Chain,
21263 N2: DAG.getSetCC(DL: SDLoc(N1), VT: N1->getValueType(ResNo: 0), LHS: S0, RHS: S1, Cond), N3: N2,
21264 Flags: N->getFlags());
21265 }
21266
21267 // If N is a constant we could fold this into a fallthrough or unconditional
21268 // branch. However that doesn't happen very often in normal code, because
21269 // Instcombine/SimplifyCFG should have handled the available opportunities.
21270 // If we did this folding here, it would be necessary to update the
21271 // MachineBasicBlock CFG, which is awkward.
21272
21273 // fold a brcond with a setcc condition into a BR_CC node if BR_CC is legal
21274 // on the target, also copy fast math flags.
21275 if (N1.getOpcode() == ISD::SETCC &&
21276 TLI.isOperationLegalOrCustom(Op: ISD::BR_CC,
21277 VT: N1.getOperand(i: 0).getValueType())) {
21278 return DAG.getNode(Opcode: ISD::BR_CC, DL: SDLoc(N), VT: MVT::Other, N1: Chain,
21279 N2: N1.getOperand(i: 2), N3: N1.getOperand(i: 0), N4: N1.getOperand(i: 1), N5: N2,
21280 Flags: N1->getFlags());
21281 }
21282
21283 if (N1.hasOneUse()) {
21284 // rebuildSetCC calls visitXor which may change the Chain when there is a
21285 // STRICT_FSETCC/STRICT_FSETCCS involved. Use a handle to track changes.
21286 HandleSDNode ChainHandle(Chain);
21287 if (SDValue NewN1 = rebuildSetCC(N: N1))
21288 return DAG.getNode(Opcode: ISD::BRCOND, DL: SDLoc(N), VT: MVT::Other,
21289 N1: ChainHandle.getValue(), N2: NewN1, N3: N2, Flags: N->getFlags());
21290 }
21291
21292 return SDValue();
21293}
21294
21295SDValue DAGCombiner::rebuildSetCC(SDValue N) {
21296 if (N.getOpcode() == ISD::SRL ||
21297 (N.getOpcode() == ISD::TRUNCATE &&
21298 (N.getOperand(i: 0).hasOneUse() &&
21299 N.getOperand(i: 0).getOpcode() == ISD::SRL))) {
21300 // Look pass the truncate.
21301 if (N.getOpcode() == ISD::TRUNCATE)
21302 N = N.getOperand(i: 0);
21303
21304 // Match this pattern so that we can generate simpler code:
21305 //
21306 // %a = ...
21307 // %b = and i32 %a, 2
21308 // %c = srl i32 %b, 1
21309 // brcond i32 %c ...
21310 //
21311 // into
21312 //
21313 // %a = ...
21314 // %b = and i32 %a, 2
21315 // %c = setcc eq %b, 0
21316 // brcond %c ...
21317 //
21318 // This applies only when the AND constant value has one bit set and the
21319 // SRL constant is equal to the log2 of the AND constant. The back-end is
21320 // smart enough to convert the result into a TEST/JMP sequence.
21321 SDValue Op0 = N.getOperand(i: 0);
21322 SDValue Op1 = N.getOperand(i: 1);
21323
21324 if (Op0.getOpcode() == ISD::AND && Op1.getOpcode() == ISD::Constant) {
21325 SDValue AndOp1 = Op0.getOperand(i: 1);
21326
21327 if (AndOp1.getOpcode() == ISD::Constant) {
21328 const APInt &AndConst = AndOp1->getAsAPIntVal();
21329
21330 if (AndConst.isPowerOf2() &&
21331 Op1->getAsAPIntVal() == AndConst.logBase2()) {
21332 SDLoc DL(N);
21333 return DAG.getSetCC(DL, VT: getSetCCResultType(VT: Op0.getValueType()),
21334 LHS: Op0, RHS: DAG.getConstant(Val: 0, DL, VT: Op0.getValueType()),
21335 Cond: ISD::SETNE);
21336 }
21337 }
21338 }
21339 }
21340
21341 // Transform (brcond (xor x, y)) -> (brcond (setcc, x, y, ne))
21342 // Transform (brcond (xor (xor x, y), -1)) -> (brcond (setcc, x, y, eq))
21343 if (N.getOpcode() == ISD::XOR) {
21344 // Because we may call this on a speculatively constructed
21345 // SimplifiedSetCC Node, we need to simplify this node first.
21346 // Ideally this should be folded into SimplifySetCC and not
21347 // here. For now, grab a handle to N so we don't lose it from
21348 // replacements interal to the visit.
21349 while (N.getOpcode() == ISD::XOR) {
21350 HandleSDNode XORHandle(N);
21351 SDValue Tmp = visitXOR(N: N.getNode());
21352 // No simplification done.
21353 if (!Tmp.getNode())
21354 break;
21355 // Returning N is form in-visit replacement that may invalidated
21356 // N. Grab value from Handle.
21357 if (Tmp.getNode() == N.getNode())
21358 N = XORHandle.getValue();
21359 else // Node simplified. Try simplifying again.
21360 N = Tmp;
21361 }
21362
21363 if (N.getOpcode() != ISD::XOR)
21364 return N;
21365
21366 SDValue Op0 = N->getOperand(Num: 0);
21367 SDValue Op1 = N->getOperand(Num: 1);
21368
21369 if (Op0.getOpcode() != ISD::SETCC && Op1.getOpcode() != ISD::SETCC) {
21370 bool Equal = false;
21371 // (brcond (xor (xor x, y), -1)) -> (brcond (setcc x, y, eq))
21372 if (isBitwiseNot(V: N) && Op0.hasOneUse() && Op0.getOpcode() == ISD::XOR &&
21373 Op0.getValueType() == MVT::i1) {
21374 N = Op0;
21375 Op0 = N->getOperand(Num: 0);
21376 Op1 = N->getOperand(Num: 1);
21377 Equal = true;
21378 }
21379
21380 EVT SetCCVT = N.getValueType();
21381 if (LegalTypes)
21382 SetCCVT = getSetCCResultType(VT: SetCCVT);
21383 // Replace the uses of XOR with SETCC. Note, avoid this transformation if
21384 // it would introduce illegal operations post-legalization as this can
21385 // result in infinite looping between converting xor->setcc here, and
21386 // expanding setcc->xor in LegalizeSetCCCondCode if requested.
21387 const ISD::CondCode CC = Equal ? ISD::SETEQ : ISD::SETNE;
21388 if (!LegalOperations || TLI.isCondCodeLegal(CC, VT: Op0.getSimpleValueType()))
21389 return DAG.getSetCC(DL: SDLoc(N), VT: SetCCVT, LHS: Op0, RHS: Op1, Cond: CC);
21390 }
21391 }
21392
21393 return SDValue();
21394}
21395
21396// Operand List for BR_CC: Chain, CondCC, CondLHS, CondRHS, DestBB.
21397//
21398SDValue DAGCombiner::visitBR_CC(SDNode *N) {
21399 CondCodeSDNode *CC = cast<CondCodeSDNode>(Val: N->getOperand(Num: 1));
21400 SDValue CondLHS = N->getOperand(Num: 2), CondRHS = N->getOperand(Num: 3);
21401
21402 // If N is a constant we could fold this into a fallthrough or unconditional
21403 // branch. However that doesn't happen very often in normal code, because
21404 // Instcombine/SimplifyCFG should have handled the available opportunities.
21405 // If we did this folding here, it would be necessary to update the
21406 // MachineBasicBlock CFG, which is awkward.
21407
21408 // Use SimplifySetCC to simplify SETCC's.
21409 SDValue Simp = SimplifySetCC(VT: getSetCCResultType(VT: CondLHS.getValueType()),
21410 N0: CondLHS, N1: CondRHS, Cond: CC->get(), DL: SDLoc(N),
21411 foldBooleans: false);
21412 if (Simp.getNode()) AddToWorklist(N: Simp.getNode());
21413
21414 // fold to a simpler setcc
21415 if (Simp.getNode() && Simp.getOpcode() == ISD::SETCC)
21416 return DAG.getNode(Opcode: ISD::BR_CC, DL: SDLoc(N), VT: MVT::Other,
21417 N1: N->getOperand(Num: 0), N2: Simp.getOperand(i: 2),
21418 N3: Simp.getOperand(i: 0), N4: Simp.getOperand(i: 1),
21419 N5: N->getOperand(Num: 4));
21420
21421 return SDValue();
21422}
21423
21424static bool getCombineLoadStoreParts(SDNode *N, unsigned Inc, unsigned Dec,
21425 bool &IsLoad, bool &IsMasked, SDValue &Ptr,
21426 const TargetLowering &TLI) {
21427 if (LoadSDNode *LD = dyn_cast<LoadSDNode>(Val: N)) {
21428 if (LD->isIndexed())
21429 return false;
21430 EVT VT = LD->getMemoryVT();
21431 if (!TLI.isIndexedLoadLegal(IdxMode: Inc, VT) && !TLI.isIndexedLoadLegal(IdxMode: Dec, VT))
21432 return false;
21433 Ptr = LD->getBasePtr();
21434 } else if (StoreSDNode *ST = dyn_cast<StoreSDNode>(Val: N)) {
21435 if (ST->isIndexed())
21436 return false;
21437 EVT VT = ST->getMemoryVT();
21438 if (!TLI.isIndexedStoreLegal(IdxMode: Inc, VT) && !TLI.isIndexedStoreLegal(IdxMode: Dec, VT))
21439 return false;
21440 Ptr = ST->getBasePtr();
21441 IsLoad = false;
21442 } else if (MaskedLoadSDNode *LD = dyn_cast<MaskedLoadSDNode>(Val: N)) {
21443 if (LD->isIndexed())
21444 return false;
21445 EVT VT = LD->getMemoryVT();
21446 if (!TLI.isIndexedMaskedLoadLegal(IdxMode: Inc, VT) &&
21447 !TLI.isIndexedMaskedLoadLegal(IdxMode: Dec, VT))
21448 return false;
21449 Ptr = LD->getBasePtr();
21450 IsMasked = true;
21451 } else if (MaskedStoreSDNode *ST = dyn_cast<MaskedStoreSDNode>(Val: N)) {
21452 if (ST->isIndexed())
21453 return false;
21454 EVT VT = ST->getMemoryVT();
21455 if (!TLI.isIndexedMaskedStoreLegal(IdxMode: Inc, VT) &&
21456 !TLI.isIndexedMaskedStoreLegal(IdxMode: Dec, VT))
21457 return false;
21458 Ptr = ST->getBasePtr();
21459 IsLoad = false;
21460 IsMasked = true;
21461 } else {
21462 return false;
21463 }
21464 return true;
21465}
21466
21467/// Try turning a load/store into a pre-indexed load/store when the base
21468/// pointer is an add or subtract and it has other uses besides the load/store.
21469/// After the transformation, the new indexed load/store has effectively folded
21470/// the add/subtract in and all of its other uses are redirected to the
21471/// new load/store.
21472bool DAGCombiner::CombineToPreIndexedLoadStore(SDNode *N) {
21473 if (Level < AfterLegalizeDAG)
21474 return false;
21475
21476 bool IsLoad = true;
21477 bool IsMasked = false;
21478 SDValue Ptr;
21479 if (!getCombineLoadStoreParts(N, Inc: ISD::PRE_INC, Dec: ISD::PRE_DEC, IsLoad, IsMasked,
21480 Ptr, TLI))
21481 return false;
21482
21483 // If the pointer is not an add/sub, or if it doesn't have multiple uses, bail
21484 // out. There is no reason to make this a preinc/predec.
21485 if ((Ptr.getOpcode() != ISD::ADD && Ptr.getOpcode() != ISD::SUB) ||
21486 Ptr->hasOneUse())
21487 return false;
21488
21489 // Ask the target to do addressing mode selection.
21490 SDValue BasePtr;
21491 SDValue Offset;
21492 ISD::MemIndexedMode AM = ISD::UNINDEXED;
21493 if (!TLI.getPreIndexedAddressParts(N, BasePtr, Offset, AM, DAG))
21494 return false;
21495
21496 // Backends without true r+i pre-indexed forms may need to pass a
21497 // constant base with a variable offset so that constant coercion
21498 // will work with the patterns in canonical form.
21499 bool Swapped = false;
21500 if (isa<ConstantSDNode>(Val: BasePtr)) {
21501 std::swap(a&: BasePtr, b&: Offset);
21502 Swapped = true;
21503 }
21504
21505 // Don't create a indexed load / store with zero offset.
21506 if (isNullConstant(V: Offset))
21507 return false;
21508
21509 // Try turning it into a pre-indexed load / store except when:
21510 // 1) The new base ptr is a frame index.
21511 // 2) If N is a store and the new base ptr is either the same as or is a
21512 // predecessor of the value being stored.
21513 // 3) Another use of old base ptr is a predecessor of N. If ptr is folded
21514 // that would create a cycle.
21515 // 4) All uses are load / store ops that use it as old base ptr.
21516
21517 // Check #1. Preinc'ing a frame index would require copying the stack pointer
21518 // (plus the implicit offset) to a register to preinc anyway.
21519 if (isa<FrameIndexSDNode>(Val: BasePtr) || isa<RegisterSDNode>(Val: BasePtr))
21520 return false;
21521
21522 // Check #2.
21523 if (!IsLoad) {
21524 SDValue Val = IsMasked ? cast<MaskedStoreSDNode>(Val: N)->getValue()
21525 : cast<StoreSDNode>(Val: N)->getValue();
21526
21527 // Would require a copy.
21528 if (Val == BasePtr)
21529 return false;
21530
21531 // Would create a cycle.
21532 if (Val == Ptr || Ptr->isPredecessorOf(N: Val.getNode()))
21533 return false;
21534 }
21535
21536 // Caches for hasPredecessorHelper.
21537 SmallPtrSet<const SDNode *, 32> Visited;
21538 SmallVector<const SDNode *, 16> Worklist;
21539 Worklist.push_back(Elt: N);
21540
21541 // If the offset is a constant, there may be other adds of constants that
21542 // can be folded with this one. We should do this to avoid having to keep
21543 // a copy of the original base pointer.
21544 SmallVector<SDNode *, 16> OtherUses;
21545 unsigned MaxSteps = SelectionDAG::getHasPredecessorMaxSteps();
21546 if (isa<ConstantSDNode>(Val: Offset))
21547 for (SDUse &Use : BasePtr->uses()) {
21548 // Skip the use that is Ptr and uses of other results from BasePtr's
21549 // node (important for nodes that return multiple results).
21550 if (Use.getUser() == Ptr.getNode() || Use != BasePtr)
21551 continue;
21552
21553 if (SDNode::hasPredecessorHelper(N: Use.getUser(), Visited, Worklist,
21554 MaxSteps))
21555 continue;
21556
21557 if (Use.getUser()->getOpcode() != ISD::ADD &&
21558 Use.getUser()->getOpcode() != ISD::SUB) {
21559 OtherUses.clear();
21560 break;
21561 }
21562
21563 SDValue Op1 = Use.getUser()->getOperand(Num: (Use.getOperandNo() + 1) & 1);
21564 if (!isa<ConstantSDNode>(Val: Op1)) {
21565 OtherUses.clear();
21566 break;
21567 }
21568
21569 // FIXME: In some cases, we can be smarter about this.
21570 if (Op1.getValueType() != Offset.getValueType()) {
21571 OtherUses.clear();
21572 break;
21573 }
21574
21575 OtherUses.push_back(Elt: Use.getUser());
21576 }
21577
21578 if (Swapped)
21579 std::swap(a&: BasePtr, b&: Offset);
21580
21581 // Now check for #3 and #4.
21582 bool RealUse = false;
21583
21584 for (SDNode *User : Ptr->users()) {
21585 if (User == N)
21586 continue;
21587 if (SDNode::hasPredecessorHelper(N: User, Visited, Worklist, MaxSteps))
21588 return false;
21589
21590 // If Ptr may be folded in addressing mode of other use, then it's
21591 // not profitable to do this transformation.
21592 if (!canFoldInAddressingMode(N: Ptr.getNode(), Use: User, DAG, TLI))
21593 RealUse = true;
21594 }
21595
21596 if (!RealUse)
21597 return false;
21598
21599 SDValue Result;
21600 if (!IsMasked) {
21601 if (IsLoad)
21602 Result = DAG.getIndexedLoad(OrigLoad: SDValue(N, 0), dl: SDLoc(N), Base: BasePtr, Offset, AM);
21603 else
21604 Result =
21605 DAG.getIndexedStore(OrigStore: SDValue(N, 0), dl: SDLoc(N), Base: BasePtr, Offset, AM);
21606 } else {
21607 if (IsLoad)
21608 Result = DAG.getIndexedMaskedLoad(OrigLoad: SDValue(N, 0), dl: SDLoc(N), Base: BasePtr,
21609 Offset, AM);
21610 else
21611 Result = DAG.getIndexedMaskedStore(OrigStore: SDValue(N, 0), dl: SDLoc(N), Base: BasePtr,
21612 Offset, AM);
21613 }
21614 ++PreIndexedNodes;
21615 ++NodesCombined;
21616 LLVM_DEBUG(dbgs() << "\nReplacing.4 "; N->dump(&DAG); dbgs() << "\nWith: ";
21617 Result.dump(&DAG); dbgs() << '\n');
21618 WorklistRemover DeadNodes(*this);
21619 if (IsLoad) {
21620 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 0), To: Result.getValue(R: 0));
21621 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 1), To: Result.getValue(R: 2));
21622 } else {
21623 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 0), To: Result.getValue(R: 1));
21624 }
21625
21626 // Finally, since the node is now dead, remove it from the graph.
21627 deleteAndRecombine(N);
21628
21629 if (Swapped)
21630 std::swap(a&: BasePtr, b&: Offset);
21631
21632 // Replace other uses of BasePtr that can be updated to use Ptr
21633 for (SDNode *OtherUse : OtherUses) {
21634 unsigned OffsetIdx = 1;
21635 if (OtherUse->getOperand(Num: OffsetIdx).getNode() == BasePtr.getNode())
21636 OffsetIdx = 0;
21637 assert(OtherUse->getOperand(!OffsetIdx).getNode() == BasePtr.getNode() &&
21638 "Expected BasePtr operand");
21639
21640 // We need to replace ptr0 in the following expression:
21641 // x0 * offset0 + y0 * ptr0 = t0
21642 // knowing that
21643 // x1 * offset1 + y1 * ptr0 = t1 (the indexed load/store)
21644 //
21645 // where x0, x1, y0 and y1 in {-1, 1} are given by the types of the
21646 // indexed load/store and the expression that needs to be re-written.
21647 //
21648 // Therefore, we have:
21649 // t0 = (x0 * offset0 - x1 * y0 * y1 *offset1) + (y0 * y1) * t1
21650
21651 auto *CN = cast<ConstantSDNode>(Val: OtherUse->getOperand(Num: OffsetIdx));
21652 const APInt &Offset0 = CN->getAPIntValue();
21653 const APInt &Offset1 = Offset->getAsAPIntVal();
21654 int X0 = (OtherUse->getOpcode() == ISD::SUB && OffsetIdx == 1) ? -1 : 1;
21655 int Y0 = (OtherUse->getOpcode() == ISD::SUB && OffsetIdx == 0) ? -1 : 1;
21656 int X1 = (AM == ISD::PRE_DEC && !Swapped) ? -1 : 1;
21657 int Y1 = (AM == ISD::PRE_DEC && Swapped) ? -1 : 1;
21658
21659 unsigned Opcode = (Y0 * Y1 < 0) ? ISD::SUB : ISD::ADD;
21660
21661 APInt CNV = Offset0;
21662 if (X0 < 0) CNV = -CNV;
21663 if (X1 * Y0 * Y1 < 0) CNV = CNV + Offset1;
21664 else CNV = CNV - Offset1;
21665
21666 SDLoc DL(OtherUse);
21667
21668 // We can now generate the new expression.
21669 SDValue NewOp1 = DAG.getConstant(Val: CNV, DL, VT: CN->getValueType(ResNo: 0));
21670 SDValue NewOp2 = Result.getValue(R: IsLoad ? 1 : 0);
21671
21672 SDValue NewUse =
21673 DAG.getNode(Opcode, DL, VT: OtherUse->getValueType(ResNo: 0), N1: NewOp1, N2: NewOp2);
21674 DAG.ReplaceAllUsesOfValueWith(From: SDValue(OtherUse, 0), To: NewUse);
21675 deleteAndRecombine(N: OtherUse);
21676 }
21677
21678 // Replace the uses of Ptr with uses of the updated base value.
21679 DAG.ReplaceAllUsesOfValueWith(From: Ptr, To: Result.getValue(R: IsLoad ? 1 : 0));
21680 deleteAndRecombine(N: Ptr.getNode());
21681 AddToWorklist(N: Result.getNode());
21682
21683 return true;
21684}
21685
21686static bool shouldCombineToPostInc(SDNode *N, SDValue Ptr, SDNode *PtrUse,
21687 SDValue &BasePtr, SDValue &Offset,
21688 ISD::MemIndexedMode &AM,
21689 SelectionDAG &DAG,
21690 const TargetLowering &TLI) {
21691 if (PtrUse == N ||
21692 (PtrUse->getOpcode() != ISD::ADD && PtrUse->getOpcode() != ISD::SUB))
21693 return false;
21694
21695 if (!TLI.getPostIndexedAddressParts(N, PtrUse, BasePtr, Offset, AM, DAG))
21696 return false;
21697
21698 // Don't create a indexed load / store with zero offset.
21699 if (isNullConstant(V: Offset))
21700 return false;
21701
21702 if (isa<FrameIndexSDNode>(Val: BasePtr) || isa<RegisterSDNode>(Val: BasePtr))
21703 return false;
21704
21705 SmallPtrSet<const SDNode *, 32> Visited;
21706 unsigned MaxSteps = SelectionDAG::getHasPredecessorMaxSteps();
21707 for (SDNode *User : BasePtr->users()) {
21708 if (User == Ptr.getNode())
21709 continue;
21710
21711 // No if there's a later user which could perform the index instead.
21712 if (isa<MemSDNode>(Val: User)) {
21713 bool IsLoad = true;
21714 bool IsMasked = false;
21715 SDValue OtherPtr;
21716 if (getCombineLoadStoreParts(N: User, Inc: ISD::POST_INC, Dec: ISD::POST_DEC, IsLoad,
21717 IsMasked, Ptr&: OtherPtr, TLI)) {
21718 SmallVector<const SDNode *, 2> Worklist;
21719 Worklist.push_back(Elt: User);
21720 if (SDNode::hasPredecessorHelper(N, Visited, Worklist, MaxSteps))
21721 return false;
21722 }
21723 }
21724
21725 // If all the uses are load / store addresses, then don't do the
21726 // transformation.
21727 if (User->getOpcode() == ISD::ADD || User->getOpcode() == ISD::SUB) {
21728 for (SDNode *UserUser : User->users())
21729 if (canFoldInAddressingMode(N: User, Use: UserUser, DAG, TLI))
21730 return false;
21731 }
21732 }
21733 return true;
21734}
21735
21736static SDNode *getPostIndexedLoadStoreOp(SDNode *N, bool &IsLoad,
21737 bool &IsMasked, SDValue &Ptr,
21738 SDValue &BasePtr, SDValue &Offset,
21739 ISD::MemIndexedMode &AM,
21740 SelectionDAG &DAG,
21741 const TargetLowering &TLI) {
21742 if (!getCombineLoadStoreParts(N, Inc: ISD::POST_INC, Dec: ISD::POST_DEC, IsLoad,
21743 IsMasked, Ptr, TLI) ||
21744 Ptr->hasOneUse())
21745 return nullptr;
21746
21747 // Try turning it into a post-indexed load / store except when
21748 // 1) All uses are load / store ops that use it as base ptr (and
21749 // it may be folded as addressing mmode).
21750 // 2) Op must be independent of N, i.e. Op is neither a predecessor
21751 // nor a successor of N. Otherwise, if Op is folded that would
21752 // create a cycle.
21753 unsigned MaxSteps = SelectionDAG::getHasPredecessorMaxSteps();
21754 for (SDUse &U : Ptr->uses()) {
21755 if (U.getResNo() != Ptr.getResNo())
21756 continue;
21757
21758 // Check for #1.
21759 SDNode *Op = U.getUser();
21760 if (!shouldCombineToPostInc(N, Ptr, PtrUse: Op, BasePtr, Offset, AM, DAG, TLI))
21761 continue;
21762
21763 // Check for #2.
21764 SmallPtrSet<const SDNode *, 32> Visited;
21765 SmallVector<const SDNode *, 8> Worklist;
21766 // Ptr is predecessor to both N and Op.
21767 Visited.insert(Ptr: Ptr.getNode());
21768 Worklist.push_back(Elt: N);
21769 Worklist.push_back(Elt: Op);
21770 if (!SDNode::hasPredecessorHelper(N, Visited, Worklist, MaxSteps) &&
21771 !SDNode::hasPredecessorHelper(N: Op, Visited, Worklist, MaxSteps))
21772 return Op;
21773 }
21774 return nullptr;
21775}
21776
21777/// Try to combine a load/store with a add/sub of the base pointer node into a
21778/// post-indexed load/store. The transformation folded the add/subtract into the
21779/// new indexed load/store effectively and all of its uses are redirected to the
21780/// new load/store.
21781bool DAGCombiner::CombineToPostIndexedLoadStore(SDNode *N) {
21782 if (Level < AfterLegalizeDAG)
21783 return false;
21784
21785 bool IsLoad = true;
21786 bool IsMasked = false;
21787 SDValue Ptr;
21788 SDValue BasePtr;
21789 SDValue Offset;
21790 ISD::MemIndexedMode AM = ISD::UNINDEXED;
21791 SDNode *Op = getPostIndexedLoadStoreOp(N, IsLoad, IsMasked, Ptr, BasePtr,
21792 Offset, AM, DAG, TLI);
21793 if (!Op)
21794 return false;
21795
21796 SDValue Result;
21797 if (!IsMasked)
21798 Result = IsLoad ? DAG.getIndexedLoad(OrigLoad: SDValue(N, 0), dl: SDLoc(N), Base: BasePtr,
21799 Offset, AM)
21800 : DAG.getIndexedStore(OrigStore: SDValue(N, 0), dl: SDLoc(N),
21801 Base: BasePtr, Offset, AM);
21802 else
21803 Result = IsLoad ? DAG.getIndexedMaskedLoad(OrigLoad: SDValue(N, 0), dl: SDLoc(N),
21804 Base: BasePtr, Offset, AM)
21805 : DAG.getIndexedMaskedStore(OrigStore: SDValue(N, 0), dl: SDLoc(N),
21806 Base: BasePtr, Offset, AM);
21807 ++PostIndexedNodes;
21808 ++NodesCombined;
21809 LLVM_DEBUG(dbgs() << "\nReplacing.5 "; N->dump(&DAG); dbgs() << "\nWith: ";
21810 Result.dump(&DAG); dbgs() << '\n');
21811 WorklistRemover DeadNodes(*this);
21812 if (IsLoad) {
21813 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 0), To: Result.getValue(R: 0));
21814 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 1), To: Result.getValue(R: 2));
21815 } else {
21816 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 0), To: Result.getValue(R: 1));
21817 }
21818
21819 // Finally, since the node is now dead, remove it from the graph.
21820 deleteAndRecombine(N);
21821
21822 // Replace the uses of Use with uses of the updated base value.
21823 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Op, 0),
21824 To: Result.getValue(R: IsLoad ? 1 : 0));
21825 deleteAndRecombine(N: Op);
21826 return true;
21827}
21828
21829/// Return the base-pointer arithmetic from an indexed \p LD.
21830SDValue DAGCombiner::SplitIndexingFromLoad(LoadSDNode *LD) {
21831 ISD::MemIndexedMode AM = LD->getAddressingMode();
21832 assert(AM != ISD::UNINDEXED);
21833 SDValue BP = LD->getOperand(Num: 1);
21834 SDValue Inc = LD->getOperand(Num: 2);
21835
21836 // Some backends use TargetConstants for load offsets, but don't expect
21837 // TargetConstants in general ADD nodes. We can convert these constants into
21838 // regular Constants (if the constant is not opaque).
21839 assert((Inc.getOpcode() != ISD::TargetConstant ||
21840 !cast<ConstantSDNode>(Inc)->isOpaque()) &&
21841 "Cannot split out indexing using opaque target constants");
21842 if (Inc.getOpcode() == ISD::TargetConstant) {
21843 ConstantSDNode *ConstInc = cast<ConstantSDNode>(Val&: Inc);
21844 Inc = DAG.getConstant(Val: *ConstInc->getConstantIntValue(), DL: SDLoc(Inc),
21845 VT: ConstInc->getValueType(ResNo: 0));
21846 }
21847
21848 unsigned Opc =
21849 (AM == ISD::PRE_INC || AM == ISD::POST_INC ? ISD::ADD : ISD::SUB);
21850 return DAG.getNode(Opcode: Opc, DL: SDLoc(LD), VT: BP.getSimpleValueType(), N1: BP, N2: Inc);
21851}
21852
21853static inline ElementCount numVectorEltsOrZero(EVT T) {
21854 return T.isVector() ? T.getVectorElementCount() : ElementCount::getFixed(MinVal: 0);
21855}
21856
21857bool DAGCombiner::getTruncatedStoreValue(StoreSDNode *ST, SDValue &Val) {
21858 EVT STType = Val.getValueType();
21859 EVT STMemType = ST->getMemoryVT();
21860 if (STType == STMemType)
21861 return true;
21862 if (isTypeLegal(VT: STMemType))
21863 return false; // fail.
21864 if (STType.isFloatingPoint() && STMemType.isFloatingPoint() &&
21865 TLI.isOperationLegal(Op: ISD::FTRUNC, VT: STMemType)) {
21866 Val = DAG.getNode(Opcode: ISD::FTRUNC, DL: SDLoc(ST), VT: STMemType, Operand: Val);
21867 return true;
21868 }
21869 if (numVectorEltsOrZero(T: STType) == numVectorEltsOrZero(T: STMemType) &&
21870 STType.isInteger() && STMemType.isInteger()) {
21871 Val = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(ST), VT: STMemType, Operand: Val);
21872 return true;
21873 }
21874 if (STType.getSizeInBits() == STMemType.getSizeInBits()) {
21875 Val = DAG.getBitcast(VT: STMemType, V: Val);
21876 return true;
21877 }
21878 return false; // fail.
21879}
21880
21881bool DAGCombiner::extendLoadedValueToExtension(LoadSDNode *LD, SDValue &Val) {
21882 EVT LDMemType = LD->getMemoryVT();
21883 EVT LDType = LD->getValueType(ResNo: 0);
21884 assert(Val.getValueType() == LDMemType &&
21885 "Attempting to extend value of non-matching type");
21886 if (LDType == LDMemType)
21887 return true;
21888 if (LDMemType.isInteger() && LDType.isInteger()) {
21889 switch (LD->getExtensionType()) {
21890 case ISD::NON_EXTLOAD:
21891 Val = DAG.getBitcast(VT: LDType, V: Val);
21892 return true;
21893 case ISD::EXTLOAD:
21894 Val = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL: SDLoc(LD), VT: LDType, Operand: Val);
21895 return true;
21896 case ISD::SEXTLOAD:
21897 Val = DAG.getNode(Opcode: ISD::SIGN_EXTEND, DL: SDLoc(LD), VT: LDType, Operand: Val);
21898 return true;
21899 case ISD::ZEXTLOAD:
21900 Val = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL: SDLoc(LD), VT: LDType, Operand: Val);
21901 return true;
21902 }
21903 }
21904 return false;
21905}
21906
21907StoreSDNode *DAGCombiner::getUniqueStoreFeeding(LoadSDNode *LD,
21908 int64_t &Offset) {
21909 SDValue Chain = LD->getOperand(Num: 0);
21910
21911 // Look through CALLSEQ_START.
21912 if (Chain.getOpcode() == ISD::CALLSEQ_START)
21913 Chain = Chain->getOperand(Num: 0);
21914
21915 StoreSDNode *ST = nullptr;
21916 SmallVector<SDValue, 8> Aliases;
21917 if (Chain.getOpcode() == ISD::TokenFactor) {
21918 // Look for unique store within the TokenFactor.
21919 for (SDValue Op : Chain->ops()) {
21920 StoreSDNode *Store = dyn_cast<StoreSDNode>(Val: Op.getNode());
21921 if (!Store)
21922 continue;
21923 BaseIndexOffset BasePtrLD = BaseIndexOffset::match(N: LD, DAG);
21924 BaseIndexOffset BasePtrST = BaseIndexOffset::match(N: Store, DAG);
21925 if (!BasePtrST.equalBaseIndex(Other: BasePtrLD, DAG, Off&: Offset))
21926 continue;
21927 // Make sure the store is not aliased with any nodes in TokenFactor.
21928 GatherAllAliases(N: Store, OriginalChain: Chain, Aliases);
21929 if (Aliases.empty() ||
21930 (Aliases.size() == 1 && Aliases.front().getNode() == Store))
21931 ST = Store;
21932 break;
21933 }
21934 } else {
21935 StoreSDNode *Store = dyn_cast<StoreSDNode>(Val: Chain.getNode());
21936 if (Store) {
21937 BaseIndexOffset BasePtrLD = BaseIndexOffset::match(N: LD, DAG);
21938 BaseIndexOffset BasePtrST = BaseIndexOffset::match(N: Store, DAG);
21939 if (BasePtrST.equalBaseIndex(Other: BasePtrLD, DAG, Off&: Offset))
21940 ST = Store;
21941 }
21942 }
21943
21944 return ST;
21945}
21946
21947SDValue DAGCombiner::ForwardStoreValueToDirectLoad(LoadSDNode *LD) {
21948 if (OptLevel == CodeGenOptLevel::None || !LD->isSimple())
21949 return SDValue();
21950 SDValue Chain = LD->getOperand(Num: 0);
21951 int64_t Offset;
21952
21953 StoreSDNode *ST = getUniqueStoreFeeding(LD, Offset);
21954 // TODO: Relax this restriction for unordered atomics (see D66309)
21955 if (!ST || !ST->isSimple() || ST->getAddressSpace() != LD->getAddressSpace())
21956 return SDValue();
21957
21958 EVT LDType = LD->getValueType(ResNo: 0);
21959 EVT LDMemType = LD->getMemoryVT();
21960 EVT STMemType = ST->getMemoryVT();
21961 EVT STType = ST->getValue().getValueType();
21962
21963 // There are two cases to consider here:
21964 // 1. The store is fixed width and the load is scalable. In this case we
21965 // don't know at compile time if the store completely envelops the load
21966 // so we abandon the optimisation.
21967 // 2. The store is scalable and the load is fixed width. We could
21968 // potentially support a limited number of cases here, but there has been
21969 // no cost-benefit analysis to prove it's worth it.
21970 bool LdStScalable = LDMemType.isScalableVT();
21971 if (LdStScalable != STMemType.isScalableVT())
21972 return SDValue();
21973
21974 // If we are dealing with scalable vectors on a big endian platform the
21975 // calculation of offsets below becomes trickier, since we do not know at
21976 // compile time the absolute size of the vector. Until we've done more
21977 // analysis on big-endian platforms it seems better to bail out for now.
21978 if (LdStScalable && DAG.getDataLayout().isBigEndian())
21979 return SDValue();
21980
21981 // Normalize for Endianness. After this Offset=0 will denote that the least
21982 // significant bit in the loaded value maps to the least significant bit in
21983 // the stored value). With Offset=n (for n > 0) the loaded value starts at the
21984 // n:th least significant byte of the stored value.
21985 int64_t OrigOffset = Offset;
21986 if (DAG.getDataLayout().isBigEndian())
21987 Offset = ((int64_t)STMemType.getStoreSizeInBits().getFixedValue() -
21988 (int64_t)LDMemType.getStoreSizeInBits().getFixedValue()) /
21989 8 -
21990 Offset;
21991
21992 // Check that the stored value cover all bits that are loaded.
21993 bool STCoversLD;
21994
21995 TypeSize LdMemSize = LDMemType.getSizeInBits();
21996 TypeSize StMemSize = STMemType.getSizeInBits();
21997 if (LdStScalable)
21998 STCoversLD = (Offset == 0) && LdMemSize == StMemSize;
21999 else
22000 STCoversLD = (Offset >= 0) && (Offset * 8 + LdMemSize.getFixedValue() <=
22001 StMemSize.getFixedValue());
22002
22003 auto ReplaceLd = [&](LoadSDNode *LD, SDValue Val, SDValue Chain) -> SDValue {
22004 if (LD->isIndexed()) {
22005 // Cannot handle opaque target constants and we must respect the user's
22006 // request not to split indexes from loads.
22007 if (!canSplitIdx(LD))
22008 return SDValue();
22009 SDValue Idx = SplitIndexingFromLoad(LD);
22010 SDValue Ops[] = {Val, Idx, Chain};
22011 return CombineTo(N: LD, To: Ops, NumTo: 3);
22012 }
22013 return CombineTo(N: LD, Res0: Val, Res1: Chain);
22014 };
22015
22016 if (!STCoversLD)
22017 return SDValue();
22018
22019 // Memory as copy space (potentially masked).
22020 if (Offset == 0 && LDType == STType && STMemType == LDMemType) {
22021 // Simple case: Direct non-truncating forwarding
22022 if (LDType.getSizeInBits() == LdMemSize)
22023 return ReplaceLd(LD, ST->getValue(), Chain);
22024 // Can we model the truncate and extension with an and mask?
22025 if (STType.isInteger() && LDMemType.isInteger() && !STType.isVector() &&
22026 !LDMemType.isVector() && LD->getExtensionType() != ISD::SEXTLOAD) {
22027 // Mask to size of LDMemType
22028 auto Mask =
22029 DAG.getConstant(Val: APInt::getLowBitsSet(numBits: STType.getFixedSizeInBits(),
22030 loBitsSet: StMemSize.getFixedValue()),
22031 DL: SDLoc(ST), VT: STType);
22032 auto Val = DAG.getNode(Opcode: ISD::AND, DL: SDLoc(LD), VT: LDType, N1: ST->getValue(), N2: Mask);
22033 return ReplaceLd(LD, Val, Chain);
22034 }
22035 }
22036
22037 // Handle some cases for big-endian that would be Offset 0 and handled for
22038 // little-endian.
22039 SDValue Val = ST->getValue();
22040 if (DAG.getDataLayout().isBigEndian() && Offset > 0 && OrigOffset == 0) {
22041 if (STType.isInteger() && !STType.isVector() && LDType.isInteger() &&
22042 !LDType.isVector() && isTypeLegal(VT: STType) &&
22043 TLI.isOperationLegal(Op: ISD::SRL, VT: STType)) {
22044 Val = DAG.getNode(
22045 Opcode: ISD::SRL, DL: SDLoc(LD), VT: STType, N1: Val,
22046 N2: DAG.getShiftAmountConstant(Val: Offset * 8, VT: STType, DL: SDLoc(LD)));
22047 Offset = 0;
22048 }
22049 }
22050
22051 // TODO: Deal with nonzero offset.
22052 if (LD->getBasePtr().isUndef() || Offset != 0)
22053 return SDValue();
22054 // Model necessary truncations / extenstions.
22055 // Truncate Value To Stored Memory Size.
22056 do {
22057 if (!getTruncatedStoreValue(ST, Val))
22058 break;
22059 if (!isTypeLegal(VT: LDMemType))
22060 break;
22061 if (STMemType != LDMemType) {
22062 if (LdMemSize == StMemSize) {
22063 if (TLI.isOperationLegal(Op: ISD::BITCAST, VT: LDMemType) &&
22064 isTypeLegal(VT: LDMemType) &&
22065 TLI.isOperationLegal(Op: ISD::BITCAST, VT: STMemType) &&
22066 isTypeLegal(VT: STMemType) &&
22067 TLI.isLoadBitCastBeneficial(LoadVT: LDMemType, BitcastVT: STMemType, DAG,
22068 MMO: *LD->getMemOperand()))
22069 Val = DAG.getBitcast(VT: LDMemType, V: Val);
22070 else
22071 break;
22072 } else if (LDMemType.isVector() && isTypeLegal(VT: STMemType)) {
22073 EVT EltVT = LDMemType.getVectorElementType();
22074 TypeSize EltSize = EltVT.getSizeInBits();
22075
22076 if (!StMemSize.isKnownMultipleOf(RHS: EltSize))
22077 break;
22078
22079 EVT InterVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: EltVT,
22080 NumElements: StMemSize.divideCoefficientBy(RHS: EltSize));
22081 if (!TLI.isOperationLegalOrCustom(Op: ISD::EXTRACT_SUBVECTOR, VT: LDMemType) ||
22082 !TLI.isTypeLegal(VT: InterVT))
22083 break;
22084
22085 // In case of big-endian the offset is normalized to zero, denoting
22086 // the last bit. For big-endian we need to transform the extraction
22087 // to the last sub-vector.
22088 unsigned ExtIdx = 0;
22089 if (DAG.getDataLayout().isBigEndian()) {
22090 ExtIdx =
22091 InterVT.getVectorNumElements() - LDMemType.getVectorNumElements();
22092 }
22093
22094 if (TLI.getExtractSubvectorCost(ResVT: LDMemType, SrcVT: InterVT, Index: ExtIdx) >
22095 TargetLowering::ExtractSubvectorCost::Cheap)
22096 break;
22097 Val = DAG.getExtractSubvector(DL: SDLoc(LD), VT: LDMemType,
22098 Vec: DAG.getBitcast(VT: InterVT, V: Val), Idx: ExtIdx);
22099 } else if (!STMemType.isVector() && !LDMemType.isVector() &&
22100 STMemType.isInteger() && LDMemType.isInteger())
22101 Val = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(LD), VT: LDMemType, Operand: Val);
22102 else
22103 break;
22104 }
22105 if (!extendLoadedValueToExtension(LD, Val))
22106 break;
22107 return ReplaceLd(LD, Val, Chain);
22108 } while (false);
22109
22110 // On failure, cleanup dead nodes we may have created.
22111 if (Val->use_empty())
22112 deleteAndRecombine(N: Val.getNode());
22113 return SDValue();
22114}
22115
22116SDValue DAGCombiner::visitLOAD(SDNode *N) {
22117 LoadSDNode *LD = cast<LoadSDNode>(Val: N);
22118 SDValue Chain = LD->getChain();
22119 SDValue Ptr = LD->getBasePtr();
22120
22121 // If load is not volatile and there are no uses of the loaded value (and
22122 // the updated indexed value in case of indexed loads), change uses of the
22123 // chain value into uses of the chain input (i.e. delete the dead load).
22124 // TODO: Allow this for unordered atomics (see D66309)
22125 if (LD->isSimple()) {
22126 if (N->getValueType(ResNo: 1) == MVT::Other) {
22127 // Unindexed loads.
22128 if (!N->hasAnyUseOfValue(Value: 0)) {
22129 // It's not safe to use the two value CombineTo variant here. e.g.
22130 // v1, chain2 = load chain1, loc
22131 // v2, chain3 = load chain2, loc
22132 // v3 = add v2, c
22133 // Now we replace use of chain2 with chain1. This makes the second load
22134 // isomorphic to the one we are deleting, and thus makes this load live.
22135 LLVM_DEBUG(dbgs() << "\nReplacing.6 "; N->dump(&DAG);
22136 dbgs() << "\nWith chain: "; Chain.dump(&DAG);
22137 dbgs() << "\n");
22138 WorklistRemover DeadNodes(*this);
22139 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 1), To: Chain);
22140 AddUsersToWorklist(N: Chain.getNode());
22141 if (N->use_empty())
22142 deleteAndRecombine(N);
22143
22144 return SDValue(N, 0); // Return N so it doesn't get rechecked!
22145 }
22146 } else {
22147 // Indexed loads.
22148 assert(N->getValueType(2) == MVT::Other && "Malformed indexed loads?");
22149
22150 // If this load has an opaque TargetConstant offset, then we cannot split
22151 // the indexing into an add/sub directly (that TargetConstant may not be
22152 // valid for a different type of node, and we cannot convert an opaque
22153 // target constant into a regular constant).
22154 bool CanSplitIdx = canSplitIdx(LD);
22155
22156 if (!N->hasAnyUseOfValue(Value: 0) && (CanSplitIdx || !N->hasAnyUseOfValue(Value: 1))) {
22157 SDValue Poison = DAG.getPOISON(VT: N->getValueType(ResNo: 0));
22158 SDValue Index;
22159 if (N->hasAnyUseOfValue(Value: 1) && CanSplitIdx) {
22160 Index = SplitIndexingFromLoad(LD);
22161 // Try to fold the base pointer arithmetic into subsequent loads and
22162 // stores.
22163 AddUsersToWorklist(N);
22164 } else
22165 Index = DAG.getPOISON(VT: N->getValueType(ResNo: 1));
22166 LLVM_DEBUG(dbgs() << "\nReplacing.7 "; N->dump(&DAG);
22167 dbgs() << "\nWith: "; Poison.dump(&DAG);
22168 dbgs() << " and 2 other values\n");
22169 WorklistRemover DeadNodes(*this);
22170 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 0), To: Poison);
22171 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 1), To: Index);
22172 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 2), To: Chain);
22173 deleteAndRecombine(N);
22174 return SDValue(N, 0); // Return N so it doesn't get rechecked!
22175 }
22176 }
22177 }
22178
22179 // If this load is directly stored, replace the load value with the stored
22180 // value.
22181 if (auto V = ForwardStoreValueToDirectLoad(LD))
22182 return V;
22183
22184 // Try to infer better alignment information than the load already has.
22185 if (OptLevel != CodeGenOptLevel::None && LD->isUnindexed() &&
22186 !LD->isAtomic()) {
22187 if (MaybeAlign Alignment = DAG.InferPtrAlign(Ptr)) {
22188 if (*Alignment > LD->getAlign() &&
22189 isAligned(Lhs: *Alignment, SizeInBytes: LD->getSrcValueOffset())) {
22190 SDValue NewLoad = DAG.getLoad(
22191 AM: LD->getAddressingMode(), ExtType: LD->getExtensionType(),
22192 VT: LD->getValueType(ResNo: 0), dl: SDLoc(N), Chain, Ptr, Offset: LD->getOffset(),
22193 PtrInfo: LD->getPointerInfo(), MemVT: LD->getMemoryVT(), Alignment: *Alignment,
22194 MMOFlags: LD->getMemOperand()->getFlags(), Metadata: LD->getAAInfo());
22195 // NewLoad will always be N as we are only refining the alignment
22196 assert(NewLoad.getNode() == N);
22197 (void)NewLoad;
22198 }
22199 }
22200 }
22201
22202 if (LD->isUnindexed()) {
22203 // Walk up chain skipping non-aliasing memory nodes.
22204 SDValue BetterChain = FindBetterChain(N: LD, Chain);
22205
22206 // If there is a better chain.
22207 if (Chain != BetterChain) {
22208 SDValue ReplLoad;
22209
22210 // Replace the chain to void dependency.
22211 if (LD->getExtensionType() == ISD::NON_EXTLOAD) {
22212 ReplLoad = DAG.getLoad(VT: N->getValueType(ResNo: 0), dl: SDLoc(LD),
22213 Chain: BetterChain, Ptr, MMO: LD->getMemOperand());
22214 } else {
22215 ReplLoad = DAG.getExtLoad(ExtType: LD->getExtensionType(), dl: SDLoc(LD),
22216 VT: LD->getValueType(ResNo: 0),
22217 Chain: BetterChain, Ptr, MemVT: LD->getMemoryVT(),
22218 MMO: LD->getMemOperand());
22219 }
22220
22221 // Create token factor to keep old chain connected.
22222 SDValue Token = DAG.getNode(Opcode: ISD::TokenFactor, DL: SDLoc(N),
22223 VT: MVT::Other, N1: Chain, N2: ReplLoad.getValue(R: 1));
22224
22225 // Replace uses with load result and token factor
22226 return CombineTo(N, Res0: ReplLoad.getValue(R: 0), Res1: Token);
22227 }
22228 }
22229
22230 // Try transforming N to an indexed load.
22231 if (CombineToPreIndexedLoadStore(N) || CombineToPostIndexedLoadStore(N))
22232 return SDValue(N, 0);
22233
22234 // Try to slice up N to more direct loads if the slices are mapped to
22235 // different register banks or pairing can take place.
22236 if (SliceUpLoad(N))
22237 return SDValue(N, 0);
22238
22239 return SDValue();
22240}
22241
22242namespace {
22243
22244/// Helper structure used to slice a load in smaller loads.
22245/// Basically a slice is obtained from the following sequence:
22246/// Origin = load Ty1, Base
22247/// Shift = srl Ty1 Origin, CstTy Amount
22248/// Inst = trunc Shift to Ty2
22249///
22250/// Then, it will be rewritten into:
22251/// Slice = load SliceTy, Base + SliceOffset
22252/// [Inst = zext Slice to Ty2], only if SliceTy <> Ty2
22253///
22254/// SliceTy is deduced from the number of bits that are actually used to
22255/// build Inst.
22256struct LoadedSlice {
22257 /// Helper structure used to compute the cost of a slice.
22258 struct Cost {
22259 /// Are we optimizing for code size.
22260 bool ForCodeSize = false;
22261
22262 /// Various cost.
22263 unsigned Loads = 0;
22264 unsigned Truncates = 0;
22265 unsigned CrossRegisterBanksCopies = 0;
22266 unsigned ZExts = 0;
22267 unsigned Shift = 0;
22268
22269 explicit Cost(bool ForCodeSize) : ForCodeSize(ForCodeSize) {}
22270
22271 /// Get the cost of one isolated slice.
22272 Cost(const LoadedSlice &LS, bool ForCodeSize)
22273 : ForCodeSize(ForCodeSize), Loads(1) {
22274 EVT TruncType = LS.Inst->getValueType(ResNo: 0);
22275 EVT LoadedType = LS.getLoadedType();
22276 if (TruncType != LoadedType &&
22277 !LS.DAG->getTargetLoweringInfo().isZExtFree(FromTy: LoadedType, ToTy: TruncType))
22278 ZExts = 1;
22279 }
22280
22281 /// Account for slicing gain in the current cost.
22282 /// Slicing provide a few gains like removing a shift or a
22283 /// truncate. This method allows to grow the cost of the original
22284 /// load with the gain from this slice.
22285 void addSliceGain(const LoadedSlice &LS) {
22286 // Each slice saves a truncate.
22287 const TargetLowering &TLI = LS.DAG->getTargetLoweringInfo();
22288 if (!TLI.isTruncateFree(Val: LS.Inst->getOperand(Num: 0), VT2: LS.Inst->getValueType(ResNo: 0)))
22289 ++Truncates;
22290 // If there is a shift amount, this slice gets rid of it.
22291 if (LS.Shift)
22292 ++Shift;
22293 // If this slice can merge a cross register bank copy, account for it.
22294 if (LS.canMergeExpensiveCrossRegisterBankCopy())
22295 ++CrossRegisterBanksCopies;
22296 }
22297
22298 Cost &operator+=(const Cost &RHS) {
22299 Loads += RHS.Loads;
22300 Truncates += RHS.Truncates;
22301 CrossRegisterBanksCopies += RHS.CrossRegisterBanksCopies;
22302 ZExts += RHS.ZExts;
22303 Shift += RHS.Shift;
22304 return *this;
22305 }
22306
22307 bool operator==(const Cost &RHS) const {
22308 return Loads == RHS.Loads && Truncates == RHS.Truncates &&
22309 CrossRegisterBanksCopies == RHS.CrossRegisterBanksCopies &&
22310 ZExts == RHS.ZExts && Shift == RHS.Shift;
22311 }
22312
22313 bool operator!=(const Cost &RHS) const { return !(*this == RHS); }
22314
22315 bool operator<(const Cost &RHS) const {
22316 // Assume cross register banks copies are as expensive as loads.
22317 // FIXME: Do we want some more target hooks?
22318 unsigned ExpensiveOpsLHS = Loads + CrossRegisterBanksCopies;
22319 unsigned ExpensiveOpsRHS = RHS.Loads + RHS.CrossRegisterBanksCopies;
22320 // Unless we are optimizing for code size, consider the
22321 // expensive operation first.
22322 if (!ForCodeSize && ExpensiveOpsLHS != ExpensiveOpsRHS)
22323 return ExpensiveOpsLHS < ExpensiveOpsRHS;
22324 return (Truncates + ZExts + Shift + ExpensiveOpsLHS) <
22325 (RHS.Truncates + RHS.ZExts + RHS.Shift + ExpensiveOpsRHS);
22326 }
22327
22328 bool operator>(const Cost &RHS) const { return RHS < *this; }
22329
22330 bool operator<=(const Cost &RHS) const { return !(RHS < *this); }
22331
22332 bool operator>=(const Cost &RHS) const { return !(*this < RHS); }
22333 };
22334
22335 // The last instruction that represent the slice. This should be a
22336 // truncate instruction.
22337 SDNode *Inst;
22338
22339 // The original load instruction.
22340 LoadSDNode *Origin;
22341
22342 // The right shift amount in bits from the original load.
22343 unsigned Shift;
22344
22345 // The DAG from which Origin came from.
22346 // This is used to get some contextual information about legal types, etc.
22347 SelectionDAG *DAG;
22348
22349 LoadedSlice(SDNode *Inst = nullptr, LoadSDNode *Origin = nullptr,
22350 unsigned Shift = 0, SelectionDAG *DAG = nullptr)
22351 : Inst(Inst), Origin(Origin), Shift(Shift), DAG(DAG) {}
22352
22353 /// Get the bits used in a chunk of bits \p BitWidth large.
22354 /// \return Result is \p BitWidth and has used bits set to 1 and
22355 /// not used bits set to 0.
22356 APInt getUsedBits() const {
22357 // Reproduce the trunc(lshr) sequence:
22358 // - Start from the truncated value.
22359 // - Zero extend to the desired bit width.
22360 // - Shift left.
22361 assert(Origin && "No original load to compare against.");
22362 unsigned BitWidth = Origin->getValueSizeInBits(ResNo: 0);
22363 assert(Inst && "This slice is not bound to an instruction");
22364 assert(Inst->getValueSizeInBits(0) <= BitWidth &&
22365 "Extracted slice is bigger than the whole type!");
22366 APInt UsedBits(Inst->getValueSizeInBits(ResNo: 0), 0);
22367 UsedBits.setAllBits();
22368 UsedBits = UsedBits.zext(width: BitWidth);
22369 UsedBits <<= Shift;
22370 return UsedBits;
22371 }
22372
22373 /// Get the size of the slice to be loaded in bytes.
22374 unsigned getLoadedSize() const {
22375 unsigned SliceSize = getUsedBits().popcount();
22376 assert(!(SliceSize & 0x7) && "Size is not a multiple of a byte.");
22377 return SliceSize / 8;
22378 }
22379
22380 /// Get the type that will be loaded for this slice.
22381 /// Note: This may not be the final type for the slice.
22382 EVT getLoadedType() const {
22383 assert(DAG && "Missing context");
22384 LLVMContext &Ctxt = *DAG->getContext();
22385 return EVT::getIntegerVT(Context&: Ctxt, BitWidth: getLoadedSize() * 8);
22386 }
22387
22388 /// Get the alignment of the load used for this slice.
22389 Align getAlign() const {
22390 Align Alignment = Origin->getAlign();
22391 uint64_t Offset = getOffsetFromBase();
22392 if (Offset != 0)
22393 Alignment = commonAlignment(A: Alignment, Offset: Alignment.value() + Offset);
22394 return Alignment;
22395 }
22396
22397 /// Check if this slice can be rewritten with legal operations.
22398 bool isLegal() const {
22399 // An invalid slice is not legal.
22400 if (!Origin || !Inst || !DAG)
22401 return false;
22402
22403 // Offsets are for indexed load only, we do not handle that.
22404 if (!Origin->getOffset().isUndef())
22405 return false;
22406
22407 const TargetLowering &TLI = DAG->getTargetLoweringInfo();
22408
22409 // Check that the type is legal.
22410 EVT SliceType = getLoadedType();
22411 if (!TLI.isTypeLegal(VT: SliceType))
22412 return false;
22413
22414 // Check that the load is legal for this type.
22415 if (!TLI.isOperationLegal(Op: ISD::LOAD, VT: SliceType))
22416 return false;
22417
22418 // Check that the offset can be computed.
22419 // 1. Check its type.
22420 EVT PtrType = Origin->getBasePtr().getValueType();
22421 if (PtrType == MVT::Untyped || PtrType.isExtended())
22422 return false;
22423
22424 // 2. Check that it fits in the immediate.
22425 if (!TLI.isLegalAddImmediate(getOffsetFromBase()))
22426 return false;
22427
22428 // 3. Check that the computation is legal.
22429 if (!TLI.isOperationLegal(Op: ISD::ADD, VT: PtrType))
22430 return false;
22431
22432 // Check that the zext is legal if it needs one.
22433 EVT TruncateType = Inst->getValueType(ResNo: 0);
22434 if (TruncateType != SliceType &&
22435 !TLI.isOperationLegal(Op: ISD::ZERO_EXTEND, VT: TruncateType))
22436 return false;
22437
22438 return true;
22439 }
22440
22441 /// Get the offset in bytes of this slice in the original chunk of
22442 /// bits.
22443 /// \pre DAG != nullptr.
22444 uint64_t getOffsetFromBase() const {
22445 assert(DAG && "Missing context.");
22446 bool IsBigEndian = DAG->getDataLayout().isBigEndian();
22447 assert(!(Shift & 0x7) && "Shifts not aligned on Bytes are not supported.");
22448 uint64_t Offset = Shift / 8;
22449 unsigned TySizeInBytes = Origin->getValueSizeInBits(ResNo: 0) / 8;
22450 assert(!(Origin->getValueSizeInBits(0) & 0x7) &&
22451 "The size of the original loaded type is not a multiple of a"
22452 " byte.");
22453 // If Offset is bigger than TySizeInBytes, it means we are loading all
22454 // zeros. This should have been optimized before in the process.
22455 assert(TySizeInBytes > Offset &&
22456 "Invalid shift amount for given loaded size");
22457 if (IsBigEndian)
22458 Offset = TySizeInBytes - Offset - getLoadedSize();
22459 return Offset;
22460 }
22461
22462 /// Generate the sequence of instructions to load the slice
22463 /// represented by this object and redirect the uses of this slice to
22464 /// this new sequence of instructions.
22465 /// \pre this->Inst && this->Origin are valid Instructions and this
22466 /// object passed the legal check: LoadedSlice::isLegal returned true.
22467 /// \return The last instruction of the sequence used to load the slice.
22468 SDValue loadSlice() const {
22469 assert(Inst && Origin && "Unable to replace a non-existing slice.");
22470 const SDValue &OldBaseAddr = Origin->getBasePtr();
22471 SDValue BaseAddr = OldBaseAddr;
22472 // Get the offset in that chunk of bytes w.r.t. the endianness.
22473 int64_t Offset = static_cast<int64_t>(getOffsetFromBase());
22474 assert(Offset >= 0 && "Offset too big to fit in int64_t!");
22475 if (Offset) {
22476 // BaseAddr = BaseAddr + Offset.
22477 EVT ArithType = BaseAddr.getValueType();
22478 SDLoc DL(Origin);
22479 BaseAddr = DAG->getNode(Opcode: ISD::ADD, DL, VT: ArithType, N1: BaseAddr,
22480 N2: DAG->getConstant(Val: Offset, DL, VT: ArithType));
22481 }
22482
22483 // Create the type of the loaded slice according to its size.
22484 EVT SliceType = getLoadedType();
22485
22486 // Create the load for the slice.
22487 SDValue LastInst =
22488 DAG->getLoad(VT: SliceType, dl: SDLoc(Origin), Chain: Origin->getChain(), Ptr: BaseAddr,
22489 PtrInfo: Origin->getPointerInfo().getWithOffset(O: Offset), Alignment: getAlign(),
22490 MMOFlags: Origin->getMemOperand()->getFlags());
22491 // If the final type is not the same as the loaded type, this means that
22492 // we have to pad with zero. Create a zero extend for that.
22493 EVT FinalType = Inst->getValueType(ResNo: 0);
22494 if (SliceType != FinalType)
22495 LastInst =
22496 DAG->getNode(Opcode: ISD::ZERO_EXTEND, DL: SDLoc(LastInst), VT: FinalType, Operand: LastInst);
22497 return LastInst;
22498 }
22499
22500 /// Check if this slice can be merged with an expensive cross register
22501 /// bank copy. E.g.,
22502 /// i = load i32
22503 /// f = bitcast i32 i to float
22504 bool canMergeExpensiveCrossRegisterBankCopy() const {
22505 if (!Inst || !Inst->hasOneUse())
22506 return false;
22507 SDNode *User = *Inst->user_begin();
22508 if (User->getOpcode() != ISD::BITCAST)
22509 return false;
22510 assert(DAG && "Missing context");
22511 const TargetLowering &TLI = DAG->getTargetLoweringInfo();
22512 EVT ResVT = User->getValueType(ResNo: 0);
22513 const TargetRegisterClass *ResRC =
22514 TLI.getRegClassFor(VT: ResVT.getSimpleVT(), isDivergent: User->isDivergent());
22515 const TargetRegisterClass *ArgRC =
22516 TLI.getRegClassFor(VT: User->getOperand(Num: 0).getValueType().getSimpleVT(),
22517 isDivergent: User->getOperand(Num: 0)->isDivergent());
22518 if (ArgRC == ResRC || !TLI.isOperationLegal(Op: ISD::LOAD, VT: ResVT))
22519 return false;
22520
22521 // At this point, we know that we perform a cross-register-bank copy.
22522 // Check if it is expensive.
22523 const TargetRegisterInfo *TRI = DAG->getSubtarget().getRegisterInfo();
22524 // Assume bitcasts are cheap, unless both register classes do not
22525 // explicitly share a common sub class.
22526 if (!TRI || TRI->getCommonSubClass(A: ArgRC, B: ResRC))
22527 return false;
22528
22529 // Check if it will be merged with the load.
22530 // 1. Check the alignment / fast memory access constraint.
22531 unsigned IsFast = 0;
22532 if (!TLI.allowsMemoryAccess(Context&: *DAG->getContext(), DL: DAG->getDataLayout(), VT: ResVT,
22533 AddrSpace: Origin->getAddressSpace(), Alignment: getAlign(),
22534 Flags: Origin->getMemOperand()->getFlags(), Fast: &IsFast) ||
22535 !IsFast)
22536 return false;
22537
22538 // 2. Check that the load is a legal operation for that type.
22539 if (!TLI.isOperationLegal(Op: ISD::LOAD, VT: ResVT))
22540 return false;
22541
22542 // 3. Check that we do not have a zext in the way.
22543 if (Inst->getValueType(ResNo: 0) != getLoadedType())
22544 return false;
22545
22546 return true;
22547 }
22548};
22549
22550} // end anonymous namespace
22551
22552/// Check that all bits set in \p UsedBits form a dense region, i.e.,
22553/// \p UsedBits looks like 0..0 1..1 0..0.
22554static bool areUsedBitsDense(const APInt &UsedBits) {
22555 // If all the bits are one, this is dense!
22556 if (UsedBits.isAllOnes())
22557 return true;
22558
22559 // Get rid of the unused bits on the right.
22560 APInt NarrowedUsedBits = UsedBits.lshr(shiftAmt: UsedBits.countr_zero());
22561 // Get rid of the unused bits on the left.
22562 if (NarrowedUsedBits.countl_zero())
22563 NarrowedUsedBits = NarrowedUsedBits.trunc(width: NarrowedUsedBits.getActiveBits());
22564 // Check that the chunk of bits is completely used.
22565 return NarrowedUsedBits.isAllOnes();
22566}
22567
22568/// Check whether or not \p First and \p Second are next to each other
22569/// in memory. This means that there is no hole between the bits loaded
22570/// by \p First and the bits loaded by \p Second.
22571static bool areSlicesNextToEachOther(const LoadedSlice &First,
22572 const LoadedSlice &Second) {
22573 assert(First.Origin == Second.Origin && First.Origin &&
22574 "Unable to match different memory origins.");
22575 APInt UsedBits = First.getUsedBits();
22576 assert((UsedBits & Second.getUsedBits()) == 0 &&
22577 "Slices are not supposed to overlap.");
22578 UsedBits |= Second.getUsedBits();
22579 return areUsedBitsDense(UsedBits);
22580}
22581
22582/// Adjust the \p GlobalLSCost according to the target
22583/// paring capabilities and the layout of the slices.
22584/// \pre \p GlobalLSCost should account for at least as many loads as
22585/// there is in the slices in \p LoadedSlices.
22586static void adjustCostForPairing(SmallVectorImpl<LoadedSlice> &LoadedSlices,
22587 LoadedSlice::Cost &GlobalLSCost) {
22588 unsigned NumberOfSlices = LoadedSlices.size();
22589 // If there is less than 2 elements, no pairing is possible.
22590 if (NumberOfSlices < 2)
22591 return;
22592
22593 // Sort the slices so that elements that are likely to be next to each
22594 // other in memory are next to each other in the list.
22595 llvm::sort(C&: LoadedSlices, Comp: [](const LoadedSlice &LHS, const LoadedSlice &RHS) {
22596 assert(LHS.Origin == RHS.Origin && "Different bases not implemented.");
22597 return LHS.getOffsetFromBase() < RHS.getOffsetFromBase();
22598 });
22599 const TargetLowering &TLI = LoadedSlices[0].DAG->getTargetLoweringInfo();
22600 // First (resp. Second) is the first (resp. Second) potentially candidate
22601 // to be placed in a paired load.
22602 const LoadedSlice *First = nullptr;
22603 const LoadedSlice *Second = nullptr;
22604 for (unsigned CurrSlice = 0; CurrSlice < NumberOfSlices; ++CurrSlice,
22605 // Set the beginning of the pair.
22606 First = Second) {
22607 Second = &LoadedSlices[CurrSlice];
22608
22609 // If First is NULL, it means we start a new pair.
22610 // Get to the next slice.
22611 if (!First)
22612 continue;
22613
22614 EVT LoadedType = First->getLoadedType();
22615
22616 // If the types of the slices are different, we cannot pair them.
22617 if (LoadedType != Second->getLoadedType())
22618 continue;
22619
22620 // Check if the target supplies paired loads for this type.
22621 Align RequiredAlignment;
22622 if (!TLI.hasPairedLoad(LoadedType, RequiredAlignment)) {
22623 // move to the next pair, this type is hopeless.
22624 Second = nullptr;
22625 continue;
22626 }
22627 // Check if we meet the alignment requirement.
22628 if (First->getAlign() < RequiredAlignment)
22629 continue;
22630
22631 // Check that both loads are next to each other in memory.
22632 if (!areSlicesNextToEachOther(First: *First, Second: *Second))
22633 continue;
22634
22635 assert(GlobalLSCost.Loads > 0 && "We save more loads than we created!");
22636 --GlobalLSCost.Loads;
22637 // Move to the next pair.
22638 Second = nullptr;
22639 }
22640}
22641
22642/// Check the profitability of all involved LoadedSlice.
22643/// Currently, it is considered profitable if there is exactly two
22644/// involved slices (1) which are (2) next to each other in memory, and
22645/// whose cost (\see LoadedSlice::Cost) is smaller than the original load (3).
22646///
22647/// Note: The order of the elements in \p LoadedSlices may be modified, but not
22648/// the elements themselves.
22649///
22650/// FIXME: When the cost model will be mature enough, we can relax
22651/// constraints (1) and (2).
22652static bool isSlicingProfitable(SmallVectorImpl<LoadedSlice> &LoadedSlices,
22653 const APInt &UsedBits, bool ForCodeSize) {
22654 unsigned NumberOfSlices = LoadedSlices.size();
22655 if (StressLoadSlicing)
22656 return NumberOfSlices > 1;
22657
22658 // Check (1).
22659 if (NumberOfSlices != 2)
22660 return false;
22661
22662 // Check (2).
22663 if (!areUsedBitsDense(UsedBits))
22664 return false;
22665
22666 // Check (3).
22667 LoadedSlice::Cost OrigCost(ForCodeSize), GlobalSlicingCost(ForCodeSize);
22668 // The original code has one big load.
22669 OrigCost.Loads = 1;
22670 for (unsigned CurrSlice = 0; CurrSlice < NumberOfSlices; ++CurrSlice) {
22671 const LoadedSlice &LS = LoadedSlices[CurrSlice];
22672 // Accumulate the cost of all the slices.
22673 LoadedSlice::Cost SliceCost(LS, ForCodeSize);
22674 GlobalSlicingCost += SliceCost;
22675
22676 // Account as cost in the original configuration the gain obtained
22677 // with the current slices.
22678 OrigCost.addSliceGain(LS);
22679 }
22680
22681 // If the target supports paired load, adjust the cost accordingly.
22682 adjustCostForPairing(LoadedSlices, GlobalLSCost&: GlobalSlicingCost);
22683 return OrigCost > GlobalSlicingCost;
22684}
22685
22686/// If the given load, \p LI, is used only by trunc or trunc(lshr)
22687/// operations, split it in the various pieces being extracted.
22688///
22689/// This sort of thing is introduced by SROA.
22690/// This slicing takes care not to insert overlapping loads.
22691/// \pre LI is a simple load (i.e., not an atomic or volatile load).
22692bool DAGCombiner::SliceUpLoad(SDNode *N) {
22693 if (Level < AfterLegalizeDAG)
22694 return false;
22695
22696 LoadSDNode *LD = cast<LoadSDNode>(Val: N);
22697 if (!LD->isSimple() || !ISD::isNormalLoad(N: LD) ||
22698 !LD->getValueType(ResNo: 0).isInteger())
22699 return false;
22700
22701 // The algorithm to split up a load of a scalable vector into individual
22702 // elements currently requires knowing the length of the loaded type,
22703 // so will need adjusting to work on scalable vectors.
22704 if (LD->getValueType(ResNo: 0).isScalableVector())
22705 return false;
22706
22707 // Keep track of already used bits to detect overlapping values.
22708 // In that case, we will just abort the transformation.
22709 APInt UsedBits(LD->getValueSizeInBits(ResNo: 0), 0);
22710
22711 SmallVector<LoadedSlice, 4> LoadedSlices;
22712
22713 // Check if this load is used as several smaller chunks of bits.
22714 // Basically, look for uses in trunc or trunc(lshr) and record a new chain
22715 // of computation for each trunc.
22716 for (SDUse &U : LD->uses()) {
22717 // Skip the uses of the chain.
22718 if (U.getResNo() != 0)
22719 continue;
22720
22721 SDNode *User = U.getUser();
22722 unsigned Shift = 0;
22723
22724 // Check if this is a trunc(lshr).
22725 if (User->getOpcode() == ISD::SRL && User->hasOneUse() &&
22726 isa<ConstantSDNode>(Val: User->getOperand(Num: 1))) {
22727 Shift = User->getConstantOperandVal(Num: 1);
22728 User = *User->user_begin();
22729 }
22730
22731 // At this point, User is a Truncate, iff we encountered, trunc or
22732 // trunc(lshr).
22733 if (User->getOpcode() != ISD::TRUNCATE)
22734 return false;
22735
22736 // The width of the type must be a power of 2 and greater than 8-bits.
22737 // Otherwise the load cannot be represented in LLVM IR.
22738 // Moreover, if we shifted with a non-8-bits multiple, the slice
22739 // will be across several bytes. We do not support that.
22740 unsigned Width = User->getValueSizeInBits(ResNo: 0);
22741 if (Width < 8 || !isPowerOf2_32(Value: Width) || (Shift & 0x7))
22742 return false;
22743
22744 // Build the slice for this chain of computations.
22745 LoadedSlice LS(User, LD, Shift, &DAG);
22746 APInt CurrentUsedBits = LS.getUsedBits();
22747
22748 // Check if this slice overlaps with another.
22749 if ((CurrentUsedBits & UsedBits) != 0)
22750 return false;
22751 // Update the bits used globally.
22752 UsedBits |= CurrentUsedBits;
22753
22754 // Check if the new slice would be legal.
22755 if (!LS.isLegal())
22756 return false;
22757
22758 // Record the slice.
22759 LoadedSlices.push_back(Elt: LS);
22760 }
22761
22762 // Abort slicing if it does not seem to be profitable.
22763 if (!isSlicingProfitable(LoadedSlices, UsedBits, ForCodeSize))
22764 return false;
22765
22766 ++SlicedLoads;
22767
22768 // Rewrite each chain to use an independent load.
22769 // By construction, each chain can be represented by a unique load.
22770
22771 // Prepare the argument for the new token factor for all the slices.
22772 SmallVector<SDValue, 8> ArgChains;
22773 for (const LoadedSlice &LS : LoadedSlices) {
22774 SDValue SliceInst = LS.loadSlice();
22775 CombineTo(N: LS.Inst, Res: SliceInst, AddTo: true);
22776 if (SliceInst.getOpcode() != ISD::LOAD)
22777 SliceInst = SliceInst.getOperand(i: 0);
22778 assert(SliceInst->getOpcode() == ISD::LOAD &&
22779 "It takes more than a zext to get to the loaded slice!!");
22780 ArgChains.push_back(Elt: SliceInst.getValue(R: 1));
22781 }
22782
22783 SDValue Chain = DAG.getNode(Opcode: ISD::TokenFactor, DL: SDLoc(LD), VT: MVT::Other,
22784 Ops: ArgChains);
22785 DAG.ReplaceAllUsesOfValueWith(From: SDValue(N, 1), To: Chain);
22786 AddToWorklist(N: Chain.getNode());
22787 return true;
22788}
22789
22790/// Check to see if V is (and load (ptr), imm), where the load is having
22791/// specific bytes cleared out. If so, return the byte size being masked out
22792/// and the shift amount.
22793static std::pair<unsigned, unsigned>
22794CheckForMaskedLoad(SDValue V, SDValue Ptr, SDValue Chain) {
22795 std::pair<unsigned, unsigned> Result(0, 0);
22796
22797 // Check for the structure we're looking for.
22798 if (V->getOpcode() != ISD::AND ||
22799 !isa<ConstantSDNode>(Val: V->getOperand(Num: 1)) ||
22800 !ISD::isNormalLoad(N: V->getOperand(Num: 0).getNode()))
22801 return Result;
22802
22803 // Check the chain and pointer.
22804 LoadSDNode *LD = cast<LoadSDNode>(Val: V->getOperand(Num: 0));
22805 if (LD->getBasePtr() != Ptr) return Result; // Not from same pointer.
22806
22807 // This only handles simple types.
22808 if (V.getValueType() != MVT::i16 &&
22809 V.getValueType() != MVT::i32 &&
22810 V.getValueType() != MVT::i64)
22811 return Result;
22812
22813 // Check the constant mask. Invert it so that the bits being masked out are
22814 // 0 and the bits being kept are 1. Use getSExtValue so that leading bits
22815 // follow the sign bit for uniformity.
22816 uint64_t NotMask = ~cast<ConstantSDNode>(Val: V->getOperand(Num: 1))->getSExtValue();
22817 unsigned NotMaskLZ = llvm::countl_zero(Val: NotMask);
22818 if (NotMaskLZ & 7) return Result; // Must be multiple of a byte.
22819 unsigned NotMaskTZ = llvm::countr_zero(Val: NotMask);
22820 if (NotMaskTZ & 7) return Result; // Must be multiple of a byte.
22821 if (NotMaskLZ == 64) return Result; // All zero mask.
22822
22823 // See if we have a continuous run of bits. If so, we have 0*1+0*
22824 if (llvm::countr_one(Value: NotMask >> NotMaskTZ) + NotMaskTZ + NotMaskLZ != 64)
22825 return Result;
22826
22827 // Adjust NotMaskLZ down to be from the actual size of the int instead of i64.
22828 if (V.getValueType() != MVT::i64 && NotMaskLZ)
22829 NotMaskLZ -= 64-V.getValueSizeInBits();
22830
22831 unsigned MaskedBytes = (V.getValueSizeInBits()-NotMaskLZ-NotMaskTZ)/8;
22832 switch (MaskedBytes) {
22833 case 1:
22834 case 2:
22835 case 4: break;
22836 default: return Result; // All one mask, or 5-byte mask.
22837 }
22838
22839 // Verify that the first bit starts at a multiple of mask so that the access
22840 // is aligned the same as the access width.
22841 if (NotMaskTZ && NotMaskTZ/8 % MaskedBytes) return Result;
22842
22843 // For narrowing to be valid, it must be the case that the load the
22844 // immediately preceding memory operation before the store.
22845 if (LD == Chain.getNode())
22846 ; // ok.
22847 else if (Chain->getOpcode() == ISD::TokenFactor &&
22848 SDValue(LD, 1).hasOneUse()) {
22849 // LD has only 1 chain use so they are no indirect dependencies.
22850 if (!LD->isOperandOf(N: Chain.getNode()))
22851 return Result;
22852 } else
22853 return Result; // Fail.
22854
22855 Result.first = MaskedBytes;
22856 Result.second = NotMaskTZ/8;
22857 return Result;
22858}
22859
22860/// Check to see if IVal is something that provides a value as specified by
22861/// MaskInfo. If so, replace the specified store with a narrower store of
22862/// truncated IVal.
22863static SDValue
22864ShrinkLoadReplaceStoreWithStore(const std::pair<unsigned, unsigned> &MaskInfo,
22865 SDValue IVal, StoreSDNode *St,
22866 DAGCombiner *DC) {
22867 unsigned NumBytes = MaskInfo.first;
22868 unsigned ByteShift = MaskInfo.second;
22869 SelectionDAG &DAG = DC->getDAG();
22870
22871 // Check to see if IVal is all zeros in the part being masked in by the 'or'
22872 // that uses this. If not, this is not a replacement.
22873 APInt Mask = ~APInt::getBitsSet(numBits: IVal.getValueSizeInBits(),
22874 loBit: ByteShift*8, hiBit: (ByteShift+NumBytes)*8);
22875 if (!DAG.MaskedValueIsZero(Op: IVal, Mask)) return SDValue();
22876
22877 // Check that it is legal on the target to do this. It is legal if the new
22878 // VT we're shrinking to (i8/i16/i32) is legal or we're still before type
22879 // legalization. If the source type is legal, but the store type isn't, see
22880 // if we can use a truncating store.
22881 MVT VT = MVT::getIntegerVT(BitWidth: NumBytes * 8);
22882 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
22883 bool UseTruncStore;
22884 if (DC->isTypeLegal(VT))
22885 UseTruncStore = false;
22886 else if (TLI.isTypeLegal(VT: IVal.getValueType()) &&
22887 TLI.isTruncStoreLegal(ValVT: IVal.getValueType(), MemVT: VT, Alignment: St->getAlign(),
22888 AddrSpace: St->getAddressSpace()))
22889 UseTruncStore = true;
22890 else
22891 return SDValue();
22892
22893 // Can't do this for indexed stores.
22894 if (St->isIndexed())
22895 return SDValue();
22896
22897 // Check that the target doesn't think this is a bad idea.
22898 if (St->getMemOperand() &&
22899 !TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT,
22900 MMO: *St->getMemOperand()))
22901 return SDValue();
22902
22903 // Okay, we can do this! Replace the 'St' store with a store of IVal that is
22904 // shifted by ByteShift and truncated down to NumBytes.
22905 if (ByteShift) {
22906 SDLoc DL(IVal);
22907 IVal = DAG.getNode(
22908 Opcode: ISD::SRL, DL, VT: IVal.getValueType(), N1: IVal,
22909 N2: DAG.getShiftAmountConstant(Val: ByteShift * 8, VT: IVal.getValueType(), DL));
22910 }
22911
22912 // Figure out the offset for the store and the alignment of the access.
22913 unsigned StOffset;
22914 if (DAG.getDataLayout().isLittleEndian())
22915 StOffset = ByteShift;
22916 else
22917 StOffset = IVal.getValueType().getStoreSize() - ByteShift - NumBytes;
22918
22919 SDValue Ptr = St->getBasePtr();
22920 if (StOffset) {
22921 SDLoc DL(IVal);
22922 Ptr = DAG.getMemBasePlusOffset(Base: Ptr, Offset: TypeSize::getFixed(ExactSize: StOffset), DL);
22923 }
22924
22925 ++OpsNarrowed;
22926 if (UseTruncStore)
22927 return DAG.getTruncStore(Chain: St->getChain(), dl: SDLoc(St), Val: IVal, Ptr,
22928 PtrInfo: St->getPointerInfo().getWithOffset(O: StOffset), SVT: VT,
22929 Alignment: St->getBaseAlign());
22930
22931 // Truncate down to the new size.
22932 IVal = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(IVal), VT, Operand: IVal);
22933
22934 return DAG.getStore(Chain: St->getChain(), dl: SDLoc(St), Val: IVal, Ptr,
22935 PtrInfo: St->getPointerInfo().getWithOffset(O: StOffset),
22936 Alignment: St->getBaseAlign());
22937}
22938
22939/// Look for sequence of load / op / store where op is one of 'or', 'xor', and
22940/// 'and' of immediates. If 'op' is only touching some of the loaded bits, try
22941/// narrowing the load and store if it would end up being a win for performance
22942/// or code size.
22943SDValue DAGCombiner::ReduceLoadOpStoreWidth(SDNode *N) {
22944 StoreSDNode *ST = cast<StoreSDNode>(Val: N);
22945 if (!ST->isSimple())
22946 return SDValue();
22947
22948 SDValue Chain = ST->getChain();
22949 SDValue Value = ST->getValue();
22950 SDValue Ptr = ST->getBasePtr();
22951 EVT VT = Value.getValueType();
22952
22953 if (ST->isTruncatingStore() || VT.isVector())
22954 return SDValue();
22955
22956 unsigned Opc = Value.getOpcode();
22957
22958 if ((Opc != ISD::OR && Opc != ISD::XOR && Opc != ISD::AND) ||
22959 !Value.hasOneUse())
22960 return SDValue();
22961
22962 // If this is "store (or X, Y), P" and X is "(and (load P), cst)", where cst
22963 // is a byte mask indicating a consecutive number of bytes, check to see if
22964 // Y is known to provide just those bytes. If so, we try to replace the
22965 // load + replace + store sequence with a single (narrower) store, which makes
22966 // the load dead.
22967 if (Opc == ISD::OR && EnableShrinkLoadReplaceStoreWithStore) {
22968 std::pair<unsigned, unsigned> MaskedLoad;
22969 MaskedLoad = CheckForMaskedLoad(V: Value.getOperand(i: 0), Ptr, Chain);
22970 if (MaskedLoad.first)
22971 if (SDValue NewST = ShrinkLoadReplaceStoreWithStore(MaskInfo: MaskedLoad,
22972 IVal: Value.getOperand(i: 1), St: ST,DC: this))
22973 return NewST;
22974
22975 // Or is commutative, so try swapping X and Y.
22976 MaskedLoad = CheckForMaskedLoad(V: Value.getOperand(i: 1), Ptr, Chain);
22977 if (MaskedLoad.first)
22978 if (SDValue NewST = ShrinkLoadReplaceStoreWithStore(MaskInfo: MaskedLoad,
22979 IVal: Value.getOperand(i: 0), St: ST,DC: this))
22980 return NewST;
22981 }
22982
22983 if (!EnableReduceLoadOpStoreWidth)
22984 return SDValue();
22985
22986 if (Value.getOperand(i: 1).getOpcode() != ISD::Constant)
22987 return SDValue();
22988
22989 SDValue N0 = Value.getOperand(i: 0);
22990 if (ISD::isNormalLoad(N: N0.getNode()) && N0.hasOneUse() &&
22991 Chain == SDValue(N0.getNode(), 1)) {
22992 LoadSDNode *LD = cast<LoadSDNode>(Val&: N0);
22993 if (LD->getBasePtr() != Ptr ||
22994 LD->getPointerInfo().getAddrSpace() !=
22995 ST->getPointerInfo().getAddrSpace())
22996 return SDValue();
22997
22998 // Find the type NewVT to narrow the load / op / store to.
22999 SDValue N1 = Value.getOperand(i: 1);
23000 unsigned BitWidth = N1.getValueSizeInBits();
23001 APInt Imm = N1->getAsAPIntVal();
23002 if (Opc == ISD::AND)
23003 Imm.flipAllBits();
23004 if (Imm == 0 || Imm.isAllOnes())
23005 return SDValue();
23006 // Find least/most significant bit that need to be part of the narrowed
23007 // operation. We assume target will need to address/access full bytes, so
23008 // we make sure to align LSB and MSB at byte boundaries.
23009 unsigned BitsPerByteMask = 7u;
23010 unsigned LSB = Imm.countr_zero() & ~BitsPerByteMask;
23011 unsigned MSB = (Imm.getActiveBits() - 1) | BitsPerByteMask;
23012 unsigned NewBW = NextPowerOf2(A: MSB - LSB);
23013 EVT NewVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: NewBW);
23014 // The narrowing should be profitable, the load/store operation should be
23015 // legal (or custom) and the store size should be equal to the NewVT width.
23016 while (NewBW < BitWidth &&
23017 (NewVT.getStoreSizeInBits() != NewBW ||
23018 !TLI.isOperationLegalOrCustom(Op: Opc, VT: NewVT) ||
23019 (!ReduceLoadOpStoreWidthForceNarrowingProfitable &&
23020 !TLI.isNarrowingProfitable(N, SrcVT: VT, DestVT: NewVT)))) {
23021 NewBW = NextPowerOf2(A: NewBW);
23022 NewVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: NewBW);
23023 }
23024 if (NewBW >= BitWidth)
23025 return SDValue();
23026
23027 // If we come this far NewVT/NewBW reflect a power-of-2 sized type that is
23028 // large enough to cover all bits that should be modified. This type might
23029 // however be larger than really needed (such as i32 while we actually only
23030 // need to modify one byte). Now we need to find our how to align the memory
23031 // accesses to satisfy preferred alignments as well as avoiding to access
23032 // memory outside the store size of the orignal access.
23033
23034 unsigned VTStoreSize = VT.getStoreSizeInBits().getFixedValue();
23035
23036 // Let ShAmt denote amount of bits to skip, counted from the least
23037 // significant bits of Imm. And let PtrOff how much the pointer needs to be
23038 // offsetted (in bytes) for the new access.
23039 unsigned ShAmt = 0;
23040 uint64_t PtrOff = 0;
23041 for (; ShAmt + NewBW <= VTStoreSize; ShAmt += 8) {
23042 // Make sure the range [ShAmt, ShAmt+NewBW) cover both LSB and MSB.
23043 if (ShAmt > LSB)
23044 return SDValue();
23045 if (ShAmt + NewBW < MSB)
23046 continue;
23047
23048 // Calculate PtrOff.
23049 unsigned PtrAdjustmentInBits = DAG.getDataLayout().isBigEndian()
23050 ? VTStoreSize - NewBW - ShAmt
23051 : ShAmt;
23052 PtrOff = PtrAdjustmentInBits / 8;
23053
23054 // Now check if narrow access is allowed and fast, considering alignments.
23055 unsigned IsFast = 0;
23056 Align NewAlign = commonAlignment(A: LD->getAlign(), Offset: PtrOff);
23057 if (TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT: NewVT,
23058 AddrSpace: LD->getAddressSpace(), Alignment: NewAlign,
23059 Flags: LD->getMemOperand()->getFlags(), Fast: &IsFast) &&
23060 IsFast)
23061 break;
23062 }
23063 // If loop above did not find any accepted ShAmt we need to exit here.
23064 if (ShAmt + NewBW > VTStoreSize)
23065 return SDValue();
23066
23067 APInt NewImm = Imm.lshr(shiftAmt: ShAmt).trunc(width: NewBW);
23068 if (Opc == ISD::AND)
23069 NewImm.flipAllBits();
23070 Align NewAlign = commonAlignment(A: LD->getAlign(), Offset: PtrOff);
23071 SDValue NewPtr =
23072 DAG.getMemBasePlusOffset(Base: Ptr, Offset: TypeSize::getFixed(ExactSize: PtrOff), DL: SDLoc(LD));
23073 SDValue NewLD =
23074 DAG.getLoad(VT: NewVT, dl: SDLoc(N0), Chain: LD->getChain(), Ptr: NewPtr,
23075 PtrInfo: LD->getPointerInfo().getWithOffset(O: PtrOff), Alignment: NewAlign,
23076 MMOFlags: LD->getMemOperand()->getFlags(), Metadata: LD->getAAInfo());
23077 SDValue NewVal = DAG.getNode(Opcode: Opc, DL: SDLoc(Value), VT: NewVT, N1: NewLD,
23078 N2: DAG.getConstant(Val: NewImm, DL: SDLoc(Value), VT: NewVT));
23079 SDValue NewST =
23080 DAG.getStore(Chain, dl: SDLoc(N), Val: NewVal, Ptr: NewPtr,
23081 PtrInfo: ST->getPointerInfo().getWithOffset(O: PtrOff), Alignment: NewAlign);
23082
23083 AddToWorklist(N: NewPtr.getNode());
23084 AddToWorklist(N: NewLD.getNode());
23085 AddToWorklist(N: NewVal.getNode());
23086 WorklistRemover DeadNodes(*this);
23087 DAG.ReplaceAllUsesOfValueWith(From: N0.getValue(R: 1), To: NewLD.getValue(R: 1));
23088 ++OpsNarrowed;
23089 return NewST;
23090 }
23091
23092 return SDValue();
23093}
23094
23095/// For a given floating point load / store pair, if the load value isn't used
23096/// by any other operations, then consider transforming the pair to integer
23097/// load / store operations if the target deems the transformation profitable.
23098SDValue DAGCombiner::TransformFPLoadStorePair(SDNode *N) {
23099 StoreSDNode *ST = cast<StoreSDNode>(Val: N);
23100 SDValue Value = ST->getValue();
23101 if (ISD::isNormalStore(N: ST) && ISD::isNormalLoad(N: Value.getNode()) &&
23102 Value.hasOneUse()) {
23103 LoadSDNode *LD = cast<LoadSDNode>(Val&: Value);
23104 EVT VT = LD->getMemoryVT();
23105 if (!VT.isSimple() || !VT.isFloatingPoint() || VT != ST->getMemoryVT() ||
23106 LD->isNonTemporal() || ST->isNonTemporal() ||
23107 LD->getPointerInfo().getAddrSpace() != 0 ||
23108 ST->getPointerInfo().getAddrSpace() != 0)
23109 return SDValue();
23110
23111 TypeSize VTSize = VT.getSizeInBits();
23112
23113 // We don't know the size of scalable types at compile time so we cannot
23114 // create an integer of the equivalent size.
23115 if (VTSize.isScalable())
23116 return SDValue();
23117
23118 unsigned FastLD = 0, FastST = 0;
23119 EVT IntVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: VTSize.getFixedValue());
23120 if (!TLI.isOperationLegal(Op: ISD::LOAD, VT: IntVT) ||
23121 !TLI.isOperationLegal(Op: ISD::STORE, VT: IntVT) ||
23122 !TLI.isDesirableToTransformToIntegerOp(ISD::LOAD, VT) ||
23123 !TLI.isDesirableToTransformToIntegerOp(ISD::STORE, VT) ||
23124 !TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT: IntVT,
23125 MMO: *LD->getMemOperand(), Fast: &FastLD) ||
23126 !TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT: IntVT,
23127 MMO: *ST->getMemOperand(), Fast: &FastST) ||
23128 !FastLD || !FastST)
23129 return SDValue();
23130
23131 SDValue NewLD = DAG.getLoad(VT: IntVT, dl: SDLoc(Value), Chain: LD->getChain(),
23132 Ptr: LD->getBasePtr(), MMO: LD->getMemOperand());
23133
23134 SDValue NewST = DAG.getStore(Chain: ST->getChain(), dl: SDLoc(N), Val: NewLD,
23135 Ptr: ST->getBasePtr(), MMO: ST->getMemOperand());
23136
23137 AddToWorklist(N: NewLD.getNode());
23138 AddToWorklist(N: NewST.getNode());
23139 WorklistRemover DeadNodes(*this);
23140 DAG.ReplaceAllUsesOfValueWith(From: Value.getValue(R: 1), To: NewLD.getValue(R: 1));
23141 ++LdStFP2Int;
23142 return NewST;
23143 }
23144
23145 return SDValue();
23146}
23147
23148// This is a helper function for visitMUL to check the profitability
23149// of folding (mul (add x, c1), c2) -> (add (mul x, c2), c1*c2).
23150// MulNode is the original multiply, AddNode is (add x, c1),
23151// and ConstNode is c2.
23152//
23153// If the (add x, c1) has multiple uses, we could increase
23154// the number of adds if we make this transformation.
23155// It would only be worth doing this if we can remove a
23156// multiply in the process. Check for that here.
23157// To illustrate:
23158// (A + c1) * c3
23159// (A + c2) * c3
23160// We're checking for cases where we have common "c3 * A" expressions.
23161bool DAGCombiner::isMulAddWithConstProfitable(SDNode *MulNode, SDValue AddNode,
23162 SDValue ConstNode) {
23163 // If the add only has one use, and the target thinks the folding is
23164 // profitable or does not lead to worse code, this would be OK to do.
23165 if (AddNode->hasOneUse() &&
23166 TLI.isMulAddWithConstProfitable(AddNode, ConstNode))
23167 return true;
23168
23169 // Walk all the users of the constant with which we're multiplying.
23170 for (SDNode *User : ConstNode->users()) {
23171 if (User == MulNode) // This use is the one we're on right now. Skip it.
23172 continue;
23173
23174 if (User->getOpcode() == ISD::MUL) { // We have another multiply use.
23175 SDNode *OtherOp;
23176 SDNode *MulVar = AddNode.getOperand(i: 0).getNode();
23177
23178 // OtherOp is what we're multiplying against the constant.
23179 if (User->getOperand(Num: 0) == ConstNode)
23180 OtherOp = User->getOperand(Num: 1).getNode();
23181 else
23182 OtherOp = User->getOperand(Num: 0).getNode();
23183
23184 // Check to see if multiply is with the same operand of our "add".
23185 //
23186 // ConstNode = CONST
23187 // User = ConstNode * A <-- visiting User. OtherOp is A.
23188 // ...
23189 // AddNode = (A + c1) <-- MulVar is A.
23190 // = AddNode * ConstNode <-- current visiting instruction.
23191 //
23192 // If we make this transformation, we will have a common
23193 // multiply (ConstNode * A) that we can save.
23194 if (OtherOp == MulVar)
23195 return true;
23196
23197 // Now check to see if a future expansion will give us a common
23198 // multiply.
23199 //
23200 // ConstNode = CONST
23201 // AddNode = (A + c1)
23202 // ... = AddNode * ConstNode <-- current visiting instruction.
23203 // ...
23204 // OtherOp = (A + c2)
23205 // User = OtherOp * ConstNode <-- visiting User.
23206 //
23207 // If we make this transformation, we will have a common
23208 // multiply (CONST * A) after we also do the same transformation
23209 // to the "t2" instruction.
23210 if (OtherOp->getOpcode() == ISD::ADD &&
23211 DAG.isConstantIntBuildVectorOrConstantInt(N: OtherOp->getOperand(Num: 1)) &&
23212 OtherOp->getOperand(Num: 0).getNode() == MulVar)
23213 return true;
23214 }
23215 }
23216
23217 // Didn't find a case where this would be profitable.
23218 return false;
23219}
23220
23221SDValue DAGCombiner::getMergeStoreChains(SmallVectorImpl<MemOpLink> &StoreNodes,
23222 unsigned NumStores) {
23223 SmallVector<SDValue, 8> Chains;
23224 SmallPtrSet<const SDNode *, 8> Visited;
23225 SDLoc StoreDL(StoreNodes[0].MemNode);
23226
23227 for (unsigned i = 0; i < NumStores; ++i) {
23228 Visited.insert(Ptr: StoreNodes[i].MemNode);
23229 }
23230
23231 // don't include nodes that are children or repeated nodes.
23232 for (unsigned i = 0; i < NumStores; ++i) {
23233 if (Visited.insert(Ptr: StoreNodes[i].MemNode->getChain().getNode()).second)
23234 Chains.push_back(Elt: StoreNodes[i].MemNode->getChain());
23235 }
23236
23237 assert(!Chains.empty() && "Chain should have generated a chain");
23238 return DAG.getTokenFactor(DL: StoreDL, Vals&: Chains);
23239}
23240
23241bool DAGCombiner::hasSameUnderlyingObj(ArrayRef<MemOpLink> StoreNodes) {
23242 const Value *UnderlyingObj = nullptr;
23243 for (const auto &MemOp : StoreNodes) {
23244 const MachineMemOperand *MMO = MemOp.MemNode->getMemOperand();
23245 // Pseudo value like stack frame has its own frame index and size, should
23246 // not use the first store's frame index for other frames.
23247 if (MMO->getPseudoValue())
23248 return false;
23249
23250 if (!MMO->getValue())
23251 return false;
23252
23253 const Value *Obj = getUnderlyingObject(V: MMO->getValue());
23254
23255 if (UnderlyingObj && UnderlyingObj != Obj)
23256 return false;
23257
23258 if (!UnderlyingObj)
23259 UnderlyingObj = Obj;
23260 }
23261
23262 return true;
23263}
23264
23265bool DAGCombiner::mergeStoresOfConstantsOrVecElts(
23266 SmallVectorImpl<MemOpLink> &StoreNodes, EVT MemVT, unsigned NumStores,
23267 bool IsConstantSrc, bool UseVector, bool UseTrunc) {
23268 // Make sure we have something to merge.
23269 if (NumStores < 2)
23270 return false;
23271
23272 assert((!UseTrunc || !UseVector) &&
23273 "This optimization cannot emit a vector truncating store");
23274
23275 // The latest Node in the DAG.
23276 SDLoc DL(StoreNodes[0].MemNode);
23277
23278 TypeSize ElementSizeBits = MemVT.getStoreSizeInBits();
23279 unsigned SizeInBits = NumStores * ElementSizeBits;
23280 unsigned NumMemElts = MemVT.isVector() ? MemVT.getVectorNumElements() : 1;
23281
23282 std::optional<MachineMemOperand::Flags> Flags;
23283 AAMDNodes AAInfo;
23284 for (unsigned I = 0; I != NumStores; ++I) {
23285 StoreSDNode *St = cast<StoreSDNode>(Val: StoreNodes[I].MemNode);
23286 if (!Flags) {
23287 Flags = St->getMemOperand()->getFlags();
23288 AAInfo = St->getAAInfo();
23289 continue;
23290 }
23291 // Skip merging if there's an inconsistent flag.
23292 if (Flags != St->getMemOperand()->getFlags())
23293 return false;
23294 // Concatenate AA metadata.
23295 AAInfo = AAInfo.concat(Other: St->getAAInfo());
23296 }
23297
23298 EVT StoreTy;
23299 if (UseVector) {
23300 unsigned Elts = NumStores * NumMemElts;
23301 // Get the type for the merged vector store.
23302 StoreTy = EVT::getVectorVT(Context&: *DAG.getContext(), VT: MemVT.getScalarType(), NumElements: Elts);
23303 } else
23304 StoreTy = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: SizeInBits);
23305
23306 SDValue StoredVal;
23307 if (UseVector) {
23308 if (IsConstantSrc) {
23309 SmallVector<SDValue, 8> BuildVector;
23310 for (unsigned I = 0; I != NumStores; ++I) {
23311 StoreSDNode *St = cast<StoreSDNode>(Val: StoreNodes[I].MemNode);
23312 SDValue Val = St->getValue();
23313 // If constant is of the wrong type, convert it now. This comes up
23314 // when one of our stores was truncating.
23315 if (MemVT != Val.getValueType()) {
23316 Val = peekThroughBitcasts(V: Val);
23317 // Deal with constants of wrong size.
23318 if (ElementSizeBits != Val.getValueSizeInBits()) {
23319 auto *C = dyn_cast<ConstantSDNode>(Val);
23320 if (!C)
23321 // Not clear how to truncate FP values.
23322 // TODO: Handle truncation of build_vector constants
23323 return false;
23324
23325 EVT IntMemVT =
23326 EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: MemVT.getSizeInBits());
23327 Val = DAG.getConstant(Val: C->getAPIntValue()
23328 .zextOrTrunc(width: Val.getValueSizeInBits())
23329 .zextOrTrunc(width: ElementSizeBits),
23330 DL: SDLoc(C), VT: IntMemVT);
23331 }
23332 // Make sure correctly size type is the correct type.
23333 Val = DAG.getBitcast(VT: MemVT, V: Val);
23334 }
23335 BuildVector.push_back(Elt: Val);
23336 }
23337 StoredVal = DAG.getNode(Opcode: MemVT.isVector() ? ISD::CONCAT_VECTORS
23338 : ISD::BUILD_VECTOR,
23339 DL, VT: StoreTy, Ops: BuildVector);
23340 } else {
23341 SmallVector<SDValue, 8> Ops;
23342 for (unsigned i = 0; i < NumStores; ++i) {
23343 StoreSDNode *St = cast<StoreSDNode>(Val: StoreNodes[i].MemNode);
23344 SDValue Val = peekThroughBitcasts(V: St->getValue());
23345 // All operands of BUILD_VECTOR / CONCAT_VECTOR must be of
23346 // type MemVT. If the underlying value is not the correct
23347 // type, but it is an extraction of an appropriate vector we
23348 // can recast Val to be of the correct type. This may require
23349 // converting between EXTRACT_VECTOR_ELT and
23350 // EXTRACT_SUBVECTOR.
23351 if ((MemVT != Val.getValueType()) &&
23352 (Val.getOpcode() == ISD::EXTRACT_VECTOR_ELT ||
23353 Val.getOpcode() == ISD::EXTRACT_SUBVECTOR)) {
23354 EVT MemVTScalarTy = MemVT.getScalarType();
23355 // We may need to add a bitcast here to get types to line up.
23356 if (MemVTScalarTy != Val.getValueType().getScalarType()) {
23357 Val = DAG.getBitcast(VT: MemVT, V: Val);
23358 } else if (MemVT.isVector() &&
23359 Val.getOpcode() == ISD::EXTRACT_VECTOR_ELT) {
23360 Val = DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT: MemVT, Operand: Val);
23361 } else {
23362 unsigned OpC = MemVT.isVector() ? ISD::EXTRACT_SUBVECTOR
23363 : ISD::EXTRACT_VECTOR_ELT;
23364 SDValue Vec = Val.getOperand(i: 0);
23365 SDValue Idx = Val.getOperand(i: 1);
23366 Val = DAG.getNode(Opcode: OpC, DL: SDLoc(Val), VT: MemVT, N1: Vec, N2: Idx);
23367 }
23368 }
23369 Ops.push_back(Elt: Val);
23370 }
23371
23372 // Build the extracted vector elements back into a vector.
23373 StoredVal = DAG.getNode(Opcode: MemVT.isVector() ? ISD::CONCAT_VECTORS
23374 : ISD::BUILD_VECTOR,
23375 DL, VT: StoreTy, Ops);
23376 }
23377 } else {
23378 // We should always use a vector store when merging extracted vector
23379 // elements, so this path implies a store of constants.
23380 assert(IsConstantSrc && "Merged vector elements should use vector store");
23381
23382 APInt StoreInt(SizeInBits, 0);
23383
23384 // Construct a single integer constant which is made of the smaller
23385 // constant inputs.
23386 bool IsLE = DAG.getDataLayout().isLittleEndian();
23387 for (unsigned i = 0; i < NumStores; ++i) {
23388 unsigned Idx = IsLE ? (NumStores - 1 - i) : i;
23389 StoreSDNode *St = cast<StoreSDNode>(Val: StoreNodes[Idx].MemNode);
23390
23391 SDValue Val = St->getValue();
23392 Val = peekThroughBitcasts(V: Val);
23393 StoreInt <<= ElementSizeBits;
23394 if (ConstantSDNode *C = dyn_cast<ConstantSDNode>(Val)) {
23395 StoreInt |= C->getAPIntValue()
23396 .zextOrTrunc(width: ElementSizeBits)
23397 .zextOrTrunc(width: SizeInBits);
23398 } else if (ConstantFPSDNode *C = dyn_cast<ConstantFPSDNode>(Val)) {
23399 StoreInt |= C->getValueAPF()
23400 .bitcastToAPInt()
23401 .zextOrTrunc(width: ElementSizeBits)
23402 .zextOrTrunc(width: SizeInBits);
23403 // If fp truncation is necessary give up for now.
23404 if (MemVT.getSizeInBits() != ElementSizeBits)
23405 return false;
23406 } else if (ISD::isBuildVectorOfConstantSDNodes(N: Val.getNode()) ||
23407 ISD::isBuildVectorOfConstantFPSDNodes(N: Val.getNode())) {
23408 // Not yet handled
23409 return false;
23410 } else {
23411 llvm_unreachable("Invalid constant element type");
23412 }
23413 }
23414
23415 // Create the new Load and Store operations.
23416 StoredVal = DAG.getConstant(Val: StoreInt, DL, VT: StoreTy);
23417 }
23418
23419 LSBaseSDNode *FirstInChain = StoreNodes[0].MemNode;
23420 SDValue NewChain = getMergeStoreChains(StoreNodes, NumStores);
23421 bool CanReusePtrInfo = hasSameUnderlyingObj(StoreNodes);
23422
23423 // make sure we use trunc store if it's necessary to be legal.
23424 // When generate the new widen store, if the first store's pointer info can
23425 // not be reused, discard the pointer info except the address space because
23426 // now the widen store can not be represented by the original pointer info
23427 // which is for the narrow memory object.
23428 SDValue NewStore;
23429 if (!UseTrunc) {
23430 NewStore = DAG.getStore(
23431 Chain: NewChain, dl: DL, Val: StoredVal, Ptr: FirstInChain->getBasePtr(),
23432 PtrInfo: CanReusePtrInfo
23433 ? FirstInChain->getPointerInfo()
23434 : MachinePointerInfo(FirstInChain->getPointerInfo().getAddrSpace()),
23435 Alignment: FirstInChain->getAlign(), MMOFlags: *Flags, Metadata: AAInfo);
23436 } else { // Must be realized as a trunc store
23437 EVT LegalizedStoredValTy =
23438 TLI.getTypeToTransformTo(Context&: *DAG.getContext(), VT: StoredVal.getValueType());
23439 unsigned LegalizedStoreSize = LegalizedStoredValTy.getSizeInBits();
23440 ConstantSDNode *C = cast<ConstantSDNode>(Val&: StoredVal);
23441 SDValue ExtendedStoreVal =
23442 DAG.getConstant(Val: C->getAPIntValue().zextOrTrunc(width: LegalizedStoreSize), DL,
23443 VT: LegalizedStoredValTy);
23444 NewStore = DAG.getTruncStore(
23445 Chain: NewChain, dl: DL, Val: ExtendedStoreVal, Ptr: FirstInChain->getBasePtr(),
23446 PtrInfo: CanReusePtrInfo
23447 ? FirstInChain->getPointerInfo()
23448 : MachinePointerInfo(FirstInChain->getPointerInfo().getAddrSpace()),
23449 SVT: StoredVal.getValueType() /*TVT*/, Alignment: FirstInChain->getAlign(), MMOFlags: *Flags,
23450 Metadata: AAInfo);
23451 }
23452
23453 // Replace all merged stores with the new store.
23454 for (unsigned i = 0; i < NumStores; ++i)
23455 CombineTo(N: StoreNodes[i].MemNode, Res: NewStore);
23456
23457 AddToWorklist(N: NewChain.getNode());
23458 return true;
23459}
23460
23461SDNode *
23462DAGCombiner::getStoreMergeCandidates(StoreSDNode *St,
23463 SmallVectorImpl<MemOpLink> &StoreNodes) {
23464 // This holds the base pointer, index, and the offset in bytes from the base
23465 // pointer. We must have a base and an offset. Do not handle stores to undef
23466 // base pointers.
23467 BaseIndexOffset BasePtr = BaseIndexOffset::match(N: St, DAG);
23468 if (!BasePtr.getBase().getNode() || BasePtr.getBase().isUndef())
23469 return nullptr;
23470
23471 SDValue Val = peekThroughBitcasts(V: St->getValue());
23472 StoreSource StoreSrc = getStoreSource(StoreVal: Val);
23473 assert(StoreSrc != StoreSource::Unknown && "Expected known source for store");
23474
23475 // Match on loadbaseptr if relevant.
23476 EVT MemVT = St->getMemoryVT();
23477 BaseIndexOffset LBasePtr;
23478 EVT LoadVT;
23479 if (StoreSrc == StoreSource::Load) {
23480 auto *Ld = cast<LoadSDNode>(Val);
23481 LBasePtr = BaseIndexOffset::match(N: Ld, DAG);
23482 LoadVT = Ld->getMemoryVT();
23483 // Load and store should be the same type.
23484 if (MemVT != LoadVT)
23485 return nullptr;
23486 // Loads must only have one use.
23487 if (!Ld->hasNUsesOfValue(NUses: 1, Value: 0))
23488 return nullptr;
23489 // The memory operands must not be volatile/indexed/atomic.
23490 // TODO: May be able to relax for unordered atomics (see D66309)
23491 if (!Ld->isSimple() || Ld->isIndexed())
23492 return nullptr;
23493 }
23494 auto CandidateMatch = [&](StoreSDNode *Other, BaseIndexOffset &Ptr,
23495 int64_t &Offset) -> bool {
23496 // The memory operands must not be volatile/indexed/atomic.
23497 // TODO: May be able to relax for unordered atomics (see D66309)
23498 if (!Other->isSimple() || Other->isIndexed())
23499 return false;
23500 // Don't mix temporal stores with non-temporal stores.
23501 if (St->isNonTemporal() != Other->isNonTemporal())
23502 return false;
23503 if (!TLI.areTwoSDNodeTargetMMOFlagsMergeable(NodeX: *St, NodeY: *Other))
23504 return false;
23505 SDValue OtherBC = peekThroughBitcasts(V: Other->getValue());
23506 // Allow merging constants of different types as integers.
23507 bool NoTypeMatch = (MemVT.isInteger()) ? !MemVT.bitsEq(VT: Other->getMemoryVT())
23508 : Other->getMemoryVT() != MemVT;
23509 switch (StoreSrc) {
23510 case StoreSource::Load: {
23511 if (NoTypeMatch)
23512 return false;
23513 // The Load's Base Ptr must also match.
23514 auto *OtherLd = dyn_cast<LoadSDNode>(Val&: OtherBC);
23515 if (!OtherLd)
23516 return false;
23517 BaseIndexOffset LPtr = BaseIndexOffset::match(N: OtherLd, DAG);
23518 if (LoadVT != OtherLd->getMemoryVT())
23519 return false;
23520 // Loads must only have one use.
23521 if (!OtherLd->hasNUsesOfValue(NUses: 1, Value: 0))
23522 return false;
23523 // The memory operands must not be volatile/indexed/atomic.
23524 // TODO: May be able to relax for unordered atomics (see D66309)
23525 if (!OtherLd->isSimple() || OtherLd->isIndexed())
23526 return false;
23527 // Don't mix temporal loads with non-temporal loads.
23528 if (cast<LoadSDNode>(Val)->isNonTemporal() != OtherLd->isNonTemporal())
23529 return false;
23530 if (!TLI.areTwoSDNodeTargetMMOFlagsMergeable(NodeX: *cast<LoadSDNode>(Val),
23531 NodeY: *OtherLd))
23532 return false;
23533 if (!(LBasePtr.equalBaseIndex(Other: LPtr, DAG)))
23534 return false;
23535 break;
23536 }
23537 case StoreSource::Constant:
23538 if (NoTypeMatch)
23539 return false;
23540 if (getStoreSource(StoreVal: OtherBC) != StoreSource::Constant)
23541 return false;
23542 break;
23543 case StoreSource::Extract:
23544 // Do not merge truncated stores here.
23545 if (Other->isTruncatingStore())
23546 return false;
23547 if (!MemVT.bitsEq(VT: OtherBC.getValueType()))
23548 return false;
23549 if (OtherBC.getOpcode() != ISD::EXTRACT_VECTOR_ELT &&
23550 OtherBC.getOpcode() != ISD::EXTRACT_SUBVECTOR)
23551 return false;
23552 break;
23553 default:
23554 llvm_unreachable("Unhandled store source for merging");
23555 }
23556 Ptr = BaseIndexOffset::match(N: Other, DAG);
23557 return (BasePtr.equalBaseIndex(Other: Ptr, DAG, Off&: Offset));
23558 };
23559
23560 // We are looking for a root node which is an ancestor to all mergable
23561 // stores. We search up through a load, to our root and then down
23562 // through all children. For instance we will find Store{1,2,3} if
23563 // St is Store1, Store2. or Store3 where the root is not a load
23564 // which always true for nonvolatile ops. TODO: Expand
23565 // the search to find all valid candidates through multiple layers of loads.
23566 //
23567 // Root
23568 // |-------|-------|
23569 // Load Load Store3
23570 // | |
23571 // Store1 Store2
23572 //
23573 // FIXME: We should be able to climb and
23574 // descend TokenFactors to find candidates as well.
23575
23576 SDNode *RootNode = St->getChain().getNode();
23577 // Bail out if we already analyzed this root node and found nothing.
23578 if (ChainsWithoutMergeableStores.contains(Ptr: RootNode))
23579 return nullptr;
23580
23581 // Check if the pair of StoreNode and the RootNode already bail out many
23582 // times which is over the limit in dependence check.
23583 auto OverLimitInDependenceCheck = [&](SDNode *StoreNode,
23584 SDNode *RootNode) -> bool {
23585 auto RootCount = StoreRootCountMap.find(Val: StoreNode);
23586 return RootCount != StoreRootCountMap.end() &&
23587 RootCount->second.first == RootNode &&
23588 RootCount->second.second > StoreMergeDependenceLimit;
23589 };
23590
23591 auto TryToAddCandidate = [&](SDUse &Use) {
23592 // This must be a chain use.
23593 if (Use.getOperandNo() != 0)
23594 return;
23595 if (auto *OtherStore = dyn_cast<StoreSDNode>(Val: Use.getUser())) {
23596 BaseIndexOffset Ptr;
23597 int64_t PtrDiff;
23598 if (CandidateMatch(OtherStore, Ptr, PtrDiff) &&
23599 !OverLimitInDependenceCheck(OtherStore, RootNode))
23600 StoreNodes.push_back(Elt: MemOpLink(OtherStore, PtrDiff));
23601 }
23602 };
23603
23604 unsigned NumNodesExplored = 0;
23605 const unsigned MaxSearchNodes = 1024;
23606 if (auto *Ldn = dyn_cast<LoadSDNode>(Val: RootNode)) {
23607 RootNode = Ldn->getChain().getNode();
23608 // Bail out if we already analyzed this root node and found nothing.
23609 if (ChainsWithoutMergeableStores.contains(Ptr: RootNode))
23610 return nullptr;
23611 for (auto I = RootNode->use_begin(), E = RootNode->use_end();
23612 I != E && NumNodesExplored < MaxSearchNodes; ++I, ++NumNodesExplored) {
23613 SDNode *User = I->getUser();
23614 if (I->getOperandNo() == 0 && isa<LoadSDNode>(Val: User)) { // walk down chain
23615 for (SDUse &U2 : User->uses())
23616 TryToAddCandidate(U2);
23617 }
23618 // Check stores that depend on the root (e.g. Store 3 in the chart above).
23619 if (I->getOperandNo() == 0 && isa<StoreSDNode>(Val: User)) {
23620 TryToAddCandidate(*I);
23621 }
23622 }
23623 } else {
23624 for (auto I = RootNode->use_begin(), E = RootNode->use_end();
23625 I != E && NumNodesExplored < MaxSearchNodes; ++I, ++NumNodesExplored)
23626 TryToAddCandidate(*I);
23627 }
23628
23629 return RootNode;
23630}
23631
23632// We need to check that merging these stores does not cause a loop in the
23633// DAG. Any store candidate may depend on another candidate indirectly through
23634// its operands. Check in parallel by searching up from operands of candidates.
23635bool DAGCombiner::checkMergeStoreCandidatesForDependencies(
23636 SmallVectorImpl<MemOpLink> &StoreNodes, unsigned NumStores,
23637 SDNode *RootNode) {
23638 // FIXME: We should be able to truncate a full search of
23639 // predecessors by doing a BFS and keeping tabs the originating
23640 // stores from which worklist nodes come from in a similar way to
23641 // TokenFactor simplfication.
23642
23643 SmallPtrSet<const SDNode *, 32> Visited;
23644 SmallVector<const SDNode *, 8> Worklist;
23645
23646 // RootNode is a predecessor to all candidates so we need not search
23647 // past it. Add RootNode (peeking through TokenFactors). Do not count
23648 // these towards size check.
23649
23650 Worklist.push_back(Elt: RootNode);
23651 while (!Worklist.empty()) {
23652 auto N = Worklist.pop_back_val();
23653 if (!Visited.insert(Ptr: N).second)
23654 continue; // Already present in Visited.
23655 if (N->getOpcode() == ISD::TokenFactor) {
23656 for (SDValue Op : N->ops())
23657 Worklist.push_back(Elt: Op.getNode());
23658 }
23659 }
23660
23661 // Don't count pruning nodes towards max.
23662 unsigned int Max = 1024 + Visited.size();
23663 // Search Ops of store candidates.
23664 for (unsigned i = 0; i < NumStores; ++i) {
23665 SDNode *N = StoreNodes[i].MemNode;
23666 // Of the 4 Store Operands:
23667 // * Chain (Op 0) -> We have already considered these
23668 // in candidate selection, but only by following the
23669 // chain dependencies. We could still have a chain
23670 // dependency to a load, that has a non-chain dep to
23671 // another load, that depends on a store, etc. So it is
23672 // possible to have dependencies that consist of a mix
23673 // of chain and non-chain deps, and we need to include
23674 // chain operands in the analysis here..
23675 // * Value (Op 1) -> Cycles may happen (e.g. through load chains)
23676 // * Address (Op 2) -> Merged addresses may only vary by a fixed constant,
23677 // but aren't necessarily fromt the same base node, so
23678 // cycles possible (e.g. via indexed store).
23679 // * (Op 3) -> Represents the pre or post-indexing offset (or undef for
23680 // non-indexed stores). Not constant on all targets (e.g. ARM)
23681 // and so can participate in a cycle.
23682 for (const SDValue &Op : N->op_values())
23683 Worklist.push_back(Elt: Op.getNode());
23684 }
23685 // Search through DAG. We can stop early if we find a store node.
23686 for (unsigned i = 0; i < NumStores; ++i)
23687 if (SDNode::hasPredecessorHelper(N: StoreNodes[i].MemNode, Visited, Worklist,
23688 MaxSteps: Max)) {
23689 // If the searching bail out, record the StoreNode and RootNode in the
23690 // StoreRootCountMap. If we have seen the pair many times over a limit,
23691 // we won't add the StoreNode into StoreNodes set again.
23692 if (Visited.size() >= Max) {
23693 auto &RootCount = StoreRootCountMap[StoreNodes[i].MemNode];
23694 if (RootCount.first == RootNode)
23695 RootCount.second++;
23696 else
23697 RootCount = {RootNode, 1};
23698 }
23699 return false;
23700 }
23701 return true;
23702}
23703
23704bool DAGCombiner::hasCallInLdStChain(StoreSDNode *St, LoadSDNode *Ld) {
23705 SmallPtrSet<const SDNode *, 32> Visited;
23706 SmallVector<std::pair<const SDNode *, bool>, 8> Worklist;
23707 Worklist.emplace_back(Args: St->getChain().getNode(), Args: false);
23708
23709 while (!Worklist.empty()) {
23710 auto [Node, FoundCall] = Worklist.pop_back_val();
23711 if (!Visited.insert(Ptr: Node).second || Node->getNumOperands() == 0)
23712 continue;
23713
23714 switch (Node->getOpcode()) {
23715 case ISD::CALLSEQ_END:
23716 Worklist.emplace_back(Args: Node->getOperand(Num: 0).getNode(), Args: true);
23717 break;
23718 case ISD::TokenFactor:
23719 for (SDValue Op : Node->ops())
23720 Worklist.emplace_back(Args: Op.getNode(), Args&: FoundCall);
23721 break;
23722 case ISD::LOAD:
23723 if (Node == Ld)
23724 return FoundCall;
23725 [[fallthrough]];
23726 default:
23727 assert(Node->getOperand(0).getValueType() == MVT::Other &&
23728 "Invalid chain type");
23729 Worklist.emplace_back(Args: Node->getOperand(Num: 0).getNode(), Args&: FoundCall);
23730 break;
23731 }
23732 }
23733 return false;
23734}
23735
23736unsigned
23737DAGCombiner::getConsecutiveStores(SmallVectorImpl<MemOpLink> &StoreNodes,
23738 int64_t ElementSizeBytes) const {
23739 while (true) {
23740 // Find a store past the width of the first store.
23741 size_t StartIdx = 0;
23742 while ((StartIdx + 1 < StoreNodes.size()) &&
23743 StoreNodes[StartIdx].OffsetFromBase + ElementSizeBytes !=
23744 StoreNodes[StartIdx + 1].OffsetFromBase)
23745 ++StartIdx;
23746
23747 // Bail if we don't have enough candidates to merge.
23748 if (StartIdx + 1 >= StoreNodes.size())
23749 return 0;
23750
23751 // Trim stores that overlapped with the first store.
23752 if (StartIdx)
23753 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + StartIdx);
23754
23755 // Scan the memory operations on the chain and find the first
23756 // non-consecutive store memory address.
23757 unsigned NumConsecutiveStores = 1;
23758 int64_t StartAddress = StoreNodes[0].OffsetFromBase;
23759 // Check that the addresses are consecutive starting from the second
23760 // element in the list of stores.
23761 for (unsigned i = 1, e = StoreNodes.size(); i < e; ++i) {
23762 int64_t CurrAddress = StoreNodes[i].OffsetFromBase;
23763 if (CurrAddress - StartAddress != (ElementSizeBytes * i))
23764 break;
23765 NumConsecutiveStores = i + 1;
23766 }
23767 if (NumConsecutiveStores > 1)
23768 return NumConsecutiveStores;
23769
23770 // There are no consecutive stores at the start of the list.
23771 // Remove the first store and try again.
23772 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + 1);
23773 }
23774}
23775
23776bool DAGCombiner::tryStoreMergeOfConstants(
23777 SmallVectorImpl<MemOpLink> &StoreNodes, unsigned NumConsecutiveStores,
23778 EVT MemVT, SDNode *RootNode, bool AllowVectors) {
23779 LLVMContext &Context = *DAG.getContext();
23780 const DataLayout &DL = DAG.getDataLayout();
23781 int64_t ElementSizeBytes = MemVT.getStoreSize();
23782 unsigned NumMemElts = MemVT.isVector() ? MemVT.getVectorNumElements() : 1;
23783 bool MadeChange = false;
23784
23785 // Store the constants into memory as one consecutive store.
23786 while (NumConsecutiveStores >= 2) {
23787 LSBaseSDNode *FirstInChain = StoreNodes[0].MemNode;
23788 unsigned FirstStoreAS = FirstInChain->getAddressSpace();
23789 Align FirstStoreAlign = FirstInChain->getAlign();
23790 unsigned LastLegalType = 1;
23791 unsigned LastLegalVectorType = 1;
23792 bool LastIntegerTrunc = false;
23793 bool NonZero = false;
23794 unsigned FirstZeroAfterNonZero = NumConsecutiveStores;
23795 for (unsigned i = 0; i < NumConsecutiveStores; ++i) {
23796 StoreSDNode *ST = cast<StoreSDNode>(Val: StoreNodes[i].MemNode);
23797 SDValue StoredVal = ST->getValue();
23798 bool IsElementZero = false;
23799 if (ConstantSDNode *C = dyn_cast<ConstantSDNode>(Val&: StoredVal))
23800 IsElementZero = C->isZero();
23801 else if (ConstantFPSDNode *C = dyn_cast<ConstantFPSDNode>(Val&: StoredVal))
23802 IsElementZero = C->getConstantFPValue()->isNullValue();
23803 else if (ISD::isBuildVectorAllZeros(N: StoredVal.getNode()))
23804 IsElementZero = true;
23805 if (IsElementZero) {
23806 if (NonZero && FirstZeroAfterNonZero == NumConsecutiveStores)
23807 FirstZeroAfterNonZero = i;
23808 }
23809 NonZero |= !IsElementZero;
23810
23811 // Find a legal type for the constant store.
23812 unsigned SizeInBits = (i + 1) * ElementSizeBytes * 8;
23813 EVT StoreTy = EVT::getIntegerVT(Context, BitWidth: SizeInBits);
23814 unsigned IsFast = 0;
23815
23816 // Break early when size is too large to be legal.
23817 if (StoreTy.getSizeInBits() > TLI.getMaximumLegalStoreInBits())
23818 break;
23819
23820 if (TLI.isTypeLegal(VT: StoreTy) &&
23821 TLI.canMergeStoresTo(AS: FirstStoreAS, MemVT: StoreTy,
23822 MF: DAG.getMachineFunction()) &&
23823 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
23824 MMO: *FirstInChain->getMemOperand(), Fast: &IsFast) &&
23825 IsFast) {
23826 LastIntegerTrunc = false;
23827 LastLegalType = i + 1;
23828 // Or check whether a truncstore is legal.
23829 } else if (TLI.getTypeAction(Context, VT: StoreTy) ==
23830 TargetLowering::TypePromoteInteger) {
23831 EVT LegalizedStoredValTy =
23832 TLI.getTypeToTransformTo(Context, VT: StoredVal.getValueType());
23833 if (TLI.isTruncStoreLegal(ValVT: LegalizedStoredValTy, MemVT: StoreTy,
23834 Alignment: FirstStoreAlign, AddrSpace: FirstStoreAS) &&
23835 TLI.canMergeStoresTo(AS: FirstStoreAS, MemVT: LegalizedStoredValTy,
23836 MF: DAG.getMachineFunction()) &&
23837 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
23838 MMO: *FirstInChain->getMemOperand(), Fast: &IsFast) &&
23839 IsFast) {
23840 LastIntegerTrunc = true;
23841 LastLegalType = i + 1;
23842 }
23843 }
23844
23845 // We only use vectors if the target allows it and the function is not
23846 // marked with the noimplicitfloat attribute.
23847 if (TLI.storeOfVectorConstantIsCheap(IsZero: !NonZero, MemVT, NumElem: i + 1, AddrSpace: FirstStoreAS) &&
23848 AllowVectors) {
23849 // Find a legal type for the vector store.
23850 unsigned Elts = (i + 1) * NumMemElts;
23851 EVT Ty = EVT::getVectorVT(Context, VT: MemVT.getScalarType(), NumElements: Elts);
23852 if (TLI.isTypeLegal(VT: Ty) && TLI.isTypeLegal(VT: MemVT) &&
23853 TLI.canMergeStoresTo(AS: FirstStoreAS, MemVT: Ty, MF: DAG.getMachineFunction()) &&
23854 TLI.allowsMemoryAccess(Context, DL, VT: Ty,
23855 MMO: *FirstInChain->getMemOperand(), Fast: &IsFast) &&
23856 IsFast)
23857 LastLegalVectorType = i + 1;
23858 }
23859 }
23860
23861 bool UseVector = (LastLegalVectorType > LastLegalType) && AllowVectors;
23862 unsigned NumElem = (UseVector) ? LastLegalVectorType : LastLegalType;
23863 bool UseTrunc = LastIntegerTrunc && !UseVector;
23864
23865 // Check if we found a legal integer type that creates a meaningful
23866 // merge.
23867 if (NumElem < 2) {
23868 // We know that candidate stores are in order and of correct
23869 // shape. While there is no mergeable sequence from the
23870 // beginning one may start later in the sequence. The only
23871 // reason a merge of size N could have failed where another of
23872 // the same size would not have, is if the alignment has
23873 // improved or we've dropped a non-zero value. Drop as many
23874 // candidates as we can here.
23875 unsigned NumSkip = 1;
23876 while ((NumSkip < NumConsecutiveStores) &&
23877 (NumSkip < FirstZeroAfterNonZero) &&
23878 (StoreNodes[NumSkip].MemNode->getAlign() <= FirstStoreAlign))
23879 NumSkip++;
23880
23881 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumSkip);
23882 NumConsecutiveStores -= NumSkip;
23883 continue;
23884 }
23885
23886 // Check that we can merge these candidates without causing a cycle.
23887 if (!checkMergeStoreCandidatesForDependencies(StoreNodes, NumStores: NumElem,
23888 RootNode)) {
23889 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumElem);
23890 NumConsecutiveStores -= NumElem;
23891 continue;
23892 }
23893
23894 MadeChange |= mergeStoresOfConstantsOrVecElts(StoreNodes, MemVT, NumStores: NumElem,
23895 /*IsConstantSrc*/ true,
23896 UseVector, UseTrunc);
23897
23898 // Remove merged stores for next iteration.
23899 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumElem);
23900 NumConsecutiveStores -= NumElem;
23901 }
23902 return MadeChange;
23903}
23904
23905bool DAGCombiner::tryStoreMergeOfExtracts(
23906 SmallVectorImpl<MemOpLink> &StoreNodes, unsigned NumConsecutiveStores,
23907 EVT MemVT, SDNode *RootNode) {
23908 LLVMContext &Context = *DAG.getContext();
23909 const DataLayout &DL = DAG.getDataLayout();
23910 unsigned NumMemElts = MemVT.isVector() ? MemVT.getVectorNumElements() : 1;
23911 bool MadeChange = false;
23912
23913 // Loop on Consecutive Stores on success.
23914 while (NumConsecutiveStores >= 2) {
23915 LSBaseSDNode *FirstInChain = StoreNodes[0].MemNode;
23916 unsigned FirstStoreAS = FirstInChain->getAddressSpace();
23917 Align FirstStoreAlign = FirstInChain->getAlign();
23918 unsigned NumStoresToMerge = 1;
23919 for (unsigned i = 0; i < NumConsecutiveStores; ++i) {
23920 // Find a legal type for the vector store.
23921 unsigned Elts = (i + 1) * NumMemElts;
23922 EVT Ty = EVT::getVectorVT(Context&: *DAG.getContext(), VT: MemVT.getScalarType(), NumElements: Elts);
23923 unsigned IsFast = 0;
23924
23925 // Break early when size is too large to be legal.
23926 if (Ty.getSizeInBits() > TLI.getMaximumLegalStoreInBits())
23927 break;
23928
23929 if (TLI.isTypeLegal(VT: Ty) &&
23930 TLI.canMergeStoresTo(AS: FirstStoreAS, MemVT: Ty, MF: DAG.getMachineFunction()) &&
23931 TLI.allowsMemoryAccess(Context, DL, VT: Ty,
23932 MMO: *FirstInChain->getMemOperand(), Fast: &IsFast) &&
23933 IsFast)
23934 NumStoresToMerge = i + 1;
23935 }
23936
23937 // Check if we found a legal integer type creating a meaningful
23938 // merge.
23939 if (NumStoresToMerge < 2) {
23940 // We know that candidate stores are in order and of correct
23941 // shape. While there is no mergeable sequence from the
23942 // beginning one may start later in the sequence. The only
23943 // reason a merge of size N could have failed where another of
23944 // the same size would not have, is if the alignment has
23945 // improved. Drop as many candidates as we can here.
23946 unsigned NumSkip = 1;
23947 while ((NumSkip < NumConsecutiveStores) &&
23948 (StoreNodes[NumSkip].MemNode->getAlign() <= FirstStoreAlign))
23949 NumSkip++;
23950
23951 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumSkip);
23952 NumConsecutiveStores -= NumSkip;
23953 continue;
23954 }
23955
23956 // Check that we can merge these candidates without causing a cycle.
23957 if (!checkMergeStoreCandidatesForDependencies(StoreNodes, NumStores: NumStoresToMerge,
23958 RootNode)) {
23959 StoreNodes.erase(CS: StoreNodes.begin(),
23960 CE: StoreNodes.begin() + NumStoresToMerge);
23961 NumConsecutiveStores -= NumStoresToMerge;
23962 continue;
23963 }
23964
23965 MadeChange |= mergeStoresOfConstantsOrVecElts(
23966 StoreNodes, MemVT, NumStores: NumStoresToMerge, /*IsConstantSrc*/ false,
23967 /*UseVector*/ true, /*UseTrunc*/ false);
23968
23969 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumStoresToMerge);
23970 NumConsecutiveStores -= NumStoresToMerge;
23971 }
23972 return MadeChange;
23973}
23974
23975bool DAGCombiner::tryStoreMergeOfLoads(SmallVectorImpl<MemOpLink> &StoreNodes,
23976 unsigned NumConsecutiveStores, EVT MemVT,
23977 SDNode *RootNode, bool AllowVectors,
23978 bool IsNonTemporalStore,
23979 bool IsNonTemporalLoad) {
23980 LLVMContext &Context = *DAG.getContext();
23981 const DataLayout &DL = DAG.getDataLayout();
23982 int64_t ElementSizeBytes = MemVT.getStoreSize();
23983 unsigned NumMemElts = MemVT.isVector() ? MemVT.getVectorNumElements() : 1;
23984 bool MadeChange = false;
23985
23986 // Look for load nodes which are used by the stored values.
23987 SmallVector<MemOpLink, 8> LoadNodes;
23988
23989 // Find acceptable loads. Loads need to have the same chain (token factor),
23990 // must not be zext, volatile, indexed, and they must be consecutive.
23991 BaseIndexOffset LdBasePtr;
23992
23993 for (unsigned i = 0; i < NumConsecutiveStores; ++i) {
23994 StoreSDNode *St = cast<StoreSDNode>(Val: StoreNodes[i].MemNode);
23995 SDValue Val = peekThroughBitcasts(V: St->getValue());
23996 LoadSDNode *Ld = cast<LoadSDNode>(Val);
23997
23998 BaseIndexOffset LdPtr = BaseIndexOffset::match(N: Ld, DAG);
23999 // If this is not the first ptr that we check.
24000 int64_t LdOffset = 0;
24001 if (LdBasePtr.getBase().getNode()) {
24002 // The base ptr must be the same.
24003 if (!LdBasePtr.equalBaseIndex(Other: LdPtr, DAG, Off&: LdOffset))
24004 break;
24005 } else {
24006 // Check that all other base pointers are the same as this one.
24007 LdBasePtr = LdPtr;
24008 }
24009
24010 // We found a potential memory operand to merge.
24011 LoadNodes.push_back(Elt: MemOpLink(Ld, LdOffset));
24012 }
24013
24014 while (NumConsecutiveStores >= 2 && LoadNodes.size() >= 2) {
24015 Align RequiredAlignment;
24016 bool NeedRotate = false;
24017 if (LoadNodes.size() == 2) {
24018 // If we have load/store pair instructions and we only have two values,
24019 // don't bother merging.
24020 if (TLI.hasPairedLoad(MemVT, RequiredAlignment) &&
24021 StoreNodes[0].MemNode->getAlign() >= RequiredAlignment) {
24022 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + 2);
24023 LoadNodes.erase(CS: LoadNodes.begin(), CE: LoadNodes.begin() + 2);
24024 break;
24025 }
24026 // If the loads are reversed, see if we can rotate the halves into place.
24027 int64_t Offset0 = LoadNodes[0].OffsetFromBase;
24028 int64_t Offset1 = LoadNodes[1].OffsetFromBase;
24029 EVT PairVT = EVT::getIntegerVT(Context, BitWidth: ElementSizeBytes * 8 * 2);
24030 if (Offset0 - Offset1 == ElementSizeBytes &&
24031 (hasOperation(Opcode: ISD::ROTL, VT: PairVT) ||
24032 hasOperation(Opcode: ISD::ROTR, VT: PairVT))) {
24033 std::swap(a&: LoadNodes[0], b&: LoadNodes[1]);
24034 NeedRotate = true;
24035 }
24036 }
24037 LSBaseSDNode *FirstInChain = StoreNodes[0].MemNode;
24038 unsigned FirstStoreAS = FirstInChain->getAddressSpace();
24039 Align FirstStoreAlign = FirstInChain->getAlign();
24040 LoadSDNode *FirstLoad = cast<LoadSDNode>(Val: LoadNodes[0].MemNode);
24041
24042 // Scan the memory operations on the chain and find the first
24043 // non-consecutive load memory address. These variables hold the index in
24044 // the store node array.
24045
24046 unsigned LastConsecutiveLoad = 1;
24047
24048 // This variable refers to the size and not index in the array.
24049 unsigned LastLegalVectorType = 1;
24050 unsigned LastLegalIntegerType = 1;
24051 bool isDereferenceable = true;
24052 bool DoIntegerTruncate = false;
24053 int64_t StartAddress = LoadNodes[0].OffsetFromBase;
24054 SDValue LoadChain = FirstLoad->getChain();
24055 for (unsigned i = 1; i < LoadNodes.size(); ++i) {
24056 // All loads must share the same chain.
24057 if (LoadNodes[i].MemNode->getChain() != LoadChain)
24058 break;
24059
24060 int64_t CurrAddress = LoadNodes[i].OffsetFromBase;
24061 if (CurrAddress - StartAddress != (ElementSizeBytes * i))
24062 break;
24063 LastConsecutiveLoad = i;
24064
24065 if (isDereferenceable && !LoadNodes[i].MemNode->isDereferenceable())
24066 isDereferenceable = false;
24067
24068 // Find a legal type for the vector store.
24069 unsigned Elts = (i + 1) * NumMemElts;
24070 EVT StoreTy = EVT::getVectorVT(Context, VT: MemVT.getScalarType(), NumElements: Elts);
24071
24072 // Break early when size is too large to be legal.
24073 if (StoreTy.getSizeInBits() > TLI.getMaximumLegalStoreInBits())
24074 break;
24075
24076 unsigned IsFastSt = 0;
24077 unsigned IsFastLd = 0;
24078 // Don't try vector types if we need a rotate. We may still fail the
24079 // legality checks for the integer type, but we can't handle the rotate
24080 // case with vectors.
24081 // FIXME: We could use a shuffle in place of the rotate.
24082 if (!NeedRotate && TLI.isTypeLegal(VT: StoreTy) &&
24083 TLI.canMergeStoresTo(AS: FirstStoreAS, MemVT: StoreTy,
24084 MF: DAG.getMachineFunction()) &&
24085 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
24086 MMO: *FirstInChain->getMemOperand(), Fast: &IsFastSt) &&
24087 IsFastSt &&
24088 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
24089 MMO: *FirstLoad->getMemOperand(), Fast: &IsFastLd) &&
24090 IsFastLd) {
24091 LastLegalVectorType = i + 1;
24092 }
24093
24094 // Find a legal type for the integer store.
24095 unsigned SizeInBits = (i + 1) * ElementSizeBytes * 8;
24096 StoreTy = EVT::getIntegerVT(Context, BitWidth: SizeInBits);
24097 if (TLI.isTypeLegal(VT: StoreTy) &&
24098 TLI.canMergeStoresTo(AS: FirstStoreAS, MemVT: StoreTy,
24099 MF: DAG.getMachineFunction()) &&
24100 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
24101 MMO: *FirstInChain->getMemOperand(), Fast: &IsFastSt) &&
24102 IsFastSt &&
24103 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
24104 MMO: *FirstLoad->getMemOperand(), Fast: &IsFastLd) &&
24105 IsFastLd) {
24106 LastLegalIntegerType = i + 1;
24107 DoIntegerTruncate = false;
24108 // Or check whether a truncstore and extload is legal.
24109 } else if (TLI.getTypeAction(Context, VT: StoreTy) ==
24110 TargetLowering::TypePromoteInteger) {
24111 EVT LegalizedStoredValTy = TLI.getTypeToTransformTo(Context, VT: StoreTy);
24112 if (TLI.isTruncStoreLegal(ValVT: LegalizedStoredValTy, MemVT: StoreTy,
24113 Alignment: FirstStoreAlign, AddrSpace: FirstStoreAS) &&
24114 TLI.canMergeStoresTo(AS: FirstStoreAS, MemVT: LegalizedStoredValTy,
24115 MF: DAG.getMachineFunction()) &&
24116 TLI.isLoadLegal(ValVT: LegalizedStoredValTy, MemVT: StoreTy,
24117 Alignment: FirstLoad->getAlign(), AddrSpace: FirstLoad->getAddressSpace(),
24118 ExtType: ISD::ZEXTLOAD, Atomic: false) &&
24119 TLI.isLoadLegal(ValVT: LegalizedStoredValTy, MemVT: StoreTy,
24120 Alignment: FirstLoad->getAlign(), AddrSpace: FirstLoad->getAddressSpace(),
24121 ExtType: ISD::SEXTLOAD, Atomic: false) &&
24122 TLI.isLoadLegal(ValVT: LegalizedStoredValTy, MemVT: StoreTy,
24123 Alignment: FirstLoad->getAlign(), AddrSpace: FirstLoad->getAddressSpace(),
24124 ExtType: ISD::EXTLOAD, Atomic: false) &&
24125 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
24126 MMO: *FirstInChain->getMemOperand(), Fast: &IsFastSt) &&
24127 IsFastSt &&
24128 TLI.allowsMemoryAccess(Context, DL, VT: StoreTy,
24129 MMO: *FirstLoad->getMemOperand(), Fast: &IsFastLd) &&
24130 IsFastLd) {
24131 LastLegalIntegerType = i + 1;
24132 DoIntegerTruncate = true;
24133 }
24134 }
24135 }
24136
24137 // Only use vector types if the vector type is larger than the integer
24138 // type. If they are the same, use integers.
24139 bool UseVectorTy =
24140 LastLegalVectorType > LastLegalIntegerType && AllowVectors;
24141 unsigned LastLegalType =
24142 std::max(a: LastLegalVectorType, b: LastLegalIntegerType);
24143
24144 // We add +1 here because the LastXXX variables refer to location while
24145 // the NumElem refers to array/index size.
24146 unsigned NumElem = std::min(a: NumConsecutiveStores, b: LastConsecutiveLoad + 1);
24147 NumElem = std::min(a: LastLegalType, b: NumElem);
24148 Align FirstLoadAlign = FirstLoad->getAlign();
24149
24150 if (NumElem < 2) {
24151 // We know that candidate stores are in order and of correct
24152 // shape. While there is no mergeable sequence from the
24153 // beginning one may start later in the sequence. The only
24154 // reason a merge of size N could have failed where another of
24155 // the same size would not have is if the alignment or either
24156 // the load or store has improved. Drop as many candidates as we
24157 // can here.
24158 unsigned NumSkip = 1;
24159 while ((NumSkip < LoadNodes.size()) &&
24160 (LoadNodes[NumSkip].MemNode->getAlign() <= FirstLoadAlign) &&
24161 (StoreNodes[NumSkip].MemNode->getAlign() <= FirstStoreAlign))
24162 NumSkip++;
24163 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumSkip);
24164 LoadNodes.erase(CS: LoadNodes.begin(), CE: LoadNodes.begin() + NumSkip);
24165 NumConsecutiveStores -= NumSkip;
24166 continue;
24167 }
24168
24169 // Check that we can merge these candidates without causing a cycle.
24170 if (!checkMergeStoreCandidatesForDependencies(StoreNodes, NumStores: NumElem,
24171 RootNode)) {
24172 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumElem);
24173 LoadNodes.erase(CS: LoadNodes.begin(), CE: LoadNodes.begin() + NumElem);
24174 NumConsecutiveStores -= NumElem;
24175 continue;
24176 }
24177
24178 // Find if it is better to use vectors or integers to load and store
24179 // to memory.
24180 EVT JointMemOpVT;
24181 if (UseVectorTy) {
24182 // Find a legal type for the vector store.
24183 unsigned Elts = NumElem * NumMemElts;
24184 JointMemOpVT = EVT::getVectorVT(Context, VT: MemVT.getScalarType(), NumElements: Elts);
24185 } else {
24186 unsigned SizeInBits = NumElem * ElementSizeBytes * 8;
24187 JointMemOpVT = EVT::getIntegerVT(Context, BitWidth: SizeInBits);
24188 }
24189
24190 // Check if there is a call in the load/store chain.
24191 if (!TLI.shouldMergeStoreOfLoadsOverCall(MemVT, JointMemOpVT) &&
24192 hasCallInLdStChain(St: cast<StoreSDNode>(Val: StoreNodes[0].MemNode),
24193 Ld: cast<LoadSDNode>(Val: LoadNodes[0].MemNode))) {
24194 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumElem);
24195 LoadNodes.erase(CS: LoadNodes.begin(), CE: LoadNodes.begin() + NumElem);
24196 NumConsecutiveStores -= NumElem;
24197 continue;
24198 }
24199
24200 SDLoc LoadDL(LoadNodes[0].MemNode);
24201 SDLoc StoreDL(StoreNodes[0].MemNode);
24202
24203 // The merged loads are required to have the same incoming chain, so
24204 // using the first's chain is acceptable.
24205
24206 SDValue NewStoreChain = getMergeStoreChains(StoreNodes, NumStores: NumElem);
24207 bool CanReusePtrInfo = hasSameUnderlyingObj(StoreNodes);
24208 AddToWorklist(N: NewStoreChain.getNode());
24209
24210 MachineMemOperand::Flags LdMMOFlags =
24211 isDereferenceable ? MachineMemOperand::MODereferenceable
24212 : MachineMemOperand::MONone;
24213 if (IsNonTemporalLoad)
24214 LdMMOFlags |= MachineMemOperand::MONonTemporal;
24215
24216 LdMMOFlags |= TLI.getTargetMMOFlags(Node: *FirstLoad);
24217
24218 MachineMemOperand::Flags StMMOFlags = IsNonTemporalStore
24219 ? MachineMemOperand::MONonTemporal
24220 : MachineMemOperand::MONone;
24221
24222 StMMOFlags |= TLI.getTargetMMOFlags(Node: *StoreNodes[0].MemNode);
24223
24224 SDValue NewLoad, NewStore;
24225 if (UseVectorTy || !DoIntegerTruncate) {
24226 NewLoad = DAG.getLoad(
24227 VT: JointMemOpVT, dl: LoadDL, Chain: FirstLoad->getChain(), Ptr: FirstLoad->getBasePtr(),
24228 PtrInfo: FirstLoad->getPointerInfo(), Alignment: FirstLoadAlign, MMOFlags: LdMMOFlags);
24229 SDValue StoreOp = NewLoad;
24230 if (NeedRotate) {
24231 unsigned LoadWidth = ElementSizeBytes * 8 * 2;
24232 assert(JointMemOpVT == EVT::getIntegerVT(Context, LoadWidth) &&
24233 "Unexpected type for rotate-able load pair");
24234 SDValue RotAmt =
24235 DAG.getShiftAmountConstant(Val: LoadWidth / 2, VT: JointMemOpVT, DL: LoadDL);
24236 // Target can convert to the identical ROTR if it does not have ROTL.
24237 StoreOp = DAG.getNode(Opcode: ISD::ROTL, DL: LoadDL, VT: JointMemOpVT, N1: NewLoad, N2: RotAmt);
24238 }
24239 NewStore = DAG.getStore(
24240 Chain: NewStoreChain, dl: StoreDL, Val: StoreOp, Ptr: FirstInChain->getBasePtr(),
24241 PtrInfo: CanReusePtrInfo ? FirstInChain->getPointerInfo()
24242 : MachinePointerInfo(FirstStoreAS),
24243 Alignment: FirstStoreAlign, MMOFlags: StMMOFlags);
24244 } else { // This must be the truncstore/extload case
24245 EVT ExtendedTy =
24246 TLI.getTypeToTransformTo(Context&: *DAG.getContext(), VT: JointMemOpVT);
24247 NewLoad = DAG.getExtLoad(ExtType: ISD::EXTLOAD, dl: LoadDL, VT: ExtendedTy,
24248 Chain: FirstLoad->getChain(), Ptr: FirstLoad->getBasePtr(),
24249 PtrInfo: FirstLoad->getPointerInfo(), MemVT: JointMemOpVT,
24250 Alignment: FirstLoadAlign, MMOFlags: LdMMOFlags);
24251 NewStore = DAG.getTruncStore(
24252 Chain: NewStoreChain, dl: StoreDL, Val: NewLoad, Ptr: FirstInChain->getBasePtr(),
24253 PtrInfo: CanReusePtrInfo ? FirstInChain->getPointerInfo()
24254 : MachinePointerInfo(FirstStoreAS),
24255 SVT: JointMemOpVT, Alignment: FirstInChain->getAlign(),
24256 MMOFlags: FirstInChain->getMemOperand()->getFlags());
24257 }
24258
24259 // Transfer chain users from old loads to the new load.
24260 for (unsigned i = 0; i < NumElem; ++i) {
24261 LoadSDNode *Ld = cast<LoadSDNode>(Val: LoadNodes[i].MemNode);
24262 DAG.ReplaceAllUsesOfValueWith(From: SDValue(Ld, 1),
24263 To: SDValue(NewLoad.getNode(), 1));
24264 }
24265
24266 // Replace all stores with the new store. Recursively remove corresponding
24267 // values if they are no longer used.
24268 for (unsigned i = 0; i < NumElem; ++i) {
24269 SDValue Val = StoreNodes[i].MemNode->getOperand(Num: 1);
24270 CombineTo(N: StoreNodes[i].MemNode, Res: NewStore);
24271 if (Val->use_empty())
24272 recursivelyDeleteUnusedNodes(N: Val.getNode());
24273 }
24274
24275 MadeChange = true;
24276 StoreNodes.erase(CS: StoreNodes.begin(), CE: StoreNodes.begin() + NumElem);
24277 LoadNodes.erase(CS: LoadNodes.begin(), CE: LoadNodes.begin() + NumElem);
24278 NumConsecutiveStores -= NumElem;
24279 }
24280 return MadeChange;
24281}
24282
24283bool DAGCombiner::mergeConsecutiveStores(StoreSDNode *St) {
24284 if (OptLevel == CodeGenOptLevel::None || !EnableStoreMerging)
24285 return false;
24286
24287 // TODO: Extend this function to merge stores of scalable vectors.
24288 // (i.e. two <vscale x 8 x i8> stores can be merged to one <vscale x 16 x i8>
24289 // store since we know <vscale x 16 x i8> is exactly twice as large as
24290 // <vscale x 8 x i8>). Until then, bail out for scalable vectors.
24291 EVT MemVT = St->getMemoryVT();
24292 if (MemVT.isScalableVT())
24293 return false;
24294 if (!MemVT.isSimple() ||
24295 MemVT.getSizeInBits() * 2 > TLI.getMaximumLegalStoreInBits())
24296 return false;
24297
24298 // This function cannot currently deal with non-byte-sized memory sizes.
24299 int64_t ElementSizeBytes = MemVT.getStoreSize();
24300 if (ElementSizeBytes * 8 != (int64_t)MemVT.getSizeInBits())
24301 return false;
24302
24303 // Do not bother looking at stored values that are not constants, loads, or
24304 // extracted vector elements.
24305 SDValue StoredVal = peekThroughBitcasts(V: St->getValue());
24306 const StoreSource StoreSrc = getStoreSource(StoreVal: StoredVal);
24307 if (StoreSrc == StoreSource::Unknown)
24308 return false;
24309
24310 SmallVector<MemOpLink, 8> StoreNodes;
24311 // Find potential store merge candidates by searching through chain sub-DAG
24312 SDNode *RootNode = getStoreMergeCandidates(St, StoreNodes);
24313
24314 // Check if there is anything to merge.
24315 if (StoreNodes.size() < 2)
24316 return false;
24317
24318 // Sort the memory operands according to their distance from the
24319 // base pointer.
24320 llvm::sort(C&: StoreNodes, Comp: [](MemOpLink LHS, MemOpLink RHS) {
24321 return LHS.OffsetFromBase < RHS.OffsetFromBase;
24322 });
24323
24324 bool AllowVectors = !DAG.getMachineFunction().getFunction().hasFnAttribute(
24325 Kind: Attribute::NoImplicitFloat);
24326 bool IsNonTemporalStore = St->isNonTemporal();
24327 bool IsNonTemporalLoad = StoreSrc == StoreSource::Load &&
24328 cast<LoadSDNode>(Val&: StoredVal)->isNonTemporal();
24329
24330 // Store Merge attempts to merge the lowest stores. This generally
24331 // works out as if successful, as the remaining stores are checked
24332 // after the first collection of stores is merged. However, in the
24333 // case that a non-mergeable store is found first, e.g., {p[-2],
24334 // p[0], p[1], p[2], p[3]}, we would fail and miss the subsequent
24335 // mergeable cases. To prevent this, we prune such stores from the
24336 // front of StoreNodes here.
24337 bool MadeChange = false;
24338 while (StoreNodes.size() > 1) {
24339 unsigned NumConsecutiveStores =
24340 getConsecutiveStores(StoreNodes, ElementSizeBytes);
24341 // There are no more stores in the list to examine.
24342 if (NumConsecutiveStores == 0)
24343 return MadeChange;
24344
24345 // We have at least 2 consecutive stores. Try to merge them.
24346 assert(NumConsecutiveStores >= 2 && "Expected at least 2 stores");
24347 switch (StoreSrc) {
24348 case StoreSource::Constant:
24349 MadeChange |= tryStoreMergeOfConstants(StoreNodes, NumConsecutiveStores,
24350 MemVT, RootNode, AllowVectors);
24351 break;
24352
24353 case StoreSource::Extract:
24354 MadeChange |= tryStoreMergeOfExtracts(StoreNodes, NumConsecutiveStores,
24355 MemVT, RootNode);
24356 break;
24357
24358 case StoreSource::Load:
24359 MadeChange |= tryStoreMergeOfLoads(StoreNodes, NumConsecutiveStores,
24360 MemVT, RootNode, AllowVectors,
24361 IsNonTemporalStore, IsNonTemporalLoad);
24362 break;
24363
24364 default:
24365 llvm_unreachable("Unhandled store source type");
24366 }
24367 }
24368
24369 // Remember if we failed to optimize, to save compile time.
24370 if (!MadeChange)
24371 ChainsWithoutMergeableStores.insert(Ptr: RootNode);
24372
24373 return MadeChange;
24374}
24375
24376SDValue DAGCombiner::replaceStoreChain(StoreSDNode *ST, SDValue BetterChain) {
24377 SDLoc SL(ST);
24378 SDValue ReplStore;
24379
24380 // Replace the chain to avoid dependency.
24381 if (ST->isTruncatingStore()) {
24382 ReplStore = DAG.getTruncStore(Chain: BetterChain, dl: SL, Val: ST->getValue(),
24383 Ptr: ST->getBasePtr(), SVT: ST->getMemoryVT(),
24384 MMO: ST->getMemOperand());
24385 } else {
24386 ReplStore = DAG.getStore(Chain: BetterChain, dl: SL, Val: ST->getValue(), Ptr: ST->getBasePtr(),
24387 MMO: ST->getMemOperand());
24388 }
24389
24390 // Create token to keep both nodes around.
24391 SDValue Token = DAG.getNode(Opcode: ISD::TokenFactor, DL: SL,
24392 VT: MVT::Other, N1: ST->getChain(), N2: ReplStore);
24393
24394 // Make sure the new and old chains are cleaned up.
24395 AddToWorklist(N: Token.getNode());
24396
24397 // Don't add users to work list.
24398 return CombineTo(N: ST, Res: Token, AddTo: false);
24399}
24400
24401SDValue DAGCombiner::replaceStoreOfFPConstant(StoreSDNode *ST) {
24402 SDValue Value = ST->getValue();
24403 if (Value.getOpcode() == ISD::TargetConstantFP)
24404 return SDValue();
24405
24406 if (!ISD::isNormalStore(N: ST))
24407 return SDValue();
24408
24409 SDLoc DL(ST);
24410
24411 SDValue Chain = ST->getChain();
24412 SDValue Ptr = ST->getBasePtr();
24413
24414 const ConstantFPSDNode *CFP = cast<ConstantFPSDNode>(Val&: Value);
24415
24416 // NOTE: If the original store is volatile, this transform must not increase
24417 // the number of stores. For example, on x86-32 an f64 can be stored in one
24418 // processor operation but an i64 (which is not legal) requires two. So the
24419 // transform should not be done in this case.
24420
24421 SDValue Tmp;
24422 switch (CFP->getSimpleValueType(ResNo: 0).SimpleTy) {
24423 default:
24424 llvm_unreachable("Unknown FP type");
24425 case MVT::f16: // We don't do this for these yet.
24426 case MVT::bf16:
24427 case MVT::f80:
24428 case MVT::f128:
24429 case MVT::ppcf128:
24430 return SDValue();
24431 case MVT::f32:
24432 if ((isTypeLegal(VT: MVT::i32) && !LegalOperations && ST->isSimple()) ||
24433 TLI.isOperationLegalOrCustom(Op: ISD::STORE, VT: MVT::i32)) {
24434 Tmp = DAG.getConstant(Val: (uint32_t)CFP->getValueAPF().
24435 bitcastToAPInt().getZExtValue(), DL: SDLoc(CFP),
24436 VT: MVT::i32);
24437 return DAG.getStore(Chain, dl: DL, Val: Tmp, Ptr, MMO: ST->getMemOperand());
24438 }
24439
24440 return SDValue();
24441 case MVT::f64:
24442 if ((TLI.isTypeLegal(VT: MVT::i64) && !LegalOperations &&
24443 ST->isSimple()) ||
24444 TLI.isOperationLegalOrCustom(Op: ISD::STORE, VT: MVT::i64)) {
24445 Tmp = DAG.getConstant(Val: CFP->getValueAPF().bitcastToAPInt().
24446 getZExtValue(), DL: SDLoc(CFP), VT: MVT::i64);
24447 return DAG.getStore(Chain, dl: DL, Val: Tmp,
24448 Ptr, MMO: ST->getMemOperand());
24449 }
24450
24451 if (ST->isSimple() && TLI.isOperationLegalOrCustom(Op: ISD::STORE, VT: MVT::i32) &&
24452 !TLI.isFPImmLegal(CFP->getValueAPF(), MVT::f64)) {
24453 // Many FP stores are not made apparent until after legalize, e.g. for
24454 // argument passing. Since this is so common, custom legalize the
24455 // 64-bit integer store into two 32-bit stores.
24456 uint64_t Val = CFP->getValueAPF().bitcastToAPInt().getZExtValue();
24457 SDValue Lo = DAG.getConstant(Val: Val & 0xFFFFFFFF, DL: SDLoc(CFP), VT: MVT::i32);
24458 SDValue Hi = DAG.getConstant(Val: Val >> 32, DL: SDLoc(CFP), VT: MVT::i32);
24459 if (DAG.getDataLayout().isBigEndian())
24460 std::swap(a&: Lo, b&: Hi);
24461
24462 MachineMemOperand::Flags MMOFlags = ST->getMemOperand()->getFlags();
24463 AAMDNodes AAInfo = ST->getAAInfo();
24464
24465 SDValue St0 = DAG.getStore(Chain, dl: DL, Val: Lo, Ptr, PtrInfo: ST->getPointerInfo(),
24466 Alignment: ST->getBaseAlign(), MMOFlags, Metadata: AAInfo);
24467 Ptr = DAG.getMemBasePlusOffset(Base: Ptr, Offset: TypeSize::getFixed(ExactSize: 4), DL);
24468 SDValue St1 = DAG.getStore(Chain, dl: DL, Val: Hi, Ptr,
24469 PtrInfo: ST->getPointerInfo().getWithOffset(O: 4),
24470 Alignment: ST->getBaseAlign(), MMOFlags, Metadata: AAInfo);
24471 return DAG.getNode(Opcode: ISD::TokenFactor, DL, VT: MVT::Other,
24472 N1: St0, N2: St1);
24473 }
24474
24475 return SDValue();
24476 }
24477}
24478
24479// (store (insert_vector_elt (load p), x, i), p) -> (store x, p+offset)
24480//
24481// If a store of a load with an element inserted into it has no other
24482// uses in between the chain, then we can consider the vector store
24483// dead and replace it with just the single scalar element store.
24484SDValue DAGCombiner::replaceStoreOfInsertLoad(StoreSDNode *ST) {
24485 SDLoc DL(ST);
24486 SDValue Value = ST->getValue();
24487 SDValue Ptr = ST->getBasePtr();
24488 SDValue Chain = ST->getChain();
24489 if (Value.getOpcode() != ISD::INSERT_VECTOR_ELT || !Value.hasOneUse())
24490 return SDValue();
24491
24492 SDValue Elt = Value.getOperand(i: 1);
24493 SDValue Idx = Value.getOperand(i: 2);
24494
24495 // If the element isn't byte sized or is implicitly truncated then we can't
24496 // compute an offset.
24497 EVT EltVT = Elt.getValueType();
24498 if (!EltVT.isByteSized() ||
24499 EltVT != Value.getOperand(i: 0).getValueType().getVectorElementType())
24500 return SDValue();
24501
24502 auto *Ld = dyn_cast<LoadSDNode>(Val: Value.getOperand(i: 0));
24503 if (!Ld || Ld->getBasePtr() != Ptr ||
24504 ST->getMemoryVT() != Ld->getMemoryVT() || !ST->isSimple() ||
24505 !ISD::isNormalStore(N: ST) ||
24506 Ld->getAddressSpace() != ST->getAddressSpace() ||
24507 !Chain.reachesChainWithoutSideEffects(Dest: SDValue(Ld, 1)))
24508 return SDValue();
24509
24510 unsigned IsFast;
24511 if (!TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(),
24512 VT: Elt.getValueType(), AddrSpace: ST->getAddressSpace(),
24513 Alignment: ST->getAlign(), Flags: ST->getMemOperand()->getFlags(),
24514 Fast: &IsFast) ||
24515 !IsFast)
24516 return SDValue();
24517
24518 MachinePointerInfo PointerInfo(ST->getAddressSpace());
24519 Align NewAlign;
24520
24521 // If the offset is a known constant then try to recover the pointer
24522 // info
24523 SDValue NewPtr;
24524 if (auto *CIdx = dyn_cast<ConstantSDNode>(Val&: Idx)) {
24525 unsigned COffset = CIdx->getSExtValue() * EltVT.getFixedSizeInBits() / 8;
24526 NewPtr = DAG.getMemBasePlusOffset(Base: Ptr, Offset: TypeSize::getFixed(ExactSize: COffset), DL);
24527 PointerInfo = ST->getPointerInfo().getWithOffset(O: COffset);
24528 NewAlign = ST->getAlign();
24529 } else {
24530 // The original DAG loaded the entire vector from memory, so arithmetic
24531 // within it must be inbounds.
24532 NewPtr = TLI.getInboundsVectorElementPointer(DAG, VecPtr: Ptr, VecVT: Value.getValueType(),
24533 Index: Idx);
24534 // MachinePointerInfo can't represent a variable offset, so use a generic
24535 // MachinePointerInfo and recompute the alignment.
24536 NewAlign = commonAlignment(A: ST->getAlign(), Offset: EltVT.getFixedSizeInBits() / 8);
24537 }
24538
24539 return DAG.getStore(Chain, dl: DL, Val: Elt, Ptr: NewPtr, PtrInfo: PointerInfo, Alignment: NewAlign,
24540 MMOFlags: ST->getMemOperand()->getFlags());
24541}
24542
24543SDValue DAGCombiner::visitATOMIC_STORE(SDNode *N) {
24544 AtomicSDNode *ST = cast<AtomicSDNode>(Val: N);
24545 SDValue Val = ST->getVal();
24546 EVT VT = Val.getValueType();
24547 EVT MemVT = ST->getMemoryVT();
24548
24549 if (MemVT.bitsLT(VT)) { // Is truncating store
24550 APInt TruncDemandedBits = APInt::getLowBitsSet(numBits: VT.getScalarSizeInBits(),
24551 loBitsSet: MemVT.getScalarSizeInBits());
24552 // See if we can simplify the operation with SimplifyDemandedBits, which
24553 // only works if the value has a single use.
24554 if (SimplifyDemandedBits(Op: Val, DemandedBits: TruncDemandedBits))
24555 return SDValue(N, 0);
24556 }
24557
24558 return SDValue();
24559}
24560
24561static SDValue foldToMaskedStore(StoreSDNode *Store, SelectionDAG &DAG,
24562 const SDLoc &Dl) {
24563 if (!Store->isSimple() || !ISD::isNormalStore(N: Store))
24564 return SDValue();
24565
24566 SDValue StoredVal = Store->getValue();
24567 SDValue StorePtr = Store->getBasePtr();
24568 SDValue StoreOffset = Store->getOffset();
24569 EVT VT = Store->getMemoryVT();
24570
24571 // Skip this combine for non-vector types and for <1 x ty> vectors, as they
24572 // will be scalarized later.
24573 if (!VT.isVector() || VT.isScalableVector() || VT.getVectorNumElements() == 1)
24574 return SDValue();
24575
24576 unsigned AddrSpace = Store->getAddressSpace();
24577 Align Alignment = Store->getAlign();
24578 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
24579
24580 // A legal masked store can still be slower than the original sequence,
24581 // e.g. on pre-AVX-512 Zen, which we avoid by checking isTypeDesirableForOp.
24582 if (!TLI.isOperationLegalOrCustom(Op: ISD::MSTORE, VT) ||
24583 !TLI.isTypeDesirableForOp(ISD::MSTORE, VT) ||
24584 !TLI.allowsMisalignedMemoryAccesses(VT, AddrSpace, Alignment))
24585 return SDValue();
24586
24587 SDValue Mask, OtherVec, LoadCh;
24588 unsigned LoadPos;
24589 if (sd_match(N: StoredVal,
24590 P: m_VSelect(Cond: m_Value(N&: Mask), T: m_Value(N&: OtherVec),
24591 F: m_Load(Ch: m_Value(N&: LoadCh), Ptr: m_Specific(N: StorePtr),
24592 Offset: m_Specific(N: StoreOffset))))) {
24593 LoadPos = 2;
24594 } else if (sd_match(N: StoredVal,
24595 P: m_VSelect(Cond: m_Value(N&: Mask),
24596 T: m_Load(Ch: m_Value(N&: LoadCh), Ptr: m_Specific(N: StorePtr),
24597 Offset: m_Specific(N: StoreOffset)),
24598 F: m_Value(N&: OtherVec)))) {
24599 LoadPos = 1;
24600 } else {
24601 return SDValue();
24602 }
24603
24604 auto *Load = cast<LoadSDNode>(Val: StoredVal.getOperand(i: LoadPos));
24605 if (!Load->isSimple() || !ISD::isNormalLoad(N: Load) ||
24606 Load->getAddressSpace() != AddrSpace)
24607 return SDValue();
24608
24609 if (!Store->getChain().reachesChainWithoutSideEffects(Dest: LoadCh))
24610 return SDValue();
24611
24612 if (LoadPos == 1)
24613 Mask = DAG.getNOT(DL: Dl, Val: Mask, VT: Mask.getValueType());
24614
24615 // A masked store follows the IR convention of a vXi1 mask (one bit per
24616 // element). A vselect condition may instead be a wider boolean vector, e.g.
24617 // a vXi32/vXi64 comparison result produced on AVX512 targets without VLX.
24618 // When the matching vXi1 type is legal, narrow the mask to it so that targets
24619 // expecting a vXi1 mask lower it correctly. Targets where vXi1 is illegal
24620 // (e.g. AVX/AVX2) keep the wide mask and lower it as a blend/vmaskmov.
24621 EVT MaskVT = Mask.getValueType();
24622 if (MaskVT.getVectorElementType() != MVT::i1) {
24623 EVT BoolVT = MaskVT.changeVectorElementType(Context&: *DAG.getContext(), EltVT: MVT::i1);
24624 if (TLI.isTypeLegal(VT: BoolVT))
24625 Mask = DAG.getNode(Opcode: ISD::TRUNCATE, DL: Dl, VT: BoolVT, Operand: Mask);
24626 }
24627
24628 return DAG.getMaskedStore(Chain: Store->getChain(), dl: Dl, Val: OtherVec, Base: StorePtr,
24629 Offset: StoreOffset, Mask, MemVT: VT, MMO: Store->getMemOperand(),
24630 AM: Store->getAddressingMode());
24631}
24632
24633// store(concat_vector(truncate, truncate))
24634// --> store(truncate)
24635// store(truncate)
24636SDValue DAGCombiner::combineStoreConcatTruncVector(StoreSDNode *ST) {
24637 if (!LegalTypes)
24638 return SDValue();
24639
24640 if (!ST->isSimple() || ST->isTruncatingStore() || ST->isIndexed())
24641 return SDValue();
24642
24643 SDValue Chain = ST->getChain();
24644 SDValue ConcatVec = ST->getValue();
24645
24646 if (ConcatVec.getOpcode() != ISD::CONCAT_VECTORS ||
24647 ConcatVec.getNumOperands() != 2 || !ConcatVec->hasOneUse())
24648 return SDValue();
24649
24650 SDValue T1 = ConcatVec.getOperand(i: 0);
24651 SDValue T2 = ConcatVec.getOperand(i: 1);
24652 if (T1.getOpcode() != ISD::TRUNCATE || T2.getOpcode() != ISD::TRUNCATE)
24653 return SDValue();
24654
24655 EVT LoMemVT = T1.getValueType();
24656 EVT HiMemVT = T2.getValueType();
24657 if (!LoMemVT.isFixedLengthVector())
24658 return SDValue();
24659
24660 if (!T1.hasOneUse() || !T2.hasOneUse())
24661 return SDValue();
24662
24663 unsigned LoBytes = LoMemVT.getStoreSize();
24664 unsigned HiBytes = HiMemVT.getStoreSize();
24665 Align LoAlign = ST->getAlign();
24666 Align HiAlign = commonAlignment(A: LoAlign, Offset: LoBytes);
24667
24668 if (!TLI.canCombineTruncStore(ValVT: T1.getOperand(i: 0).getValueType(), MemVT: LoMemVT,
24669 Alignment: LoAlign, AddrSpace: ST->getAddressSpace(),
24670 LegalOnly: LegalOperations) ||
24671 !TLI.canCombineTruncStore(ValVT: T2.getOperand(i: 0).getValueType(), MemVT: HiMemVT,
24672 Alignment: HiAlign, AddrSpace: ST->getAddressSpace(),
24673 LegalOnly: LegalOperations))
24674 return SDValue();
24675
24676 SDLoc DL(ST);
24677 SDValue LoPtr = ST->getBasePtr();
24678 SDValue HiPtr =
24679 DAG.getObjectPtrOffset(SL: DL, Ptr: LoPtr, Offset: TypeSize::getFixed(ExactSize: LoBytes));
24680
24681 MachineFunction &MF = DAG.getMachineFunction();
24682 MachineMemOperand *LoMMO =
24683 MF.getMachineMemOperand(MMO: ST->getMemOperand(), Offset: 0, Size: LoBytes);
24684 MachineMemOperand *HiMMO =
24685 MF.getMachineMemOperand(MMO: ST->getMemOperand(), Offset: LoBytes, Size: HiBytes);
24686
24687 SDValue LoSt =
24688 DAG.getTruncStore(Chain, dl: DL, Val: T1.getOperand(i: 0), Ptr: LoPtr, SVT: LoMemVT, MMO: LoMMO);
24689 SDValue HiSt =
24690 DAG.getTruncStore(Chain, dl: DL, Val: T2.getOperand(i: 0), Ptr: HiPtr, SVT: HiMemVT, MMO: HiMMO);
24691
24692 return DAG.getNode(Opcode: ISD::TokenFactor, DL, VT: MVT::Other, N1: LoSt, N2: HiSt);
24693}
24694
24695SDValue DAGCombiner::visitSTORE(SDNode *N) {
24696 StoreSDNode *ST = cast<StoreSDNode>(Val: N);
24697 SDValue Chain = ST->getChain();
24698 SDValue Value = ST->getValue();
24699 SDValue Ptr = ST->getBasePtr();
24700
24701 // If this is a store of a bit convert, store the input value if the
24702 // resultant store does not need a higher alignment than the original.
24703 if (Value.getOpcode() == ISD::BITCAST && !ST->isTruncatingStore() &&
24704 ST->isUnindexed()) {
24705 EVT SVT = Value.getOperand(i: 0).getValueType();
24706 // If the store is volatile, we only want to change the store type if the
24707 // resulting store is legal. Otherwise we might increase the number of
24708 // memory accesses. We don't care if the original type was legal or not
24709 // as we assume software couldn't rely on the number of accesses of an
24710 // illegal type.
24711 // TODO: May be able to relax for unordered atomics (see D66309)
24712 if (((!LegalOperations && ST->isSimple()) ||
24713 TLI.isOperationLegal(Op: ISD::STORE, VT: SVT)) &&
24714 TLI.isStoreBitCastBeneficial(StoreVT: Value.getValueType(), BitcastVT: SVT,
24715 DAG, MMO: *ST->getMemOperand())) {
24716 return DAG.getStore(Chain, dl: SDLoc(N), Val: Value.getOperand(i: 0), Ptr,
24717 MMO: ST->getMemOperand());
24718 }
24719 }
24720
24721 // Turn 'store undef, Ptr' -> nothing.
24722 if (Value.isUndef() && ST->isUnindexed() && !ST->isVolatile())
24723 return Chain;
24724
24725 // Try to infer better alignment information than the store already has.
24726 if (OptLevel != CodeGenOptLevel::None && ST->isUnindexed() &&
24727 !ST->isAtomic()) {
24728 if (MaybeAlign Alignment = DAG.InferPtrAlign(Ptr)) {
24729 if (*Alignment > ST->getAlign() &&
24730 isAligned(Lhs: *Alignment, SizeInBytes: ST->getSrcValueOffset())) {
24731 SDValue NewStore = DAG.getTruncStore(
24732 Chain, dl: SDLoc(N), Val: Value, Ptr, Offset: ST->getOffset(), PtrInfo: ST->getPointerInfo(),
24733 SVT: ST->getMemoryVT(), Alignment: *Alignment, MMOFlags: ST->getMemOperand()->getFlags(),
24734 Metadata: ST->getAAInfo());
24735 // NewStore will always be N as we are only refining the alignment
24736 assert(NewStore.getNode() == N);
24737 (void)NewStore;
24738 }
24739 }
24740 }
24741
24742 // Try transforming a pair floating point load / store ops to integer
24743 // load / store ops.
24744 if (SDValue NewST = TransformFPLoadStorePair(N))
24745 return NewST;
24746
24747 // Try transforming several stores into STORE (BSWAP).
24748 if (SDValue Store = mergeTruncStores(N: ST))
24749 return Store;
24750
24751 if (ST->isUnindexed()) {
24752 // Walk up chain skipping non-aliasing memory nodes, on this store and any
24753 // adjacent stores.
24754 if (findBetterNeighborChains(St: ST)) {
24755 // replaceStoreChain uses CombineTo, which handled all of the worklist
24756 // manipulation. Return the original node to not do anything else.
24757 return SDValue(ST, 0);
24758 }
24759 Chain = ST->getChain();
24760 }
24761
24762 if (SDValue R = combineStoreConcatTruncVector(ST))
24763 return R;
24764
24765 // FIXME: is there such a thing as a truncating indexed store?
24766 if (ST->isTruncatingStore() && ST->isUnindexed() &&
24767 Value.getValueType().isInteger() &&
24768 (!isa<ConstantSDNode>(Val: Value) ||
24769 !cast<ConstantSDNode>(Val&: Value)->isOpaque())) {
24770 // Convert a truncating store of a extension into a standard store.
24771 if ((Value.getOpcode() == ISD::ZERO_EXTEND ||
24772 Value.getOpcode() == ISD::SIGN_EXTEND ||
24773 Value.getOpcode() == ISD::ANY_EXTEND) &&
24774 Value.getOperand(i: 0).getValueType() == ST->getMemoryVT() &&
24775 TLI.isOperationLegalOrCustom(Op: ISD::STORE, VT: ST->getMemoryVT()))
24776 return DAG.getStore(Chain, dl: SDLoc(N), Val: Value.getOperand(i: 0), Ptr,
24777 MMO: ST->getMemOperand());
24778
24779 APInt TruncDemandedBits =
24780 APInt::getLowBitsSet(numBits: Value.getScalarValueSizeInBits(),
24781 loBitsSet: ST->getMemoryVT().getScalarSizeInBits());
24782
24783 // See if we can simplify the operation with SimplifyDemandedBits, which
24784 // only works if the value has a single use.
24785 AddToWorklist(N: Value.getNode());
24786 if (SimplifyDemandedBits(Op: Value, DemandedBits: TruncDemandedBits)) {
24787 // Re-visit the store if anything changed and the store hasn't been merged
24788 // with another node (N is deleted) SimplifyDemandedBits will add Value's
24789 // node back to the worklist if necessary, but we also need to re-visit
24790 // the Store node itself.
24791 if (N->getOpcode() != ISD::DELETED_NODE)
24792 AddToWorklist(N);
24793 return SDValue(N, 0);
24794 }
24795
24796 // Otherwise, see if we can simplify the input to this truncstore with
24797 // knowledge that only the low bits are being used. For example:
24798 // "truncstore (or (shl x, 8), y), i8" -> "truncstore y, i8"
24799 if (SDValue Shorter =
24800 TLI.SimplifyMultipleUseDemandedBits(Op: Value, DemandedBits: TruncDemandedBits, DAG))
24801 return DAG.getTruncStore(Chain, dl: SDLoc(N), Val: Shorter, Ptr, SVT: ST->getMemoryVT(),
24802 MMO: ST->getMemOperand());
24803
24804 // If we're storing a truncated constant, see if we can simplify it.
24805 // TODO: Move this to targetShrinkDemandedConstant?
24806 if (auto *Cst = dyn_cast<ConstantSDNode>(Val&: Value))
24807 if (!Cst->isOpaque()) {
24808 const APInt &CValue = Cst->getAPIntValue();
24809 APInt NewVal = CValue & TruncDemandedBits;
24810 if (NewVal != CValue) {
24811 SDValue Shorter =
24812 DAG.getConstant(Val: NewVal, DL: SDLoc(N), VT: Value.getValueType());
24813 return DAG.getTruncStore(Chain, dl: SDLoc(N), Val: Shorter, Ptr,
24814 SVT: ST->getMemoryVT(), MMO: ST->getMemOperand());
24815 }
24816 }
24817 }
24818
24819 // If this is a load followed by a store to the same location, then the store
24820 // is dead/noop. Peek through any truncates if canCombineTruncStore failed.
24821 // TODO: Add big-endian truncate support with test coverage.
24822 // TODO: Can relax for unordered atomics (see D66309)
24823 SDValue TruncVal = DAG.getDataLayout().isLittleEndian()
24824 ? peekThroughTruncates(V: Value)
24825 : Value;
24826 if (auto *Ld = dyn_cast<LoadSDNode>(Val&: TruncVal)) {
24827 if (Ld->getBasePtr() == Ptr && ST->getMemoryVT() == Ld->getMemoryVT() &&
24828 ST->isUnindexed() && ST->isSimple() &&
24829 Ld->getAddressSpace() == ST->getAddressSpace() &&
24830 // There can't be any side effects between the load and store, such as
24831 // a call or store.
24832 Chain.reachesChainWithoutSideEffects(Dest: SDValue(Ld, 1))) {
24833 // The store is dead, remove it.
24834 return Chain;
24835 }
24836 }
24837
24838 // Try scalarizing vector stores of loads where we only change one element
24839 if (SDValue NewST = replaceStoreOfInsertLoad(ST))
24840 return NewST;
24841
24842 // TODO: Can relax for unordered atomics (see D66309)
24843 if (StoreSDNode *ST1 = dyn_cast<StoreSDNode>(Val&: Chain)) {
24844 if (ST->isUnindexed() && ST->isSimple() &&
24845 ST1->isUnindexed() && ST1->isSimple()) {
24846 if (OptLevel != CodeGenOptLevel::None && ST1->getBasePtr() == Ptr &&
24847 ST1->getValue() == Value && ST->getMemoryVT() == ST1->getMemoryVT() &&
24848 ST->getAddressSpace() == ST1->getAddressSpace()) {
24849 // If this is a store followed by a store with the same value to the
24850 // same location, then the store is dead/noop.
24851 return Chain;
24852 }
24853
24854 if (OptLevel != CodeGenOptLevel::None && ST1->hasOneUse() &&
24855 !ST1->getBasePtr().isUndef() &&
24856 ST->getAddressSpace() == ST1->getAddressSpace()) {
24857 // If we consider two stores and one smaller in size is a scalable
24858 // vector type and another one a bigger size store with a fixed type,
24859 // then we could not allow the scalable store removal because we don't
24860 // know its final size in the end.
24861 if (ST->getMemoryVT().isScalableVector() ||
24862 ST1->getMemoryVT().isScalableVector()) {
24863 if (ST1->getBasePtr() == Ptr &&
24864 TypeSize::isKnownLE(LHS: ST1->getMemoryVT().getStoreSize(),
24865 RHS: ST->getMemoryVT().getStoreSize())) {
24866 CombineTo(N: ST1, Res: ST1->getChain());
24867 return SDValue(N, 0);
24868 }
24869 } else {
24870 const BaseIndexOffset STBase = BaseIndexOffset::match(N: ST, DAG);
24871 const BaseIndexOffset ChainBase = BaseIndexOffset::match(N: ST1, DAG);
24872 // If this is a store who's preceding store to a subset of the current
24873 // location and no one other node is chained to that store we can
24874 // effectively drop the store. Do not remove stores to undef as they
24875 // may be used as data sinks.
24876 if (STBase.contains(DAG, BitSize: ST->getMemoryVT().getFixedSizeInBits(),
24877 Other: ChainBase,
24878 OtherBitSize: ST1->getMemoryVT().getFixedSizeInBits())) {
24879 CombineTo(N: ST1, Res: ST1->getChain());
24880 return SDValue(N, 0);
24881 }
24882 }
24883 }
24884 }
24885 }
24886
24887 // If this is an FP_ROUND or TRUNC followed by a store, fold this into a
24888 // truncating store. We can do this even if this is already a truncstore.
24889 if ((Value.getOpcode() == ISD::FP_ROUND ||
24890 Value.getOpcode() == ISD::TRUNCATE) &&
24891 Value->hasOneUse() && ST->isUnindexed() &&
24892 TLI.canCombineTruncStore(ValVT: Value.getOperand(i: 0).getValueType(),
24893 MemVT: ST->getMemoryVT(), Alignment: ST->getAlign(),
24894 AddrSpace: ST->getAddressSpace(), LegalOnly: LegalOperations)) {
24895 return DAG.getTruncStore(Chain, dl: SDLoc(N), Val: Value.getOperand(i: 0), Ptr,
24896 SVT: ST->getMemoryVT(), MMO: ST->getMemOperand());
24897 }
24898
24899 // Always perform this optimization before types are legal. If the target
24900 // prefers, also try this after legalization to catch stores that were created
24901 // by intrinsics or other nodes.
24902 if (!LegalTypes || (TLI.mergeStoresAfterLegalization(MemVT: ST->getMemoryVT()))) {
24903 while (true) {
24904 // There can be multiple store sequences on the same chain.
24905 // Keep trying to merge store sequences until we are unable to do so
24906 // or until we merge the last store on the chain.
24907 bool Changed = mergeConsecutiveStores(St: ST);
24908 if (!Changed) break;
24909 // Return N as merge only uses CombineTo and no worklist clean
24910 // up is necessary.
24911 if (N->getOpcode() == ISD::DELETED_NODE || !isa<StoreSDNode>(Val: N))
24912 return SDValue(N, 0);
24913 }
24914 }
24915
24916 // Try transforming N to an indexed store.
24917 if (CombineToPreIndexedLoadStore(N) || CombineToPostIndexedLoadStore(N))
24918 return SDValue(N, 0);
24919
24920 // Turn 'store float 1.0, Ptr' -> 'store int 0x12345678, Ptr'
24921 //
24922 // Make sure to do this only after attempting to merge stores in order to
24923 // avoid changing the types of some subset of stores due to visit order,
24924 // preventing their merging.
24925 if (isa<ConstantFPSDNode>(Val: ST->getValue())) {
24926 if (SDValue NewSt = replaceStoreOfFPConstant(ST))
24927 return NewSt;
24928 }
24929
24930 if (SDValue NewSt = splitMergedValStore(ST))
24931 return NewSt;
24932
24933 if (SDValue MaskedStore = foldToMaskedStore(Store: ST, DAG, Dl: SDLoc(N)))
24934 return MaskedStore;
24935
24936 return ReduceLoadOpStoreWidth(N);
24937}
24938
24939SDValue DAGCombiner::visitLIFETIME_END(SDNode *N) {
24940 const auto *LifetimeEnd = cast<LifetimeSDNode>(Val: N);
24941 const BaseIndexOffset LifetimeEndBase(N->getOperand(Num: 1), SDValue(), 0, false);
24942
24943 // We walk up the chains to find stores.
24944 SmallVector<SDValue, 8> Chains = {N->getOperand(Num: 0)};
24945 while (!Chains.empty()) {
24946 SDValue Chain = Chains.pop_back_val();
24947 if (!Chain.hasOneUse())
24948 continue;
24949 switch (Chain.getOpcode()) {
24950 case ISD::TokenFactor:
24951 for (unsigned Nops = Chain.getNumOperands(); Nops;)
24952 Chains.push_back(Elt: Chain.getOperand(i: --Nops));
24953 break;
24954 case ISD::LIFETIME_START:
24955 case ISD::LIFETIME_END:
24956 // We can forward past any lifetime start/end that can be proven not to
24957 // alias the node.
24958 if (!mayAlias(Op0: Chain.getNode(), Op1: N))
24959 Chains.push_back(Elt: Chain.getOperand(i: 0));
24960 break;
24961 case ISD::STORE: {
24962 StoreSDNode *ST = dyn_cast<StoreSDNode>(Val&: Chain);
24963 // TODO: Can relax for unordered atomics (see D66309)
24964 if (!ST->isSimple() || ST->isIndexed())
24965 continue;
24966 const TypeSize StoreSize = ST->getMemoryVT().getStoreSize();
24967 // The bounds of a scalable store are not known until runtime, so this
24968 // store cannot be elided.
24969 if (StoreSize.isScalable())
24970 continue;
24971 const BaseIndexOffset StoreBase = BaseIndexOffset::match(N: ST, DAG);
24972 // If we store purely within object bounds just before its lifetime ends,
24973 // we can remove the store.
24974 MachineFrameInfo &MFI = DAG.getMachineFunction().getFrameInfo();
24975 if (LifetimeEndBase.contains(
24976 DAG, BitSize: MFI.getObjectSize(ObjectIdx: LifetimeEnd->getFrameIndex()) * 8,
24977 Other: StoreBase, OtherBitSize: StoreSize.getFixedValue() * 8)) {
24978 LLVM_DEBUG(dbgs() << "\nRemoving store:"; StoreBase.dump();
24979 dbgs() << "\nwithin LIFETIME_END of : ";
24980 LifetimeEndBase.dump(); dbgs() << "\n");
24981 CombineTo(N: ST, Res: ST->getChain());
24982 return SDValue(N, 0);
24983 }
24984 }
24985 }
24986 }
24987 return SDValue();
24988}
24989
24990/// For the instruction sequence of store below, F and I values
24991/// are bundled together as an i64 value before being stored into memory.
24992/// Sometimes it is more efficent to generate separate stores for F and I,
24993/// which can remove the bitwise instructions or sink them to colder places.
24994///
24995/// (store (or (zext (bitcast F to i32) to i64),
24996/// (shl (zext I to i64), 32)), addr) -->
24997/// (store F, addr) and (store I, addr+4)
24998///
24999/// Similarly, splitting for other merged store can also be beneficial, like:
25000/// For pair of {i32, i32}, i64 store --> two i32 stores.
25001/// For pair of {i32, i16}, i64 store --> two i32 stores.
25002/// For pair of {i16, i16}, i32 store --> two i16 stores.
25003/// For pair of {i16, i8}, i32 store --> two i16 stores.
25004/// For pair of {i8, i8}, i16 store --> two i8 stores.
25005///
25006/// We allow each target to determine specifically which kind of splitting is
25007/// supported.
25008///
25009/// The store patterns are commonly seen from the simple code snippet below
25010/// if only std::make_pair(...) is sroa transformed before inlined into hoo.
25011/// void goo(const std::pair<int, float> &);
25012/// hoo() {
25013/// ...
25014/// goo(std::make_pair(tmp, ftmp));
25015/// ...
25016/// }
25017///
25018SDValue DAGCombiner::splitMergedValStore(StoreSDNode *ST) {
25019 if (OptLevel == CodeGenOptLevel::None)
25020 return SDValue();
25021
25022 // Can't change the number of memory accesses for a volatile store or break
25023 // atomicity for an atomic one.
25024 if (!ST->isSimple())
25025 return SDValue();
25026
25027 SDValue Val = ST->getValue();
25028 SDLoc DL(ST);
25029
25030 // Match OR operand.
25031 if (!Val.getValueType().isScalarInteger() || Val.getOpcode() != ISD::OR)
25032 return SDValue();
25033
25034 // Match SHL operand and get Lower and Higher parts of Val.
25035 SDValue Op1 = Val.getOperand(i: 0);
25036 SDValue Op2 = Val.getOperand(i: 1);
25037 SDValue Lo, Hi;
25038 if (Op1.getOpcode() != ISD::SHL) {
25039 std::swap(a&: Op1, b&: Op2);
25040 if (Op1.getOpcode() != ISD::SHL)
25041 return SDValue();
25042 }
25043 Lo = Op2;
25044 Hi = Op1.getOperand(i: 0);
25045 if (!Op1.hasOneUse())
25046 return SDValue();
25047
25048 // Match shift amount to HalfValBitSize.
25049 unsigned HalfValBitSize = Val.getValueSizeInBits() / 2;
25050 ConstantSDNode *ShAmt = dyn_cast<ConstantSDNode>(Val: Op1.getOperand(i: 1));
25051 if (!ShAmt || ShAmt->getAPIntValue() != HalfValBitSize)
25052 return SDValue();
25053
25054 // Lo and Hi are zero-extended from int with size less equal than 32
25055 // to i64.
25056 if (Lo.getOpcode() != ISD::ZERO_EXTEND || !Lo.hasOneUse() ||
25057 !Lo.getOperand(i: 0).getValueType().isScalarInteger() ||
25058 Lo.getOperand(i: 0).getValueSizeInBits() > HalfValBitSize ||
25059 Hi.getOpcode() != ISD::ZERO_EXTEND || !Hi.hasOneUse() ||
25060 !Hi.getOperand(i: 0).getValueType().isScalarInteger() ||
25061 Hi.getOperand(i: 0).getValueSizeInBits() > HalfValBitSize)
25062 return SDValue();
25063
25064 // Use the EVT of low and high parts before bitcast as the input
25065 // of target query.
25066 EVT LowTy = (Lo.getOperand(i: 0).getOpcode() == ISD::BITCAST)
25067 ? Lo.getOperand(i: 0).getValueType()
25068 : Lo.getValueType();
25069 EVT HighTy = (Hi.getOperand(i: 0).getOpcode() == ISD::BITCAST)
25070 ? Hi.getOperand(i: 0).getValueType()
25071 : Hi.getValueType();
25072 if (!TLI.isMultiStoresCheaperThanBitsMerge(LTy: LowTy, HTy: HighTy))
25073 return SDValue();
25074
25075 // Start to split store.
25076 MachineMemOperand::Flags MMOFlags = ST->getMemOperand()->getFlags();
25077 AAMDNodes AAInfo = ST->getAAInfo();
25078
25079 // Change the sizes of Lo and Hi's value types to HalfValBitSize.
25080 EVT VT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: HalfValBitSize);
25081 Lo = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: Lo.getOperand(i: 0));
25082 Hi = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT, Operand: Hi.getOperand(i: 0));
25083
25084 SDValue Chain = ST->getChain();
25085 SDValue Ptr = ST->getBasePtr();
25086 // Lower value store.
25087 SDValue St0 = DAG.getStore(Chain, dl: DL, Val: Lo, Ptr, PtrInfo: ST->getPointerInfo(),
25088 Alignment: ST->getBaseAlign(), MMOFlags, Metadata: AAInfo);
25089 Ptr =
25090 DAG.getMemBasePlusOffset(Base: Ptr, Offset: TypeSize::getFixed(ExactSize: HalfValBitSize / 8), DL);
25091 // Higher value store.
25092 SDValue St1 = DAG.getStore(
25093 Chain: St0, dl: DL, Val: Hi, Ptr, PtrInfo: ST->getPointerInfo().getWithOffset(O: HalfValBitSize / 8),
25094 Alignment: ST->getBaseAlign(), MMOFlags, Metadata: AAInfo);
25095 return St1;
25096}
25097
25098// Merge an insertion into an existing shuffle:
25099// (insert_vector_elt (vector_shuffle X, Y, Mask),
25100// .(extract_vector_elt X, N), InsIndex)
25101// --> (vector_shuffle X, Y, NewMask)
25102// and variations where shuffle operands may be CONCAT_VECTORS.
25103static bool mergeEltWithShuffle(SDValue &X, SDValue &Y, ArrayRef<int> Mask,
25104 SmallVectorImpl<int> &NewMask, SDValue Elt,
25105 unsigned InsIndex) {
25106 if (Elt.getOpcode() != ISD::EXTRACT_VECTOR_ELT ||
25107 !isa<ConstantSDNode>(Val: Elt.getOperand(i: 1)))
25108 return false;
25109
25110 // Vec's operand 0 is using indices from 0 to N-1 and
25111 // operand 1 from N to 2N - 1, where N is the number of
25112 // elements in the vectors.
25113 SDValue InsertVal0 = Elt.getOperand(i: 0);
25114 int ElementOffset = -1;
25115
25116 // We explore the inputs of the shuffle in order to see if we find the
25117 // source of the extract_vector_elt. If so, we can use it to modify the
25118 // shuffle rather than perform an insert_vector_elt.
25119 SmallVector<std::pair<int, SDValue>, 8> ArgWorkList;
25120 ArgWorkList.emplace_back(Args: Mask.size(), Args&: Y);
25121 ArgWorkList.emplace_back(Args: 0, Args&: X);
25122
25123 while (!ArgWorkList.empty()) {
25124 int ArgOffset;
25125 SDValue ArgVal;
25126 std::tie(args&: ArgOffset, args&: ArgVal) = ArgWorkList.pop_back_val();
25127
25128 if (ArgVal == InsertVal0) {
25129 ElementOffset = ArgOffset;
25130 break;
25131 }
25132
25133 // Peek through concat_vector.
25134 if (ArgVal.getOpcode() == ISD::CONCAT_VECTORS) {
25135 int CurrentArgOffset =
25136 ArgOffset + ArgVal.getValueType().getVectorNumElements();
25137 int Step = ArgVal.getOperand(i: 0).getValueType().getVectorNumElements();
25138 for (SDValue Op : reverse(C: ArgVal->ops())) {
25139 CurrentArgOffset -= Step;
25140 ArgWorkList.emplace_back(Args&: CurrentArgOffset, Args&: Op);
25141 }
25142
25143 // Make sure we went through all the elements and did not screw up index
25144 // computation.
25145 assert(CurrentArgOffset == ArgOffset);
25146 }
25147 }
25148
25149 // If we failed to find a match, see if we can replace an UNDEF shuffle
25150 // operand.
25151 if (ElementOffset == -1) {
25152 if (!Y.isUndef() || InsertVal0.getValueType() != Y.getValueType())
25153 return false;
25154 ElementOffset = Mask.size();
25155 Y = InsertVal0;
25156 }
25157
25158 NewMask.assign(in_start: Mask.begin(), in_end: Mask.end());
25159 NewMask[InsIndex] = ElementOffset + Elt.getConstantOperandVal(i: 1);
25160 assert(NewMask[InsIndex] < (int)(2 * Mask.size()) && NewMask[InsIndex] >= 0 &&
25161 "NewMask[InsIndex] is out of bound");
25162 return true;
25163}
25164
25165// Merge an insertion into an existing shuffle:
25166// (insert_vector_elt (vector_shuffle X, Y), (extract_vector_elt X, N),
25167// InsIndex)
25168// --> (vector_shuffle X, Y) and variations where shuffle operands may be
25169// CONCAT_VECTORS.
25170SDValue DAGCombiner::mergeInsertEltWithShuffle(SDNode *N, unsigned InsIndex) {
25171 assert(N->getOpcode() == ISD::INSERT_VECTOR_ELT &&
25172 "Expected extract_vector_elt");
25173 SDValue InsertVal = N->getOperand(Num: 1);
25174 SDValue Vec = N->getOperand(Num: 0);
25175
25176 auto *SVN = dyn_cast<ShuffleVectorSDNode>(Val&: Vec);
25177 if (!SVN || !Vec.hasOneUse())
25178 return SDValue();
25179
25180 ArrayRef<int> Mask = SVN->getMask();
25181 SDValue X = Vec.getOperand(i: 0);
25182 SDValue Y = Vec.getOperand(i: 1);
25183
25184 SmallVector<int, 16> NewMask(Mask);
25185 if (mergeEltWithShuffle(X, Y, Mask, NewMask, Elt: InsertVal, InsIndex)) {
25186 SDValue LegalShuffle = TLI.buildLegalVectorShuffle(
25187 VT: Vec.getValueType(), DL: SDLoc(N), N0: X, N1: Y, Mask: NewMask, DAG);
25188 if (LegalShuffle)
25189 return LegalShuffle;
25190 }
25191
25192 return SDValue();
25193}
25194
25195// Convert a disguised subvector insertion into a shuffle:
25196// insert_vector_elt V, (bitcast X from vector type), IdxC -->
25197// bitcast(shuffle (bitcast V), (extended X), Mask)
25198// Note: We do not use an insert_subvector node because that requires a
25199// legal subvector type.
25200SDValue DAGCombiner::combineInsertEltToShuffle(SDNode *N, unsigned InsIndex) {
25201 assert(N->getOpcode() == ISD::INSERT_VECTOR_ELT &&
25202 "Expected extract_vector_elt");
25203 SDValue InsertVal = N->getOperand(Num: 1);
25204
25205 if (InsertVal.getOpcode() != ISD::BITCAST || !InsertVal.hasOneUse() ||
25206 !InsertVal.getOperand(i: 0).getValueType().isVector())
25207 return SDValue();
25208
25209 SDValue SubVec = InsertVal.getOperand(i: 0);
25210 SDValue DestVec = N->getOperand(Num: 0);
25211 EVT SubVecVT = SubVec.getValueType();
25212 EVT VT = DestVec.getValueType();
25213 unsigned NumSrcElts = SubVecVT.getVectorNumElements();
25214 // Bail out if the inserted value is larger than the vector element, as
25215 // insert_vector_elt performs an implicit truncation in this case.
25216 if (InsertVal.getValueType() != VT.getVectorElementType())
25217 return SDValue();
25218 // If the source only has a single vector element, the cost of creating adding
25219 // it to a vector is likely to exceed the cost of a insert_vector_elt.
25220 if (NumSrcElts == 1)
25221 return SDValue();
25222 unsigned ExtendRatio = VT.getSizeInBits() / SubVecVT.getSizeInBits();
25223 unsigned NumMaskVals = ExtendRatio * NumSrcElts;
25224
25225 // Step 1: Create a shuffle mask that implements this insert operation. The
25226 // vector that we are inserting into will be operand 0 of the shuffle, so
25227 // those elements are just 'i'. The inserted subvector is in the first
25228 // positions of operand 1 of the shuffle. Example:
25229 // insert v4i32 V, (v2i16 X), 2 --> shuffle v8i16 V', X', {0,1,2,3,8,9,6,7}
25230 SmallVector<int, 16> Mask(NumMaskVals);
25231 for (unsigned i = 0; i != NumMaskVals; ++i) {
25232 if (i / NumSrcElts == InsIndex)
25233 Mask[i] = (i % NumSrcElts) + NumMaskVals;
25234 else
25235 Mask[i] = i;
25236 }
25237
25238 // Bail out if the target can not handle the shuffle we want to create, or
25239 // would create an illegal-typed shuffle after type legalization.
25240 EVT SubVecEltVT = SubVecVT.getVectorElementType();
25241 EVT ShufVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: SubVecEltVT, NumElements: NumMaskVals);
25242 if ((LegalTypes && !TLI.isTypeLegal(VT: ShufVT)) ||
25243 !TLI.isShuffleMaskLegal(Mask, ShufVT))
25244 return SDValue();
25245
25246 // Step 2: Create a wide vector from the inserted source vector by appending
25247 // poison elements. This is the same size as our destination vector.
25248 SDLoc DL(N);
25249 SmallVector<SDValue, 8> ConcatOps(ExtendRatio, DAG.getPOISON(VT: SubVecVT));
25250 ConcatOps[0] = SubVec;
25251 SDValue PaddedSubV = DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT: ShufVT, Ops: ConcatOps);
25252
25253 // Step 3: Shuffle in the padded subvector.
25254 SDValue DestVecBC = DAG.getBitcast(VT: ShufVT, V: DestVec);
25255 SDValue Shuf = DAG.getVectorShuffle(VT: ShufVT, dl: DL, N1: DestVecBC, N2: PaddedSubV, Mask);
25256 AddToWorklist(N: PaddedSubV.getNode());
25257 AddToWorklist(N: DestVecBC.getNode());
25258 AddToWorklist(N: Shuf.getNode());
25259 return DAG.getBitcast(VT, V: Shuf);
25260}
25261
25262// Combine insert(shuffle(load, <u,0,1,2>), load, 0) into a single load if
25263// possible and the new load will be quick. We use more loads but less shuffles
25264// and inserts.
25265SDValue DAGCombiner::combineInsertEltToLoad(SDNode *N, unsigned InsIndex) {
25266 EVT VT = N->getValueType(ResNo: 0);
25267
25268 // InsIndex is expected to be the first of last lane.
25269 if (!VT.isFixedLengthVector() ||
25270 (InsIndex != 0 && InsIndex != VT.getVectorNumElements() - 1))
25271 return SDValue();
25272
25273 // Look for a shuffle with the mask u,0,1,2,3,4,5,6 or 1,2,3,4,5,6,7,u
25274 // depending on the InsIndex.
25275 auto *Shuffle = dyn_cast<ShuffleVectorSDNode>(Val: N->getOperand(Num: 0));
25276 SDValue Scalar = N->getOperand(Num: 1);
25277 if (!Shuffle || !all_of(Range: enumerate(First: Shuffle->getMask()), P: [&](auto P) {
25278 return InsIndex == P.index() || P.value() < 0 ||
25279 (InsIndex == 0 && P.value() == (int)P.index() - 1) ||
25280 (InsIndex == VT.getVectorNumElements() - 1 &&
25281 P.value() == (int)P.index() + 1);
25282 }))
25283 return SDValue();
25284
25285 // We optionally skip over an extend so long as both loads are extended in the
25286 // same way from the same type.
25287 unsigned Extend = 0;
25288 if (Scalar.getOpcode() == ISD::ZERO_EXTEND ||
25289 Scalar.getOpcode() == ISD::SIGN_EXTEND ||
25290 Scalar.getOpcode() == ISD::ANY_EXTEND) {
25291 Extend = Scalar.getOpcode();
25292 Scalar = Scalar.getOperand(i: 0);
25293 }
25294
25295 auto *ScalarLoad = dyn_cast<LoadSDNode>(Val&: Scalar);
25296 if (!ScalarLoad)
25297 return SDValue();
25298
25299 SDValue Vec = Shuffle->getOperand(Num: 0);
25300 if (Extend) {
25301 if (Vec.getOpcode() != Extend)
25302 return SDValue();
25303 Vec = Vec.getOperand(i: 0);
25304 }
25305 auto *VecLoad = dyn_cast<LoadSDNode>(Val&: Vec);
25306 if (!VecLoad || Vec.getValueType().getScalarType() != Scalar.getValueType())
25307 return SDValue();
25308
25309 int EltSize = ScalarLoad->getValueType(ResNo: 0).getScalarSizeInBits();
25310 if (EltSize == 0 || EltSize % 8 != 0 || !ScalarLoad->isSimple() ||
25311 !VecLoad->isSimple() || VecLoad->getExtensionType() != ISD::NON_EXTLOAD ||
25312 ScalarLoad->getExtensionType() != ISD::NON_EXTLOAD ||
25313 ScalarLoad->getAddressSpace() != VecLoad->getAddressSpace())
25314 return SDValue();
25315
25316 // Check that the offset between the pointers to produce a single continuous
25317 // load.
25318 if (InsIndex == 0) {
25319 if (!DAG.areNonVolatileConsecutiveLoads(LD: ScalarLoad, Base: VecLoad, Bytes: EltSize / 8,
25320 Dist: -1))
25321 return SDValue();
25322 } else {
25323 if (!DAG.areNonVolatileConsecutiveLoads(
25324 LD: VecLoad, Base: ScalarLoad, Bytes: VT.getVectorNumElements() * EltSize / 8, Dist: -1))
25325 return SDValue();
25326 }
25327
25328 // And that the new unaligned load will be fast.
25329 unsigned IsFast = 0;
25330 Align NewAlign = commonAlignment(A: VecLoad->getAlign(), Offset: EltSize / 8);
25331 if (!TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(),
25332 VT: Vec.getValueType(), AddrSpace: VecLoad->getAddressSpace(),
25333 Alignment: NewAlign, Flags: VecLoad->getMemOperand()->getFlags(),
25334 Fast: &IsFast) ||
25335 !IsFast)
25336 return SDValue();
25337
25338 // Calculate the new Ptr and create the new load.
25339 SDLoc DL(N);
25340 SDValue Ptr = ScalarLoad->getBasePtr();
25341 if (InsIndex != 0)
25342 Ptr = DAG.getNode(Opcode: ISD::ADD, DL, VT: Ptr.getValueType(), N1: VecLoad->getBasePtr(),
25343 N2: DAG.getConstant(Val: EltSize / 8, DL, VT: Ptr.getValueType()));
25344 MachinePointerInfo PtrInfo =
25345 InsIndex == 0 ? ScalarLoad->getPointerInfo()
25346 : VecLoad->getPointerInfo().getWithOffset(O: EltSize / 8);
25347
25348 SDValue Load = DAG.getLoad(VT: VecLoad->getValueType(ResNo: 0), dl: DL,
25349 Chain: ScalarLoad->getChain(), Ptr, PtrInfo, Alignment: NewAlign);
25350 DAG.makeEquivalentMemoryOrdering(OldLoad: ScalarLoad, NewMemOp: Load.getValue(R: 1));
25351 DAG.makeEquivalentMemoryOrdering(OldLoad: VecLoad, NewMemOp: Load.getValue(R: 1));
25352 return Extend ? DAG.getNode(Opcode: Extend, DL, VT, Operand: Load) : Load;
25353}
25354
25355SDValue DAGCombiner::visitINSERT_VECTOR_ELT(SDNode *N) {
25356 SDValue InVec = N->getOperand(Num: 0);
25357 SDValue InVal = N->getOperand(Num: 1);
25358 SDValue EltNo = N->getOperand(Num: 2);
25359 SDLoc DL(N);
25360
25361 EVT VT = InVec.getValueType();
25362 auto *IndexC = dyn_cast<ConstantSDNode>(Val&: EltNo);
25363
25364 // Insert into out-of-bounds element is poison.
25365 if (IndexC && VT.isFixedLengthVector() &&
25366 IndexC->getZExtValue() >= VT.getVectorNumElements())
25367 return DAG.getPOISON(VT);
25368
25369 // Remove redundant insertions:
25370 // (insert_vector_elt x (extract_vector_elt x idx) idx) -> x
25371 if (InVal.getOpcode() == ISD::EXTRACT_VECTOR_ELT &&
25372 InVec == InVal.getOperand(i: 0) && EltNo == InVal.getOperand(i: 1))
25373 return InVec;
25374
25375 // Remove insert of UNDEF/POISON elements.
25376 if (InVal.isUndef()) {
25377 if (InVal.getOpcode() == ISD::POISON || InVec.getOpcode() == ISD::UNDEF)
25378 return InVec;
25379 return DAG.getFreeze(V: InVec);
25380 }
25381
25382 if (!IndexC) {
25383 // If this is variable insert to undef vector, it might be better to splat:
25384 // inselt undef, InVal, EltNo --> build_vector < InVal, InVal, ... >
25385 if (InVec.isUndef() && TLI.shouldSplatInsEltVarIndex(VT))
25386 return DAG.getSplat(VT, DL, Op: InVal);
25387
25388 // Extend this type to be byte-addressable
25389 EVT OldVT = VT;
25390 EVT EltVT = VT.getVectorElementType();
25391 bool IsByteSized = EltVT.isByteSized();
25392 if (!IsByteSized) {
25393 EltVT =
25394 EltVT.changeTypeToInteger().getRoundIntegerType(Context&: *DAG.getContext());
25395 VT = VT.changeElementType(Context&: *DAG.getContext(), EltVT);
25396 }
25397
25398 // Check if this operation will be handled the default way for its type.
25399 auto IsTypeDefaultHandled = [this](EVT VT) {
25400 return TLI.getTypeAction(Context&: *DAG.getContext(), VT) ==
25401 TargetLowering::TypeSplitVector ||
25402 TLI.isOperationExpand(Op: ISD::INSERT_VECTOR_ELT, VT);
25403 };
25404
25405 // Check if this operation is illegal and will be handled the default way,
25406 // even after extending the type to be byte-addressable.
25407 if (IsTypeDefaultHandled(OldVT) && IsTypeDefaultHandled(VT)) {
25408 // For each dynamic insertelt, the default way will save the vector to
25409 // the stack, store at an offset, and load the modified vector. This can
25410 // dramatically increase code size if we have a chain of insertelts on a
25411 // large vector: requiring O(V*C) stores/loads where V = length of
25412 // vector and C is length of chain. If each insertelt is only fed into the
25413 // next, the vector is write-only across this chain, and we can just
25414 // save once before the chain and load after in O(V + C) operations.
25415 SmallVector<SDNode *> Seq{N};
25416 unsigned NumDynamic = 1;
25417 while (true) {
25418 SDValue InVec = Seq.back()->getOperand(Num: 0);
25419 if (InVec.getOpcode() != ISD::INSERT_VECTOR_ELT)
25420 break;
25421 Seq.push_back(Elt: InVec.getNode());
25422 NumDynamic += !isa<ConstantSDNode>(Val: InVec.getOperand(i: 2));
25423 }
25424
25425 // It always and only makes sense to lower this sequence when we have more
25426 // than one dynamic insertelt, since we will not have more than V constant
25427 // insertelts, so we will be reducing the total number of stores+loads.
25428 if (NumDynamic > 1) {
25429 // In cases where the vector is illegal it will be broken down into
25430 // parts and stored in parts - we should use the alignment for the
25431 // smallest part.
25432 Align SmallestAlign = DAG.getReducedAlign(VT, /*UseABI=*/false);
25433 SDValue StackPtr =
25434 DAG.CreateStackTemporary(Bytes: VT.getStoreSize(), Alignment: SmallestAlign);
25435 auto &MF = DAG.getMachineFunction();
25436 int FrameIndex = cast<FrameIndexSDNode>(Val: StackPtr.getNode())->getIndex();
25437 auto PtrInfo = MachinePointerInfo::getFixedStack(MF, FI: FrameIndex);
25438
25439 // Save the vector to the stack
25440 SDValue InVec = Seq.back()->getOperand(Num: 0);
25441 if (!IsByteSized)
25442 InVec = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT, Operand: InVec);
25443 SDValue Store = DAG.getStore(Chain: DAG.getEntryNode(), dl: DL, Val: InVec, Ptr: StackPtr,
25444 PtrInfo, Alignment: SmallestAlign);
25445
25446 // Lower each dynamic insertelt to a store
25447 for (SDNode *N : reverse(C&: Seq)) {
25448 SDValue Elmnt = N->getOperand(Num: 1);
25449 SDValue Index = N->getOperand(Num: 2);
25450
25451 // Check if we have to extend the element type
25452 if (!IsByteSized && Elmnt.getValueType().bitsLT(VT: EltVT))
25453 Elmnt = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT: EltVT, Operand: Elmnt);
25454
25455 // Store the new element. This may be larger than the vector element
25456 // type, so use a truncating store.
25457 SDValue EltPtr =
25458 TLI.getVectorElementPointer(DAG, VecPtr: StackPtr, VecVT: VT, Index);
25459 EVT EltVT = Elmnt.getValueType();
25460 Store = DAG.getTruncStore(
25461 Chain: Store, dl: DL, Val: Elmnt, Ptr: EltPtr, PtrInfo: MachinePointerInfo::getUnknownStack(MF),
25462 SVT: EltVT,
25463 Alignment: commonAlignment(A: SmallestAlign, Offset: EltVT.getFixedSizeInBits() / 8));
25464 }
25465
25466 // Load the saved vector from the stack
25467 SDValue Load =
25468 DAG.getLoad(VT, dl: DL, Chain: Store, Ptr: StackPtr, PtrInfo, Alignment: SmallestAlign);
25469 SDValue LoadV = Load.getValue(R: 0);
25470 return IsByteSized ? LoadV : DAG.getAnyExtOrTrunc(Op: LoadV, DL, VT: OldVT);
25471 }
25472 }
25473
25474 return SDValue();
25475 }
25476
25477 if (VT.isScalableVector())
25478 return SDValue();
25479
25480 unsigned NumElts = VT.getVectorNumElements();
25481
25482 // We must know which element is being inserted for folds below here.
25483 unsigned Elt = IndexC->getZExtValue();
25484
25485 // Handle <1 x ???> vector insertion special cases.
25486 if (NumElts == 1) {
25487 // insert_vector_elt(x, extract_vector_elt(y, 0), 0) -> y
25488 if (InVal.getOpcode() == ISD::EXTRACT_VECTOR_ELT &&
25489 InVal.getOperand(i: 0).getValueType() == VT &&
25490 isNullConstant(V: InVal.getOperand(i: 1)))
25491 return InVal.getOperand(i: 0);
25492 }
25493
25494 // Canonicalize insert_vector_elt dag nodes.
25495 // Example:
25496 // (insert_vector_elt (insert_vector_elt A, Idx0), Idx1)
25497 // -> (insert_vector_elt (insert_vector_elt A, Idx1), Idx0)
25498 //
25499 // Do this only if the child insert_vector node has one use; also
25500 // do this only if indices are both constants and Idx1 < Idx0.
25501 if (InVec.getOpcode() == ISD::INSERT_VECTOR_ELT && InVec.hasOneUse()
25502 && isa<ConstantSDNode>(Val: InVec.getOperand(i: 2))) {
25503 unsigned OtherElt = InVec.getConstantOperandVal(i: 2);
25504 if (Elt < OtherElt) {
25505 // Swap nodes.
25506 SDValue NewOp = DAG.getNode(Opcode: ISD::INSERT_VECTOR_ELT, DL, VT,
25507 N1: InVec.getOperand(i: 0), N2: InVal, N3: EltNo);
25508 AddToWorklist(N: NewOp.getNode());
25509 return DAG.getNode(Opcode: ISD::INSERT_VECTOR_ELT, DL: SDLoc(InVec.getNode()),
25510 VT, N1: NewOp, N2: InVec.getOperand(i: 1), N3: InVec.getOperand(i: 2));
25511 }
25512 }
25513
25514 if (SDValue Shuf = mergeInsertEltWithShuffle(N, InsIndex: Elt))
25515 return Shuf;
25516
25517 if (SDValue Shuf = combineInsertEltToShuffle(N, InsIndex: Elt))
25518 return Shuf;
25519
25520 if (SDValue Shuf = combineInsertEltToLoad(N, InsIndex: Elt))
25521 return Shuf;
25522
25523 // Attempt to convert an insert_vector_elt chain into a legal build_vector.
25524 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::BUILD_VECTOR, VT)) {
25525 // vXi1 vector - we don't need to recurse.
25526 if (NumElts == 1)
25527 return DAG.getBuildVector(VT, DL, Ops: {InVal});
25528
25529 // If we haven't already collected the element, insert into the op list.
25530 EVT MaxEltVT = InVal.getValueType();
25531 auto AddBuildVectorOp = [&](SmallVectorImpl<SDValue> &Ops, SDValue Elt,
25532 unsigned Idx) {
25533 if (!Ops[Idx]) {
25534 Ops[Idx] = Elt;
25535 if (VT.isInteger()) {
25536 EVT EltVT = Elt.getValueType();
25537 MaxEltVT = MaxEltVT.bitsGE(VT: EltVT) ? MaxEltVT : EltVT;
25538 }
25539 }
25540 };
25541
25542 // Ensure all the operands are the same value type, fill any missing
25543 // operands with UNDEF and create the BUILD_VECTOR.
25544 auto CanonicalizeBuildVector = [&](SmallVectorImpl<SDValue> &Ops,
25545 bool FreezeUndef = false) {
25546 assert(Ops.size() == NumElts && "Unexpected vector size");
25547 SDValue UndefOp = FreezeUndef ? DAG.getFreeze(V: DAG.getUNDEF(VT: MaxEltVT))
25548 : DAG.getUNDEF(VT: MaxEltVT);
25549 for (SDValue &Op : Ops) {
25550 if (Op)
25551 Op = VT.isInteger() ? DAG.getAnyExtOrTrunc(Op, DL, VT: MaxEltVT) : Op;
25552 else
25553 Op = UndefOp;
25554 }
25555 return DAG.getBuildVector(VT, DL, Ops);
25556 };
25557
25558 SmallVector<SDValue, 8> Ops(NumElts, SDValue());
25559 Ops[Elt] = InVal;
25560
25561 // Recurse up a INSERT_VECTOR_ELT chain to build a BUILD_VECTOR.
25562 for (SDValue CurVec = InVec; CurVec;) {
25563 // UNDEF - build new BUILD_VECTOR from already inserted operands.
25564 if (CurVec.isUndef())
25565 return CanonicalizeBuildVector(Ops);
25566
25567 // FREEZE(UNDEF) - build new BUILD_VECTOR from already inserted operands.
25568 if (ISD::isFreezeUndef(N: CurVec.getNode()) && CurVec.hasOneUse())
25569 return CanonicalizeBuildVector(Ops, /*FreezeUndef=*/true);
25570
25571 // BUILD_VECTOR - insert unused operands and build new BUILD_VECTOR.
25572 // A multi-use base is allowed when the target prefers build_vector
25573 // sources: rebuilding only re-references the base's scalar operands, and
25574 // it un-shares the base so each vector is built independently.
25575 if (CurVec.getOpcode() == ISD::BUILD_VECTOR &&
25576 (CurVec.hasOneUse() ||
25577 TLI.aggressivelyPreferBuildVectorSources(VecVT: VT))) {
25578 for (unsigned I = 0; I != NumElts; ++I)
25579 AddBuildVectorOp(Ops, CurVec.getOperand(i: I), I);
25580 return CanonicalizeBuildVector(Ops);
25581 }
25582
25583 // SCALAR_TO_VECTOR - insert unused scalar and build new BUILD_VECTOR.
25584 if (CurVec.getOpcode() == ISD::SCALAR_TO_VECTOR && CurVec.hasOneUse()) {
25585 AddBuildVectorOp(Ops, CurVec.getOperand(i: 0), 0);
25586 return CanonicalizeBuildVector(Ops);
25587 }
25588
25589 // INSERT_VECTOR_ELT - insert operand and continue up the chain.
25590 if (CurVec.getOpcode() == ISD::INSERT_VECTOR_ELT && CurVec.hasOneUse())
25591 if (auto *CurIdx = dyn_cast<ConstantSDNode>(Val: CurVec.getOperand(i: 2)))
25592 if (CurIdx->getAPIntValue().ult(RHS: NumElts)) {
25593 unsigned Idx = CurIdx->getZExtValue();
25594 AddBuildVectorOp(Ops, CurVec.getOperand(i: 1), Idx);
25595
25596 // Found entire BUILD_VECTOR.
25597 if (all_of(Range&: Ops, P: [](SDValue Op) { return !!Op; }))
25598 return CanonicalizeBuildVector(Ops);
25599
25600 CurVec = CurVec->getOperand(Num: 0);
25601 continue;
25602 }
25603
25604 // VECTOR_SHUFFLE - if all the operands match the shuffle's sources,
25605 // update the shuffle mask (and second operand if we started with unary
25606 // shuffle) and create a new legal shuffle.
25607 if (CurVec.getOpcode() == ISD::VECTOR_SHUFFLE && CurVec.hasOneUse()) {
25608 auto *SVN = cast<ShuffleVectorSDNode>(Val&: CurVec);
25609 SDValue LHS = SVN->getOperand(Num: 0);
25610 SDValue RHS = SVN->getOperand(Num: 1);
25611 SmallVector<int, 16> Mask(SVN->getMask());
25612 bool Merged = true;
25613 for (auto I : enumerate(First&: Ops)) {
25614 SDValue &Op = I.value();
25615 if (Op) {
25616 SmallVector<int, 16> NewMask;
25617 if (!mergeEltWithShuffle(X&: LHS, Y&: RHS, Mask, NewMask, Elt: Op, InsIndex: I.index())) {
25618 Merged = false;
25619 break;
25620 }
25621 Mask = std::move(NewMask);
25622 }
25623 }
25624 if (Merged)
25625 if (SDValue NewShuffle =
25626 TLI.buildLegalVectorShuffle(VT, DL, N0: LHS, N1: RHS, Mask, DAG))
25627 return NewShuffle;
25628 }
25629
25630 if (!LegalOperations) {
25631 bool IsNull = llvm::isNullConstant(V: InVal);
25632 // We can convert to AND/OR mask if all insertions are zero or -1
25633 // respectively.
25634 if ((IsNull || llvm::isAllOnesConstant(V: InVal)) &&
25635 all_of(Range&: Ops, P: [InVal](SDValue Op) { return !Op || Op == InVal; }) &&
25636 count_if(Range&: Ops, P: [InVal](SDValue Op) { return Op == InVal; }) >= 2) {
25637 SDValue Zero = DAG.getConstant(Val: 0, DL, VT: MaxEltVT);
25638 SDValue AllOnes = DAG.getAllOnesConstant(DL, VT: MaxEltVT);
25639 SmallVector<SDValue, 8> Mask(NumElts);
25640
25641 // Build the mask and return the corresponding DAG node.
25642 auto BuildMaskAndNode = [&](SDValue TrueVal, SDValue FalseVal,
25643 unsigned MaskOpcode) {
25644 APInt InsertedEltMask = APInt::getZero(numBits: NumElts);
25645 for (unsigned I = 0; I != NumElts; ++I) {
25646 Mask[I] = Ops[I] ? TrueVal : FalseVal;
25647 if (Ops[I])
25648 InsertedEltMask.setBit(I);
25649 }
25650 // Make sure to freeze the source vector in case any of the elements
25651 // overwritten by the insert may be poison. Otherwise those elements
25652 // could end up being poison instead of 0/-1 after the AND/OR.
25653 CurVec = DAG.getFreeze(V: CurVec, DemandedElts: InsertedEltMask,
25654 Kind: UndefPoisonKind::PoisonOnly);
25655 return DAG.getNode(Opcode: MaskOpcode, DL, VT, N1: CurVec,
25656 N2: DAG.getBuildVector(VT, DL, Ops: Mask));
25657 };
25658
25659 // If all elements are zero, we can use AND with all ones.
25660 if (IsNull)
25661 return BuildMaskAndNode(Zero, AllOnes, ISD::AND);
25662
25663 // If all elements are -1, we can use OR with zero.
25664 return BuildMaskAndNode(AllOnes, Zero, ISD::OR);
25665 }
25666 }
25667
25668 // Failed to find a match in the chain - bail.
25669 break;
25670 }
25671
25672 // See if we can fill in the missing constant elements as zeros.
25673 // TODO: Should we do this for any constant?
25674 APInt DemandedZeroElts = APInt::getZero(numBits: NumElts);
25675 for (unsigned I = 0; I != NumElts; ++I)
25676 if (!Ops[I])
25677 DemandedZeroElts.setBit(I);
25678
25679 if (DAG.MaskedVectorIsZero(Op: InVec, DemandedElts: DemandedZeroElts)) {
25680 SDValue Zero = VT.isInteger() ? DAG.getConstant(Val: 0, DL, VT: MaxEltVT)
25681 : DAG.getConstantFP(Val: 0, DL, VT: MaxEltVT);
25682 for (unsigned I = 0; I != NumElts; ++I)
25683 if (!Ops[I])
25684 Ops[I] = Zero;
25685
25686 return CanonicalizeBuildVector(Ops);
25687 }
25688 }
25689
25690 return SDValue();
25691}
25692
25693/// Transform a vector binary operation into a scalar binary operation by moving
25694/// the math/logic after an extract element of a vector.
25695static SDValue scalarizeExtractedBinOp(SDNode *ExtElt, SelectionDAG &DAG,
25696 const SDLoc &DL, bool LegalTypes) {
25697 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
25698 SDValue Vec = ExtElt->getOperand(Num: 0);
25699 SDValue Index = ExtElt->getOperand(Num: 1);
25700 auto *IndexC = dyn_cast<ConstantSDNode>(Val&: Index);
25701 unsigned Opc = Vec.getOpcode();
25702 if (!IndexC || !Vec.hasOneUse() || (!TLI.isBinOp(Opcode: Opc) && Opc != ISD::SETCC) ||
25703 Vec->getNumValues() != 1)
25704 return SDValue();
25705
25706 // Targets may want to avoid this to prevent an expensive register transfer.
25707 if (!TLI.shouldScalarizeBinop(VecOp: Vec))
25708 return SDValue();
25709
25710 EVT ResVT = ExtElt->getValueType(ResNo: 0);
25711 if (Opc == ISD::SETCC &&
25712 (ResVT != Vec.getValueType().getVectorElementType() || LegalTypes))
25713 return SDValue();
25714
25715 // Extracting an element of a vector constant is constant-folded, so this
25716 // transform is just replacing a vector op with a scalar op while moving the
25717 // extract.
25718 auto IsExtractFree = [](SDValue Op) {
25719 APInt SplatVal;
25720 return isAnyConstantBuildVector(V: Op, NoOpaques: true) ||
25721 ISD::isConstantSplatVector(N: Op.getNode(), SplatValue&: SplatVal) ||
25722 (Op.getOpcode() == ISD::BUILD_VECTOR && Op.hasOneUse());
25723 };
25724 SDValue Op0 = Vec.getOperand(i: 0);
25725 SDValue Op1 = Vec.getOperand(i: 1);
25726 if (!IsExtractFree(Op0) && !IsExtractFree(Op1))
25727 return SDValue();
25728
25729 // extractelt (op X, C), IndexC --> op (extractelt X, IndexC), C'
25730 // extractelt (op C, X), IndexC --> op C', (extractelt X, IndexC)
25731 if (Opc == ISD::SETCC) {
25732 EVT OpVT = Op0.getValueType().getVectorElementType();
25733 Op0 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: OpVT, N1: Op0, N2: Index);
25734 Op1 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: OpVT, N1: Op1, N2: Index);
25735 SDValue NewVal = DAG.getSetCC(
25736 DL, VT: ResVT, LHS: Op0, RHS: Op1, Cond: cast<CondCodeSDNode>(Val: Vec->getOperand(Num: 2))->get());
25737 // We may need to sign- or zero-extend the result to match the same
25738 // behaviour as the vector version of SETCC.
25739 unsigned VecBoolContents = TLI.getBooleanContents(Type: Vec.getValueType());
25740 if (ResVT != MVT::i1 &&
25741 VecBoolContents != TargetLowering::UndefinedBooleanContent &&
25742 VecBoolContents != TLI.getBooleanContents(Type: ResVT)) {
25743 if (VecBoolContents == TargetLowering::ZeroOrNegativeOneBooleanContent)
25744 NewVal = DAG.getNode(Opcode: ISD::SIGN_EXTEND_INREG, DL, VT: ResVT, N1: NewVal,
25745 N2: DAG.getValueType(MVT::i1));
25746 else
25747 NewVal = DAG.getZeroExtendInReg(Op: NewVal, DL, VT: MVT::i1);
25748 }
25749 return NewVal;
25750 }
25751 Op0 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: ResVT, N1: Op0, N2: Index);
25752 Op1 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: ResVT, N1: Op1, N2: Index);
25753 return DAG.getNode(Opcode: Opc, DL, VT: ResVT, N1: Op0, N2: Op1);
25754}
25755
25756// Given a ISD::EXTRACT_VECTOR_ELT, which is a glorified bit sequence extract,
25757// recursively analyse all of it's users. and try to model themselves as
25758// bit sequence extractions. If all of them agree on the new, narrower element
25759// type, and all of them can be modelled as ISD::EXTRACT_VECTOR_ELT's of that
25760// new element type, do so now.
25761// This is mainly useful to recover from legalization that scalarized
25762// the vector as wide elements, but tries to rebuild it with narrower elements.
25763//
25764// Some more nodes could be modelled if that helps cover interesting patterns.
25765bool DAGCombiner::refineExtractVectorEltIntoMultipleNarrowExtractVectorElts(
25766 SDNode *N) {
25767 // We perform this optimization post type-legalization because
25768 // the type-legalizer often scalarizes integer-promoted vectors.
25769 // Performing this optimization before may cause legalizaton cycles.
25770 if (Level != AfterLegalizeVectorOps && Level != AfterLegalizeTypes)
25771 return false;
25772
25773 // TODO: Add support for big-endian.
25774 if (DAG.getDataLayout().isBigEndian())
25775 return false;
25776
25777 SDValue VecOp = N->getOperand(Num: 0);
25778 EVT VecVT = VecOp.getValueType();
25779 assert(!VecVT.isScalableVector() && "Only for fixed vectors.");
25780
25781 // We must start with a constant extraction index.
25782 auto *IndexC = dyn_cast<ConstantSDNode>(Val: N->getOperand(Num: 1));
25783 if (!IndexC)
25784 return false;
25785
25786 assert(IndexC->getZExtValue() < VecVT.getVectorNumElements() &&
25787 "Original ISD::EXTRACT_VECTOR_ELT is undefinend?");
25788
25789 // TODO: deal with the case of implicit anyext of the extraction.
25790 unsigned VecEltBitWidth = VecVT.getScalarSizeInBits();
25791 EVT ScalarVT = N->getValueType(ResNo: 0);
25792 if (VecVT.getScalarType() != ScalarVT)
25793 return false;
25794
25795 // TODO: deal with the cases other than everything being integer-typed.
25796 if (!ScalarVT.isScalarInteger())
25797 return false;
25798
25799 struct Entry {
25800 SDNode *Producer;
25801
25802 // Which bits of VecOp does it contain?
25803 unsigned BitPos;
25804 int NumBits;
25805 // NOTE: the actual width of \p Producer may be wider than NumBits!
25806
25807 Entry(Entry &&) = default;
25808 Entry(SDNode *Producer_, unsigned BitPos_, int NumBits_)
25809 : Producer(Producer_), BitPos(BitPos_), NumBits(NumBits_) {}
25810
25811 Entry() = delete;
25812 Entry(const Entry &) = delete;
25813 Entry &operator=(const Entry &) = delete;
25814 Entry &operator=(Entry &&) = delete;
25815 };
25816 SmallVector<Entry, 32> Worklist;
25817 SmallVector<Entry, 32> Leafs;
25818
25819 // We start at the "root" ISD::EXTRACT_VECTOR_ELT.
25820 Worklist.emplace_back(Args&: N, /*BitPos=*/Args: VecEltBitWidth * IndexC->getZExtValue(),
25821 /*NumBits=*/Args&: VecEltBitWidth);
25822
25823 while (!Worklist.empty()) {
25824 Entry E = Worklist.pop_back_val();
25825 // Does the node not even use any of the VecOp bits?
25826 if (!(E.NumBits > 0 && E.BitPos < VecVT.getSizeInBits() &&
25827 E.BitPos + E.NumBits <= VecVT.getSizeInBits()))
25828 return false; // Let's allow the other combines clean this up first.
25829 // Did we fail to model any of the users of the Producer?
25830 bool ProducerIsLeaf = false;
25831 // Look at each user of this Producer.
25832 for (SDNode *User : E.Producer->users()) {
25833 switch (User->getOpcode()) {
25834 // TODO: support ISD::BITCAST
25835 // TODO: support ISD::ANY_EXTEND
25836 // TODO: support ISD::ZERO_EXTEND
25837 // TODO: support ISD::SIGN_EXTEND
25838 case ISD::TRUNCATE:
25839 // Truncation simply means we keep position, but extract less bits.
25840 Worklist.emplace_back(Args&: User, Args&: E.BitPos,
25841 /*NumBits=*/Args: User->getValueSizeInBits(ResNo: 0));
25842 break;
25843 // TODO: support ISD::SRA
25844 // TODO: support ISD::SHL
25845 case ISD::SRL:
25846 // We should be shifting the Producer by a constant amount.
25847 if (auto *ShAmtC = dyn_cast<ConstantSDNode>(Val: User->getOperand(Num: 1));
25848 User->getOperand(Num: 0).getNode() == E.Producer && ShAmtC) {
25849 // Logical right-shift means that we start extraction later,
25850 // but stop it at the same position we did previously.
25851 unsigned ShAmt = ShAmtC->getZExtValue();
25852 Worklist.emplace_back(Args&: User, Args: E.BitPos + ShAmt, Args: E.NumBits - ShAmt);
25853 break;
25854 }
25855 [[fallthrough]];
25856 default:
25857 // We can not model this user of the Producer.
25858 // Which means the current Producer will be a ISD::EXTRACT_VECTOR_ELT.
25859 ProducerIsLeaf = true;
25860 // Profitability check: all users that we can not model
25861 // must be ISD::BUILD_VECTOR's.
25862 if (User->getOpcode() != ISD::BUILD_VECTOR)
25863 return false;
25864 break;
25865 }
25866 }
25867 if (ProducerIsLeaf)
25868 Leafs.emplace_back(Args: std::move(E));
25869 }
25870
25871 unsigned NewVecEltBitWidth = Leafs.front().NumBits;
25872
25873 // If we are still at the same element granularity, give up,
25874 if (NewVecEltBitWidth == VecEltBitWidth)
25875 return false;
25876
25877 // The vector width must be a multiple of the new element width.
25878 if (VecVT.getSizeInBits() % NewVecEltBitWidth != 0)
25879 return false;
25880
25881 // All leafs must agree on the new element width.
25882 // All leafs must not expect any "padding" bits ontop of that width.
25883 // All leafs must start extraction from multiple of that width.
25884 if (!all_of(Range&: Leafs, P: [NewVecEltBitWidth](const Entry &E) {
25885 return (unsigned)E.NumBits == NewVecEltBitWidth &&
25886 E.Producer->getValueSizeInBits(ResNo: 0) == NewVecEltBitWidth &&
25887 E.BitPos % NewVecEltBitWidth == 0;
25888 }))
25889 return false;
25890
25891 EVT NewScalarVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: NewVecEltBitWidth);
25892 EVT NewVecVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: NewScalarVT,
25893 NumElements: VecVT.getSizeInBits() / NewVecEltBitWidth);
25894
25895 if (LegalTypes &&
25896 !(TLI.isTypeLegal(VT: NewScalarVT) && TLI.isTypeLegal(VT: NewVecVT)))
25897 return false;
25898
25899 if (LegalOperations &&
25900 !(TLI.isOperationLegalOrCustom(Op: ISD::BITCAST, VT: NewVecVT) &&
25901 TLI.isOperationLegalOrCustom(Op: ISD::EXTRACT_VECTOR_ELT, VT: NewVecVT)))
25902 return false;
25903
25904 SDValue NewVecOp = DAG.getBitcast(VT: NewVecVT, V: VecOp);
25905 for (const Entry &E : Leafs) {
25906 SDLoc DL(E.Producer);
25907 unsigned NewIndex = E.BitPos / NewVecEltBitWidth;
25908 assert(NewIndex < NewVecVT.getVectorNumElements() &&
25909 "Creating out-of-bounds ISD::EXTRACT_VECTOR_ELT?");
25910 SDValue V = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: NewScalarVT, N1: NewVecOp,
25911 N2: DAG.getVectorIdxConstant(Val: NewIndex, DL));
25912 CombineTo(N: E.Producer, Res: V);
25913 }
25914
25915 return true;
25916}
25917
25918SDValue DAGCombiner::visitEXTRACT_VECTOR_ELT(SDNode *N) {
25919 SDValue VecOp = N->getOperand(Num: 0);
25920 SDValue Index = N->getOperand(Num: 1);
25921 EVT ScalarVT = N->getValueType(ResNo: 0);
25922 EVT VecVT = VecOp.getValueType();
25923 if (VecOp.getOpcode() == ISD::POISON)
25924 return DAG.getPOISON(VT: ScalarVT);
25925
25926 if (VecOp.getOpcode() == ISD::UNDEF)
25927 return DAG.getUNDEF(VT: ScalarVT);
25928
25929 // extract_vector_elt (insert_vector_elt vec, val, idx), idx) -> val
25930 //
25931 // This only really matters if the index is non-constant since other combines
25932 // on the constant elements already work.
25933 SDLoc DL(N);
25934 if (VecOp.getOpcode() == ISD::INSERT_VECTOR_ELT &&
25935 Index == VecOp.getOperand(i: 2)) {
25936 SDValue Elt = VecOp.getOperand(i: 1);
25937 AddUsersToWorklist(N: VecOp.getNode());
25938 return VecVT.isInteger() ? DAG.getAnyExtOrTrunc(Op: Elt, DL, VT: ScalarVT) : Elt;
25939 }
25940
25941 // (vextract (scalar_to_vector val, 0) -> val
25942 if (VecOp.getOpcode() == ISD::SCALAR_TO_VECTOR) {
25943 // Only 0'th element of SCALAR_TO_VECTOR is defined.
25944 if (DAG.isKnownNeverZero(Op: Index))
25945 return DAG.getPOISON(VT: ScalarVT);
25946
25947 // Check if the result type doesn't match the inserted element type.
25948 // The inserted element and extracted element may have mismatched bitwidth.
25949 // As a result, EXTRACT_VECTOR_ELT may extend or truncate the extracted vector.
25950 SDValue InOp = VecOp.getOperand(i: 0);
25951 if (InOp.getValueType() != ScalarVT) {
25952 assert(InOp.getValueType().isInteger() && ScalarVT.isInteger());
25953 if (InOp.getValueType().bitsGT(VT: ScalarVT))
25954 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: ScalarVT, Operand: InOp);
25955 return DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT: ScalarVT, Operand: InOp);
25956 }
25957 return InOp;
25958 }
25959
25960 // extract_vector_elt of out-of-bounds element -> POISON
25961 auto *IndexC = dyn_cast<ConstantSDNode>(Val&: Index);
25962 if (IndexC && VecVT.isFixedLengthVector() &&
25963 IndexC->getAPIntValue().uge(RHS: VecVT.getVectorNumElements()))
25964 return DAG.getPOISON(VT: ScalarVT);
25965
25966 // extract_vector_elt (build_vector x, y), 1 -> y
25967 if (((IndexC && VecOp.getOpcode() == ISD::BUILD_VECTOR) ||
25968 VecOp.getOpcode() == ISD::SPLAT_VECTOR) &&
25969 TLI.isTypeLegal(VT: VecVT)) {
25970 assert((VecOp.getOpcode() != ISD::BUILD_VECTOR ||
25971 VecVT.isFixedLengthVector()) &&
25972 "BUILD_VECTOR used for scalable vectors");
25973 unsigned IndexVal =
25974 VecOp.getOpcode() == ISD::BUILD_VECTOR ? IndexC->getZExtValue() : 0;
25975 SDValue Elt = VecOp.getOperand(i: IndexVal);
25976 EVT InEltVT = Elt.getValueType();
25977
25978 if (VecOp.hasOneUse() || TLI.aggressivelyPreferBuildVectorSources(VecVT) ||
25979 isNullConstant(V: Elt)) {
25980 // Sometimes build_vector's scalar input types do not match result type.
25981 if (ScalarVT == InEltVT)
25982 return Elt;
25983
25984 // TODO: It may be useful to truncate if free if the build_vector
25985 // implicitly converts.
25986 }
25987 }
25988
25989 if (SDValue BO = scalarizeExtractedBinOp(ExtElt: N, DAG, DL, LegalTypes))
25990 return BO;
25991
25992 if (VecVT.isScalableVector())
25993 return SDValue();
25994
25995 // All the code from this point onwards assumes fixed width vectors, but it's
25996 // possible that some of the combinations could be made to work for scalable
25997 // vectors too.
25998 unsigned NumElts = VecVT.getVectorNumElements();
25999 unsigned VecEltBitWidth = VecVT.getScalarSizeInBits();
26000
26001 // See if the extracted element is constant, in which case fold it if its
26002 // a legal fp immediate.
26003 if (IndexC && ScalarVT.isFloatingPoint()) {
26004 APInt EltMask = APInt::getOneBitSet(numBits: NumElts, BitNo: IndexC->getZExtValue());
26005 KnownBits KnownElt = DAG.computeKnownBits(Op: VecOp, DemandedElts: EltMask);
26006 if (KnownElt.isConstant()) {
26007 APFloat CstFP =
26008 APFloat(ScalarVT.getFltSemantics(), KnownElt.getConstant());
26009 if (TLI.isFPImmLegal(CstFP, ScalarVT))
26010 return DAG.getConstantFP(Val: CstFP, DL, VT: ScalarVT);
26011 }
26012 }
26013
26014 // TODO: These transforms should not require the 'hasOneUse' restriction, but
26015 // there are regressions on multiple targets without it. We can end up with a
26016 // mess of scalar and vector code if we reduce only part of the DAG to scalar.
26017 if (IndexC && VecOp.getOpcode() == ISD::BITCAST && VecVT.isInteger() &&
26018 VecOp.hasOneUse()) {
26019 // The vector index of the LSBs of the source depend on the endian-ness.
26020 bool IsLE = DAG.getDataLayout().isLittleEndian();
26021 unsigned ExtractIndex = IndexC->getZExtValue();
26022 // extract_elt (v2i32 (bitcast i64:x)), BCTruncElt -> i32 (trunc i64:x)
26023 unsigned BCTruncElt = IsLE ? 0 : NumElts - 1;
26024 SDValue BCSrc = VecOp.getOperand(i: 0);
26025 if (ExtractIndex == BCTruncElt && BCSrc.getValueType().isScalarInteger())
26026 return DAG.getAnyExtOrTrunc(Op: BCSrc, DL, VT: ScalarVT);
26027
26028 // TODO: Add support for SCALAR_TO_VECTOR implicit truncation.
26029 if (LegalTypes && BCSrc.getValueType().isInteger() &&
26030 BCSrc.getOpcode() == ISD::SCALAR_TO_VECTOR &&
26031 BCSrc.getScalarValueSizeInBits() ==
26032 BCSrc.getOperand(i: 0).getScalarValueSizeInBits()) {
26033 // ext_elt (bitcast (scalar_to_vec i64 X to v2i64) to v4i32), TruncElt -->
26034 // trunc i64 X to i32
26035 SDValue X = BCSrc.getOperand(i: 0);
26036 EVT XVT = X.getValueType();
26037 assert(XVT.isScalarInteger() && ScalarVT.isScalarInteger() &&
26038 "Extract element and scalar to vector can't change element type "
26039 "from FP to integer.");
26040 unsigned XBitWidth = X.getValueSizeInBits();
26041 unsigned Scale = XBitWidth / VecEltBitWidth;
26042 BCTruncElt = IsLE ? 0 : Scale - 1;
26043
26044 // An extract element return value type can be wider than its vector
26045 // operand element type. In that case, the high bits are undefined, so
26046 // it's possible that we may need to extend rather than truncate.
26047 if (ExtractIndex < Scale && XBitWidth > VecEltBitWidth) {
26048 assert(XBitWidth % VecEltBitWidth == 0 &&
26049 "Scalar bitwidth must be a multiple of vector element bitwidth");
26050
26051 if (ExtractIndex != BCTruncElt) {
26052 unsigned ShiftIndex =
26053 IsLE ? ExtractIndex : (Scale - 1) - ExtractIndex;
26054 X = DAG.getNode(
26055 Opcode: ISD::SRL, DL, VT: XVT, N1: X,
26056 N2: DAG.getShiftAmountConstant(Val: ShiftIndex * VecEltBitWidth, VT: XVT, DL));
26057 }
26058
26059 return DAG.getAnyExtOrTrunc(Op: X, DL, VT: ScalarVT);
26060 }
26061 }
26062 }
26063
26064 // Transform: (EXTRACT_VECTOR_ELT( VECTOR_SHUFFLE )) -> EXTRACT_VECTOR_ELT.
26065 // We only perform this optimization before the op legalization phase because
26066 // we may introduce new vector instructions which are not backed by TD
26067 // patterns. For example on AVX, extracting elements from a wide vector
26068 // without using extract_subvector. However, if we can find an underlying
26069 // scalar value, then we can always use that.
26070 if (IndexC && VecOp.getOpcode() == ISD::VECTOR_SHUFFLE) {
26071 auto *Shuf = cast<ShuffleVectorSDNode>(Val&: VecOp);
26072 // Find the new index to extract from.
26073 int OrigElt = Shuf->getMaskElt(Idx: IndexC->getZExtValue());
26074
26075 // Extracting an undef index is poison.
26076 if (OrigElt == -1)
26077 return DAG.getPOISON(VT: ScalarVT);
26078
26079 // Select the right vector half to extract from.
26080 SDValue SVInVec;
26081 if (OrigElt < (int)NumElts) {
26082 SVInVec = VecOp.getOperand(i: 0);
26083 } else {
26084 SVInVec = VecOp.getOperand(i: 1);
26085 OrigElt -= NumElts;
26086 }
26087
26088 if (SVInVec.getOpcode() == ISD::BUILD_VECTOR) {
26089 // TODO: Check if shuffle mask is legal?
26090 if (LegalOperations && TLI.isOperationLegal(Op: ISD::VECTOR_SHUFFLE, VT: VecVT) &&
26091 !VecOp.hasOneUse())
26092 return SDValue();
26093
26094 SDValue InOp = SVInVec.getOperand(i: OrigElt);
26095 if (InOp.getValueType() != ScalarVT) {
26096 assert(InOp.getValueType().isInteger() && ScalarVT.isInteger());
26097 InOp = DAG.getSExtOrTrunc(Op: InOp, DL, VT: ScalarVT);
26098 }
26099
26100 return InOp;
26101 }
26102
26103 // FIXME: We should handle recursing on other vector shuffles and
26104 // scalar_to_vector here as well.
26105
26106 if (!LegalOperations ||
26107 // FIXME: Should really be just isOperationLegalOrCustom.
26108 TLI.isOperationLegal(Op: ISD::EXTRACT_VECTOR_ELT, VT: VecVT) ||
26109 TLI.isOperationExpand(Op: ISD::VECTOR_SHUFFLE, VT: VecVT)) {
26110 return DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: ScalarVT, N1: SVInVec,
26111 N2: DAG.getVectorIdxConstant(Val: OrigElt, DL));
26112 }
26113 }
26114
26115 // If only EXTRACT_VECTOR_ELT nodes use the source vector we can
26116 // simplify it based on the (valid) extraction indices.
26117 if (llvm::all_of(Range: VecOp->users(), P: [&](SDNode *Use) {
26118 return Use->getOpcode() == ISD::EXTRACT_VECTOR_ELT &&
26119 Use->getOperand(Num: 0) == VecOp &&
26120 isa<ConstantSDNode>(Val: Use->getOperand(Num: 1));
26121 })) {
26122 APInt DemandedElts = APInt::getZero(numBits: NumElts);
26123 for (SDNode *User : VecOp->users()) {
26124 auto *CstElt = cast<ConstantSDNode>(Val: User->getOperand(Num: 1));
26125 if (CstElt->getAPIntValue().ult(RHS: NumElts))
26126 DemandedElts.setBit(CstElt->getZExtValue());
26127 }
26128 if (SimplifyDemandedVectorElts(Op: VecOp, DemandedElts, AssumeSingleUse: true)) {
26129 // We simplified the vector operand of this extract element. If this
26130 // extract is not dead, visit it again so it is folded properly.
26131 if (N->getOpcode() != ISD::DELETED_NODE)
26132 AddToWorklist(N);
26133 return SDValue(N, 0);
26134 }
26135 APInt DemandedBits = APInt::getAllOnes(numBits: VecEltBitWidth);
26136 if (SimplifyDemandedBits(Op: VecOp, DemandedBits, DemandedElts, AssumeSingleUse: true)) {
26137 // We simplified the vector operand of this extract element. If this
26138 // extract is not dead, visit it again so it is folded properly.
26139 if (N->getOpcode() != ISD::DELETED_NODE)
26140 AddToWorklist(N);
26141 return SDValue(N, 0);
26142 }
26143 }
26144
26145 if (refineExtractVectorEltIntoMultipleNarrowExtractVectorElts(N))
26146 return SDValue(N, 0);
26147
26148 // Everything under here is trying to match an extract of a loaded value.
26149 // If the result of load has to be truncated, then it's not necessarily
26150 // profitable.
26151 bool BCNumEltsChanged = false;
26152 EVT ExtVT = VecVT.getVectorElementType();
26153 EVT LVT = ExtVT;
26154 if (ScalarVT.bitsLT(VT: LVT) && !TLI.isTruncateFree(FromVT: LVT, ToVT: ScalarVT))
26155 return SDValue();
26156
26157 if (VecOp.getOpcode() == ISD::BITCAST) {
26158 // Don't duplicate a load with other uses.
26159 if (!VecOp.hasOneUse())
26160 return SDValue();
26161
26162 EVT BCVT = VecOp.getOperand(i: 0).getValueType();
26163 if (!BCVT.isVector() || ExtVT.bitsGT(VT: BCVT.getVectorElementType()))
26164 return SDValue();
26165 if (NumElts != BCVT.getVectorNumElements())
26166 BCNumEltsChanged = true;
26167 VecOp = VecOp.getOperand(i: 0);
26168 ExtVT = BCVT.getVectorElementType();
26169 }
26170
26171 // extract (vector load $addr), i --> load $addr + i * size
26172 if (!LegalOperations && !IndexC && VecOp.hasOneUse() &&
26173 ISD::isNormalLoad(N: VecOp.getNode()) &&
26174 !Index->hasPredecessor(N: VecOp.getNode())) {
26175 auto *VecLoad = dyn_cast<LoadSDNode>(Val&: VecOp);
26176 if (VecLoad && VecLoad->isSimple()) {
26177 if (SDValue Scalarized = TLI.scalarizeExtractedVectorLoad(
26178 ResultVT: ScalarVT, DL: SDLoc(N), InVecVT: VecVT, EltNo: Index, OriginalLoad: VecLoad, DAG)) {
26179 ++OpsNarrowed;
26180 return Scalarized;
26181 }
26182 }
26183 }
26184
26185 // Perform only after legalization to ensure build_vector / vector_shuffle
26186 // optimizations have already been done.
26187 if (!LegalOperations || !IndexC)
26188 return SDValue();
26189
26190 bool IsFrozen = false;
26191 if (VecOp.getOpcode() == ISD::FREEZE && VecOp.hasOneUse()) {
26192 VecOp = VecOp.getOperand(i: 0);
26193 IsFrozen = true;
26194 }
26195
26196 // (vextract (v4f32 load $addr), c) -> (f32 load $addr+c*size)
26197 // (vextract (v4f32 s2v (f32 load $addr)), c) -> (f32 load $addr+c*size)
26198 // (vextract (v4f32 shuffle (load $addr), <1,u,u,u>), 0) -> (f32 load $addr)
26199 int Elt = IndexC->getZExtValue();
26200 LoadSDNode *LN0 = nullptr;
26201 if (ISD::isNormalLoad(N: VecOp.getNode())) {
26202 LN0 = cast<LoadSDNode>(Val&: VecOp);
26203 } else if (VecOp.getOpcode() == ISD::SCALAR_TO_VECTOR &&
26204 VecOp.getOperand(i: 0).getValueType() == ExtVT &&
26205 ISD::isNormalLoad(N: VecOp.getOperand(i: 0).getNode())) {
26206 // Don't duplicate a load with other uses.
26207 if (!VecOp.hasOneUse())
26208 return SDValue();
26209
26210 LN0 = cast<LoadSDNode>(Val: VecOp.getOperand(i: 0));
26211 }
26212 if (auto *Shuf = dyn_cast<ShuffleVectorSDNode>(Val&: VecOp)) {
26213 // (vextract (vector_shuffle (load $addr), v2, <1, u, u, u>), 1)
26214 // =>
26215 // (load $addr+1*size)
26216
26217 // Don't duplicate a load with other uses.
26218 if (!VecOp.hasOneUse())
26219 return SDValue();
26220
26221 // If the bit convert changed the number of elements, it is unsafe
26222 // to examine the mask.
26223 if (BCNumEltsChanged)
26224 return SDValue();
26225
26226 // Select the input vector, guarding against out of range extract vector.
26227 int Idx = (Elt > (int)NumElts) ? -1 : Shuf->getMaskElt(Idx: Elt);
26228 VecOp = (Idx < (int)NumElts) ? VecOp.getOperand(i: 0) : VecOp.getOperand(i: 1);
26229
26230 if (VecOp.getOpcode() == ISD::BITCAST) {
26231 // Don't duplicate a load with other uses.
26232 if (!VecOp.hasOneUse())
26233 return SDValue();
26234
26235 VecOp = VecOp.getOperand(i: 0);
26236 }
26237 if (ISD::isNormalLoad(N: VecOp.getNode())) {
26238 LN0 = cast<LoadSDNode>(Val&: VecOp);
26239 Elt = (Idx < (int)NumElts) ? Idx : Idx - (int)NumElts;
26240 Index = DAG.getConstant(Val: Elt, DL, VT: Index.getValueType());
26241 }
26242 } else if (VecOp.getOpcode() == ISD::CONCAT_VECTORS && !BCNumEltsChanged &&
26243 VecVT.getVectorElementType() == ScalarVT &&
26244 (!LegalTypes ||
26245 TLI.isTypeLegal(
26246 VT: VecOp.getOperand(i: 0).getValueType().getVectorElementType()))) {
26247 // extract_vector_elt (concat_vectors v2i16:a, v2i16:b), 0
26248 // -> extract_vector_elt a, 0
26249 // extract_vector_elt (concat_vectors v2i16:a, v2i16:b), 1
26250 // -> extract_vector_elt a, 1
26251 // extract_vector_elt (concat_vectors v2i16:a, v2i16:b), 2
26252 // -> extract_vector_elt b, 0
26253 // extract_vector_elt (concat_vectors v2i16:a, v2i16:b), 3
26254 // -> extract_vector_elt b, 1
26255 EVT ConcatVT = VecOp.getOperand(i: 0).getValueType();
26256 unsigned ConcatNumElts = ConcatVT.getVectorNumElements();
26257 SDValue NewIdx = DAG.getConstant(Val: Elt % ConcatNumElts, DL,
26258 VT: Index.getValueType());
26259
26260 SDValue ConcatOp = VecOp.getOperand(i: Elt / ConcatNumElts);
26261 SDValue Elt = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL,
26262 VT: ConcatVT.getVectorElementType(),
26263 N1: ConcatOp, N2: NewIdx);
26264 return DAG.getNode(Opcode: ISD::BITCAST, DL, VT: ScalarVT, Operand: Elt);
26265 }
26266
26267 // Make sure we found a non-volatile load and the extractelement is
26268 // the only use.
26269 if (!LN0 || !LN0->hasNUsesOfValue(NUses: 1,Value: 0) || !LN0->isSimple())
26270 return SDValue();
26271
26272 // If Idx was -1 above, Elt is going to be -1, so just return poison.
26273 if (Elt == -1)
26274 return DAG.getPOISON(VT: LVT);
26275
26276 if (SDValue Scalarized =
26277 TLI.scalarizeExtractedVectorLoad(ResultVT: LVT, DL, InVecVT: VecVT, EltNo: Index, OriginalLoad: LN0, DAG)) {
26278 ++OpsNarrowed;
26279 if (IsFrozen)
26280 return DAG.getFreeze(V: Scalarized);
26281 return Scalarized;
26282 }
26283
26284 return SDValue();
26285}
26286
26287// Simplify (build_vec (ext )) to (bitcast (build_vec ))
26288SDValue DAGCombiner::reduceBuildVecExtToExtBuildVec(SDNode *N) {
26289 // We perform this optimization post type-legalization because
26290 // the type-legalizer often scalarizes integer-promoted vectors.
26291 // Performing this optimization before may create bit-casts which
26292 // will be type-legalized to complex code sequences.
26293 // We perform this optimization only before the operation legalizer because we
26294 // may introduce illegal operations.
26295 if (Level != AfterLegalizeVectorOps && Level != AfterLegalizeTypes)
26296 return SDValue();
26297
26298 unsigned NumInScalars = N->getNumOperands();
26299 SDLoc DL(N);
26300 EVT VT = N->getValueType(ResNo: 0);
26301
26302 // Check to see if this is a BUILD_VECTOR of a bunch of values
26303 // which come from any_extend or zero_extend nodes. If so, we can create
26304 // a new BUILD_VECTOR using bit-casts which may enable other BUILD_VECTOR
26305 // optimizations. We do not handle sign-extend because we can't fill the sign
26306 // using shuffles.
26307 EVT SourceType = MVT::Other;
26308 bool AllAnyExt = true;
26309
26310 for (unsigned i = 0; i != NumInScalars; ++i) {
26311 SDValue In = N->getOperand(Num: i);
26312 // Ignore undef inputs.
26313 if (In.isUndef()) continue;
26314
26315 bool AnyExt = In.getOpcode() == ISD::ANY_EXTEND;
26316 bool ZeroExt = In.getOpcode() == ISD::ZERO_EXTEND;
26317
26318 // Abort if the element is not an extension.
26319 if (!ZeroExt && !AnyExt) {
26320 SourceType = MVT::Other;
26321 break;
26322 }
26323
26324 // The input is a ZeroExt or AnyExt. Check the original type.
26325 EVT InTy = In.getOperand(i: 0).getValueType();
26326
26327 // Check that all of the widened source types are the same.
26328 if (SourceType == MVT::Other)
26329 // First time.
26330 SourceType = InTy;
26331 else if (InTy != SourceType) {
26332 // Multiple income types. Abort.
26333 SourceType = MVT::Other;
26334 break;
26335 }
26336
26337 // Check if all of the extends are ANY_EXTENDs.
26338 AllAnyExt &= AnyExt;
26339 }
26340
26341 // In order to have valid types, all of the inputs must be extended from the
26342 // same source type and all of the inputs must be any or zero extend.
26343 // Scalar sizes must be a power of two.
26344 EVT OutScalarTy = VT.getScalarType();
26345 bool ValidTypes =
26346 SourceType != MVT::Other &&
26347 llvm::has_single_bit<uint32_t>(Value: OutScalarTy.getSizeInBits()) &&
26348 llvm::has_single_bit<uint32_t>(Value: SourceType.getSizeInBits());
26349
26350 // Create a new simpler BUILD_VECTOR sequence which other optimizations can
26351 // turn into a single shuffle instruction.
26352 if (!ValidTypes)
26353 return SDValue();
26354
26355 // If we already have a splat buildvector, then don't fold it if it means
26356 // introducing zeros.
26357 if (!AllAnyExt && DAG.isSplatValue(V: SDValue(N, 0), /*AllowUndefs*/ true))
26358 return SDValue();
26359
26360 bool isLE = DAG.getDataLayout().isLittleEndian();
26361 unsigned ElemRatio = OutScalarTy.getSizeInBits()/SourceType.getSizeInBits();
26362 assert(ElemRatio > 1 && "Invalid element size ratio");
26363 SDValue Filler = AllAnyExt ? DAG.getPOISON(VT: SourceType)
26364 : DAG.getConstant(Val: 0, DL, VT: SourceType);
26365
26366 unsigned NewBVElems = ElemRatio * VT.getVectorNumElements();
26367 SmallVector<SDValue, 8> Ops(NewBVElems, Filler);
26368
26369 // Populate the new build_vector
26370 for (unsigned i = 0, e = N->getNumOperands(); i != e; ++i) {
26371 SDValue Cast = N->getOperand(Num: i);
26372 assert((Cast.getOpcode() == ISD::ANY_EXTEND ||
26373 Cast.getOpcode() == ISD::ZERO_EXTEND ||
26374 Cast.isUndef()) && "Invalid cast opcode");
26375 SDValue In;
26376 if (Cast.isUndef())
26377 In = DAG.getUNDEF(VT: SourceType);
26378 else
26379 In = Cast->getOperand(Num: 0);
26380 unsigned Index = isLE ? (i * ElemRatio) :
26381 (i * ElemRatio + (ElemRatio - 1));
26382
26383 assert(Index < Ops.size() && "Invalid index");
26384 Ops[Index] = In;
26385 }
26386
26387 // The type of the new BUILD_VECTOR node.
26388 EVT VecVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: SourceType, NumElements: NewBVElems);
26389 assert(VecVT.getSizeInBits() == VT.getSizeInBits() &&
26390 "Invalid vector size");
26391 // Check if the new vector type is legal.
26392 if (!isTypeLegal(VT: VecVT) ||
26393 (!TLI.isOperationLegal(Op: ISD::BUILD_VECTOR, VT: VecVT) &&
26394 TLI.isOperationLegal(Op: ISD::BUILD_VECTOR, VT)))
26395 return SDValue();
26396
26397 // Make the new BUILD_VECTOR.
26398 SDValue BV = DAG.getBuildVector(VT: VecVT, DL, Ops);
26399
26400 // The new BUILD_VECTOR node has the potential to be further optimized.
26401 AddToWorklist(N: BV.getNode());
26402 // Bitcast to the desired type.
26403 return DAG.getBitcast(VT, V: BV);
26404}
26405
26406// Simplify (build_vec (trunc $1)
26407// (trunc (srl $1 half-width))
26408// (trunc (srl $1 (2 * half-width))))
26409// to (bitcast $1)
26410SDValue DAGCombiner::reduceBuildVecTruncToBitCast(SDNode *N) {
26411 assert(N->getOpcode() == ISD::BUILD_VECTOR && "Expected build vector");
26412
26413 EVT VT = N->getValueType(ResNo: 0);
26414
26415 // Don't run this before LegalizeTypes if VT is legal.
26416 // Targets may have other preferences.
26417 if (Level < AfterLegalizeTypes && TLI.isTypeLegal(VT))
26418 return SDValue();
26419
26420 // Only for little endian
26421 if (!DAG.getDataLayout().isLittleEndian())
26422 return SDValue();
26423
26424 EVT OutScalarTy = VT.getScalarType();
26425 uint64_t ScalarTypeBitsize = OutScalarTy.getSizeInBits();
26426
26427 // Only for power of two types to be sure that bitcast works well
26428 if (!isPowerOf2_64(Value: ScalarTypeBitsize))
26429 return SDValue();
26430
26431 unsigned NumInScalars = N->getNumOperands();
26432
26433 // Look through bitcasts
26434 auto PeekThroughBitcast = [](SDValue Op) {
26435 if (Op.getOpcode() == ISD::BITCAST)
26436 return Op.getOperand(i: 0);
26437 return Op;
26438 };
26439
26440 // The source value where all the parts are extracted.
26441 SDValue Src;
26442 for (unsigned i = 0; i != NumInScalars; ++i) {
26443 SDValue In = PeekThroughBitcast(N->getOperand(Num: i));
26444 // Ignore undef inputs.
26445 if (In.isUndef()) continue;
26446
26447 if (In.getOpcode() != ISD::TRUNCATE)
26448 return SDValue();
26449
26450 In = PeekThroughBitcast(In.getOperand(i: 0));
26451
26452 if (In.getOpcode() != ISD::SRL) {
26453 // For now only build_vec without shuffling, handle shifts here in the
26454 // future.
26455 if (i != 0)
26456 return SDValue();
26457
26458 Src = In;
26459 } else {
26460 // In is SRL
26461 SDValue part = PeekThroughBitcast(In.getOperand(i: 0));
26462
26463 if (!Src) {
26464 Src = part;
26465 } else if (Src != part) {
26466 // Vector parts do not stem from the same variable
26467 return SDValue();
26468 }
26469
26470 SDValue ShiftAmtVal = In.getOperand(i: 1);
26471 if (!isa<ConstantSDNode>(Val: ShiftAmtVal))
26472 return SDValue();
26473
26474 uint64_t ShiftAmt = In.getConstantOperandVal(i: 1);
26475
26476 // The extracted value is not extracted at the right position
26477 if (ShiftAmt != i * ScalarTypeBitsize)
26478 return SDValue();
26479 }
26480 }
26481
26482 // Only cast if the size is the same
26483 if (!Src || Src.getValueType().getSizeInBits() != VT.getSizeInBits())
26484 return SDValue();
26485
26486 return DAG.getBitcast(VT, V: Src);
26487}
26488
26489SDValue DAGCombiner::createBuildVecShuffle(const SDLoc &DL, SDNode *N,
26490 ArrayRef<int> VectorMask,
26491 SDValue VecIn1, SDValue VecIn2,
26492 unsigned LeftIdx, bool DidSplitVec) {
26493 EVT VT = N->getValueType(ResNo: 0);
26494 EVT InVT1 = VecIn1.getValueType();
26495 EVT InVT2 = VecIn2.getNode() ? VecIn2.getValueType() : InVT1;
26496
26497 unsigned NumElems = VT.getVectorNumElements();
26498 unsigned ShuffleNumElems = NumElems;
26499
26500 // If we artificially split a vector in two already, then the offsets in the
26501 // operands will all be based off of VecIn1, even those in VecIn2.
26502 unsigned Vec2Offset = DidSplitVec ? 0 : InVT1.getVectorNumElements();
26503
26504 uint64_t VTSize = VT.getFixedSizeInBits();
26505 uint64_t InVT1Size = InVT1.getFixedSizeInBits();
26506 uint64_t InVT2Size = InVT2.getFixedSizeInBits();
26507
26508 assert(InVT2Size <= InVT1Size &&
26509 "Inputs must be sorted to be in non-increasing vector size order.");
26510
26511 // We can't generate a shuffle node with mismatched input and output types.
26512 // Try to make the types match the type of the output.
26513 if (InVT1 != VT || InVT2 != VT) {
26514 if ((VTSize % InVT1Size == 0) && InVT1 == InVT2) {
26515 // If the output vector length is a multiple of both input lengths,
26516 // we can concatenate them and pad the rest with poison.
26517 unsigned NumConcats = VTSize / InVT1Size;
26518 assert(NumConcats >= 2 && "Concat needs at least two inputs!");
26519 SmallVector<SDValue, 2> ConcatOps(NumConcats, DAG.getPOISON(VT: InVT1));
26520 ConcatOps[0] = VecIn1;
26521 ConcatOps[1] = VecIn2 ? VecIn2 : DAG.getPOISON(VT: InVT1);
26522 VecIn1 = DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT, Ops: ConcatOps);
26523 VecIn2 = SDValue();
26524 } else if (InVT1Size == VTSize * 2) {
26525 if (TLI.getExtractSubvectorCost(ResVT: VT, SrcVT: InVT1, Index: NumElems) >
26526 TargetLowering::ExtractSubvectorCost::Cheap)
26527 return SDValue();
26528
26529 if (!VecIn2.getNode()) {
26530 // If we only have one input vector, and it's twice the size of the
26531 // output, split it in two.
26532 VecIn2 = DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT, N1: VecIn1,
26533 N2: DAG.getVectorIdxConstant(Val: NumElems, DL));
26534 VecIn1 = DAG.getExtractSubvector(DL, VT, Vec: VecIn1, Idx: 0);
26535 // Since we now have shorter input vectors, adjust the offset of the
26536 // second vector's start.
26537 Vec2Offset = NumElems;
26538 } else {
26539 assert(InVT2Size <= InVT1Size &&
26540 "Second input is not going to be larger than the first one.");
26541
26542 // VecIn1 is wider than the output, and we have another, possibly
26543 // smaller input. Pad the smaller input with undefs, shuffle at the
26544 // input vector width, and extract the output.
26545 // The shuffle type is different than VT, so check legality again.
26546 if (LegalOperations &&
26547 !TLI.isOperationLegal(Op: ISD::VECTOR_SHUFFLE, VT: InVT1))
26548 return SDValue();
26549
26550 // Legalizing INSERT_SUBVECTOR is tricky - you basically have to
26551 // lower it back into a BUILD_VECTOR. So if the inserted type is
26552 // illegal, don't even try.
26553 if (InVT1 != InVT2) {
26554 if (!TLI.isTypeLegal(VT: InVT2))
26555 return SDValue();
26556 VecIn2 = DAG.getInsertSubvector(DL, Vec: DAG.getPOISON(VT: InVT1), SubVec: VecIn2, Idx: 0);
26557 }
26558 ShuffleNumElems = NumElems * 2;
26559 }
26560 } else if (InVT2Size * 2 == VTSize && InVT1Size == VTSize) {
26561 SmallVector<SDValue, 2> ConcatOps(2, DAG.getPOISON(VT: InVT2));
26562 ConcatOps[0] = VecIn2;
26563 VecIn2 = DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT, Ops: ConcatOps);
26564 } else if (InVT1Size / VTSize > 1 && InVT1Size % VTSize == 0) {
26565 if (TLI.getExtractSubvectorCost(ResVT: VT, SrcVT: InVT1, Index: NumElems) >
26566 TargetLowering::ExtractSubvectorCost::Cheap ||
26567 !TLI.isTypeLegal(VT: InVT1) || !TLI.isTypeLegal(VT: InVT2))
26568 return SDValue();
26569 // If dest vector has less than two elements, then use shuffle and extract
26570 // from larger regs will cost even more.
26571 if (VT.getVectorNumElements() <= 2 || !VecIn2.getNode())
26572 return SDValue();
26573 assert(InVT2Size <= InVT1Size &&
26574 "Second input is not going to be larger than the first one.");
26575
26576 // VecIn1 is wider than the output, and we have another, possibly
26577 // smaller input. Pad the smaller input with undefs, shuffle at the
26578 // input vector width, and extract the output.
26579 // The shuffle type is different than VT, so check legality again.
26580 if (LegalOperations && !TLI.isOperationLegal(Op: ISD::VECTOR_SHUFFLE, VT: InVT1))
26581 return SDValue();
26582
26583 if (InVT1 != InVT2) {
26584 VecIn2 = DAG.getInsertSubvector(DL, Vec: DAG.getPOISON(VT: InVT1), SubVec: VecIn2, Idx: 0);
26585 }
26586 ShuffleNumElems = InVT1Size / VTSize * NumElems;
26587 } else {
26588 // TODO: Support cases where the length mismatch isn't exactly by a
26589 // factor of 2.
26590 // TODO: Move this check upwards, so that if we have bad type
26591 // mismatches, we don't create any DAG nodes.
26592 return SDValue();
26593 }
26594 }
26595
26596 // Initialize mask to undef.
26597 SmallVector<int, 8> Mask(ShuffleNumElems, -1);
26598
26599 // Only need to run up to the number of elements actually used, not the
26600 // total number of elements in the shuffle - if we are shuffling a wider
26601 // vector, the high lanes should be set to undef.
26602 for (unsigned i = 0; i != NumElems; ++i) {
26603 if (VectorMask[i] <= 0)
26604 continue;
26605
26606 unsigned ExtIndex = N->getOperand(Num: i).getConstantOperandVal(i: 1);
26607 if (VectorMask[i] == (int)LeftIdx) {
26608 Mask[i] = ExtIndex;
26609 } else if (VectorMask[i] == (int)LeftIdx + 1) {
26610 Mask[i] = Vec2Offset + ExtIndex;
26611 }
26612 }
26613
26614 // The type the input vectors may have changed above.
26615 InVT1 = VecIn1.getValueType();
26616
26617 // If we already have a VecIn2, it should have the same type as VecIn1.
26618 // If we don't, get an poison/zero vector of the appropriate type.
26619 VecIn2 = VecIn2.getNode() ? VecIn2 : DAG.getPOISON(VT: InVT1);
26620 assert(InVT1 == VecIn2.getValueType() && "Unexpected second input type.");
26621
26622 SDValue Shuffle = DAG.getVectorShuffle(VT: InVT1, dl: DL, N1: VecIn1, N2: VecIn2, Mask);
26623 if (ShuffleNumElems > NumElems)
26624 Shuffle = DAG.getExtractSubvector(DL, VT, Vec: Shuffle, Idx: 0);
26625
26626 return Shuffle;
26627}
26628
26629static SDValue reduceBuildVecToShuffleWithZero(SDNode *BV, SelectionDAG &DAG) {
26630 assert(BV->getOpcode() == ISD::BUILD_VECTOR && "Expected build vector");
26631
26632 // First, determine where the build vector is not undef.
26633 // TODO: We could extend this to handle zero elements as well as undefs.
26634 int NumBVOps = BV->getNumOperands();
26635 int ZextElt = -1;
26636 for (int i = 0; i != NumBVOps; ++i) {
26637 SDValue Op = BV->getOperand(Num: i);
26638 if (Op.isUndef())
26639 continue;
26640 if (ZextElt == -1)
26641 ZextElt = i;
26642 else
26643 return SDValue();
26644 }
26645 // Bail out if there's no non-undef element.
26646 if (ZextElt == -1)
26647 return SDValue();
26648
26649 // The build vector contains some number of undef elements and exactly
26650 // one other element. That other element must be a zero-extended scalar
26651 // extracted from a vector at a constant index to turn this into a shuffle.
26652 // Also, require that the build vector does not implicitly truncate/extend
26653 // its elements.
26654 // TODO: This could be enhanced to allow ANY_EXTEND as well as ZERO_EXTEND.
26655 EVT VT = BV->getValueType(ResNo: 0);
26656 SDValue Zext = BV->getOperand(Num: ZextElt);
26657 if (Zext.getOpcode() != ISD::ZERO_EXTEND || !Zext.hasOneUse() ||
26658 Zext.getOperand(i: 0).getOpcode() != ISD::EXTRACT_VECTOR_ELT ||
26659 !isa<ConstantSDNode>(Val: Zext.getOperand(i: 0).getOperand(i: 1)) ||
26660 Zext.getValueSizeInBits() != VT.getScalarSizeInBits())
26661 return SDValue();
26662
26663 // The zero-extend must be a multiple of the source size, and we must be
26664 // building a vector of the same size as the source of the extract element.
26665 SDValue Extract = Zext.getOperand(i: 0);
26666 unsigned DestSize = Zext.getValueSizeInBits();
26667 unsigned SrcSize = Extract.getValueSizeInBits();
26668 if (DestSize % SrcSize != 0 ||
26669 Extract.getOperand(i: 0).getValueSizeInBits() != VT.getSizeInBits())
26670 return SDValue();
26671
26672 // Create a shuffle mask that will combine the extracted element with zeros
26673 // and undefs.
26674 int ZextRatio = DestSize / SrcSize;
26675 int NumMaskElts = NumBVOps * ZextRatio;
26676 SmallVector<int, 32> ShufMask(NumMaskElts, -1);
26677 for (int i = 0; i != NumMaskElts; ++i) {
26678 if (i / ZextRatio == ZextElt) {
26679 // The low bits of the (potentially translated) extracted element map to
26680 // the source vector. The high bits map to zero. We will use a zero vector
26681 // as the 2nd source operand of the shuffle, so use the 1st element of
26682 // that vector (mask value is number-of-elements) for the high bits.
26683 int Low = DAG.getDataLayout().isBigEndian() ? (ZextRatio - 1) : 0;
26684 ShufMask[i] = (i % ZextRatio == Low) ? Extract.getConstantOperandVal(i: 1)
26685 : NumMaskElts;
26686 }
26687
26688 // Undef elements of the build vector remain undef because we initialize
26689 // the shuffle mask with -1.
26690 }
26691
26692 // buildvec undef, ..., (zext (extractelt V, IndexC)), undef... -->
26693 // bitcast (shuffle V, ZeroVec, VectorMask)
26694 SDLoc DL(BV);
26695 EVT VecVT = Extract.getOperand(i: 0).getValueType();
26696 SDValue ZeroVec = DAG.getConstant(Val: 0, DL, VT: VecVT);
26697 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
26698 SDValue Shuf = TLI.buildLegalVectorShuffle(VT: VecVT, DL, N0: Extract.getOperand(i: 0),
26699 N1: ZeroVec, Mask: ShufMask, DAG);
26700 if (!Shuf)
26701 return SDValue();
26702 return DAG.getBitcast(VT, V: Shuf);
26703}
26704
26705// FIXME: promote to STLExtras.
26706template <typename R, typename T>
26707static auto getFirstIndexOf(R &&Range, const T &Val) {
26708 auto I = find(Range, Val);
26709 if (I == Range.end())
26710 return static_cast<decltype(std::distance(Range.begin(), I))>(-1);
26711 return std::distance(Range.begin(), I);
26712}
26713
26714// Check to see if this is a BUILD_VECTOR of a bunch of EXTRACT_VECTOR_ELT
26715// operations. If the types of the vectors we're extracting from allow it,
26716// turn this into a vector_shuffle node.
26717SDValue DAGCombiner::reduceBuildVecToShuffle(SDNode *N) {
26718 SDLoc DL(N);
26719 EVT VT = N->getValueType(ResNo: 0);
26720
26721 // Only type-legal BUILD_VECTOR nodes are converted to shuffle nodes.
26722 if (!isTypeLegal(VT))
26723 return SDValue();
26724
26725 if (SDValue V = reduceBuildVecToShuffleWithZero(BV: N, DAG))
26726 return V;
26727
26728 // May only combine to shuffle after legalize if shuffle is legal.
26729 if (LegalOperations && !TLI.isOperationLegal(Op: ISD::VECTOR_SHUFFLE, VT))
26730 return SDValue();
26731
26732 bool UsesZeroVector = false;
26733 unsigned NumElems = N->getNumOperands();
26734
26735 // Record, for each element of the newly built vector, which input vector
26736 // that element comes from. -1 stands for undef, 0 for the zero vector,
26737 // and positive values for the input vectors.
26738 // VectorMask maps each element to its vector number, and VecIn maps vector
26739 // numbers to their initial SDValues.
26740
26741 SmallVector<int, 8> VectorMask(NumElems, -1);
26742 SmallVector<SDValue, 8> VecIn;
26743 VecIn.push_back(Elt: SDValue());
26744
26745 // If we have a single extract_element with a constant index, track the index
26746 // value.
26747 unsigned OneConstExtractIndex = ~0u;
26748
26749 // Count the number of extract_vector_elt sources (i.e. non-constant or undef)
26750 unsigned NumExtracts = 0;
26751
26752 for (unsigned i = 0; i != NumElems; ++i) {
26753 SDValue Op = N->getOperand(Num: i);
26754
26755 if (Op.isUndef())
26756 continue;
26757
26758 // See if we can use a blend with a zero vector.
26759 // TODO: Should we generalize this to a blend with an arbitrary constant
26760 // vector?
26761 if (isNullConstant(V: Op) || isNullFPConstant(V: Op)) {
26762 UsesZeroVector = true;
26763 VectorMask[i] = 0;
26764 continue;
26765 }
26766
26767 // Not an undef or zero. If the input is something other than an
26768 // EXTRACT_VECTOR_ELT with an in-range constant index, bail out.
26769 if (Op.getOpcode() != ISD::EXTRACT_VECTOR_ELT)
26770 return SDValue();
26771
26772 SDValue ExtractedFromVec = Op.getOperand(i: 0);
26773 if (ExtractedFromVec.getValueType().isScalableVector())
26774 return SDValue();
26775 auto *ExtractIdx = dyn_cast<ConstantSDNode>(Val: Op.getOperand(i: 1));
26776 if (!ExtractIdx)
26777 return SDValue();
26778
26779 if (ExtractIdx->getAsAPIntVal().uge(
26780 RHS: ExtractedFromVec.getValueType().getVectorNumElements()))
26781 return SDValue();
26782
26783 // All inputs must have the same element type as the output.
26784 if (VT.getVectorElementType() !=
26785 ExtractedFromVec.getValueType().getVectorElementType())
26786 return SDValue();
26787
26788 OneConstExtractIndex = ExtractIdx->getZExtValue();
26789 ++NumExtracts;
26790
26791 // Have we seen this input vector before?
26792 // The vectors are expected to be tiny (usually 1 or 2 elements), so using
26793 // a map back from SDValues to numbers isn't worth it.
26794 int Idx = getFirstIndexOf(Range&: VecIn, Val: ExtractedFromVec);
26795 if (Idx == -1) { // A new source vector?
26796 Idx = VecIn.size();
26797 VecIn.push_back(Elt: ExtractedFromVec);
26798 }
26799
26800 VectorMask[i] = Idx;
26801 }
26802
26803 // If we didn't find at least one input vector, bail out.
26804 if (VecIn.size() < 2)
26805 return SDValue();
26806
26807 // If all the Operands of BUILD_VECTOR extract from same
26808 // vector, then split the vector efficiently based on the maximum
26809 // vector access index and adjust the VectorMask and
26810 // VecIn accordingly.
26811 bool DidSplitVec = false;
26812 if (VecIn.size() == 2) {
26813 // If we only found a single constant indexed extract_vector_elt feeding the
26814 // build_vector, do not produce a more complicated shuffle if the extract is
26815 // cheap with other constant/undef elements. Skip broadcast patterns with
26816 // multiple uses in the build_vector.
26817
26818 // TODO: This should be more aggressive about skipping the shuffle
26819 // formation, particularly if VecIn[1].hasOneUse(), and regardless of the
26820 // index.
26821 if (NumExtracts == 1 &&
26822 TLI.isOperationLegalOrCustom(Op: ISD::EXTRACT_VECTOR_ELT, VT) &&
26823 TLI.isTypeLegal(VT: VT.getVectorElementType()) &&
26824 TLI.isExtractVecEltCheap(VT, Index: OneConstExtractIndex))
26825 return SDValue();
26826
26827 unsigned MaxIndex = 0;
26828 unsigned NearestPow2 = 0;
26829 SDValue Vec = VecIn.back();
26830 EVT InVT = Vec.getValueType();
26831 SmallVector<unsigned, 8> IndexVec(NumElems, 0);
26832
26833 for (unsigned i = 0; i < NumElems; i++) {
26834 if (VectorMask[i] <= 0)
26835 continue;
26836 unsigned Index = N->getOperand(Num: i).getConstantOperandVal(i: 1);
26837 IndexVec[i] = Index;
26838 MaxIndex = std::max(a: MaxIndex, b: Index);
26839 }
26840
26841 NearestPow2 = PowerOf2Ceil(A: MaxIndex);
26842 if (InVT.isSimple() && NearestPow2 > 2 && MaxIndex < NearestPow2 &&
26843 NumElems * 2 < NearestPow2) {
26844 unsigned SplitSize = NearestPow2 / 2;
26845 EVT SplitVT = EVT::getVectorVT(Context&: *DAG.getContext(),
26846 VT: InVT.getVectorElementType(), NumElements: SplitSize);
26847 if (TLI.isTypeLegal(VT: SplitVT) &&
26848 SplitSize + SplitVT.getVectorNumElements() <=
26849 InVT.getVectorNumElements()) {
26850 SDValue VecIn2 = DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: SplitVT, N1: Vec,
26851 N2: DAG.getVectorIdxConstant(Val: SplitSize, DL));
26852 SDValue VecIn1 = DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: SplitVT, N1: Vec,
26853 N2: DAG.getVectorIdxConstant(Val: 0, DL));
26854 VecIn.pop_back();
26855 VecIn.push_back(Elt: VecIn1);
26856 VecIn.push_back(Elt: VecIn2);
26857 DidSplitVec = true;
26858
26859 for (unsigned i = 0; i < NumElems; i++) {
26860 if (VectorMask[i] <= 0)
26861 continue;
26862 VectorMask[i] = (IndexVec[i] < SplitSize) ? 1 : 2;
26863 }
26864 }
26865 }
26866 }
26867
26868 // Sort input vectors by decreasing vector element count,
26869 // while preserving the relative order of equally-sized vectors.
26870 // Note that we keep the first "implicit zero vector as-is.
26871 SmallVector<SDValue, 8> SortedVecIn(VecIn);
26872 llvm::stable_sort(Range: MutableArrayRef<SDValue>(SortedVecIn).drop_front(),
26873 C: [](const SDValue &a, const SDValue &b) {
26874 return a.getValueType().getVectorNumElements() >
26875 b.getValueType().getVectorNumElements();
26876 });
26877
26878 // We now also need to rebuild the VectorMask, because it referenced element
26879 // order in VecIn, and we just sorted them.
26880 for (int &SourceVectorIndex : VectorMask) {
26881 if (SourceVectorIndex <= 0)
26882 continue;
26883 unsigned Idx = getFirstIndexOf(Range&: SortedVecIn, Val: VecIn[SourceVectorIndex]);
26884 assert(Idx > 0 && Idx < SortedVecIn.size() &&
26885 VecIn[SourceVectorIndex] == SortedVecIn[Idx] && "Remapping failure");
26886 SourceVectorIndex = Idx;
26887 }
26888
26889 VecIn = std::move(SortedVecIn);
26890
26891 // TODO: Should this fire if some of the input vectors has illegal type (like
26892 // it does now), or should we let legalization run its course first?
26893
26894 // Shuffle phase:
26895 // Take pairs of vectors, and shuffle them so that the result has elements
26896 // from these vectors in the correct places.
26897 // For example, given:
26898 // t10: i32 = extract_vector_elt t1, Constant:i64<0>
26899 // t11: i32 = extract_vector_elt t2, Constant:i64<0>
26900 // t12: i32 = extract_vector_elt t3, Constant:i64<0>
26901 // t13: i32 = extract_vector_elt t1, Constant:i64<1>
26902 // t14: v4i32 = BUILD_VECTOR t10, t11, t12, t13
26903 // We will generate:
26904 // t20: v4i32 = vector_shuffle<0,4,u,1> t1, t2
26905 // t21: v4i32 = vector_shuffle<u,u,0,u> t3, undef
26906 SmallVector<SDValue, 4> Shuffles;
26907 for (unsigned In = 0, Len = (VecIn.size() / 2); In < Len; ++In) {
26908 unsigned LeftIdx = 2 * In + 1;
26909 SDValue VecLeft = VecIn[LeftIdx];
26910 SDValue VecRight =
26911 (LeftIdx + 1) < VecIn.size() ? VecIn[LeftIdx + 1] : SDValue();
26912
26913 if (SDValue Shuffle = createBuildVecShuffle(DL, N, VectorMask, VecIn1: VecLeft,
26914 VecIn2: VecRight, LeftIdx, DidSplitVec))
26915 Shuffles.push_back(Elt: Shuffle);
26916 else
26917 return SDValue();
26918 }
26919
26920 // If we need the zero vector as an "ingredient" in the blend tree, add it
26921 // to the list of shuffles.
26922 if (UsesZeroVector)
26923 Shuffles.push_back(Elt: VT.isInteger() ? DAG.getConstant(Val: 0, DL, VT)
26924 : DAG.getConstantFP(Val: 0.0, DL, VT));
26925
26926 // If we only have one shuffle, we're done.
26927 if (Shuffles.size() == 1)
26928 return Shuffles[0];
26929
26930 // Update the vector mask to point to the post-shuffle vectors.
26931 for (int &Vec : VectorMask)
26932 if (Vec == 0)
26933 Vec = Shuffles.size() - 1;
26934 else
26935 Vec = (Vec - 1) / 2;
26936
26937 // More than one shuffle. Generate a binary tree of blends, e.g. if from
26938 // the previous step we got the set of shuffles t10, t11, t12, t13, we will
26939 // generate:
26940 // t10: v8i32 = vector_shuffle<0,8,u,u,u,u,u,u> t1, t2
26941 // t11: v8i32 = vector_shuffle<u,u,0,8,u,u,u,u> t3, t4
26942 // t12: v8i32 = vector_shuffle<u,u,u,u,0,8,u,u> t5, t6
26943 // t13: v8i32 = vector_shuffle<u,u,u,u,u,u,0,8> t7, t8
26944 // t20: v8i32 = vector_shuffle<0,1,10,11,u,u,u,u> t10, t11
26945 // t21: v8i32 = vector_shuffle<u,u,u,u,4,5,14,15> t12, t13
26946 // t30: v8i32 = vector_shuffle<0,1,2,3,12,13,14,15> t20, t21
26947
26948 // Make sure the initial size of the shuffle list is even.
26949 if (Shuffles.size() % 2)
26950 Shuffles.push_back(Elt: DAG.getPOISON(VT));
26951
26952 for (unsigned CurSize = Shuffles.size(); CurSize > 1; CurSize /= 2) {
26953 if (CurSize % 2) {
26954 Shuffles[CurSize] = DAG.getPOISON(VT);
26955 CurSize++;
26956 }
26957 for (unsigned In = 0, Len = CurSize / 2; In < Len; ++In) {
26958 int Left = 2 * In;
26959 int Right = 2 * In + 1;
26960 SmallVector<int, 8> Mask(NumElems, -1);
26961 SDValue L = Shuffles[Left];
26962 ArrayRef<int> LMask;
26963 bool IsLeftShuffle = L.getOpcode() == ISD::VECTOR_SHUFFLE &&
26964 L.use_empty() && L.getOperand(i: 1).isUndef() &&
26965 L.getOperand(i: 0).getValueType() == L.getValueType();
26966 if (IsLeftShuffle) {
26967 LMask = cast<ShuffleVectorSDNode>(Val: L.getNode())->getMask();
26968 L = L.getOperand(i: 0);
26969 }
26970 SDValue R = Shuffles[Right];
26971 ArrayRef<int> RMask;
26972 bool IsRightShuffle = R.getOpcode() == ISD::VECTOR_SHUFFLE &&
26973 R.use_empty() && R.getOperand(i: 1).isUndef() &&
26974 R.getOperand(i: 0).getValueType() == R.getValueType();
26975 if (IsRightShuffle) {
26976 RMask = cast<ShuffleVectorSDNode>(Val: R.getNode())->getMask();
26977 R = R.getOperand(i: 0);
26978 }
26979 for (unsigned I = 0; I != NumElems; ++I) {
26980 if (VectorMask[I] == Left) {
26981 Mask[I] = I;
26982 if (IsLeftShuffle)
26983 Mask[I] = LMask[I];
26984 VectorMask[I] = In;
26985 } else if (VectorMask[I] == Right) {
26986 Mask[I] = I + NumElems;
26987 if (IsRightShuffle)
26988 Mask[I] = RMask[I] + NumElems;
26989 VectorMask[I] = In;
26990 }
26991 }
26992
26993 Shuffles[In] = DAG.getVectorShuffle(VT, dl: DL, N1: L, N2: R, Mask);
26994 }
26995 }
26996 return Shuffles[0];
26997}
26998
26999// Try to turn a build vector of zero/sign extends of extract vector elts into
27000// a vector zero/sign extend and possibly an extract subvector.
27001// TODO: Allow undef elements?
27002SDValue DAGCombiner::convertBuildVecExtToExt(SDNode *N) {
27003 if (LegalOperations)
27004 return SDValue();
27005
27006 EVT VT = N->getValueType(ResNo: 0);
27007
27008 bool FoundZeroExtend = false;
27009 bool FoundSignExtend = false;
27010 SDValue Op0 = N->getOperand(Num: 0);
27011 auto checkElem = [&](SDValue Op) -> int64_t {
27012 unsigned Opc = Op.getOpcode();
27013 FoundZeroExtend |= (Opc == ISD::ZERO_EXTEND);
27014 FoundSignExtend |= (Opc == ISD::SIGN_EXTEND);
27015 if ((Opc == ISD::ZERO_EXTEND || Opc == ISD::SIGN_EXTEND ||
27016 Opc == ISD::ANY_EXTEND) &&
27017 Op.getOperand(i: 0).getOpcode() == ISD::EXTRACT_VECTOR_ELT &&
27018 Op0.getOperand(i: 0).getOperand(i: 0) == Op.getOperand(i: 0).getOperand(i: 0))
27019 if (auto *C = dyn_cast<ConstantSDNode>(Val: Op.getOperand(i: 0).getOperand(i: 1)))
27020 return C->getZExtValue();
27021 return -1;
27022 };
27023
27024 // Make sure the first element matches
27025 // (zext (extract_vector_elt X, C))
27026 // Offset must be a constant multiple of the
27027 // known-minimum vector length of the result type.
27028 int64_t Offset = checkElem(Op0);
27029 if (Offset < 0 || (Offset % VT.getVectorNumElements()) != 0)
27030 return SDValue();
27031
27032 unsigned NumElems = N->getNumOperands();
27033 SDValue In = Op0.getOperand(i: 0).getOperand(i: 0);
27034 EVT InSVT = In.getValueType().getScalarType();
27035 EVT InVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: InSVT, NumElements: NumElems);
27036
27037 // Don't create an illegal input type after type legalization.
27038 if (LegalTypes && !TLI.isTypeLegal(VT: InVT))
27039 return SDValue();
27040
27041 // Ensure all the elements come from the same vector and are adjacent.
27042 for (unsigned i = 1; i != NumElems; ++i) {
27043 if ((Offset + i) != checkElem(N->getOperand(Num: i)))
27044 return SDValue();
27045 }
27046
27047 // Can't mix zero and sign extends in the same build_vector.
27048 if (FoundZeroExtend && FoundSignExtend)
27049 return SDValue();
27050
27051 unsigned ExtOpc = ISD::ANY_EXTEND;
27052 if (FoundSignExtend)
27053 ExtOpc = ISD::SIGN_EXTEND;
27054 else if (FoundZeroExtend)
27055 ExtOpc = ISD::ZERO_EXTEND;
27056
27057 SDLoc DL(N);
27058 In = DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: InVT, N1: In,
27059 N2: Op0.getOperand(i: 0).getOperand(i: 1));
27060 return DAG.getNode(Opcode: ExtOpc, DL, VT, Operand: In);
27061}
27062
27063// If this is a very simple BUILD_VECTOR with first element being a ZERO_EXTEND,
27064// and all other elements being constant zero's, granularize the BUILD_VECTOR's
27065// element width, absorbing the ZERO_EXTEND, turning it into a constant zero op.
27066// This patten can appear during legalization.
27067//
27068// NOTE: This can be generalized to allow more than a single
27069// non-constant-zero op, UNDEF's, and to be KnownBits-based,
27070SDValue DAGCombiner::convertBuildVecZextToBuildVecWithZeros(SDNode *N) {
27071 // Don't run this after legalization. Targets may have other preferences.
27072 if (Level >= AfterLegalizeDAG)
27073 return SDValue();
27074
27075 // FIXME: support big-endian.
27076 if (DAG.getDataLayout().isBigEndian())
27077 return SDValue();
27078
27079 EVT VT = N->getValueType(ResNo: 0);
27080 EVT OpVT = N->getOperand(Num: 0).getValueType();
27081 assert(!VT.isScalableVector() && "Encountered scalable BUILD_VECTOR?");
27082
27083 EVT OpIntVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: OpVT.getSizeInBits());
27084
27085 if (!TLI.isTypeLegal(VT: OpIntVT) ||
27086 (LegalOperations && !TLI.isOperationLegalOrCustom(Op: ISD::BITCAST, VT: OpIntVT)))
27087 return SDValue();
27088
27089 unsigned EltBitwidth = VT.getScalarSizeInBits();
27090 // NOTE: the actual width of operands may be wider than that!
27091
27092 // Analyze all operands of this BUILD_VECTOR. What is the largest number of
27093 // active bits they all have? We'll want to truncate them all to that width.
27094 unsigned ActiveBits = 0;
27095 APInt KnownZeroOps(VT.getVectorNumElements(), 0);
27096 for (auto I : enumerate(First: N->ops())) {
27097 SDValue Op = I.value();
27098 // FIXME: support UNDEF elements?
27099 if (auto *Cst = dyn_cast<ConstantSDNode>(Val&: Op)) {
27100 unsigned OpActiveBits =
27101 Cst->getAPIntValue().trunc(width: EltBitwidth).getActiveBits();
27102 if (OpActiveBits == 0) {
27103 KnownZeroOps.setBit(I.index());
27104 continue;
27105 }
27106 // Profitability check: don't allow non-zero constant operands.
27107 return SDValue();
27108 }
27109 // Profitability check: there must only be a single non-zero operand,
27110 // and it must be the first operand of the BUILD_VECTOR.
27111 if (I.index() != 0)
27112 return SDValue();
27113 // The operand must be a zero-extension itself.
27114 // FIXME: this could be generalized to known leading zeros check.
27115 if (Op.getOpcode() != ISD::ZERO_EXTEND)
27116 return SDValue();
27117 unsigned CurrActiveBits =
27118 Op.getOperand(i: 0).getValueSizeInBits().getFixedValue();
27119 assert(!ActiveBits && "Already encountered non-constant-zero operand?");
27120 ActiveBits = CurrActiveBits;
27121 // We want to at least halve the element size.
27122 if (2 * ActiveBits > EltBitwidth)
27123 return SDValue();
27124 }
27125
27126 // This BUILD_VECTOR must have at least one non-constant-zero operand.
27127 if (ActiveBits == 0)
27128 return SDValue();
27129
27130 // We have EltBitwidth bits, the *minimal* chunk size is ActiveBits,
27131 // into how many chunks can we split our element width?
27132 EVT NewScalarIntVT, NewIntVT;
27133 std::optional<unsigned> Factor;
27134 // We can split the element into at least two chunks, but not into more
27135 // than |_ EltBitwidth / ActiveBits _| chunks. Find a largest split factor
27136 // for which the element width is a multiple of it,
27137 // and the resulting types/operations on that chunk width are legal.
27138 assert(2 * ActiveBits <= EltBitwidth &&
27139 "We know that half or less bits of the element are active.");
27140 for (unsigned Scale = EltBitwidth / ActiveBits; Scale >= 2; --Scale) {
27141 if (EltBitwidth % Scale != 0)
27142 continue;
27143 unsigned ChunkBitwidth = EltBitwidth / Scale;
27144 assert(ChunkBitwidth >= ActiveBits && "As per starting point.");
27145 NewScalarIntVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: ChunkBitwidth);
27146 NewIntVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: NewScalarIntVT,
27147 NumElements: Scale * N->getNumOperands());
27148 if (!TLI.isTypeLegal(VT: NewScalarIntVT) || !TLI.isTypeLegal(VT: NewIntVT) ||
27149 (LegalOperations &&
27150 !(TLI.isOperationLegalOrCustom(Op: ISD::TRUNCATE, VT: NewScalarIntVT) &&
27151 TLI.isOperationLegalOrCustom(Op: ISD::BUILD_VECTOR, VT: NewIntVT))))
27152 continue;
27153 Factor = Scale;
27154 break;
27155 }
27156 if (!Factor)
27157 return SDValue();
27158
27159 SDLoc DL(N);
27160 SDValue ZeroOp = DAG.getConstant(Val: 0, DL, VT: NewScalarIntVT);
27161
27162 // Recreate the BUILD_VECTOR, with elements now being Factor times smaller.
27163 SmallVector<SDValue, 16> NewOps;
27164 NewOps.reserve(N: NewIntVT.getVectorNumElements());
27165 for (auto I : enumerate(First: N->ops())) {
27166 SDValue Op = I.value();
27167 assert(!Op.isUndef() && "FIXME: after allowing UNDEF's, handle them here.");
27168 unsigned SrcOpIdx = I.index();
27169 if (KnownZeroOps[SrcOpIdx]) {
27170 NewOps.append(NumInputs: *Factor, Elt: ZeroOp);
27171 continue;
27172 }
27173 Op = DAG.getBitcast(VT: OpIntVT, V: Op);
27174 Op = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: NewScalarIntVT, Operand: Op);
27175 NewOps.emplace_back(Args&: Op);
27176 NewOps.append(NumInputs: *Factor - 1, Elt: ZeroOp);
27177 }
27178 assert(NewOps.size() == NewIntVT.getVectorNumElements());
27179 SDValue NewBV = DAG.getBuildVector(VT: NewIntVT, DL, Ops: NewOps);
27180 NewBV = DAG.getBitcast(VT, V: NewBV);
27181 return NewBV;
27182}
27183
27184SDValue DAGCombiner::visitBUILD_VECTOR(SDNode *N) {
27185 EVT VT = N->getValueType(ResNo: 0);
27186
27187 // A vector built entirely of undefs is undef.
27188 if (ISD::allOperandsUndef(N))
27189 return DAG.getUNDEF(VT);
27190
27191 // If this is a splat of a bitcast from another vector, change to a
27192 // concat_vector.
27193 // For example:
27194 // (build_vector (i64 (bitcast (v2i32 X))), (i64 (bitcast (v2i32 X)))) ->
27195 // (v2i64 (bitcast (concat_vectors (v2i32 X), (v2i32 X))))
27196 //
27197 // If X is a build_vector itself, the concat can become a larger build_vector.
27198 // TODO: Maybe this is useful for non-splat too?
27199 if (!LegalOperations) {
27200 SDValue Splat = cast<BuildVectorSDNode>(Val: N)->getSplatValue();
27201 // Only change build_vector to a concat_vector if the splat value type is
27202 // same as the vector element type.
27203 if (Splat && Splat.getValueType() == VT.getVectorElementType()) {
27204 Splat = peekThroughBitcasts(V: Splat);
27205 EVT SrcVT = Splat.getValueType();
27206 if (SrcVT.isVector()) {
27207 unsigned NumElts = N->getNumOperands() * SrcVT.getVectorNumElements();
27208 EVT NewVT = EVT::getVectorVT(Context&: *DAG.getContext(),
27209 VT: SrcVT.getVectorElementType(), NumElements: NumElts);
27210 if (!LegalTypes || TLI.isTypeLegal(VT: NewVT)) {
27211 SmallVector<SDValue, 8> Ops(N->getNumOperands(), Splat);
27212 SDValue Concat =
27213 DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT: NewVT, Ops);
27214 return DAG.getBitcast(VT, V: Concat);
27215 }
27216 }
27217 }
27218 }
27219
27220 // Check if we can express BUILD VECTOR via subvector extract.
27221 if (!LegalTypes && (N->getNumOperands() > 1)) {
27222 SDValue Op0 = N->getOperand(Num: 0);
27223 auto checkElem = [&](SDValue Op) -> uint64_t {
27224 if ((Op.getOpcode() == ISD::EXTRACT_VECTOR_ELT) &&
27225 (Op0.getOperand(i: 0) == Op.getOperand(i: 0)))
27226 if (auto CNode = dyn_cast<ConstantSDNode>(Val: Op.getOperand(i: 1)))
27227 return CNode->getZExtValue();
27228 return -1;
27229 };
27230
27231 int Offset = checkElem(Op0);
27232 for (unsigned i = 0; i < N->getNumOperands(); ++i) {
27233 if (Offset + i != checkElem(N->getOperand(Num: i))) {
27234 Offset = -1;
27235 break;
27236 }
27237 }
27238
27239 if ((Offset == 0) &&
27240 (Op0.getOperand(i: 0).getValueType() == N->getValueType(ResNo: 0)))
27241 return Op0.getOperand(i: 0);
27242 if ((Offset != -1) &&
27243 ((Offset % N->getValueType(ResNo: 0).getVectorNumElements()) ==
27244 0)) // IDX must be multiple of output size.
27245 return DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL: SDLoc(N), VT: N->getValueType(ResNo: 0),
27246 N1: Op0.getOperand(i: 0), N2: Op0.getOperand(i: 1));
27247 }
27248
27249 if (SDValue V = convertBuildVecExtToExt(N))
27250 return V;
27251
27252 if (SDValue V = convertBuildVecZextToBuildVecWithZeros(N))
27253 return V;
27254
27255 if (SDValue V = reduceBuildVecExtToExtBuildVec(N))
27256 return V;
27257
27258 if (SDValue V = reduceBuildVecTruncToBitCast(N))
27259 return V;
27260
27261 if (SDValue V = reduceBuildVecToShuffle(N))
27262 return V;
27263
27264 // A splat of a single element is a SPLAT_VECTOR if supported on the target.
27265 // Do this late as some of the above may replace the splat.
27266 if (TLI.getOperationAction(Op: ISD::SPLAT_VECTOR, VT) != TargetLowering::Expand)
27267 if (SDValue V = cast<BuildVectorSDNode>(Val: N)->getSplatValue()) {
27268 assert(!V.isUndef() && "Splat of undef should have been handled earlier");
27269 return DAG.getNode(Opcode: ISD::SPLAT_VECTOR, DL: SDLoc(N), VT, Operand: V);
27270 }
27271
27272 return SDValue();
27273}
27274
27275static SDValue combineConcatVectorOfScalars(SDNode *N, SelectionDAG &DAG) {
27276 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
27277 EVT OpVT = N->getOperand(Num: 0).getValueType();
27278
27279 // If the operands are legal vectors, leave them alone.
27280 if (TLI.isTypeLegal(VT: OpVT) || OpVT.isScalableVector())
27281 return SDValue();
27282
27283 SDLoc DL(N);
27284 EVT VT = N->getValueType(ResNo: 0);
27285 SmallVector<SDValue, 8> Ops;
27286 EVT SVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: OpVT.getSizeInBits());
27287
27288 // Keep track of what we encounter.
27289 EVT AnyFPVT;
27290
27291 for (const SDValue &Op : N->ops()) {
27292 if (ISD::BITCAST == Op.getOpcode() &&
27293 !Op.getOperand(i: 0).getValueType().isVector())
27294 Ops.push_back(Elt: Op.getOperand(i: 0));
27295 else if (Op.isUndef())
27296 Ops.push_back(Elt: DAG.getNode(Opcode: Op.getOpcode(), DL, VT: SVT));
27297 else
27298 return SDValue();
27299
27300 // Note whether we encounter an integer or floating point scalar.
27301 // If it's neither, bail out, it could be something weird like x86mmx.
27302 EVT LastOpVT = Ops.back().getValueType();
27303 if (LastOpVT.isFloatingPoint())
27304 AnyFPVT = LastOpVT;
27305 else if (!LastOpVT.isInteger())
27306 return SDValue();
27307 }
27308
27309 // If any of the operands is a floating point scalar bitcast to a vector,
27310 // use floating point types throughout, and bitcast everything.
27311 // Replace UNDEFs by another scalar UNDEF node, of the final desired type.
27312 if (AnyFPVT != EVT()) {
27313 SVT = AnyFPVT;
27314 for (SDValue &Op : Ops) {
27315 if (Op.getValueType() == SVT)
27316 continue;
27317 if (Op.isUndef())
27318 Op = DAG.getNode(Opcode: Op.getOpcode(), DL, VT: SVT);
27319 else
27320 Op = DAG.getBitcast(VT: SVT, V: Op);
27321 }
27322 }
27323
27324 EVT VecVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: SVT,
27325 NumElements: VT.getSizeInBits() / SVT.getSizeInBits());
27326 return DAG.getBitcast(VT, V: DAG.getBuildVector(VT: VecVT, DL, Ops));
27327}
27328
27329// Attempt to merge nested concat_vectors/undefs.
27330// Fold concat_vectors(concat_vectors(x,y,z,w),u,u,concat_vectors(a,b,c,d))
27331// --> concat_vectors(x,y,z,w,u,u,u,u,u,u,u,u,a,b,c,d)
27332static SDValue combineConcatVectorOfConcatVectors(SDNode *N,
27333 SelectionDAG &DAG) {
27334 EVT VT = N->getValueType(ResNo: 0);
27335
27336 // Ensure we're concatenating UNDEF and CONCAT_VECTORS nodes of similar types.
27337 EVT SubVT;
27338 SDValue FirstConcat;
27339 for (const SDValue &Op : N->ops()) {
27340 if (Op.isUndef())
27341 continue;
27342 if (Op.getOpcode() != ISD::CONCAT_VECTORS)
27343 return SDValue();
27344 if (!FirstConcat) {
27345 SubVT = Op.getOperand(i: 0).getValueType();
27346 if (!DAG.getTargetLoweringInfo().isTypeLegal(VT: SubVT))
27347 return SDValue();
27348 FirstConcat = Op;
27349 continue;
27350 }
27351 if (SubVT != Op.getOperand(i: 0).getValueType())
27352 return SDValue();
27353 }
27354 assert(FirstConcat && "Concat of all-undefs found");
27355
27356 SmallVector<SDValue> ConcatOps;
27357 for (const SDValue &Op : N->ops()) {
27358 if (Op.isUndef()) {
27359 ConcatOps.append(NumInputs: FirstConcat->getNumOperands(),
27360 Elt: DAG.getNode(Opcode: Op.getOpcode(), DL: SDLoc(), VT: SubVT));
27361 continue;
27362 }
27363 ConcatOps.append(in_start: Op->op_begin(), in_end: Op->op_end());
27364 }
27365 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT, Ops: ConcatOps);
27366}
27367
27368// Check to see if this is a CONCAT_VECTORS of a bunch of EXTRACT_SUBVECTOR
27369// operations. If so, and if the EXTRACT_SUBVECTOR vector inputs come from at
27370// most two distinct vectors the same size as the result, attempt to turn this
27371// into a legal shuffle.
27372static SDValue combineConcatVectorOfExtracts(SDNode *N, SelectionDAG &DAG) {
27373 EVT VT = N->getValueType(ResNo: 0);
27374 EVT OpVT = N->getOperand(Num: 0).getValueType();
27375
27376 // We currently can't generate an appropriate shuffle for a scalable vector.
27377 if (VT.isScalableVector())
27378 return SDValue();
27379
27380 int NumElts = VT.getVectorNumElements();
27381 int NumOpElts = OpVT.getVectorNumElements();
27382
27383 SDValue SV0 = DAG.getPOISON(VT), SV1 = DAG.getPOISON(VT);
27384 SmallVector<int, 8> Mask;
27385
27386 for (SDValue Op : N->ops()) {
27387 Op = peekThroughBitcasts(V: Op);
27388
27389 // UNDEF nodes convert to UNDEF shuffle mask values.
27390 if (Op.isUndef()) {
27391 Mask.append(NumInputs: (unsigned)NumOpElts, Elt: -1);
27392 continue;
27393 }
27394
27395 if (Op.getOpcode() != ISD::EXTRACT_SUBVECTOR)
27396 return SDValue();
27397
27398 // What vector are we extracting the subvector from and at what index?
27399 SDValue ExtVec = Op.getOperand(i: 0);
27400 int ExtIdx = Op.getConstantOperandVal(i: 1);
27401
27402 // We want the EVT of the original extraction to correctly scale the
27403 // extraction index.
27404 EVT ExtVT = ExtVec.getValueType();
27405 ExtVec = peekThroughBitcasts(V: ExtVec);
27406
27407 // UNDEF nodes convert to UNDEF shuffle mask values.
27408 if (ExtVec.isUndef()) {
27409 Mask.append(NumInputs: (unsigned)NumOpElts, Elt: -1);
27410 continue;
27411 }
27412
27413 // Ensure that we are extracting a subvector from a vector the same
27414 // size as the result.
27415 if (ExtVT.getSizeInBits() != VT.getSizeInBits())
27416 return SDValue();
27417
27418 // Scale the subvector index to account for any bitcast.
27419 int NumExtElts = ExtVT.getVectorNumElements();
27420 if (0 == (NumExtElts % NumElts))
27421 ExtIdx /= (NumExtElts / NumElts);
27422 else if (0 == (NumElts % NumExtElts))
27423 ExtIdx *= (NumElts / NumExtElts);
27424 else
27425 return SDValue();
27426
27427 // At most we can reference 2 inputs in the final shuffle.
27428 if (SV0.isUndef() || SV0 == ExtVec) {
27429 SV0 = ExtVec;
27430 for (int i = 0; i != NumOpElts; ++i)
27431 Mask.push_back(Elt: i + ExtIdx);
27432 } else if (SV1.isUndef() || SV1 == ExtVec) {
27433 SV1 = ExtVec;
27434 for (int i = 0; i != NumOpElts; ++i)
27435 Mask.push_back(Elt: i + ExtIdx + NumElts);
27436 } else {
27437 return SDValue();
27438 }
27439 }
27440
27441 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
27442 return TLI.buildLegalVectorShuffle(VT, DL: SDLoc(N), N0: DAG.getBitcast(VT, V: SV0),
27443 N1: DAG.getBitcast(VT, V: SV1), Mask, DAG);
27444}
27445
27446static SDValue combineConcatVectorOfCasts(SDNode *N, SelectionDAG &DAG) {
27447 unsigned CastOpcode = N->getOperand(Num: 0).getOpcode();
27448 switch (CastOpcode) {
27449 case ISD::SINT_TO_FP:
27450 case ISD::UINT_TO_FP:
27451 case ISD::FP_TO_SINT:
27452 case ISD::FP_TO_UINT:
27453 // TODO: Allow more opcodes?
27454 // case ISD::BITCAST:
27455 // case ISD::TRUNCATE:
27456 // case ISD::ZERO_EXTEND:
27457 // case ISD::SIGN_EXTEND:
27458 // case ISD::FP_EXTEND:
27459 break;
27460 default:
27461 return SDValue();
27462 }
27463
27464 EVT SrcVT = N->getOperand(Num: 0).getOperand(i: 0).getValueType();
27465 if (!SrcVT.isVector())
27466 return SDValue();
27467
27468 // All operands of the concat must be the same kind of cast from the same
27469 // source type.
27470 SmallVector<SDValue, 4> SrcOps;
27471 for (SDValue Op : N->ops()) {
27472 if (Op.getOpcode() != CastOpcode || !Op.hasOneUse() ||
27473 Op.getOperand(i: 0).getValueType() != SrcVT)
27474 return SDValue();
27475 SrcOps.push_back(Elt: Op.getOperand(i: 0));
27476 }
27477
27478 // The wider cast must be supported by the target. This is unusual because
27479 // the operation support type parameter depends on the opcode. In addition,
27480 // check the other type in the cast to make sure this is really legal.
27481 EVT VT = N->getValueType(ResNo: 0);
27482 ElementCount NumElts = SrcVT.getVectorElementCount() * N->getNumOperands();
27483 EVT ConcatSrcVT = SrcVT.changeVectorElementCount(Context&: *DAG.getContext(), EC: NumElts);
27484 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
27485 switch (CastOpcode) {
27486 case ISD::SINT_TO_FP:
27487 case ISD::UINT_TO_FP:
27488 if (!TLI.isOperationLegalOrCustom(Op: CastOpcode, VT: ConcatSrcVT) ||
27489 !TLI.isTypeLegal(VT))
27490 return SDValue();
27491 break;
27492 case ISD::FP_TO_SINT:
27493 case ISD::FP_TO_UINT:
27494 if (!TLI.isOperationLegalOrCustom(Op: CastOpcode, VT) ||
27495 !TLI.isTypeLegal(VT: ConcatSrcVT))
27496 return SDValue();
27497 break;
27498 default:
27499 llvm_unreachable("Unexpected cast opcode");
27500 }
27501
27502 // concat (cast X), (cast Y)... -> cast (concat X, Y...)
27503 SDLoc DL(N);
27504 SDValue NewConcat = DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT: ConcatSrcVT, Ops: SrcOps);
27505 return DAG.getNode(Opcode: CastOpcode, DL, VT, Operand: NewConcat);
27506}
27507
27508// See if this is a simple CONCAT_VECTORS with no UNDEF operands, and if one of
27509// the operands is a SHUFFLE_VECTOR, and all other operands are also operands
27510// to that SHUFFLE_VECTOR, create wider SHUFFLE_VECTOR.
27511static SDValue combineConcatVectorOfShuffleAndItsOperands(
27512 SDNode *N, SelectionDAG &DAG, const TargetLowering &TLI, bool LegalTypes,
27513 bool LegalOperations) {
27514 EVT VT = N->getValueType(ResNo: 0);
27515 EVT OpVT = N->getOperand(Num: 0).getValueType();
27516 if (VT.isScalableVector())
27517 return SDValue();
27518
27519 // For now, only allow simple 2-operand concatenations.
27520 if (N->getNumOperands() != 2)
27521 return SDValue();
27522
27523 // Don't create illegal types/shuffles when not allowed to.
27524 if ((LegalTypes && !TLI.isTypeLegal(VT)) ||
27525 (LegalOperations &&
27526 !TLI.isOperationLegalOrCustom(Op: ISD::VECTOR_SHUFFLE, VT)))
27527 return SDValue();
27528
27529 // Analyze all of the operands of the CONCAT_VECTORS. Out of all of them,
27530 // we want to find one that is: (1) a SHUFFLE_VECTOR (2) only used by us,
27531 // and (3) all operands of CONCAT_VECTORS must be either that SHUFFLE_VECTOR,
27532 // or one of the operands of that SHUFFLE_VECTOR (but not UNDEF!).
27533 // (4) and for now, the SHUFFLE_VECTOR must be unary.
27534 ShuffleVectorSDNode *SVN = nullptr;
27535 for (SDValue Op : N->ops()) {
27536 if (auto *CurSVN = dyn_cast<ShuffleVectorSDNode>(Val&: Op);
27537 CurSVN && CurSVN->getOperand(Num: 1).isUndef() && N->isOnlyUserOf(N: CurSVN) &&
27538 all_of(Range: N->ops(), P: [CurSVN](SDValue Op) {
27539 // FIXME: can we allow UNDEF operands?
27540 return !Op.isUndef() &&
27541 (Op.getNode() == CurSVN || is_contained(Range: CurSVN->ops(), Element: Op));
27542 })) {
27543 SVN = CurSVN;
27544 break;
27545 }
27546 }
27547 if (!SVN)
27548 return SDValue();
27549
27550 // We are going to pad the shuffle operands, so any indice, that was picking
27551 // from the second operand, must be adjusted.
27552 SmallVector<int, 16> AdjustedMask(SVN->getMask());
27553 assert(SVN->getOperand(1).isUndef() && "Expected unary shuffle!");
27554
27555 // Identity masks for the operands of the (padded) shuffle.
27556 SmallVector<int, 32> IdentityMask(2 * OpVT.getVectorNumElements());
27557 MutableArrayRef<int> FirstShufOpIdentityMask =
27558 MutableArrayRef<int>(IdentityMask)
27559 .take_front(N: OpVT.getVectorNumElements());
27560 MutableArrayRef<int> SecondShufOpIdentityMask =
27561 MutableArrayRef<int>(IdentityMask).take_back(N: OpVT.getVectorNumElements());
27562 std::iota(first: FirstShufOpIdentityMask.begin(), last: FirstShufOpIdentityMask.end(), value: 0);
27563 std::iota(first: SecondShufOpIdentityMask.begin(), last: SecondShufOpIdentityMask.end(),
27564 value: VT.getVectorNumElements());
27565
27566 // New combined shuffle mask.
27567 SmallVector<int, 32> Mask;
27568 Mask.reserve(N: VT.getVectorNumElements());
27569 for (SDValue Op : N->ops()) {
27570 assert(!Op.isUndef() && "Not expecting to concatenate UNDEF.");
27571 if (Op.getNode() == SVN) {
27572 append_range(C&: Mask, R&: AdjustedMask);
27573 continue;
27574 }
27575 if (Op == SVN->getOperand(Num: 0)) {
27576 append_range(C&: Mask, R&: FirstShufOpIdentityMask);
27577 continue;
27578 }
27579 if (Op == SVN->getOperand(Num: 1)) {
27580 append_range(C&: Mask, R&: SecondShufOpIdentityMask);
27581 continue;
27582 }
27583 llvm_unreachable("Unexpected operand!");
27584 }
27585
27586 // Don't create illegal shuffle masks.
27587 if (!TLI.isShuffleMaskLegal(Mask, VT))
27588 return SDValue();
27589
27590 // Pad the shuffle operands with poison.
27591 SDLoc dl(N);
27592 std::array<SDValue, 2> ShufOps;
27593 for (auto I : zip(t: SVN->ops(), u&: ShufOps)) {
27594 SDValue ShufOp = std::get<0>(t&: I);
27595 SDValue &NewShufOp = std::get<1>(t&: I);
27596 if (ShufOp.isUndef())
27597 NewShufOp = DAG.getNode(Opcode: ShufOp.getOpcode(), DL: SDLoc(), VT);
27598 else {
27599 SmallVector<SDValue, 2> ShufOpParts(N->getNumOperands(),
27600 DAG.getPOISON(VT: OpVT));
27601 ShufOpParts[0] = ShufOp;
27602 NewShufOp = DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: dl, VT, Ops: ShufOpParts);
27603 }
27604 }
27605 // Finally, create the new wide shuffle.
27606 return DAG.getVectorShuffle(VT, dl, N1: ShufOps[0], N2: ShufOps[1], Mask);
27607}
27608
27609// concat(shuffle(loadA, loadB, mask0), shuffle(loadA, loadB, mask1))
27610// -> shuffle(loadAB, poison, concat(mask0, mask1))
27611// only if loadA and loadB can be proven consecutive.
27612static SDValue combineConcatVectorOfShuffles(SDNode *N, SelectionDAG &DAG,
27613 const TargetLowering &TLI,
27614 bool LegalOperations) {
27615 SDValue A, B;
27616 ArrayRef<int> M0, M1;
27617 if (!sd_match(N, P: m_Node<ISD::CONCAT_VECTORS>(
27618 Preds: m_OneUse(P: m_Shuffle(v1: m_NUses<2>(P: m_Value(N&: A)),
27619 v2: m_NUses<2>(P: m_Value(N&: B)), mask: m_Mask(M0))),
27620 Preds: m_OneUse(P: m_Shuffle(v1: m_Deferred(V&: A), v2: m_Deferred(V&: B),
27621 mask: m_Mask(M1))))))
27622 return SDValue();
27623 auto *LoadA = dyn_cast<LoadSDNode>(Val: A.getNode());
27624 auto *LoadB = dyn_cast<LoadSDNode>(Val: B.getNode());
27625 if (!LoadA || !LoadB || !ISD::isNON_EXTLoad(N: LoadA) ||
27626 !ISD::isNON_EXTLoad(N: LoadB))
27627 return SDValue();
27628
27629 // Check if the address spaces of both loads are the same.
27630 if (LoadA->getAddressSpace() != LoadB->getAddressSpace())
27631 return SDValue();
27632
27633 // Check if the loads are consecutive.
27634 LoadSDNode *Base = nullptr;
27635 if (DAG.areNonVolatileConsecutiveLoads(
27636 LD: LoadB, Base: LoadA, Bytes: LoadB->getMemoryVT().getStoreSize(), /*Dist=*/1)) {
27637 Base = LoadA;
27638 } else if (DAG.areNonVolatileConsecutiveLoads(
27639 LD: LoadA, Base: LoadB, Bytes: LoadA->getMemoryVT().getStoreSize(),
27640 /*Dist=*/1)) {
27641 Base = LoadB;
27642 } else {
27643 return SDValue(); // not adjacent
27644 }
27645
27646 unsigned Fast = 0;
27647 Align NewAlign = Base->getAlign();
27648 EVT WideVT =
27649 LoadA->getMemoryVT().getDoubleNumVectorElementsVT(Context&: *DAG.getContext());
27650 if (!TLI.allowsMemoryAccess(Context&: *DAG.getContext(), DL: DAG.getDataLayout(), VT: WideVT,
27651 AddrSpace: Base->getAddressSpace(), Alignment: NewAlign,
27652 Flags: Base->getMemOperand()->getFlags(), Fast: &Fast) ||
27653 !Fast)
27654 return SDValue();
27655
27656 // Create a shuffle of the wide load.
27657 SmallVector<int, 32> Mask;
27658 if (Base == LoadA) {
27659 llvm::append_range(C&: Mask, R&: M0);
27660 llvm::append_range(C&: Mask, R&: M1);
27661 } else {
27662 SmallVector<int, 16> C0(M0), C1(M1);
27663 ShuffleVectorSDNode::commuteMask(Mask: C0);
27664 ShuffleVectorSDNode::commuteMask(Mask: C1);
27665 llvm::append_range(C&: Mask, R&: C0);
27666 llvm::append_range(C&: Mask, R&: C1);
27667 }
27668
27669 // Check if the wide load, new shuffle and it's mask is legal.
27670 if (LegalOperations &&
27671 (!TLI.isOperationLegal(Op: ISD::LOAD, VT: WideVT) ||
27672 !TLI.isOperationLegalOrCustom(Op: ISD::VECTOR_SHUFFLE, VT: WideVT) ||
27673 !TLI.isShuffleMaskLegal(Mask, WideVT)))
27674 return SDValue();
27675
27676 // Create a wide load of twice the size of the original load.
27677 MachineFunction &MF = DAG.getMachineFunction();
27678 MachineMemOperand *WideMMO = MF.getMachineMemOperand(
27679 MMO: Base->getMemOperand(), /*Offset=*/0, Size: WideVT.getStoreSize());
27680 SDValue WideLoad = DAG.getLoad(VT: WideVT, dl: SDLoc(N), Chain: Base->getChain(),
27681 Ptr: Base->getBasePtr(), MMO: WideMMO);
27682 // Redirect old chain users to the new chain.
27683 DAG.makeEquivalentMemoryOrdering(OldLoad: LoadA, NewMemOp: WideLoad);
27684 DAG.makeEquivalentMemoryOrdering(OldLoad: LoadB, NewMemOp: WideLoad);
27685
27686 // Create a new shuffle with the new mask.
27687 return DAG.getVectorShuffle(VT: WideVT, dl: SDLoc(N), N1: WideLoad, N2: DAG.getPOISON(VT: WideVT),
27688 Mask);
27689}
27690
27691static SDValue combineConcatVectorOfSplats(SDNode *N, SelectionDAG &DAG,
27692 const TargetLowering &TLI,
27693 bool LegalTypes,
27694 bool LegalOperations) {
27695 EVT VT = N->getValueType(ResNo: 0);
27696
27697 // Post-legalization we can only create wider SPLAT_VECTOR operations if both
27698 // the type and operation is legal. The Hexagon target has custom
27699 // legalization for SPLAT_VECTOR that splits the operation into two parts and
27700 // concatenates them. Therefore, custom lowering must also be rejected in
27701 // order to avoid an infinite loop.
27702 if ((LegalTypes && !TLI.isTypeLegal(VT)) ||
27703 (LegalOperations && !TLI.isOperationLegal(Op: ISD::SPLAT_VECTOR, VT)))
27704 return SDValue();
27705
27706 SDValue Op0 = N->getOperand(Num: 0);
27707 if (!llvm::all_equal(Range: N->op_values()) || Op0.getOpcode() != ISD::SPLAT_VECTOR)
27708 return SDValue();
27709
27710 return DAG.getNode(Opcode: ISD::SPLAT_VECTOR, DL: SDLoc(N), VT, Operand: Op0.getOperand(i: 0));
27711}
27712
27713SDValue DAGCombiner::visitCONCAT_VECTORS(SDNode *N) {
27714 // If we only have one input vector, we don't need to do any concatenation.
27715 if (N->getNumOperands() == 1)
27716 return N->getOperand(Num: 0);
27717
27718 // Check if all of the operands are undefs.
27719 EVT VT = N->getValueType(ResNo: 0);
27720 if (ISD::allOperandsUndef(N))
27721 return DAG.getUNDEF(VT);
27722
27723 // Optimize concat_vectors where all but the first of the vectors are undef.
27724 if (all_of(Range: drop_begin(RangeOrContainer: N->ops()),
27725 P: [](const SDValue &Op) { return Op.isUndef(); })) {
27726 SDValue In = N->getOperand(Num: 0);
27727 assert(In.getValueType().isVector() && "Must concat vectors");
27728
27729 // If the input is a concat_vectors, just make a larger concat by padding
27730 // with smaller undefs.
27731 //
27732 // Legalizing in AArch64TargetLowering::LowerCONCAT_VECTORS() and combining
27733 // here could cause an infinite loop. That legalizing happens when LegalDAG
27734 // is true and input of AArch64TargetLowering::LowerCONCAT_VECTORS() is
27735 // scalable.
27736 if (In.getOpcode() == ISD::CONCAT_VECTORS && In.hasOneUse() &&
27737 !(LegalDAG && In.getValueType().isScalableVector())) {
27738 unsigned NumOps = N->getNumOperands() * In.getNumOperands();
27739 SmallVector<SDValue, 4> Ops(In->ops());
27740 Ops.resize(N: NumOps, NV: DAG.getPOISON(VT: Ops[0].getValueType()));
27741 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT, Ops);
27742 }
27743
27744 SDValue Scalar = peekThroughOneUseBitcasts(V: In);
27745
27746 // concat_vectors(scalar_to_vector(scalar), undef) ->
27747 // scalar_to_vector(scalar)
27748 if (!LegalOperations && Scalar.getOpcode() == ISD::SCALAR_TO_VECTOR &&
27749 Scalar.hasOneUse()) {
27750 EVT SVT = Scalar.getValueType().getVectorElementType();
27751 if (SVT == Scalar.getOperand(i: 0).getValueType())
27752 Scalar = Scalar.getOperand(i: 0);
27753 }
27754
27755 // concat_vectors(scalar, undef) -> scalar_to_vector(scalar)
27756 if (!Scalar.getValueType().isVector() && In.hasOneUse()) {
27757 // If the bitcast type isn't legal, it might be a trunc of a legal type;
27758 // look through the trunc so we can still do the transform:
27759 // concat_vectors(trunc(scalar), undef) -> scalar_to_vector(scalar)
27760 // However, this is only equivalent on little-endian targets.
27761 if (Scalar->getOpcode() == ISD::TRUNCATE &&
27762 !TLI.isTypeLegal(VT: Scalar.getValueType()) &&
27763 TLI.isTypeLegal(VT: Scalar->getOperand(Num: 0).getValueType()) &&
27764 DAG.getDataLayout().isLittleEndian())
27765 Scalar = Scalar->getOperand(Num: 0);
27766
27767 EVT SclTy = Scalar.getValueType();
27768
27769 if (!SclTy.isFloatingPoint() && !SclTy.isInteger())
27770 return SDValue();
27771
27772 // Bail out if the vector size is not a multiple of the scalar size.
27773 if (VT.getSizeInBits() % SclTy.getSizeInBits())
27774 return SDValue();
27775
27776 unsigned VNTNumElms = VT.getSizeInBits() / SclTy.getSizeInBits();
27777 if (VNTNumElms < 2)
27778 return SDValue();
27779
27780 EVT NVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: SclTy, NumElements: VNTNumElms);
27781 if (!TLI.isTypeLegal(VT: NVT) || !TLI.isTypeLegal(VT: Scalar.getValueType()))
27782 return SDValue();
27783
27784 SDValue Res = DAG.getNode(Opcode: ISD::SCALAR_TO_VECTOR, DL: SDLoc(N), VT: NVT, Operand: Scalar);
27785 return DAG.getBitcast(VT, V: Res);
27786 }
27787 }
27788
27789 // Fold any combination of BUILD_VECTOR or UNDEF nodes into one BUILD_VECTOR.
27790 // We have already tested above for an UNDEF only concatenation.
27791 // fold (concat_vectors (BUILD_VECTOR A, B, ...), (BUILD_VECTOR C, D, ...))
27792 // -> (BUILD_VECTOR A, B, ..., C, D, ...)
27793 auto IsBuildVectorOrUndef = [](const SDValue &Op) {
27794 return Op.isUndef() || ISD::BUILD_VECTOR == Op.getOpcode();
27795 };
27796 if (llvm::all_of(Range: N->ops(), P: IsBuildVectorOrUndef)) {
27797 SmallVector<SDValue, 8> Opnds;
27798 EVT SVT = VT.getScalarType();
27799
27800 EVT MinVT = SVT;
27801 if (!SVT.isFloatingPoint()) {
27802 // If BUILD_VECTOR are from built from integer, they may have different
27803 // operand types. Get the smallest type and truncate all operands to it.
27804 bool FoundMinVT = false;
27805 for (const SDValue &Op : N->ops())
27806 if (ISD::BUILD_VECTOR == Op.getOpcode()) {
27807 EVT OpSVT = Op.getOperand(i: 0).getValueType();
27808 MinVT = (!FoundMinVT || OpSVT.bitsLE(VT: MinVT)) ? OpSVT : MinVT;
27809 FoundMinVT = true;
27810 }
27811 assert(FoundMinVT && "Concat vector type mismatch");
27812 }
27813
27814 for (const SDValue &Op : N->ops()) {
27815 EVT OpVT = Op.getValueType();
27816 unsigned NumElts = OpVT.getVectorNumElements();
27817
27818 if (Op.isUndef())
27819 Opnds.append(NumInputs: NumElts, Elt: DAG.getNode(Opcode: Op.getOpcode(), DL: SDLoc(), VT: MinVT));
27820
27821 if (ISD::BUILD_VECTOR == Op.getOpcode()) {
27822 if (SVT.isFloatingPoint()) {
27823 assert(SVT == OpVT.getScalarType() && "Concat vector type mismatch");
27824 Opnds.append(in_start: Op->op_begin(), in_end: Op->op_begin() + NumElts);
27825 } else {
27826 for (unsigned i = 0; i != NumElts; ++i)
27827 Opnds.push_back(
27828 Elt: DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(N), VT: MinVT, Operand: Op.getOperand(i)));
27829 }
27830 }
27831 }
27832
27833 assert(VT.getVectorNumElements() == Opnds.size() &&
27834 "Concat vector type mismatch");
27835 return DAG.getBuildVector(VT, DL: SDLoc(N), Ops: Opnds);
27836 }
27837
27838 if (SDValue V =
27839 combineConcatVectorOfSplats(N, DAG, TLI, LegalTypes, LegalOperations))
27840 return V;
27841
27842 // Fold CONCAT_VECTORS of only bitcast scalars (or undef) to BUILD_VECTOR.
27843 // FIXME: Add support for concat_vectors(bitcast(vec0),bitcast(vec1),...).
27844 if (SDValue V = combineConcatVectorOfScalars(N, DAG))
27845 return V;
27846
27847 if (Level <= AfterLegalizeVectorOps && TLI.isTypeLegal(VT)) {
27848 // Fold CONCAT_VECTORS of CONCAT_VECTORS (or undef) to VECTOR_SHUFFLE.
27849 if (SDValue V = combineConcatVectorOfConcatVectors(N, DAG))
27850 return V;
27851
27852 // Fold CONCAT_VECTORS of EXTRACT_SUBVECTOR (or undef) to VECTOR_SHUFFLE.
27853 if (SDValue V = combineConcatVectorOfExtracts(N, DAG))
27854 return V;
27855 }
27856
27857 if (SDValue V = combineConcatVectorOfCasts(N, DAG))
27858 return V;
27859
27860 if (SDValue V = combineConcatVectorOfShuffleAndItsOperands(
27861 N, DAG, TLI, LegalTypes, LegalOperations))
27862 return V;
27863
27864 if (SDValue V = combineConcatVectorOfShuffles(N, DAG, TLI, LegalOperations))
27865 return V;
27866
27867 // Type legalization of vectors and DAG canonicalization of SHUFFLE_VECTOR
27868 // nodes often generate nop CONCAT_VECTOR nodes. Scan the CONCAT_VECTOR
27869 // operands and look for a CONCAT operations that place the incoming vectors
27870 // at the exact same location.
27871 //
27872 // For scalable vectors, EXTRACT_SUBVECTOR indexes are implicitly scaled.
27873 SDValue SingleSource = SDValue();
27874 unsigned PartNumElem =
27875 N->getOperand(Num: 0).getValueType().getVectorMinNumElements();
27876
27877 for (unsigned i = 0, e = N->getNumOperands(); i != e; ++i) {
27878 SDValue Op = N->getOperand(Num: i);
27879
27880 if (Op.isUndef())
27881 continue;
27882
27883 // Check if this is the identity extract:
27884 if (Op.getOpcode() != ISD::EXTRACT_SUBVECTOR)
27885 return SDValue();
27886
27887 // Find the single incoming vector for the extract_subvector.
27888 if (SingleSource.getNode()) {
27889 if (Op.getOperand(i: 0) != SingleSource)
27890 return SDValue();
27891 } else {
27892 SingleSource = Op.getOperand(i: 0);
27893
27894 // Check the source type is the same as the type of the result.
27895 // If not, this concat may extend the vector, so we can not
27896 // optimize it away.
27897 if (SingleSource.getValueType() != N->getValueType(ResNo: 0))
27898 return SDValue();
27899 }
27900
27901 // Check that we are reading from the identity index.
27902 unsigned IdentityIndex = i * PartNumElem;
27903 if (Op.getConstantOperandAPInt(i: 1) != IdentityIndex)
27904 return SDValue();
27905 }
27906
27907 if (SingleSource.getNode())
27908 return SingleSource;
27909
27910 return SDValue();
27911}
27912
27913SDValue DAGCombiner::visitVECTOR_INTERLEAVE(SDNode *N) {
27914 EVT VT = N->getValueType(ResNo: 0);
27915 SDValue Op0 = N->getOperand(Num: 0);
27916 unsigned Factor = N->getNumOperands();
27917
27918 // Canonicalize shuffle undef, undef -> undef
27919 if (all_of(Range: N->op_values(), P: [](SDValue Op) { return Op.isUndef(); })) {
27920 SDLoc DL(N);
27921 SmallVector<SDValue> Ops(N->getNumValues(), DAG.getUNDEF(VT));
27922 return DAG.getMergeValues(Ops, dl: DL);
27923 }
27924
27925 // Fold an interleave of fixed-length BUILD_VECTORs by rearranging their
27926 // scalar operands directly.
27927 if (Op0.getOpcode() == ISD::BUILD_VECTOR) {
27928 EVT EltVT = Op0.getOperand(i: 0).getValueType();
27929 if (llvm::all_of(Range: N->op_values(), P: [&](SDValue Op) {
27930 return Op.getOpcode() == ISD::BUILD_VECTOR &&
27931 Op.getOperand(i: 0).getValueType() == EltVT;
27932 })) {
27933 unsigned NumElts = VT.getVectorNumElements();
27934 SDLoc DL(N);
27935 SmallVector<SDValue, 4> Results;
27936 SmallVector<SDValue, 16> InterleavedElts;
27937 for (unsigned I = 0; I != NumElts; ++I) {
27938 for (SDValue op : N->op_values())
27939 InterleavedElts.push_back(Elt: op.getOperand(i: I));
27940 }
27941 for (unsigned I = 0; I < Factor; I++)
27942 Results.push_back(Elt: DAG.getBuildVector(
27943 VT, DL, Ops: ArrayRef(InterleavedElts).slice(N: I * NumElts, M: NumElts)));
27944 return CombineTo(N, To: &Results);
27945 }
27946 }
27947
27948 // Fold interleave(splat(S[J]), ..., splat(S[J + Factor - 1])) to shuffles
27949 // of S.
27950 if (Op0.getOpcode() == ISD::VECTOR_SHUFFLE &&
27951 VT.getVectorElementCount().isKnownMultipleOf(RHS: Factor)) {
27952 int FirstIndex;
27953 unsigned NumElts = VT.getVectorNumElements();
27954 SDValue Source = DAG.getSplatSourceVector(V: Op0, SplatIndex&: FirstIndex);
27955 if (Source && llvm::all_of(Range: llvm::enumerate(First: N->op_values()), P: [&](auto Item) {
27956 int SplatIndex;
27957 return DAG.getSplatSourceVector(V: Item.value(), SplatIndex) == Source &&
27958 SplatIndex == FirstIndex + static_cast<int>(Item.index());
27959 })) {
27960 SmallVector<int, 16> Mask;
27961 for (unsigned I = 0; I != NumElts; ++I)
27962 Mask.push_back(Elt: FirstIndex + I % Factor);
27963 SDValue Shuffle =
27964 DAG.getVectorShuffle(VT, dl: SDLoc(N), N1: Source, N2: DAG.getPOISON(VT), Mask);
27965 SmallVector<SDValue, 4> Results(Factor, Shuffle);
27966 return CombineTo(N, To: &Results);
27967 }
27968 }
27969
27970 // Check to see if all operands are identical.
27971 if (!llvm::all_equal(Range: N->op_values()))
27972 return SDValue();
27973
27974 // Check to see if the identical operand is a splat.
27975 if (!DAG.isSplatValue(V: N->getOperand(Num: 0)))
27976 return SDValue();
27977
27978 // interleave splat(X), splat(X).... --> splat(X), splat(X)....
27979 SmallVector<SDValue, 4> Ops;
27980 Ops.append(in_start: N->op_values().begin(), in_end: N->op_values().end());
27981 return CombineTo(N, To: &Ops);
27982}
27983
27984SDValue DAGCombiner::visitVECTOR_DEINTERLEAVE(SDNode *N) {
27985 EVT VT = N->getValueType(ResNo: 0);
27986 SDValue Op0 = N->getOperand(Num: 0);
27987
27988 // Canonicalize shuffle undef -> {undef, undef, ..}
27989 if (Op0.isUndef()) {
27990 SDLoc DL(N);
27991 SmallVector<SDValue> Ops(N->getNumValues(), DAG.getUNDEF(VT));
27992 return DAG.getMergeValues(Ops, dl: DL);
27993 }
27994
27995 return SDValue();
27996}
27997
27998// Helper that peeks through INSERT_SUBVECTOR/CONCAT_VECTORS to find
27999// if the subvector can be sourced for free.
28000static SDValue getSubVectorSrc(SDValue V, unsigned Index, EVT SubVT) {
28001 if (V.getOpcode() == ISD::INSERT_SUBVECTOR &&
28002 V.getOperand(i: 1).getValueType() == SubVT &&
28003 V.getConstantOperandAPInt(i: 2) == Index) {
28004 return V.getOperand(i: 1);
28005 }
28006 if (V.getOpcode() == ISD::CONCAT_VECTORS &&
28007 V.getOperand(i: 0).getValueType() == SubVT &&
28008 (Index % SubVT.getVectorMinNumElements()) == 0) {
28009 uint64_t SubIdx = Index / SubVT.getVectorMinNumElements();
28010 return V.getOperand(i: SubIdx);
28011 }
28012 return SDValue();
28013}
28014
28015static SDValue narrowInsertExtractVectorBinOp(EVT SubVT, SDValue BinOp,
28016 unsigned Index, const SDLoc &DL,
28017 SelectionDAG &DAG,
28018 bool LegalOperations) {
28019 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
28020 unsigned BinOpcode = BinOp.getOpcode();
28021 if (!TLI.isBinOp(Opcode: BinOpcode) || BinOp->getNumValues() != 1)
28022 return SDValue();
28023
28024 EVT VecVT = BinOp.getValueType();
28025 SDValue Bop0 = BinOp.getOperand(i: 0), Bop1 = BinOp.getOperand(i: 1);
28026 if (VecVT != Bop0.getValueType() || VecVT != Bop1.getValueType())
28027 return SDValue();
28028 if (!TLI.isOperationLegalOrCustom(Op: BinOpcode, VT: SubVT, LegalOnly: LegalOperations))
28029 return SDValue();
28030
28031 SDValue Sub0 = getSubVectorSrc(V: Bop0, Index, SubVT);
28032 SDValue Sub1 = getSubVectorSrc(V: Bop1, Index, SubVT);
28033
28034 // TODO: We could handle the case where only 1 operand is being inserted by
28035 // creating an extract of the other operand, but that requires checking
28036 // number of uses and/or costs.
28037 if (!Sub0 || !Sub1)
28038 return SDValue();
28039
28040 // We are inserting both operands of the wide binop only to extract back
28041 // to the narrow vector size. Eliminate all of the insert/extract:
28042 // ext (binop (ins ?, X, Index), (ins ?, Y, Index)), Index --> binop X, Y
28043 return DAG.getNode(Opcode: BinOpcode, DL, VT: SubVT, N1: Sub0, N2: Sub1, Flags: BinOp->getFlags());
28044}
28045
28046/// If we are extracting a subvector produced by a wide binary operator try
28047/// to use a narrow binary operator and/or avoid concatenation and extraction.
28048static SDValue narrowExtractedVectorBinOp(EVT VT, SDValue Src, unsigned Index,
28049 const SDLoc &DL, SelectionDAG &DAG,
28050 bool LegalOperations) {
28051 // TODO: Refactor with the caller (visitEXTRACT_SUBVECTOR), so we can share
28052 // some of these bailouts with other transforms.
28053
28054 if (SDValue V = narrowInsertExtractVectorBinOp(SubVT: VT, BinOp: Src, Index, DL, DAG,
28055 LegalOperations))
28056 return V;
28057
28058 // We are looking for an optionally bitcasted wide vector binary operator
28059 // feeding an extract subvector.
28060 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
28061 SDValue BinOp = peekThroughBitcasts(V: Src);
28062 unsigned BOpcode = BinOp.getOpcode();
28063 if (!TLI.isBinOp(Opcode: BOpcode) || BinOp->getNumValues() != 1)
28064 return SDValue();
28065
28066 // Exclude the fake form of fneg (fsub -0.0, x) because that is likely to be
28067 // reduced to the unary fneg when it is visited, and we probably want to deal
28068 // with fneg in a target-specific way.
28069 if (BOpcode == ISD::FSUB) {
28070 auto *C = isConstOrConstSplatFP(N: BinOp.getOperand(i: 0), /*AllowUndefs*/ true);
28071 if (C && C->getValueAPF().isNegZero())
28072 return SDValue();
28073 }
28074
28075 // The binop must be a vector type, so we can extract some fraction of it.
28076 EVT WideBVT = BinOp.getValueType();
28077 // The optimisations below currently assume we are dealing with fixed length
28078 // vectors. It is possible to add support for scalable vectors, but at the
28079 // moment we've done no analysis to prove whether they are profitable or not.
28080 if (!WideBVT.isFixedLengthVector())
28081 return SDValue();
28082
28083 assert((Index % VT.getVectorNumElements()) == 0 &&
28084 "Extract index is not a multiple of the vector length.");
28085
28086 // Bail out if this is not a proper multiple width extraction.
28087 unsigned WideWidth = WideBVT.getSizeInBits();
28088 unsigned NarrowWidth = VT.getSizeInBits();
28089 if (WideWidth % NarrowWidth != 0)
28090 return SDValue();
28091
28092 // Bail out if we are extracting a fraction of a single operation. This can
28093 // occur because we potentially looked through a bitcast of the binop.
28094 unsigned NarrowingRatio = WideWidth / NarrowWidth;
28095 unsigned WideNumElts = WideBVT.getVectorNumElements();
28096 if (WideNumElts % NarrowingRatio != 0)
28097 return SDValue();
28098
28099 // Bail out if the target does not support a narrower version of the binop.
28100 EVT NarrowBVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: WideBVT.getScalarType(),
28101 NumElements: WideNumElts / NarrowingRatio);
28102 if (!TLI.isOperationLegalOrCustomOrPromote(Op: BOpcode, VT: NarrowBVT,
28103 LegalOnly: LegalOperations))
28104 return SDValue();
28105
28106 // If extraction is cheap, we don't need to look at the binop operands
28107 // for concat ops. The narrow binop alone makes this transform profitable.
28108 // We can't just reuse the original extract index operand because we may have
28109 // bitcasted.
28110 unsigned ConcatOpNum = Index / VT.getVectorNumElements();
28111 unsigned ExtBOIdx = ConcatOpNum * NarrowBVT.getVectorNumElements();
28112 if (TLI.getExtractSubvectorCost(ResVT: NarrowBVT, SrcVT: WideBVT, Index: ExtBOIdx) <=
28113 TargetLowering::ExtractSubvectorCost::Cheap &&
28114 BinOp.hasOneUse() && Src->hasOneUse()) {
28115 // extract (binop B0, B1), N --> binop (extract B0, N), (extract B1, N)
28116 SDValue NewExtIndex = DAG.getVectorIdxConstant(Val: ExtBOIdx, DL);
28117 SDValue X = DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NarrowBVT,
28118 N1: BinOp.getOperand(i: 0), N2: NewExtIndex);
28119 SDValue Y = DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NarrowBVT,
28120 N1: BinOp.getOperand(i: 1), N2: NewExtIndex);
28121 SDValue NarrowBinOp =
28122 DAG.getNode(Opcode: BOpcode, DL, VT: NarrowBVT, N1: X, N2: Y, Flags: BinOp->getFlags());
28123 return DAG.getBitcast(VT, V: NarrowBinOp);
28124 }
28125
28126 // Only handle the case where we are doubling and then halving. A larger ratio
28127 // may require more than two narrow binops to replace the wide binop.
28128 if (NarrowingRatio != 2)
28129 return SDValue();
28130
28131 // TODO: The motivating case for this transform is an x86 AVX1 target. That
28132 // target has temptingly almost legal versions of bitwise logic ops in 256-bit
28133 // flavors, but no other 256-bit integer support. This could be extended to
28134 // handle any binop, but that may require fixing/adding other folds to avoid
28135 // codegen regressions.
28136 if (BOpcode != ISD::AND && BOpcode != ISD::OR && BOpcode != ISD::XOR)
28137 return SDValue();
28138
28139 // We need at least one concatenation operation of a binop operand to make
28140 // this transform worthwhile. The concat must double the input vector sizes.
28141 auto GetSubVector = [ConcatOpNum](SDValue V) -> SDValue {
28142 if (V.getOpcode() == ISD::CONCAT_VECTORS && V.getNumOperands() == 2)
28143 return V.getOperand(i: ConcatOpNum);
28144 return SDValue();
28145 };
28146 SDValue SubVecL = GetSubVector(peekThroughBitcasts(V: BinOp.getOperand(i: 0)));
28147 SDValue SubVecR = GetSubVector(peekThroughBitcasts(V: BinOp.getOperand(i: 1)));
28148
28149 if (SubVecL || SubVecR) {
28150 // If a binop operand was not the result of a concat, we must extract a
28151 // half-sized operand for our new narrow binop:
28152 // extract (binop (concat X1, X2), (concat Y1, Y2)), N --> binop XN, YN
28153 // extract (binop (concat X1, X2), Y), N --> binop XN, (extract Y, IndexC)
28154 // extract (binop X, (concat Y1, Y2)), N --> binop (extract X, IndexC), YN
28155 SDValue IndexC = DAG.getVectorIdxConstant(Val: ExtBOIdx, DL);
28156 SDValue X = SubVecL ? DAG.getBitcast(VT: NarrowBVT, V: SubVecL)
28157 : DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NarrowBVT,
28158 N1: BinOp.getOperand(i: 0), N2: IndexC);
28159
28160 SDValue Y = SubVecR ? DAG.getBitcast(VT: NarrowBVT, V: SubVecR)
28161 : DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NarrowBVT,
28162 N1: BinOp.getOperand(i: 1), N2: IndexC);
28163
28164 SDValue NarrowBinOp =
28165 DAG.getNode(Opcode: BOpcode, DL, VT: NarrowBVT, N1: X, N2: Y, Flags: BinOp->getFlags());
28166 return DAG.getBitcast(VT, V: NarrowBinOp);
28167 }
28168
28169 return SDValue();
28170}
28171
28172/// If we are extracting a subvector from a wide vector load, convert to a
28173/// narrow load to eliminate the extraction:
28174/// (extract_subvector (load wide vector)) --> (load narrow vector)
28175static SDValue narrowExtractedVectorLoad(EVT VT, SDValue Src, unsigned Index,
28176 const SDLoc &DL, SelectionDAG &DAG) {
28177 // TODO: Add support for big-endian. The offset calculation must be adjusted.
28178 if (DAG.getDataLayout().isBigEndian())
28179 return SDValue();
28180
28181 auto *Ld = dyn_cast<LoadSDNode>(Val&: Src);
28182 if (!Ld || !ISD::isNormalLoad(N: Ld) || !Ld->isSimple())
28183 return SDValue();
28184
28185 // We can only create byte sized loads.
28186 if (!VT.isByteSized())
28187 return SDValue();
28188
28189 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
28190 if (!TLI.isOperationLegalOrCustomOrPromote(Op: ISD::LOAD, VT))
28191 return SDValue();
28192
28193 unsigned NumElts = VT.getVectorMinNumElements();
28194 // A fixed length vector being extracted from a scalable vector
28195 // may not be any *smaller* than the scalable one.
28196 if (Index == 0 && NumElts >= Ld->getValueType(ResNo: 0).getVectorMinNumElements())
28197 return SDValue();
28198
28199 // The definition of EXTRACT_SUBVECTOR states that the index must be a
28200 // multiple of the minimum number of elements in the result type.
28201 assert(Index % NumElts == 0 && "The extract subvector index is not a "
28202 "multiple of the result's element count");
28203
28204 // It's fine to use TypeSize here as we know the offset will not be negative.
28205 TypeSize Offset = VT.getStoreSize() * (Index / NumElts);
28206 std::optional<unsigned> ByteOffset;
28207 if (Offset.isFixed())
28208 ByteOffset = Offset.getFixedValue();
28209
28210 if (!TLI.shouldReduceLoadWidth(Load: Ld, ExtTy: Ld->getExtensionType(), NewVT: VT, ByteOffset))
28211 return SDValue();
28212
28213 // The narrow load will be offset from the base address of the old load if
28214 // we are extracting from something besides index 0 (little-endian).
28215 // TODO: Use "BaseIndexOffset" to make this more effective.
28216 SDValue NewAddr = DAG.getMemBasePlusOffset(Base: Ld->getBasePtr(), Offset, DL);
28217
28218 MachineFunction &MF = DAG.getMachineFunction();
28219 MachineMemOperand *MMO;
28220 if (Offset.isScalable()) {
28221 MachinePointerInfo MPI =
28222 MachinePointerInfo(Ld->getPointerInfo().getAddrSpace());
28223 MMO = MF.getMachineMemOperand(MMO: Ld->getMemOperand(), PtrInfo: MPI, Size: VT.getStoreSize());
28224 } else
28225 MMO = MF.getMachineMemOperand(MMO: Ld->getMemOperand(), Offset: Offset.getFixedValue(),
28226 Size: VT.getStoreSize());
28227
28228 SDValue NewLd = DAG.getLoad(VT, dl: DL, Chain: Ld->getChain(), Ptr: NewAddr, MMO);
28229 DAG.makeEquivalentMemoryOrdering(OldLoad: Ld, NewMemOp: NewLd);
28230 return NewLd;
28231}
28232
28233/// Given EXTRACT_SUBVECTOR(VECTOR_SHUFFLE(Op0, Op1, Mask)),
28234/// try to produce VECTOR_SHUFFLE(EXTRACT_SUBVECTOR(Op?, ?),
28235/// EXTRACT_SUBVECTOR(Op?, ?),
28236/// Mask'))
28237/// iff it is legal and profitable to do so. Notably, the trimmed mask
28238/// (containing only the elements that are extracted)
28239/// must reference at most two subvectors.
28240static SDValue foldExtractSubvectorFromShuffleVector(EVT NarrowVT, SDValue Src,
28241 unsigned Index,
28242 const SDLoc &DL,
28243 SelectionDAG &DAG,
28244 bool LegalOperations) {
28245 // Only deal with non-scalable vectors.
28246 EVT WideVT = Src.getValueType();
28247 if (!NarrowVT.isFixedLengthVector() || !WideVT.isFixedLengthVector())
28248 return SDValue();
28249
28250 // The operand must be a shufflevector.
28251 auto *WideShuffleVector = dyn_cast<ShuffleVectorSDNode>(Val&: Src);
28252 if (!WideShuffleVector)
28253 return SDValue();
28254
28255 // The old shuffleneeds to go away.
28256 if (!WideShuffleVector->hasOneUse())
28257 return SDValue();
28258
28259 // And the narrow shufflevector that we'll form must be legal.
28260 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
28261 if (LegalOperations &&
28262 !TLI.isOperationLegalOrCustom(Op: ISD::VECTOR_SHUFFLE, VT: NarrowVT))
28263 return SDValue();
28264
28265 int NumEltsExtracted = NarrowVT.getVectorNumElements();
28266 assert((Index % NumEltsExtracted) == 0 &&
28267 "Extract index is not a multiple of the output vector length.");
28268
28269 int WideNumElts = WideVT.getVectorNumElements();
28270
28271 SmallVector<int, 16> NewMask;
28272 NewMask.reserve(N: NumEltsExtracted);
28273 SmallSetVector<std::pair<SDValue /*Op*/, int /*SubvectorIndex*/>, 2>
28274 DemandedSubvectors;
28275
28276 // Try to decode the wide mask into narrow mask from at most two subvectors.
28277 for (int M : WideShuffleVector->getMask().slice(N: Index, M: NumEltsExtracted)) {
28278 assert((M >= -1) && (M < (2 * WideNumElts)) &&
28279 "Out-of-bounds shuffle mask?");
28280
28281 if (M < 0) {
28282 // Does not depend on operands, does not require adjustment.
28283 NewMask.emplace_back(Args&: M);
28284 continue;
28285 }
28286
28287 // From which operand of the shuffle does this shuffle mask element pick?
28288 int WideShufOpIdx = M / WideNumElts;
28289 // Which element of that operand is picked?
28290 int OpEltIdx = M % WideNumElts;
28291
28292 assert((OpEltIdx + WideShufOpIdx * WideNumElts) == M &&
28293 "Shuffle mask vector decomposition failure.");
28294
28295 // And which NumEltsExtracted-sized subvector of that operand is that?
28296 int OpSubvecIdx = OpEltIdx / NumEltsExtracted;
28297 // And which element within that subvector of that operand is that?
28298 int OpEltIdxInSubvec = OpEltIdx % NumEltsExtracted;
28299
28300 assert((OpEltIdxInSubvec + OpSubvecIdx * NumEltsExtracted) == OpEltIdx &&
28301 "Shuffle mask subvector decomposition failure.");
28302
28303 assert((OpEltIdxInSubvec + OpSubvecIdx * NumEltsExtracted +
28304 WideShufOpIdx * WideNumElts) == M &&
28305 "Shuffle mask full decomposition failure.");
28306
28307 SDValue Op = WideShuffleVector->getOperand(Num: WideShufOpIdx);
28308
28309 if (Op.isUndef()) {
28310 // Picking from an undef operand. Let's adjust mask instead.
28311 NewMask.emplace_back(Args: -1);
28312 continue;
28313 }
28314
28315 const std::pair<SDValue, int> DemandedSubvector =
28316 std::make_pair(x&: Op, y&: OpSubvecIdx);
28317
28318 if (DemandedSubvectors.insert(X: DemandedSubvector)) {
28319 if (DemandedSubvectors.size() > 2)
28320 return SDValue(); // We can't handle more than two subvectors.
28321 // How many elements into the WideVT does this subvector start?
28322 int Index = NumEltsExtracted * OpSubvecIdx;
28323 // Bail out if the extraction isn't going to be cheap.
28324 if (TLI.getExtractSubvectorCost(ResVT: NarrowVT, SrcVT: WideVT, Index) >
28325 TargetLowering::ExtractSubvectorCost::Cheap)
28326 return SDValue();
28327 }
28328
28329 // Ok, but from which operand of the new shuffle will this element pick?
28330 int NewOpIdx =
28331 getFirstIndexOf(Range: DemandedSubvectors.getArrayRef(), Val: DemandedSubvector);
28332 assert((NewOpIdx == 0 || NewOpIdx == 1) && "Unexpected operand index.");
28333
28334 int AdjM = OpEltIdxInSubvec + NewOpIdx * NumEltsExtracted;
28335 NewMask.emplace_back(Args&: AdjM);
28336 }
28337 assert(NewMask.size() == (unsigned)NumEltsExtracted && "Produced bad mask.");
28338 assert(DemandedSubvectors.size() <= 2 &&
28339 "Should have ended up demanding at most two subvectors.");
28340
28341 // Did we discover that the shuffle does not actually depend on operands?
28342 if (DemandedSubvectors.empty())
28343 return DAG.getPOISON(VT: NarrowVT);
28344
28345 // Profitability check: only deal with extractions from the first subvector
28346 // unless the mask becomes an identity mask.
28347 if (!ShuffleVectorInst::isIdentityMask(Mask: NewMask, NumSrcElts: NewMask.size()) ||
28348 any_of(Range&: NewMask, P: [](int M) { return M < 0; }))
28349 for (auto &DemandedSubvector : DemandedSubvectors)
28350 if (DemandedSubvector.second != 0)
28351 return SDValue();
28352
28353 // We still perform the exact same EXTRACT_SUBVECTOR, just on different
28354 // operand[s]/index[es], so there is no point in checking for it's legality.
28355
28356 // Do not turn a legal shuffle into an illegal one.
28357 if (TLI.isShuffleMaskLegal(WideShuffleVector->getMask(), WideVT) &&
28358 !TLI.isShuffleMaskLegal(NewMask, NarrowVT))
28359 return SDValue();
28360
28361 SmallVector<SDValue, 2> NewOps;
28362 for (const std::pair<SDValue /*Op*/, int /*SubvectorIndex*/>
28363 &DemandedSubvector : DemandedSubvectors) {
28364 // How many elements into the WideVT does this subvector start?
28365 int Index = NumEltsExtracted * DemandedSubvector.second;
28366 SDValue IndexC = DAG.getVectorIdxConstant(Val: Index, DL);
28367 NewOps.emplace_back(Args: DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NarrowVT,
28368 N1: DemandedSubvector.first, N2: IndexC));
28369 }
28370 assert((NewOps.size() == 1 || NewOps.size() == 2) &&
28371 "Should end up with either one or two ops");
28372
28373 // If we ended up with only one operand, pad with poison.
28374 if (NewOps.size() == 1)
28375 NewOps.emplace_back(Args: DAG.getPOISON(VT: NarrowVT));
28376
28377 return DAG.getVectorShuffle(VT: NarrowVT, dl: DL, N1: NewOps[0], N2: NewOps[1], Mask: NewMask);
28378}
28379
28380SDValue DAGCombiner::foldExtractSubvectorFromConcatVectors(EVT VT, SDValue V,
28381 uint64_t ExtIdx,
28382 const SDLoc &DL) {
28383 assert(V.getOpcode() == ISD::CONCAT_VECTORS &&
28384 "Expected a CONCAT_VECTORS operand");
28385 ElementCount ExtNumElts = VT.getVectorElementCount();
28386 assert(ExtIdx % ExtNumElts.getKnownMinValue() == 0 &&
28387 "subvector extract is alligned");
28388 EVT ConcatSrcVT = V.getOperand(i: 0).getValueType();
28389
28390 ElementCount ConcatSrcNumElts = ConcatSrcVT.getVectorElementCount();
28391 unsigned ConcatOpIdx = ExtIdx / ConcatSrcNumElts.getKnownMinValue();
28392 if (ConcatOpIdx >= V.getNumOperands())
28393 return SDValue();
28394
28395 // If the concatenated source types match this extract, it's a direct
28396 // simplification:
28397 // extract_subvector (concat V1, V2, ...), i --> Vi
28398 if (VT.getVectorElementCount() == ConcatSrcVT.getVectorElementCount())
28399 return V.getOperand(i: ConcatOpIdx);
28400
28401 // If the concatenated source vectors are a multiple length of this extract,
28402 // then extract a fraction of one of those source vectors directly from a
28403 // concat operand. Example:
28404 // v2i8 extract_subvector (v16i8 concat_subvector v8i8:X, v8i8:Y), 14 -->
28405 // v2i8 extract_subvector v8i8:Y, 6
28406 if (ConcatSrcNumElts.hasKnownScalarFactor(RHS: ExtNumElts)) {
28407 uint64_t NewExtIdx =
28408 ExtIdx - ConcatOpIdx * ConcatSrcNumElts.getKnownMinValue();
28409 return DAG.getExtractSubvector(DL, VT, Vec: V.getOperand(i: ConcatOpIdx),
28410 Idx: NewExtIdx);
28411 }
28412
28413 // If the extract covers multiple whole concat operands, rebuild that smaller
28414 // concat directly.
28415 if (ExtNumElts.hasKnownScalarFactor(RHS: ConcatSrcNumElts) &&
28416 ExtIdx % ConcatSrcNumElts.getKnownMinValue() == 0 &&
28417 (!LegalOperations || hasOperation(Opcode: ISD::CONCAT_VECTORS, VT))) {
28418 unsigned NumConcatOps = ExtNumElts.getKnownScalarFactor(RHS: ConcatSrcNumElts);
28419 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT,
28420 Ops: V->ops().slice(N: ConcatOpIdx, M: NumConcatOps));
28421 }
28422
28423 return SDValue();
28424}
28425
28426SDValue DAGCombiner::visitEXTRACT_SUBVECTOR(SDNode *N) {
28427 EVT NVT = N->getValueType(ResNo: 0);
28428 SDValue V = N->getOperand(Num: 0);
28429 uint64_t ExtIdx = N->getConstantOperandVal(Num: 1);
28430 SDLoc DL(N);
28431
28432 // Extract from UNDEF is UNDEF.
28433 if (V.isUndef())
28434 return DAG.getUNDEF(VT: NVT);
28435
28436 if (SDValue NarrowLoad = narrowExtractedVectorLoad(VT: NVT, Src: V, Index: ExtIdx, DL, DAG))
28437 return NarrowLoad;
28438
28439 // Peek through frozen loads, but ensure the load has a single use.
28440 if (V.getOpcode() == ISD::FREEZE && V.hasOneUse() &&
28441 V.getOperand(i: 0).hasOneUse())
28442 if (SDValue NarrowLoad =
28443 narrowExtractedVectorLoad(VT: NVT, Src: V.getOperand(i: 0), Index: ExtIdx, DL, DAG))
28444 return DAG.getFreeze(V: NarrowLoad);
28445
28446 // Combine an extract of an extract into a single extract_subvector.
28447 // ext (ext X, C1), C2 --> ext X, C1 + C2
28448 if (V.getOpcode() == ISD::EXTRACT_SUBVECTOR && V.hasOneUse()) {
28449 // Both indices must have the same scaling factor and C has to be a
28450 // multiple of the new result type's known minimum vector length.
28451 uint64_t InnerExtIdx = V.getConstantOperandVal(i: 1);
28452 uint64_t NewExtIdx = InnerExtIdx + ExtIdx;
28453 if (V.getValueType().isScalableVector() == NVT.isScalableVector() &&
28454 NewExtIdx % NVT.getVectorMinNumElements() == 0 &&
28455 TLI.getExtractSubvectorCost(ResVT: NVT, SrcVT: V.getOperand(i: 0).getValueType(),
28456 Index: NewExtIdx) <=
28457 TargetLowering::ExtractSubvectorCost::Cheap &&
28458 TLI.isOperationLegalOrCustom(Op: ISD::EXTRACT_SUBVECTOR, VT: NVT))
28459 return DAG.getExtractSubvector(DL, VT: NVT, Vec: V.getOperand(i: 0), Idx: NewExtIdx);
28460 }
28461
28462 // ty1 extract_vector(ty2 splat(V))) -> ty1 splat(V)
28463 if (V.getOpcode() == ISD::SPLAT_VECTOR)
28464 if ((DAG.isConstantValueOfAnyType(N: V.getOperand(i: 0)) &&
28465 !(NVT.isScalableVector() &&
28466 TLI.getExtractSubvectorCost(ResVT: NVT, SrcVT: V.getValueType(), Index: ExtIdx) <=
28467 TargetLowering::ExtractSubvectorCost::Cheap)) ||
28468 V.hasOneUse())
28469 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::SPLAT_VECTOR, VT: NVT))
28470 return DAG.getSplatVector(VT: NVT, DL, Op: V.getOperand(i: 0));
28471
28472 // ty1 extract_vector(ty2 get_active_lane_mask(X, Y), 0) --> ty1
28473 // get_active_lane_mask(X, Y)
28474 if (ExtIdx == 0 && V.getOpcode() == ISD::GET_ACTIVE_LANE_MASK &&
28475 V.hasOneUse() &&
28476 (!LegalOperations ||
28477 TLI.isOperationLegal(Op: ISD::GET_ACTIVE_LANE_MASK, VT: NVT)))
28478 return DAG.getNode(Opcode: ISD::GET_ACTIVE_LANE_MASK, DL, VT: NVT, N1: V.getOperand(i: 0),
28479 N2: V.getOperand(i: 1));
28480
28481 // extract_subvector(insert_subvector(x,y,c1),c2)
28482 // --> extract_subvector(y,c2-c1)
28483 // iff we're just extracting from the inserted subvector.
28484 if (V.getOpcode() == ISD::INSERT_SUBVECTOR) {
28485 SDValue InsSub = V.getOperand(i: 1);
28486 EVT InsSubVT = InsSub.getValueType();
28487 unsigned NumInsElts = InsSubVT.getVectorMinNumElements();
28488 unsigned InsIdx = V.getConstantOperandVal(i: 2);
28489 unsigned NumSubElts = NVT.getVectorMinNumElements();
28490 if (InsIdx <= ExtIdx && (ExtIdx + NumSubElts) <= (InsIdx + NumInsElts) &&
28491 TLI.getExtractSubvectorCost(ResVT: NVT, SrcVT: InsSubVT, Index: ExtIdx - InsIdx) <=
28492 TargetLowering::ExtractSubvectorCost::Cheap &&
28493 InsSubVT.isFixedLengthVector() && NVT.isFixedLengthVector() &&
28494 V.getValueType().isFixedLengthVector())
28495 return DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NVT, N1: InsSub,
28496 N2: DAG.getVectorIdxConstant(Val: ExtIdx - InsIdx, DL));
28497 }
28498
28499 // Try to move vector bitcast after extract_subv by scaling extraction index:
28500 // extract_subv (bitcast X), Index --> bitcast (extract_subv X, Index')
28501 if (V.getOpcode() == ISD::BITCAST &&
28502 V.getOperand(i: 0).getValueType().isVector() &&
28503 (!LegalOperations || TLI.isOperationLegal(Op: ISD::BITCAST, VT: NVT))) {
28504 SDValue SrcOp = V.getOperand(i: 0);
28505 EVT SrcVT = SrcOp.getValueType();
28506 unsigned SrcNumElts = SrcVT.getVectorMinNumElements();
28507 unsigned DestNumElts = V.getValueType().getVectorMinNumElements();
28508 if ((SrcNumElts % DestNumElts) == 0) {
28509 unsigned SrcDestRatio = SrcNumElts / DestNumElts;
28510 ElementCount NewExtEC = NVT.getVectorElementCount() * SrcDestRatio;
28511 EVT NewExtVT =
28512 EVT::getVectorVT(Context&: *DAG.getContext(), VT: SrcVT.getScalarType(), EC: NewExtEC);
28513 if (TLI.isOperationLegalOrCustom(Op: ISD::EXTRACT_SUBVECTOR, VT: NewExtVT)) {
28514 SDValue NewIndex = DAG.getVectorIdxConstant(Val: ExtIdx * SrcDestRatio, DL);
28515 SDValue NewExtract = DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NewExtVT,
28516 N1: V.getOperand(i: 0), N2: NewIndex);
28517 return DAG.getBitcast(VT: NVT, V: NewExtract);
28518 }
28519 }
28520 if ((DestNumElts % SrcNumElts) == 0) {
28521 unsigned DestSrcRatio = DestNumElts / SrcNumElts;
28522 if (NVT.getVectorElementCount().isKnownMultipleOf(RHS: DestSrcRatio)) {
28523 ElementCount NewExtEC =
28524 NVT.getVectorElementCount().divideCoefficientBy(RHS: DestSrcRatio);
28525 EVT ScalarVT = SrcVT.getScalarType();
28526 if ((ExtIdx % DestSrcRatio) == 0) {
28527 unsigned IndexValScaled = ExtIdx / DestSrcRatio;
28528 EVT NewExtVT =
28529 EVT::getVectorVT(Context&: *DAG.getContext(), VT: ScalarVT, EC: NewExtEC);
28530 if (TLI.isOperationLegalOrCustom(Op: ISD::EXTRACT_SUBVECTOR, VT: NewExtVT)) {
28531 SDValue NewIndex = DAG.getVectorIdxConstant(Val: IndexValScaled, DL);
28532 SDValue NewExtract =
28533 DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NewExtVT,
28534 N1: V.getOperand(i: 0), N2: NewIndex);
28535 return DAG.getBitcast(VT: NVT, V: NewExtract);
28536 }
28537 if (NewExtEC.isScalar() &&
28538 TLI.isOperationLegalOrCustom(Op: ISD::EXTRACT_VECTOR_ELT, VT: ScalarVT)) {
28539 SDValue NewIndex = DAG.getVectorIdxConstant(Val: IndexValScaled, DL);
28540 SDValue NewExtract =
28541 DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: ScalarVT,
28542 N1: V.getOperand(i: 0), N2: NewIndex);
28543 return DAG.getBitcast(VT: NVT, V: NewExtract);
28544 }
28545 }
28546 }
28547 }
28548 }
28549
28550 if (V.getOpcode() == ISD::CONCAT_VECTORS) {
28551 if (SDValue Folded =
28552 foldExtractSubvectorFromConcatVectors(VT: NVT, V, ExtIdx, DL))
28553 return Folded;
28554 }
28555
28556 if (SDValue Shuffle = foldExtractSubvectorFromShuffleVector(
28557 NarrowVT: NVT, Src: V, Index: ExtIdx, DL, DAG, LegalOperations))
28558 return Shuffle;
28559
28560 if (SDValue NarrowBOp =
28561 narrowExtractedVectorBinOp(VT: NVT, Src: V, Index: ExtIdx, DL, DAG, LegalOperations))
28562 return NarrowBOp;
28563
28564 V = peekThroughBitcasts(V);
28565
28566 // If the input is a build vector. Try to make a smaller build vector.
28567 if (V.getOpcode() == ISD::BUILD_VECTOR) {
28568 EVT InVT = V.getValueType();
28569 unsigned ExtractSize = NVT.getSizeInBits();
28570 unsigned EltSize = InVT.getScalarSizeInBits();
28571 // Only do this if we won't split any elements.
28572 if (ExtractSize % EltSize == 0) {
28573 unsigned NumElems = ExtractSize / EltSize;
28574 EVT EltVT = InVT.getVectorElementType();
28575 EVT ExtractVT =
28576 NumElems == 1 ? EltVT
28577 : EVT::getVectorVT(Context&: *DAG.getContext(), VT: EltVT, NumElements: NumElems);
28578 if ((Level < AfterLegalizeDAG ||
28579 (NumElems == 1 ||
28580 TLI.isOperationLegal(Op: ISD::BUILD_VECTOR, VT: ExtractVT))) &&
28581 (!LegalTypes || TLI.isTypeLegal(VT: ExtractVT))) {
28582 unsigned IdxVal = (ExtIdx * NVT.getScalarSizeInBits()) / EltSize;
28583
28584 if (NumElems == 1) {
28585 SDValue Src = V->getOperand(Num: IdxVal);
28586 if (EltVT != Src.getValueType())
28587 Src = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: EltVT, Operand: Src);
28588 return DAG.getBitcast(VT: NVT, V: Src);
28589 }
28590
28591 // Extract the pieces from the original build_vector.
28592 SDValue BuildVec =
28593 DAG.getBuildVector(VT: ExtractVT, DL, Ops: V->ops().slice(N: IdxVal, M: NumElems));
28594 return DAG.getBitcast(VT: NVT, V: BuildVec);
28595 }
28596 }
28597 }
28598
28599 if (V.getOpcode() == ISD::INSERT_SUBVECTOR) {
28600 // Handle only simple case where vector being inserted and vector
28601 // being extracted are of same size.
28602 EVT SmallVT = V.getOperand(i: 1).getValueType();
28603 if (NVT.bitsEq(VT: SmallVT)) {
28604 // Combine:
28605 // (extract_subvec (insert_subvec V1, V2, InsIdx), ExtIdx)
28606 // Into:
28607 // indices are equal or bit offsets are equal => V1
28608 // otherwise => (extract_subvec V1, ExtIdx)
28609 uint64_t InsIdx = V.getConstantOperandVal(i: 2);
28610 if (InsIdx * SmallVT.getScalarSizeInBits() ==
28611 ExtIdx * NVT.getScalarSizeInBits()) {
28612 if (!LegalOperations || TLI.isOperationLegal(Op: ISD::BITCAST, VT: NVT))
28613 return DAG.getBitcast(VT: NVT, V: V.getOperand(i: 1));
28614 } else {
28615 return DAG.getNode(
28616 Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: NVT,
28617 N1: DAG.getBitcast(VT: N->getOperand(Num: 0).getValueType(), V: V.getOperand(i: 0)),
28618 N2: N->getOperand(Num: 1));
28619 }
28620 }
28621 }
28622
28623 // If only EXTRACT_SUBVECTOR nodes use the source vector we can
28624 // simplify it based on the (valid) extractions.
28625 if (!V.getValueType().isScalableVector() &&
28626 llvm::all_of(Range: V->users(), P: [&](SDNode *Use) {
28627 return Use->getOpcode() == ISD::EXTRACT_SUBVECTOR &&
28628 Use->getOperand(Num: 0) == V;
28629 })) {
28630 unsigned NumElts = V.getValueType().getVectorNumElements();
28631 APInt DemandedElts = APInt::getZero(numBits: NumElts);
28632 for (SDNode *User : V->users()) {
28633 unsigned ExtIdx = User->getConstantOperandVal(Num: 1);
28634 unsigned NumSubElts = User->getValueType(ResNo: 0).getVectorNumElements();
28635 DemandedElts.setBits(loBit: ExtIdx, hiBit: ExtIdx + NumSubElts);
28636 }
28637 if (SimplifyDemandedVectorElts(Op: V, DemandedElts, /*AssumeSingleUse=*/true)) {
28638 // We simplified the vector operand of this extract subvector. If this
28639 // extract is not dead, visit it again so it is folded properly.
28640 if (N->getOpcode() != ISD::DELETED_NODE)
28641 AddToWorklist(N);
28642 return SDValue(N, 0);
28643 }
28644 } else {
28645 if (SimplifyDemandedVectorElts(Op: SDValue(N, 0)))
28646 return SDValue(N, 0);
28647 }
28648
28649 return SDValue();
28650}
28651
28652/// Try to convert a wide shuffle of concatenated vectors into 2 narrow shuffles
28653/// followed by concatenation. Narrow vector ops may have better performance
28654/// than wide ops, and this can unlock further narrowing of other vector ops.
28655/// Targets can invert this transform later if it is not profitable.
28656static SDValue foldShuffleOfConcatUndefs(ShuffleVectorSDNode *Shuf,
28657 SelectionDAG &DAG) {
28658 SDValue N0 = Shuf->getOperand(Num: 0), N1 = Shuf->getOperand(Num: 1);
28659 if (N0.getOpcode() != ISD::CONCAT_VECTORS || N0.getNumOperands() != 2 ||
28660 N1.getOpcode() != ISD::CONCAT_VECTORS || N1.getNumOperands() != 2 ||
28661 !N0.getOperand(i: 1).isUndef() || !N1.getOperand(i: 1).isUndef())
28662 return SDValue();
28663
28664 // Split the wide shuffle mask into halves. Any mask element that is accessing
28665 // operand 1 is offset down to account for narrowing of the vectors.
28666 ArrayRef<int> Mask = Shuf->getMask();
28667 EVT VT = Shuf->getValueType(ResNo: 0);
28668 unsigned NumElts = VT.getVectorNumElements();
28669 unsigned HalfNumElts = NumElts / 2;
28670 SmallVector<int, 16> Mask0(HalfNumElts, -1);
28671 SmallVector<int, 16> Mask1(HalfNumElts, -1);
28672 for (unsigned i = 0; i != NumElts; ++i) {
28673 if (Mask[i] == -1)
28674 continue;
28675 // If we reference the upper (undef) subvector then the element is undef.
28676 if ((Mask[i] % NumElts) >= HalfNumElts)
28677 continue;
28678 int M = Mask[i] < (int)NumElts ? Mask[i] : Mask[i] - (int)HalfNumElts;
28679 if (i < HalfNumElts)
28680 Mask0[i] = M;
28681 else
28682 Mask1[i - HalfNumElts] = M;
28683 }
28684
28685 // Ask the target if this is a valid transform.
28686 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
28687 EVT HalfVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: VT.getScalarType(),
28688 NumElements: HalfNumElts);
28689 if (!TLI.isShuffleMaskLegal(Mask0, HalfVT) ||
28690 !TLI.isShuffleMaskLegal(Mask1, HalfVT))
28691 return SDValue();
28692
28693 // shuffle (concat X, undef), (concat Y, undef), Mask -->
28694 // concat (shuffle X, Y, Mask0), (shuffle X, Y, Mask1)
28695 SDValue X = N0.getOperand(i: 0), Y = N1.getOperand(i: 0);
28696 SDLoc DL(Shuf);
28697 SDValue Shuf0 = DAG.getVectorShuffle(VT: HalfVT, dl: DL, N1: X, N2: Y, Mask: Mask0);
28698 SDValue Shuf1 = DAG.getVectorShuffle(VT: HalfVT, dl: DL, N1: X, N2: Y, Mask: Mask1);
28699 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT, N1: Shuf0, N2: Shuf1);
28700}
28701
28702// Tries to turn a shuffle of two CONCAT_VECTORS into a single concat,
28703// or turn a shuffle of a single concat into simpler shuffle then concat.
28704static SDValue partitionShuffleOfConcats(SDNode *N, SelectionDAG &DAG) {
28705 EVT VT = N->getValueType(ResNo: 0);
28706 unsigned NumElts = VT.getVectorNumElements();
28707
28708 SDValue N0 = N->getOperand(Num: 0);
28709 SDValue N1 = N->getOperand(Num: 1);
28710 ShuffleVectorSDNode *SVN = cast<ShuffleVectorSDNode>(Val: N);
28711 ArrayRef<int> Mask = SVN->getMask();
28712
28713 SmallVector<SDValue, 4> Ops;
28714 EVT ConcatVT = N0.getOperand(i: 0).getValueType();
28715 unsigned NumElemsPerConcat = ConcatVT.getVectorNumElements();
28716 unsigned NumConcats = NumElts / NumElemsPerConcat;
28717
28718 auto IsUndefMaskElt = [](int i) { return i == -1; };
28719
28720 // Special case: shuffle(concat(A,B)) can be more efficiently represented
28721 // as concat(shuffle(A,B),UNDEF) if the shuffle doesn't set any of the high
28722 // half vector elements.
28723 if (NumElemsPerConcat * 2 == NumElts && N1.isUndef() &&
28724 llvm::all_of(Range: Mask.slice(N: NumElemsPerConcat, M: NumElemsPerConcat),
28725 P: IsUndefMaskElt)) {
28726 N0 = DAG.getVectorShuffle(VT: ConcatVT, dl: SDLoc(N), N1: N0.getOperand(i: 0),
28727 N2: N0.getOperand(i: 1),
28728 Mask: Mask.slice(N: 0, M: NumElemsPerConcat));
28729 N1 = DAG.getPOISON(VT: ConcatVT);
28730 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT, N1: N0, N2: N1);
28731 }
28732
28733 // Look at every vector that's inserted. We're looking for exact
28734 // subvector-sized copies from a concatenated vector
28735 for (unsigned I = 0; I != NumConcats; ++I) {
28736 unsigned Begin = I * NumElemsPerConcat;
28737 ArrayRef<int> SubMask = Mask.slice(N: Begin, M: NumElemsPerConcat);
28738
28739 // Make sure we're dealing with a copy.
28740 if (llvm::all_of(Range&: SubMask, P: IsUndefMaskElt)) {
28741 Ops.push_back(Elt: DAG.getUNDEF(VT: ConcatVT));
28742 continue;
28743 }
28744
28745 int OpIdx = -1;
28746 for (int i = 0; i != (int)NumElemsPerConcat; ++i) {
28747 if (IsUndefMaskElt(SubMask[i]))
28748 continue;
28749 if ((SubMask[i] % (int)NumElemsPerConcat) != i)
28750 return SDValue();
28751 int EltOpIdx = SubMask[i] / NumElemsPerConcat;
28752 if (0 <= OpIdx && EltOpIdx != OpIdx)
28753 return SDValue();
28754 OpIdx = EltOpIdx;
28755 }
28756 assert(0 <= OpIdx && "Unknown concat_vectors op");
28757
28758 if (OpIdx < (int)N0.getNumOperands())
28759 Ops.push_back(Elt: N0.getOperand(i: OpIdx));
28760 else
28761 Ops.push_back(Elt: N1.getOperand(i: OpIdx - N0.getNumOperands()));
28762 }
28763
28764 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT, Ops);
28765}
28766
28767// Attempt to combine a shuffle of 2 inputs of 'scalar sources' -
28768// BUILD_VECTOR or SCALAR_TO_VECTOR into a single BUILD_VECTOR.
28769//
28770// SHUFFLE(BUILD_VECTOR(), BUILD_VECTOR()) -> BUILD_VECTOR() is always
28771// a simplification in some sense, but it isn't appropriate in general: some
28772// BUILD_VECTORs are substantially cheaper than others. The general case
28773// of a BUILD_VECTOR requires inserting each element individually (or
28774// performing the equivalent in a temporary stack variable). A BUILD_VECTOR of
28775// all constants is a single constant pool load. A BUILD_VECTOR where each
28776// element is identical is a splat. A BUILD_VECTOR where most of the operands
28777// are undef lowers to a small number of element insertions.
28778//
28779// To deal with this, we currently use a bunch of mostly arbitrary heuristics.
28780// We don't fold shuffles where one side is a non-zero constant, and we don't
28781// fold shuffles if the resulting (non-splat) BUILD_VECTOR would have duplicate
28782// non-constant operands. This seems to work out reasonably well in practice.
28783static SDValue combineShuffleOfScalars(ShuffleVectorSDNode *SVN,
28784 SelectionDAG &DAG,
28785 const TargetLowering &TLI) {
28786 EVT VT = SVN->getValueType(ResNo: 0);
28787 unsigned NumElts = VT.getVectorNumElements();
28788 SDValue N0 = SVN->getOperand(Num: 0);
28789 SDValue N1 = SVN->getOperand(Num: 1);
28790
28791 if (!N0->hasOneUse())
28792 return SDValue();
28793
28794 // If only one of N1,N2 is constant, bail out if it is not ALL_ZEROS as
28795 // discussed above.
28796 if (!N1.isUndef()) {
28797 if (!N1->hasOneUse())
28798 return SDValue();
28799
28800 bool N0AnyConst = isAnyConstantBuildVector(V: N0);
28801 bool N1AnyConst = isAnyConstantBuildVector(V: N1);
28802 if (N0AnyConst && !N1AnyConst && !ISD::isBuildVectorAllZeros(N: N0.getNode()))
28803 return SDValue();
28804 if (!N0AnyConst && N1AnyConst && !ISD::isBuildVectorAllZeros(N: N1.getNode()))
28805 return SDValue();
28806 }
28807
28808 // If both inputs are splats of the same value then we can safely merge this
28809 // to a single BUILD_VECTOR with undef elements based on the shuffle mask.
28810 bool IsSplat = false;
28811 auto *BV0 = dyn_cast<BuildVectorSDNode>(Val&: N0);
28812 auto *BV1 = dyn_cast<BuildVectorSDNode>(Val&: N1);
28813 if (BV0 && BV1)
28814 if (SDValue Splat0 = BV0->getSplatValue())
28815 IsSplat = (Splat0 == BV1->getSplatValue());
28816
28817 SmallVector<SDValue, 8> Ops;
28818 SmallSet<SDValue, 16> DuplicateOps;
28819 for (int M : SVN->getMask()) {
28820 SDValue Op = DAG.getPOISON(VT: VT.getScalarType());
28821 if (M >= 0) {
28822 int Idx = M < (int)NumElts ? M : M - NumElts;
28823 SDValue &S = (M < (int)NumElts ? N0 : N1);
28824 if (S.getOpcode() == ISD::BUILD_VECTOR) {
28825 Op = S.getOperand(i: Idx);
28826 } else if (S.getOpcode() == ISD::SCALAR_TO_VECTOR) {
28827 SDValue Op0 = S.getOperand(i: 0);
28828 Op = Idx == 0 ? Op0 : DAG.getPOISON(VT: Op0.getValueType());
28829 } else {
28830 // Operand can't be combined - bail out.
28831 return SDValue();
28832 }
28833 }
28834
28835 // Don't duplicate a non-constant BUILD_VECTOR operand unless we're
28836 // generating a splat; semantically, this is fine, but it's likely to
28837 // generate low-quality code if the target can't reconstruct an appropriate
28838 // shuffle.
28839 if (!Op.isUndef() && !isIntOrFPConstant(V: Op))
28840 if (!IsSplat && !DuplicateOps.insert(V: Op).second)
28841 return SDValue();
28842
28843 Ops.push_back(Elt: Op);
28844 }
28845
28846 // BUILD_VECTOR requires all inputs to be of the same type, find the
28847 // maximum type and extend them all.
28848 EVT SVT = VT.getScalarType();
28849 if (SVT.isInteger())
28850 for (SDValue &Op : Ops)
28851 SVT = (SVT.bitsLT(VT: Op.getValueType()) ? Op.getValueType() : SVT);
28852 if (SVT != VT.getScalarType())
28853 for (SDValue &Op : Ops)
28854 Op = Op.isUndef() ? DAG.getUNDEF(VT: SVT)
28855 : (TLI.isZExtFree(FromTy: Op.getValueType(), ToTy: SVT)
28856 ? DAG.getZExtOrTrunc(Op, DL: SDLoc(SVN), VT: SVT)
28857 : DAG.getSExtOrTrunc(Op, DL: SDLoc(SVN), VT: SVT));
28858 return DAG.getBuildVector(VT, DL: SDLoc(SVN), Ops);
28859}
28860
28861// Match shuffles that can be converted to *_vector_extend_in_reg.
28862// This is often generated during legalization.
28863// e.g. v4i32 <0,u,1,u> -> (v2i64 any_vector_extend_in_reg(v4i32 src)),
28864// and returns the EVT to which the extension should be performed.
28865// NOTE: this assumes that the src is the first operand of the shuffle.
28866static std::optional<EVT> canCombineShuffleToExtendVectorInreg(
28867 unsigned Opcode, EVT VT, std::function<bool(unsigned)> Match,
28868 SelectionDAG &DAG, const TargetLowering &TLI, bool LegalTypes,
28869 bool LegalOperations) {
28870 bool IsBigEndian = DAG.getDataLayout().isBigEndian();
28871
28872 // TODO Add support for big-endian when we have a test case.
28873 if (!VT.isInteger() || IsBigEndian)
28874 return std::nullopt;
28875
28876 unsigned NumElts = VT.getVectorNumElements();
28877 unsigned EltSizeInBits = VT.getScalarSizeInBits();
28878
28879 // Attempt to match a '*_extend_vector_inreg' shuffle, we just search for
28880 // power-of-2 extensions as they are the most likely.
28881 // FIXME: should try Scale == NumElts case too,
28882 for (unsigned Scale = 2; Scale < NumElts; Scale *= 2) {
28883 // The vector width must be a multiple of Scale.
28884 if (NumElts % Scale != 0)
28885 continue;
28886
28887 EVT OutSVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: EltSizeInBits * Scale);
28888 EVT OutVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: OutSVT, NumElements: NumElts / Scale);
28889
28890 if ((LegalTypes && !TLI.isTypeLegal(VT: OutVT)) ||
28891 (LegalOperations && !TLI.isOperationLegalOrCustom(Op: Opcode, VT: OutVT)))
28892 continue;
28893
28894 if (Match(Scale))
28895 return OutVT;
28896 }
28897
28898 return std::nullopt;
28899}
28900
28901// Match shuffles that can be converted to any_vector_extend_in_reg.
28902// This is often generated during legalization.
28903// e.g. v4i32 <0,u,1,u> -> (v2i64 any_vector_extend_in_reg(v4i32 src))
28904static SDValue combineShuffleToAnyExtendVectorInreg(ShuffleVectorSDNode *SVN,
28905 SelectionDAG &DAG,
28906 const TargetLowering &TLI,
28907 bool LegalOperations) {
28908 EVT VT = SVN->getValueType(ResNo: 0);
28909 bool IsBigEndian = DAG.getDataLayout().isBigEndian();
28910
28911 // TODO Add support for big-endian when we have a test case.
28912 if (!VT.isInteger() || IsBigEndian)
28913 return SDValue();
28914
28915 // shuffle<0,-1,1,-1> == (v2i64 anyextend_vector_inreg(v4i32))
28916 auto isAnyExtend = [NumElts = VT.getVectorNumElements(),
28917 Mask = SVN->getMask()](unsigned Scale) {
28918 for (unsigned i = 0; i != NumElts; ++i) {
28919 if (Mask[i] < 0)
28920 continue;
28921 if ((i % Scale) == 0 && Mask[i] == (int)(i / Scale))
28922 continue;
28923 return false;
28924 }
28925 return true;
28926 };
28927
28928 unsigned Opcode = ISD::ANY_EXTEND_VECTOR_INREG;
28929 SDValue N0 = SVN->getOperand(Num: 0);
28930 // Never create an illegal type. Only create unsupported operations if we
28931 // are pre-legalization.
28932 std::optional<EVT> OutVT = canCombineShuffleToExtendVectorInreg(
28933 Opcode, VT, Match: isAnyExtend, DAG, TLI, /*LegalTypes=*/true, LegalOperations);
28934 if (!OutVT)
28935 return SDValue();
28936 return DAG.getBitcast(VT, V: DAG.getNode(Opcode, DL: SDLoc(SVN), VT: *OutVT, Operand: N0));
28937}
28938
28939// Match shuffles that can be converted to zero_extend_vector_inreg.
28940// This is often generated during legalization.
28941// e.g. v4i32 <0,z,1,u> -> (v2i64 zero_extend_vector_inreg(v4i32 src))
28942static SDValue combineShuffleToZeroExtendVectorInReg(ShuffleVectorSDNode *SVN,
28943 SelectionDAG &DAG,
28944 const TargetLowering &TLI,
28945 bool LegalOperations) {
28946 bool LegalTypes = true;
28947 EVT VT = SVN->getValueType(ResNo: 0);
28948 assert(!VT.isScalableVector() && "Encountered scalable shuffle?");
28949 unsigned NumElts = VT.getVectorNumElements();
28950 unsigned EltSizeInBits = VT.getScalarSizeInBits();
28951
28952 // TODO: add support for big-endian when we have a test case.
28953 bool IsBigEndian = DAG.getDataLayout().isBigEndian();
28954 if (!VT.isInteger() || IsBigEndian)
28955 return SDValue();
28956
28957 SmallVector<int, 16> Mask(SVN->getMask());
28958 auto ForEachDecomposedIndice = [NumElts, &Mask](auto Fn) {
28959 for (int &Indice : Mask) {
28960 if (Indice < 0)
28961 continue;
28962 int OpIdx = (unsigned)Indice < NumElts ? 0 : 1;
28963 int OpEltIdx = (unsigned)Indice < NumElts ? Indice : Indice - NumElts;
28964 Fn(Indice, OpIdx, OpEltIdx);
28965 }
28966 };
28967
28968 // Which elements of which operand does this shuffle demand?
28969 std::array<APInt, 2> OpsDemandedElts;
28970 for (APInt &OpDemandedElts : OpsDemandedElts)
28971 OpDemandedElts = APInt::getZero(numBits: NumElts);
28972 ForEachDecomposedIndice(
28973 [&OpsDemandedElts](int &Indice, int OpIdx, int OpEltIdx) {
28974 OpsDemandedElts[OpIdx].setBit(OpEltIdx);
28975 });
28976
28977 // Element-wise(!), which of these demanded elements are know to be zero?
28978 std::array<APInt, 2> OpsKnownZeroElts;
28979 for (auto I : zip(t: SVN->ops(), u&: OpsDemandedElts, args&: OpsKnownZeroElts))
28980 std::get<2>(t&: I) =
28981 DAG.computeVectorKnownZeroElements(Op: std::get<0>(t&: I), DemandedElts: std::get<1>(t&: I));
28982
28983 // Manifest zeroable element knowledge in the shuffle mask.
28984 // NOTE: we don't have 'zeroable' sentinel value in generic DAG,
28985 // this is a local invention, but it won't leak into DAG.
28986 // FIXME: should we not manifest them, but just check when matching?
28987 bool HadZeroableElts = false;
28988 ForEachDecomposedIndice([&OpsKnownZeroElts, &HadZeroableElts](
28989 int &Indice, int OpIdx, int OpEltIdx) {
28990 if (OpsKnownZeroElts[OpIdx][OpEltIdx]) {
28991 Indice = -2; // Zeroable element.
28992 HadZeroableElts = true;
28993 }
28994 });
28995
28996 // Don't proceed unless we've refined at least one zeroable mask indice.
28997 // If we didn't, then we are still trying to match the same shuffle mask
28998 // we previously tried to match as ISD::ANY_EXTEND_VECTOR_INREG,
28999 // and evidently failed. Proceeding will lead to endless combine loops.
29000 if (!HadZeroableElts)
29001 return SDValue();
29002
29003 // The shuffle may be more fine-grained than we want. Widen elements first.
29004 // FIXME: should we do this before manifesting zeroable shuffle mask indices?
29005 SmallVector<int, 16> ScaledMask;
29006 getShuffleMaskWithWidestElts(Mask, ScaledMask);
29007 assert(Mask.size() >= ScaledMask.size() &&
29008 Mask.size() % ScaledMask.size() == 0 && "Unexpected mask widening.");
29009 int Prescale = Mask.size() / ScaledMask.size();
29010
29011 NumElts = ScaledMask.size();
29012 EltSizeInBits *= Prescale;
29013
29014 EVT PrescaledVT = EVT::getVectorVT(
29015 Context&: *DAG.getContext(), VT: EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: EltSizeInBits),
29016 NumElements: NumElts);
29017
29018 if (LegalTypes && !TLI.isTypeLegal(VT: PrescaledVT) && TLI.isTypeLegal(VT))
29019 return SDValue();
29020
29021 // For example,
29022 // shuffle<0,z,1,-1> == (v2i64 zero_extend_vector_inreg(v4i32))
29023 // But not shuffle<z,z,1,-1> and not shuffle<0,z,z,-1> ! (for same types)
29024 auto isZeroExtend = [NumElts, &ScaledMask](unsigned Scale) {
29025 assert(Scale >= 2 && Scale <= NumElts && NumElts % Scale == 0 &&
29026 "Unexpected mask scaling factor.");
29027 ArrayRef<int> Mask = ScaledMask;
29028 for (unsigned SrcElt = 0, NumSrcElts = NumElts / Scale;
29029 SrcElt != NumSrcElts; ++SrcElt) {
29030 // Analyze the shuffle mask in Scale-sized chunks.
29031 ArrayRef<int> MaskChunk = Mask.take_front(N: Scale);
29032 assert(MaskChunk.size() == Scale && "Unexpected mask size.");
29033 Mask = Mask.drop_front(N: MaskChunk.size());
29034 // The first indice in this chunk must be SrcElt, but not zero!
29035 // FIXME: undef should be fine, but that results in more-defined result.
29036 if (int FirstIndice = MaskChunk[0]; (unsigned)FirstIndice != SrcElt)
29037 return false;
29038 // The rest of the indices in this chunk must be zeros.
29039 // FIXME: undef should be fine, but that results in more-defined result.
29040 if (!all_of(Range: MaskChunk.drop_front(N: 1),
29041 P: [](int Indice) { return Indice == -2; }))
29042 return false;
29043 }
29044 assert(Mask.empty() && "Did not process the whole mask?");
29045 return true;
29046 };
29047
29048 unsigned Opcode = ISD::ZERO_EXTEND_VECTOR_INREG;
29049 for (bool Commuted : {false, true}) {
29050 SDValue Op = SVN->getOperand(Num: !Commuted ? 0 : 1);
29051 if (Commuted)
29052 ShuffleVectorSDNode::commuteMask(Mask: ScaledMask);
29053 std::optional<EVT> OutVT = canCombineShuffleToExtendVectorInreg(
29054 Opcode, VT: PrescaledVT, Match: isZeroExtend, DAG, TLI, LegalTypes,
29055 LegalOperations);
29056 if (OutVT)
29057 return DAG.getBitcast(VT, V: DAG.getNode(Opcode, DL: SDLoc(SVN), VT: *OutVT,
29058 Operand: DAG.getBitcast(VT: PrescaledVT, V: Op)));
29059 }
29060 return SDValue();
29061}
29062
29063// Detect 'truncate_vector_inreg' style shuffles that pack the lower parts of
29064// each source element of a large type into the lowest elements of a smaller
29065// destination type. This is often generated during legalization.
29066// If the source node itself was a '*_extend_vector_inreg' node then we should
29067// then be able to remove it.
29068static SDValue combineTruncationShuffle(ShuffleVectorSDNode *SVN,
29069 SelectionDAG &DAG) {
29070 EVT VT = SVN->getValueType(ResNo: 0);
29071 bool IsBigEndian = DAG.getDataLayout().isBigEndian();
29072
29073 // TODO Add support for big-endian when we have a test case.
29074 if (!VT.isInteger() || IsBigEndian)
29075 return SDValue();
29076
29077 SDValue N0 = peekThroughBitcasts(V: SVN->getOperand(Num: 0));
29078
29079 unsigned Opcode = N0.getOpcode();
29080 if (!ISD::isExtVecInRegOpcode(Opcode))
29081 return SDValue();
29082
29083 SDValue N00 = N0.getOperand(i: 0);
29084 ArrayRef<int> Mask = SVN->getMask();
29085 unsigned NumElts = VT.getVectorNumElements();
29086 unsigned EltSizeInBits = VT.getScalarSizeInBits();
29087 unsigned ExtSrcSizeInBits = N00.getScalarValueSizeInBits();
29088 unsigned ExtDstSizeInBits = N0.getScalarValueSizeInBits();
29089
29090 if (ExtDstSizeInBits % ExtSrcSizeInBits != 0)
29091 return SDValue();
29092 unsigned ExtScale = ExtDstSizeInBits / ExtSrcSizeInBits;
29093
29094 // (v4i32 truncate_vector_inreg(v2i64)) == shuffle<0,2-1,-1>
29095 // (v8i16 truncate_vector_inreg(v4i32)) == shuffle<0,2,4,6,-1,-1,-1,-1>
29096 // (v8i16 truncate_vector_inreg(v2i64)) == shuffle<0,4,-1,-1,-1,-1,-1,-1>
29097 auto isTruncate = [&Mask, &NumElts](unsigned Scale) {
29098 for (unsigned i = 0; i != NumElts; ++i) {
29099 if (Mask[i] < 0)
29100 continue;
29101 if ((i * Scale) < NumElts && Mask[i] == (int)(i * Scale))
29102 continue;
29103 return false;
29104 }
29105 return true;
29106 };
29107
29108 // At the moment we just handle the case where we've truncated back to the
29109 // same size as before the extension.
29110 // TODO: handle more extension/truncation cases as cases arise.
29111 if (EltSizeInBits != ExtSrcSizeInBits)
29112 return SDValue();
29113 if (VT.getSizeInBits() != N00.getValueSizeInBits())
29114 return SDValue();
29115
29116 // We can remove *extend_vector_inreg only if the truncation happens at
29117 // the same scale as the extension.
29118 if (isTruncate(ExtScale))
29119 return DAG.getBitcast(VT, V: N00);
29120
29121 return SDValue();
29122}
29123
29124// Combine shuffles of splat-shuffles of the form:
29125// shuffle (shuffle V, undef, splat-mask), undef, M
29126// If splat-mask contains undef elements, we need to be careful about
29127// introducing undef's in the folded mask which are not the result of composing
29128// the masks of the shuffles.
29129static SDValue combineShuffleOfSplatVal(ShuffleVectorSDNode *Shuf,
29130 SelectionDAG &DAG) {
29131 EVT VT = Shuf->getValueType(ResNo: 0);
29132 unsigned NumElts = VT.getVectorNumElements();
29133
29134 if (!Shuf->getOperand(Num: 1).isUndef())
29135 return SDValue();
29136
29137 // See if this unary non-splat shuffle actually *is* a splat shuffle,
29138 // in disguise, with all demanded elements being identical.
29139 // FIXME: this can be done per-operand.
29140 if (!Shuf->isSplat()) {
29141 APInt DemandedElts(NumElts, 0);
29142 for (int Idx : Shuf->getMask()) {
29143 if (Idx < 0)
29144 continue; // Ignore sentinel indices.
29145 assert((unsigned)Idx < NumElts && "Out-of-bounds shuffle indice?");
29146 DemandedElts.setBit(Idx);
29147 }
29148 assert(DemandedElts.popcount() > 1 && "Is a splat shuffle already?");
29149 APInt UndefElts;
29150 if (DAG.isSplatValue(V: Shuf->getOperand(Num: 0), DemandedElts, UndefElts)) {
29151 // Even if all demanded elements are splat, some of them could be undef.
29152 // Which lowest demanded element is *not* known-undef?
29153 std::optional<unsigned> MinNonUndefIdx;
29154 for (int Idx : Shuf->getMask()) {
29155 if (Idx < 0 || UndefElts[Idx])
29156 continue; // Ignore sentinel indices, and undef elements.
29157 MinNonUndefIdx = std::min<unsigned>(a: Idx, b: MinNonUndefIdx.value_or(u: ~0U));
29158 }
29159 if (!MinNonUndefIdx)
29160 return DAG.getUNDEF(VT); // All undef - result is undef.
29161 assert(*MinNonUndefIdx < NumElts && "Expected valid element index.");
29162 SmallVector<int, 8> SplatMask(Shuf->getMask());
29163 for (int &Idx : SplatMask) {
29164 if (Idx < 0)
29165 continue; // Passthrough sentinel indices.
29166 // Otherwise, just pick the lowest demanded non-undef element.
29167 // Or sentinel undef, if we know we'd pick a known-undef element.
29168 Idx = UndefElts[Idx] ? -1 : *MinNonUndefIdx;
29169 }
29170 assert(SplatMask != Shuf->getMask() && "Expected mask to change!");
29171 return DAG.getVectorShuffle(VT, dl: SDLoc(Shuf), N1: Shuf->getOperand(Num: 0),
29172 N2: Shuf->getOperand(Num: 1), Mask: SplatMask);
29173 }
29174 }
29175
29176 // If the inner operand is a known splat with no undefs, just return that directly.
29177 // TODO: Create DemandedElts mask from Shuf's mask.
29178 // TODO: Allow undef elements and merge with the shuffle code below.
29179 if (DAG.isSplatValue(V: Shuf->getOperand(Num: 0), /*AllowUndefs*/ false))
29180 return Shuf->getOperand(Num: 0);
29181
29182 auto *Splat = dyn_cast<ShuffleVectorSDNode>(Val: Shuf->getOperand(Num: 0));
29183 if (!Splat || !Splat->isSplat())
29184 return SDValue();
29185
29186 ArrayRef<int> ShufMask = Shuf->getMask();
29187 ArrayRef<int> SplatMask = Splat->getMask();
29188 assert(ShufMask.size() == SplatMask.size() && "Mask length mismatch");
29189
29190 // Prefer simplifying to the splat-shuffle, if possible. This is legal if
29191 // every undef mask element in the splat-shuffle has a corresponding undef
29192 // element in the user-shuffle's mask or if the composition of mask elements
29193 // would result in undef.
29194 // Examples for (shuffle (shuffle v, undef, SplatMask), undef, UserMask):
29195 // * UserMask=[0,2,u,u], SplatMask=[2,u,2,u] -> [2,2,u,u]
29196 // In this case it is not legal to simplify to the splat-shuffle because we
29197 // may be exposing the users of the shuffle an undef element at index 1
29198 // which was not there before the combine.
29199 // * UserMask=[0,u,2,u], SplatMask=[2,u,2,u] -> [2,u,2,u]
29200 // In this case the composition of masks yields SplatMask, so it's ok to
29201 // simplify to the splat-shuffle.
29202 // * UserMask=[3,u,2,u], SplatMask=[2,u,2,u] -> [u,u,2,u]
29203 // In this case the composed mask includes all undef elements of SplatMask
29204 // and in addition sets element zero to undef. It is safe to simplify to
29205 // the splat-shuffle.
29206 auto CanSimplifyToExistingSplat = [](ArrayRef<int> UserMask,
29207 ArrayRef<int> SplatMask) {
29208 for (unsigned i = 0, e = UserMask.size(); i != e; ++i)
29209 if (UserMask[i] != -1 && SplatMask[i] == -1 &&
29210 SplatMask[UserMask[i]] != -1)
29211 return false;
29212 return true;
29213 };
29214 if (CanSimplifyToExistingSplat(ShufMask, SplatMask))
29215 return Shuf->getOperand(Num: 0);
29216
29217 // Create a new shuffle with a mask that is composed of the two shuffles'
29218 // masks.
29219 SmallVector<int, 32> NewMask;
29220 for (int Idx : ShufMask)
29221 NewMask.push_back(Elt: Idx == -1 ? -1 : SplatMask[Idx]);
29222
29223 return DAG.getVectorShuffle(VT: Splat->getValueType(ResNo: 0), dl: SDLoc(Splat),
29224 N1: Splat->getOperand(Num: 0), N2: Splat->getOperand(Num: 1),
29225 Mask: NewMask);
29226}
29227
29228// Combine shuffles of bitcasts into a shuffle of the bitcast type, providing
29229// the mask can be treated as a larger type.
29230static SDValue combineShuffleOfBitcast(ShuffleVectorSDNode *SVN,
29231 SelectionDAG &DAG,
29232 const TargetLowering &TLI,
29233 bool LegalOperations) {
29234 SDValue Op0 = SVN->getOperand(Num: 0);
29235 SDValue Op1 = SVN->getOperand(Num: 1);
29236 EVT VT = SVN->getValueType(ResNo: 0);
29237 if (Op0.getOpcode() != ISD::BITCAST)
29238 return SDValue();
29239 EVT InVT = Op0.getOperand(i: 0).getValueType();
29240 if (!InVT.isVector() ||
29241 (!Op1.isUndef() && (Op1.getOpcode() != ISD::BITCAST ||
29242 Op1.getOperand(i: 0).getValueType() != InVT)))
29243 return SDValue();
29244 if (isAnyConstantBuildVector(V: Op0.getOperand(i: 0)) &&
29245 (Op1.isUndef() || isAnyConstantBuildVector(V: Op1.getOperand(i: 0))))
29246 return SDValue();
29247
29248 int VTLanes = VT.getVectorNumElements();
29249 int InLanes = InVT.getVectorNumElements();
29250 if (VTLanes <= InLanes || VTLanes % InLanes != 0 ||
29251 (LegalOperations &&
29252 !TLI.isOperationLegalOrCustom(Op: ISD::VECTOR_SHUFFLE, VT: InVT)))
29253 return SDValue();
29254 int Factor = VTLanes / InLanes;
29255
29256 // Check that each group of lanes in the mask are either undef or make a valid
29257 // mask for the wider lane type.
29258 ArrayRef<int> Mask = SVN->getMask();
29259 SmallVector<int> NewMask;
29260 if (!widenShuffleMaskElts(Scale: Factor, Mask, ScaledMask&: NewMask))
29261 return SDValue();
29262
29263 if (!TLI.isShuffleMaskLegal(NewMask, InVT))
29264 return SDValue();
29265
29266 // Create the new shuffle with the new mask and bitcast it back to the
29267 // original type.
29268 SDLoc DL(SVN);
29269 Op0 = Op0.getOperand(i: 0);
29270 Op1 = Op1.isUndef() ? DAG.getUNDEF(VT: InVT) : Op1.getOperand(i: 0);
29271 SDValue NewShuf = DAG.getVectorShuffle(VT: InVT, dl: DL, N1: Op0, N2: Op1, Mask: NewMask);
29272 return DAG.getBitcast(VT, V: NewShuf);
29273}
29274
29275/// Combine shuffle of shuffle of the form:
29276/// shuf (shuf X, undef, InnerMask), undef, OuterMask --> splat X
29277static SDValue formSplatFromShuffles(ShuffleVectorSDNode *OuterShuf,
29278 SelectionDAG &DAG) {
29279 if (!OuterShuf->getOperand(Num: 1).isUndef())
29280 return SDValue();
29281 auto *InnerShuf = dyn_cast<ShuffleVectorSDNode>(Val: OuterShuf->getOperand(Num: 0));
29282 if (!InnerShuf || !InnerShuf->getOperand(Num: 1).isUndef())
29283 return SDValue();
29284
29285 ArrayRef<int> OuterMask = OuterShuf->getMask();
29286 ArrayRef<int> InnerMask = InnerShuf->getMask();
29287 unsigned NumElts = OuterMask.size();
29288 assert(NumElts == InnerMask.size() && "Mask length mismatch");
29289 SmallVector<int, 32> CombinedMask(NumElts, -1);
29290 int SplatIndex = -1;
29291 for (unsigned i = 0; i != NumElts; ++i) {
29292 // Undef lanes remain undef.
29293 int OuterMaskElt = OuterMask[i];
29294 if (OuterMaskElt == -1)
29295 continue;
29296
29297 // Peek through the shuffle masks to get the underlying source element.
29298 int InnerMaskElt = InnerMask[OuterMaskElt];
29299 if (InnerMaskElt == -1)
29300 continue;
29301
29302 // Initialize the splatted element.
29303 if (SplatIndex == -1)
29304 SplatIndex = InnerMaskElt;
29305
29306 // Non-matching index - this is not a splat.
29307 if (SplatIndex != InnerMaskElt)
29308 return SDValue();
29309
29310 CombinedMask[i] = InnerMaskElt;
29311 }
29312 assert((all_of(CombinedMask, equal_to(-1)) ||
29313 getSplatIndex(CombinedMask) != -1) &&
29314 "Expected a splat mask");
29315
29316 // TODO: The transform may be a win even if the mask is not legal.
29317 EVT VT = OuterShuf->getValueType(ResNo: 0);
29318 assert(VT == InnerShuf->getValueType(0) && "Expected matching shuffle types");
29319 if (!DAG.getTargetLoweringInfo().isShuffleMaskLegal(CombinedMask, VT))
29320 return SDValue();
29321
29322 return DAG.getVectorShuffle(VT, dl: SDLoc(OuterShuf), N1: InnerShuf->getOperand(Num: 0),
29323 N2: InnerShuf->getOperand(Num: 1), Mask: CombinedMask);
29324}
29325
29326/// If the shuffle mask is taking exactly one element from the first vector
29327/// operand and passing through all other elements from the second vector
29328/// operand, return the index of the mask element that is choosing an element
29329/// from the first operand. Otherwise, return -1.
29330static int getShuffleMaskIndexOfOneElementFromOp0IntoOp1(ArrayRef<int> Mask) {
29331 int MaskSize = Mask.size();
29332 int EltFromOp0 = -1;
29333 // TODO: This does not match if there are undef elements in the shuffle mask.
29334 // Should we ignore undefs in the shuffle mask instead? The trade-off is
29335 // removing an instruction (a shuffle), but losing the knowledge that some
29336 // vector lanes are not needed.
29337 for (int i = 0; i != MaskSize; ++i) {
29338 if (Mask[i] >= 0 && Mask[i] < MaskSize) {
29339 // We're looking for a shuffle of exactly one element from operand 0.
29340 if (EltFromOp0 != -1)
29341 return -1;
29342 EltFromOp0 = i;
29343 } else if (Mask[i] != i + MaskSize) {
29344 // Nothing from operand 1 can change lanes.
29345 return -1;
29346 }
29347 }
29348 return EltFromOp0;
29349}
29350
29351/// If a shuffle inserts exactly one element from a source vector operand into
29352/// another vector operand and we can access the specified element as a scalar,
29353/// then we can eliminate the shuffle.
29354SDValue DAGCombiner::replaceShuffleOfInsert(ShuffleVectorSDNode *Shuf) {
29355 // First, check if we are taking one element of a vector and shuffling that
29356 // element into another vector.
29357 ArrayRef<int> Mask = Shuf->getMask();
29358 SmallVector<int, 16> CommutedMask(Mask);
29359 SDValue Op0 = Shuf->getOperand(Num: 0);
29360 SDValue Op1 = Shuf->getOperand(Num: 1);
29361 int ShufOp0Index = getShuffleMaskIndexOfOneElementFromOp0IntoOp1(Mask);
29362 if (ShufOp0Index == -1) {
29363 // Commute mask and check again.
29364 ShuffleVectorSDNode::commuteMask(Mask: CommutedMask);
29365 ShufOp0Index = getShuffleMaskIndexOfOneElementFromOp0IntoOp1(Mask: CommutedMask);
29366 if (ShufOp0Index == -1)
29367 return SDValue();
29368 // Commute operands to match the commuted shuffle mask.
29369 std::swap(a&: Op0, b&: Op1);
29370 Mask = CommutedMask;
29371 }
29372
29373 // The shuffle inserts exactly one element from operand 0 into operand 1.
29374 // Now see if we can access that element as a scalar via a real insert element
29375 // instruction.
29376 // TODO: We can try harder to locate the element as a scalar. Examples: it
29377 // could be an operand of BUILD_VECTOR, or a constant.
29378 assert(Mask[ShufOp0Index] >= 0 && Mask[ShufOp0Index] < (int)Mask.size() &&
29379 "Shuffle mask value must be from operand 0");
29380
29381 SDValue Elt;
29382 if (sd_match(N: Op0, P: m_InsertElt(Vec: m_Value(), Val: m_Value(N&: Elt),
29383 Idx: m_SpecificInt(V: Mask[ShufOp0Index])))) {
29384 // There's an existing insertelement with constant insertion index, so we
29385 // don't need to check the legality/profitability of a replacement operation
29386 // that differs at most in the constant value. The target should be able to
29387 // lower any of those in a similar way. If not, legalization will expand
29388 // this to a scalar-to-vector plus shuffle.
29389 //
29390 // Note that the shuffle may move the scalar from the position that the
29391 // insert element used. Therefore, our new insert element occurs at the
29392 // shuffle's mask index value, not the insert's index value.
29393 //
29394 // shuffle (insertelt v1, x, C), v2, mask --> insertelt v2, x, C'
29395 SDValue NewInsIndex = DAG.getVectorIdxConstant(Val: ShufOp0Index, DL: SDLoc(Shuf));
29396 return DAG.getNode(Opcode: ISD::INSERT_VECTOR_ELT, DL: SDLoc(Shuf), VT: Op0.getValueType(),
29397 N1: Op1, N2: Elt, N3: NewInsIndex);
29398 }
29399
29400 if (!hasOperation(Opcode: ISD::INSERT_VECTOR_ELT, VT: Op0.getValueType()))
29401 return SDValue();
29402
29403 if (sd_match(N: Op0, P: m_UnaryOp(Opc: ISD::SCALAR_TO_VECTOR, Op: m_Value(N&: Elt))) &&
29404 Mask[ShufOp0Index] == 0) {
29405 SDValue NewInsIndex = DAG.getVectorIdxConstant(Val: ShufOp0Index, DL: SDLoc(Shuf));
29406 return DAG.getNode(Opcode: ISD::INSERT_VECTOR_ELT, DL: SDLoc(Shuf), VT: Op0.getValueType(),
29407 N1: Op1, N2: Elt, N3: NewInsIndex);
29408 }
29409
29410 return SDValue();
29411}
29412
29413/// If we have a unary shuffle of a shuffle, see if it can be folded away
29414/// completely. This has the potential to lose undef knowledge because the first
29415/// shuffle may not have an undef mask element where the second one does. So
29416/// only call this after doing simplifications based on demanded elements.
29417static SDValue simplifyShuffleOfShuffle(ShuffleVectorSDNode *Shuf) {
29418 // shuf (shuf0 X, Y, Mask0), undef, Mask
29419 auto *Shuf0 = dyn_cast<ShuffleVectorSDNode>(Val: Shuf->getOperand(Num: 0));
29420 if (!Shuf0 || !Shuf->getOperand(Num: 1).isUndef())
29421 return SDValue();
29422
29423 ArrayRef<int> Mask = Shuf->getMask();
29424 ArrayRef<int> Mask0 = Shuf0->getMask();
29425 for (int i = 0, e = (int)Mask.size(); i != e; ++i) {
29426 // Ignore undef elements.
29427 if (Mask[i] == -1)
29428 continue;
29429 assert(Mask[i] >= 0 && Mask[i] < e && "Unexpected shuffle mask value");
29430
29431 // Is the element of the shuffle operand chosen by this shuffle the same as
29432 // the element chosen by the shuffle operand itself?
29433 if (Mask0[Mask[i]] != Mask0[i])
29434 return SDValue();
29435 }
29436 // Every element of this shuffle is identical to the result of the previous
29437 // shuffle, so we can replace this value.
29438 return Shuf->getOperand(Num: 0);
29439}
29440
29441SDValue DAGCombiner::visitVECTOR_SHUFFLE(SDNode *N) {
29442 EVT VT = N->getValueType(ResNo: 0);
29443 unsigned NumElts = VT.getVectorNumElements();
29444
29445 SDValue N0 = N->getOperand(Num: 0);
29446 SDValue N1 = N->getOperand(Num: 1);
29447
29448 assert(N0.getValueType() == VT && "Vector shuffle must be normalized in DAG");
29449
29450 // Canonicalize shuffle undef, undef -> undef
29451 if (N0.isUndef() && N1.isUndef())
29452 return DAG.getUNDEF(VT);
29453
29454 ShuffleVectorSDNode *SVN = cast<ShuffleVectorSDNode>(Val: N);
29455
29456 // Canonicalize shuffle v, v -> v, poison
29457 if (N0 == N1)
29458 return DAG.getVectorShuffle(VT, dl: SDLoc(N), N1: N0, N2: DAG.getPOISON(VT),
29459 Mask: createUnaryMask(Mask: SVN->getMask(), NumElts));
29460
29461 // Canonicalize shuffle undef, v -> v, undef. Commute the shuffle mask.
29462 if (N0.isUndef())
29463 return DAG.getCommutedVectorShuffle(SV: *SVN);
29464
29465 // Remove references to rhs if it is undef
29466 if (N1.isUndef()) {
29467 bool Changed = false;
29468 SmallVector<int, 8> NewMask;
29469 for (unsigned i = 0; i != NumElts; ++i) {
29470 int Idx = SVN->getMaskElt(Idx: i);
29471 if (Idx >= (int)NumElts) {
29472 Idx = -1;
29473 Changed = true;
29474 }
29475 NewMask.push_back(Elt: Idx);
29476 }
29477 if (Changed)
29478 return DAG.getVectorShuffle(VT, dl: SDLoc(N), N1: N0, N2: N1, Mask: NewMask);
29479 }
29480
29481 if (SDValue InsElt = replaceShuffleOfInsert(Shuf: SVN))
29482 return InsElt;
29483
29484 // A shuffle of a single vector that is a splatted value can always be folded.
29485 if (SDValue V = combineShuffleOfSplatVal(Shuf: SVN, DAG))
29486 return V;
29487
29488 if (SDValue V = formSplatFromShuffles(OuterShuf: SVN, DAG))
29489 return V;
29490
29491 // If it is a splat, check if the argument vector is another splat or a
29492 // build_vector.
29493 if (SVN->isSplat() && SVN->getSplatIndex() < (int)NumElts) {
29494 int SplatIndex = SVN->getSplatIndex();
29495 if (N0.hasOneUse() && TLI.isExtractVecEltCheap(VT, Index: SplatIndex) &&
29496 TLI.isBinOp(Opcode: N0.getOpcode()) && N0->getNumValues() == 1) {
29497 // splat (vector_bo L, R), Index -->
29498 // splat (scalar_bo (extelt L, Index), (extelt R, Index))
29499 SDValue L = N0.getOperand(i: 0), R = N0.getOperand(i: 1);
29500 SDLoc DL(N);
29501 EVT EltVT = VT.getScalarType();
29502 SDValue Index = DAG.getVectorIdxConstant(Val: SplatIndex, DL);
29503 SDValue ExtL = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: L, N2: Index);
29504 SDValue ExtR = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: R, N2: Index);
29505 SDValue NewBO =
29506 DAG.getNode(Opcode: N0.getOpcode(), DL, VT: EltVT, N1: ExtL, N2: ExtR, Flags: N0->getFlags());
29507 SDValue Insert = DAG.getNode(Opcode: ISD::SCALAR_TO_VECTOR, DL, VT, Operand: NewBO);
29508 SmallVector<int, 16> ZeroMask(VT.getVectorNumElements(), 0);
29509 return DAG.getVectorShuffle(VT, dl: DL, N1: Insert, N2: DAG.getPOISON(VT), Mask: ZeroMask);
29510 }
29511
29512 // splat(scalar_to_vector(x), 0) -> build_vector(x,...,x)
29513 // splat(insert_vector_elt(v, x, c), c) -> build_vector(x,...,x)
29514 if ((!LegalOperations || TLI.isOperationLegal(Op: ISD::BUILD_VECTOR, VT)) &&
29515 N0.hasOneUse()) {
29516 if (N0.getOpcode() == ISD::SCALAR_TO_VECTOR && SplatIndex == 0)
29517 return DAG.getSplatBuildVector(VT, DL: SDLoc(N), Op: N0.getOperand(i: 0));
29518
29519 if (N0.getOpcode() == ISD::INSERT_VECTOR_ELT)
29520 if (auto *Idx = dyn_cast<ConstantSDNode>(Val: N0.getOperand(i: 2)))
29521 if (Idx->getAPIntValue() == SplatIndex)
29522 return DAG.getSplatBuildVector(VT, DL: SDLoc(N), Op: N0.getOperand(i: 1));
29523
29524 // Look through a bitcast if LE and splatting lane 0, through to a
29525 // scalar_to_vector or a build_vector.
29526 if (N0.getOpcode() == ISD::BITCAST && N0.getOperand(i: 0).hasOneUse() &&
29527 SplatIndex == 0 && DAG.getDataLayout().isLittleEndian() &&
29528 (N0.getOperand(i: 0).getOpcode() == ISD::SCALAR_TO_VECTOR ||
29529 N0.getOperand(i: 0).getOpcode() == ISD::BUILD_VECTOR)) {
29530 EVT N00VT = N0.getOperand(i: 0).getValueType();
29531 if (VT.getScalarSizeInBits() <= N00VT.getScalarSizeInBits() &&
29532 VT.isInteger() && N00VT.isInteger()) {
29533 EVT InVT =
29534 TLI.getTypeToTransformTo(Context&: *DAG.getContext(), VT: VT.getScalarType());
29535 SDValue Op = DAG.getZExtOrTrunc(Op: N0.getOperand(i: 0).getOperand(i: 0),
29536 DL: SDLoc(N), VT: InVT);
29537 return DAG.getSplatBuildVector(VT, DL: SDLoc(N), Op);
29538 }
29539 }
29540 }
29541
29542 // If this is a bit convert that changes the element type of the vector but
29543 // not the number of vector elements, look through it. Be careful not to
29544 // look though conversions that change things like v4f32 to v2f64.
29545 SDNode *V = N0.getNode();
29546 if (V->getOpcode() == ISD::BITCAST) {
29547 SDValue ConvInput = V->getOperand(Num: 0);
29548 if (ConvInput.getValueType().isVector() &&
29549 ConvInput.getValueType().getVectorNumElements() == NumElts)
29550 V = ConvInput.getNode();
29551 }
29552
29553 if (V->getOpcode() == ISD::BUILD_VECTOR) {
29554 assert(V->getNumOperands() == NumElts &&
29555 "BUILD_VECTOR has wrong number of operands");
29556 SDValue Base;
29557 bool AllSame = true;
29558 for (unsigned i = 0; i != NumElts; ++i) {
29559 if (!V->getOperand(Num: i).isUndef()) {
29560 Base = V->getOperand(Num: i);
29561 break;
29562 }
29563 }
29564 // Splat of <u, u, u, u>, return <u, u, u, u>
29565 if (!Base.getNode())
29566 return N0;
29567 for (unsigned i = 0; i != NumElts; ++i) {
29568 if (V->getOperand(Num: i) != Base) {
29569 AllSame = false;
29570 break;
29571 }
29572 }
29573 // Splat of <x, x, x, x>, return <x, x, x, x>
29574 if (AllSame)
29575 return N0;
29576
29577 // Canonicalize any other splat as a build_vector, but avoid defining any
29578 // undefined elements in the mask.
29579 SDValue Splatted = V->getOperand(Num: SplatIndex);
29580 SmallVector<SDValue, 8> Ops(NumElts, Splatted);
29581 EVT EltVT = Splatted.getValueType();
29582
29583 for (unsigned i = 0; i != NumElts; ++i) {
29584 if (SVN->getMaskElt(Idx: i) < 0)
29585 Ops[i] = DAG.getPOISON(VT: EltVT);
29586 }
29587
29588 SDValue NewBV = DAG.getBuildVector(VT: V->getValueType(ResNo: 0), DL: SDLoc(N), Ops);
29589
29590 // We may have jumped through bitcasts, so the type of the
29591 // BUILD_VECTOR may not match the type of the shuffle.
29592 if (V->getValueType(ResNo: 0) != VT)
29593 NewBV = DAG.getBitcast(VT, V: NewBV);
29594 return NewBV;
29595 }
29596 }
29597
29598 // Simplify source operands based on shuffle mask.
29599 if (SimplifyDemandedVectorElts(Op: SDValue(N, 0)))
29600 return SDValue(N, 0);
29601
29602 // This is intentionally placed after demanded elements simplification because
29603 // it could eliminate knowledge of undef elements created by this shuffle.
29604 if (SDValue ShufOp = simplifyShuffleOfShuffle(Shuf: SVN))
29605 return ShufOp;
29606
29607 // Match shuffles that can be converted to any_vector_extend_in_reg.
29608 if (SDValue V =
29609 combineShuffleToAnyExtendVectorInreg(SVN, DAG, TLI, LegalOperations))
29610 return V;
29611
29612 // Combine "truncate_vector_in_reg" style shuffles.
29613 if (SDValue V = combineTruncationShuffle(SVN, DAG))
29614 return V;
29615
29616 if (N0.getOpcode() == ISD::CONCAT_VECTORS &&
29617 Level < AfterLegalizeVectorOps &&
29618 (N1.isUndef() ||
29619 (N1.getOpcode() == ISD::CONCAT_VECTORS &&
29620 N0.getOperand(i: 0).getValueType() == N1.getOperand(i: 0).getValueType()))) {
29621 if (SDValue V = partitionShuffleOfConcats(N, DAG))
29622 return V;
29623 }
29624
29625 // A shuffle of a concat of the same narrow vector can be reduced to use
29626 // only low-half elements of a concat with undef:
29627 // shuf (concat X, X), undef, Mask --> shuf (concat X, undef), undef, Mask'
29628 if (N0.getOpcode() == ISD::CONCAT_VECTORS && N1.isUndef() &&
29629 N0.getNumOperands() == 2 &&
29630 N0.getOperand(i: 0) == N0.getOperand(i: 1)) {
29631 int HalfNumElts = (int)NumElts / 2;
29632 SmallVector<int, 8> NewMask;
29633 for (unsigned i = 0; i != NumElts; ++i) {
29634 int Idx = SVN->getMaskElt(Idx: i);
29635 if (Idx >= HalfNumElts) {
29636 assert(Idx < (int)NumElts && "Shuffle mask chooses undef op");
29637 Idx -= HalfNumElts;
29638 }
29639 NewMask.push_back(Elt: Idx);
29640 }
29641 if (TLI.isShuffleMaskLegal(NewMask, VT)) {
29642 SDValue UndefVec = DAG.getPOISON(VT: N0.getOperand(i: 0).getValueType());
29643 SDValue NewCat = DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT,
29644 N1: N0.getOperand(i: 0), N2: UndefVec);
29645 return DAG.getVectorShuffle(VT, dl: SDLoc(N), N1: NewCat, N2: N1, Mask: NewMask);
29646 }
29647 }
29648
29649 // See if we can replace a shuffle with an insert_subvector.
29650 // e.g. v2i32 into v8i32:
29651 // shuffle(lhs,concat(rhs0,rhs1,rhs2,rhs3),0,1,2,3,10,11,6,7).
29652 // --> insert_subvector(lhs,rhs1,4).
29653 if (Level < AfterLegalizeVectorOps && TLI.isTypeLegal(VT) &&
29654 TLI.isOperationLegalOrCustom(Op: ISD::INSERT_SUBVECTOR, VT)) {
29655 auto ShuffleToInsert = [&](SDValue LHS, SDValue RHS, ArrayRef<int> Mask) {
29656 // Ensure RHS subvectors are legal.
29657 assert(RHS.getOpcode() == ISD::CONCAT_VECTORS && "Can't find subvectors");
29658 EVT SubVT = RHS.getOperand(i: 0).getValueType();
29659 int NumSubVecs = RHS.getNumOperands();
29660 int NumSubElts = SubVT.getVectorNumElements();
29661 assert((NumElts % NumSubElts) == 0 && "Subvector mismatch");
29662 if (!TLI.isTypeLegal(VT: SubVT))
29663 return SDValue();
29664
29665 // Don't bother if we have an unary shuffle (matches undef + LHS elts).
29666 if (all_of(Range&: Mask, P: [NumElts](int M) { return M < (int)NumElts; }))
29667 return SDValue();
29668
29669 // Search [NumSubElts] spans for RHS sequence.
29670 // TODO: Can we avoid nested loops to increase performance?
29671 SmallVector<int> InsertionMask(NumElts);
29672 for (int SubVec = 0; SubVec != NumSubVecs; ++SubVec) {
29673 for (int SubIdx = 0; SubIdx != (int)NumElts; SubIdx += NumSubElts) {
29674 // Reset mask to identity.
29675 std::iota(first: InsertionMask.begin(), last: InsertionMask.end(), value: 0);
29676
29677 // Add subvector insertion.
29678 std::iota(first: InsertionMask.begin() + SubIdx,
29679 last: InsertionMask.begin() + SubIdx + NumSubElts,
29680 value: NumElts + (SubVec * NumSubElts));
29681
29682 // See if the shuffle mask matches the reference insertion mask.
29683 bool MatchingShuffle = true;
29684 for (int i = 0; i != (int)NumElts; ++i) {
29685 int ExpectIdx = InsertionMask[i];
29686 int ActualIdx = Mask[i];
29687 if (0 <= ActualIdx && ExpectIdx != ActualIdx) {
29688 MatchingShuffle = false;
29689 break;
29690 }
29691 }
29692
29693 if (MatchingShuffle)
29694 return DAG.getInsertSubvector(DL: SDLoc(N), Vec: LHS, SubVec: RHS.getOperand(i: SubVec),
29695 Idx: SubIdx);
29696 }
29697 }
29698 return SDValue();
29699 };
29700 ArrayRef<int> Mask = SVN->getMask();
29701 if (N1.getOpcode() == ISD::CONCAT_VECTORS)
29702 if (SDValue InsertN1 = ShuffleToInsert(N0, N1, Mask))
29703 return InsertN1;
29704 if (N0.getOpcode() == ISD::CONCAT_VECTORS) {
29705 SmallVector<int> CommuteMask(Mask);
29706 ShuffleVectorSDNode::commuteMask(Mask: CommuteMask);
29707 if (SDValue InsertN0 = ShuffleToInsert(N1, N0, CommuteMask))
29708 return InsertN0;
29709 }
29710 }
29711
29712 // If we're not performing a select/blend shuffle, see if we can convert the
29713 // shuffle into a AND node, with all the out-of-lane elements are known zero.
29714 if (Level < AfterLegalizeDAG && TLI.isTypeLegal(VT)) {
29715 bool IsInLaneMask = true;
29716 ArrayRef<int> Mask = SVN->getMask();
29717 SmallVector<int, 16> ClearMask(NumElts, -1);
29718 APInt DemandedLHS = APInt::getZero(numBits: NumElts);
29719 APInt DemandedRHS = APInt::getZero(numBits: NumElts);
29720 for (int I = 0; I != (int)NumElts; ++I) {
29721 int M = Mask[I];
29722 if (M < 0)
29723 continue;
29724 ClearMask[I] = M == I ? I : (I + NumElts);
29725 IsInLaneMask &= (M == I) || (M == (int)(I + NumElts));
29726 if (M != I) {
29727 APInt &Demanded = M < (int)NumElts ? DemandedLHS : DemandedRHS;
29728 Demanded.setBit(M % NumElts);
29729 }
29730 }
29731 // TODO: Should we try to mask with N1 as well?
29732 if (!IsInLaneMask && (!DemandedLHS.isZero() || !DemandedRHS.isZero()) &&
29733 (DemandedLHS.isZero() || DAG.MaskedVectorIsZero(Op: N0, DemandedElts: DemandedLHS)) &&
29734 (DemandedRHS.isZero() || DAG.MaskedVectorIsZero(Op: N1, DemandedElts: DemandedRHS))) {
29735 SDLoc DL(N);
29736 EVT IntVT = VT.changeVectorElementTypeToInteger();
29737 EVT IntSVT = VT.getVectorElementType().changeTypeToInteger();
29738 // Transform the type to a legal type so that the buildvector constant
29739 // elements are not illegal. Make sure that the result is larger than the
29740 // original type, incase the value is split into two (eg i64->i32).
29741 if (!TLI.isTypeLegal(VT: IntSVT) && LegalTypes)
29742 IntSVT = TLI.getTypeToTransformTo(Context&: *DAG.getContext(), VT: IntSVT);
29743 if (IntSVT.getSizeInBits() >= IntVT.getScalarSizeInBits()) {
29744 SDValue ZeroElt = DAG.getConstant(Val: 0, DL, VT: IntSVT);
29745 SDValue AllOnesElt = DAG.getAllOnesConstant(DL, VT: IntSVT);
29746 SmallVector<SDValue, 16> AndMask(NumElts, DAG.getPOISON(VT: IntSVT));
29747 for (int I = 0; I != (int)NumElts; ++I)
29748 if (0 <= Mask[I])
29749 AndMask[I] = Mask[I] == I ? AllOnesElt : ZeroElt;
29750
29751 // See if a clear mask is legal instead of going via
29752 // XformToShuffleWithZero which loses UNDEF mask elements.
29753 if (TLI.isVectorClearMaskLegal(ClearMask, IntVT))
29754 return DAG.getBitcast(
29755 VT, V: DAG.getVectorShuffle(VT: IntVT, dl: DL, N1: DAG.getBitcast(VT: IntVT, V: N0),
29756 N2: DAG.getConstant(Val: 0, DL, VT: IntVT), Mask: ClearMask));
29757
29758 if (TLI.isOperationLegalOrCustom(Op: ISD::AND, VT: IntVT))
29759 return DAG.getBitcast(
29760 VT, V: DAG.getNode(Opcode: ISD::AND, DL, VT: IntVT, N1: DAG.getBitcast(VT: IntVT, V: N0),
29761 N2: DAG.getBuildVector(VT: IntVT, DL, Ops: AndMask)));
29762 }
29763 }
29764 }
29765
29766 // Attempt to combine a shuffle of 2 inputs of 'scalar sources' -
29767 // BUILD_VECTOR or SCALAR_TO_VECTOR into a single BUILD_VECTOR.
29768 if (Level < AfterLegalizeDAG && TLI.isTypeLegal(VT))
29769 if (SDValue Res = combineShuffleOfScalars(SVN, DAG, TLI))
29770 return Res;
29771
29772 // If this shuffle only has a single input that is a bitcasted shuffle,
29773 // attempt to merge the 2 shuffles and suitably bitcast the inputs/output
29774 // back to their original types.
29775 if (N0.getOpcode() == ISD::BITCAST && N0.hasOneUse() &&
29776 N1.isUndef() && Level < AfterLegalizeVectorOps &&
29777 TLI.isTypeLegal(VT)) {
29778
29779 SDValue BC0 = peekThroughOneUseBitcasts(V: N0);
29780 if (BC0.getOpcode() == ISD::VECTOR_SHUFFLE && BC0.hasOneUse()) {
29781 EVT SVT = VT.getScalarType();
29782 EVT InnerVT = BC0->getValueType(ResNo: 0);
29783 EVT InnerSVT = InnerVT.getScalarType();
29784
29785 // Determine which shuffle works with the smaller scalar type.
29786 EVT ScaleVT = SVT.bitsLT(VT: InnerSVT) ? VT : InnerVT;
29787 EVT ScaleSVT = ScaleVT.getScalarType();
29788
29789 if (TLI.isTypeLegal(VT: ScaleVT) &&
29790 0 == (InnerSVT.getSizeInBits() % ScaleSVT.getSizeInBits()) &&
29791 0 == (SVT.getSizeInBits() % ScaleSVT.getSizeInBits())) {
29792 int InnerScale = InnerSVT.getSizeInBits() / ScaleSVT.getSizeInBits();
29793 int OuterScale = SVT.getSizeInBits() / ScaleSVT.getSizeInBits();
29794
29795 // Scale the shuffle masks to the smaller scalar type.
29796 ShuffleVectorSDNode *InnerSVN = cast<ShuffleVectorSDNode>(Val&: BC0);
29797 SmallVector<int, 8> InnerMask;
29798 SmallVector<int, 8> OuterMask;
29799 narrowShuffleMaskElts(Scale: InnerScale, Mask: InnerSVN->getMask(), ScaledMask&: InnerMask);
29800 narrowShuffleMaskElts(Scale: OuterScale, Mask: SVN->getMask(), ScaledMask&: OuterMask);
29801
29802 // Merge the shuffle masks.
29803 SmallVector<int, 8> NewMask;
29804 for (int M : OuterMask)
29805 NewMask.push_back(Elt: M < 0 ? -1 : InnerMask[M]);
29806
29807 // Test for shuffle mask legality over both commutations.
29808 SDValue SV0 = BC0->getOperand(Num: 0);
29809 SDValue SV1 = BC0->getOperand(Num: 1);
29810 bool LegalMask = TLI.isShuffleMaskLegal(NewMask, ScaleVT);
29811 if (!LegalMask) {
29812 std::swap(a&: SV0, b&: SV1);
29813 ShuffleVectorSDNode::commuteMask(Mask: NewMask);
29814 LegalMask = TLI.isShuffleMaskLegal(NewMask, ScaleVT);
29815 }
29816
29817 if (LegalMask) {
29818 SV0 = DAG.getBitcast(VT: ScaleVT, V: SV0);
29819 SV1 = DAG.getBitcast(VT: ScaleVT, V: SV1);
29820 return DAG.getBitcast(
29821 VT, V: DAG.getVectorShuffle(VT: ScaleVT, dl: SDLoc(N), N1: SV0, N2: SV1, Mask: NewMask));
29822 }
29823 }
29824 }
29825 }
29826
29827 // Match shuffles of bitcasts, so long as the mask can be treated as the
29828 // larger type.
29829 if (SDValue V = combineShuffleOfBitcast(SVN, DAG, TLI, LegalOperations))
29830 return V;
29831
29832 // Compute the combined shuffle mask for a shuffle with SV0 as the first
29833 // operand, and SV1 as the second operand.
29834 // i.e. Merge SVN(OtherSVN, N1) -> shuffle(SV0, SV1, Mask) iff Commute = false
29835 // Merge SVN(N1, OtherSVN) -> shuffle(SV0, SV1, Mask') iff Commute = true
29836 auto MergeInnerShuffle =
29837 [NumElts, &VT](bool Commute, ShuffleVectorSDNode *SVN,
29838 ShuffleVectorSDNode *OtherSVN, SDValue N1,
29839 const TargetLowering &TLI, SDValue &SV0, SDValue &SV1,
29840 SmallVectorImpl<int> &Mask) -> bool {
29841 // Don't try to fold splats; they're likely to simplify somehow, or they
29842 // might be free.
29843 if (OtherSVN->isSplat())
29844 return false;
29845
29846 SV0 = SV1 = SDValue();
29847 Mask.clear();
29848
29849 for (unsigned i = 0; i != NumElts; ++i) {
29850 int Idx = SVN->getMaskElt(Idx: i);
29851 if (Idx < 0) {
29852 // Propagate Undef.
29853 Mask.push_back(Elt: Idx);
29854 continue;
29855 }
29856
29857 if (Commute)
29858 Idx = (Idx < (int)NumElts) ? (Idx + NumElts) : (Idx - NumElts);
29859
29860 SDValue CurrentVec;
29861 if (Idx < (int)NumElts) {
29862 // This shuffle index refers to the inner shuffle N0. Lookup the inner
29863 // shuffle mask to identify which vector is actually referenced.
29864 Idx = OtherSVN->getMaskElt(Idx);
29865 if (Idx < 0) {
29866 // Propagate Undef.
29867 Mask.push_back(Elt: Idx);
29868 continue;
29869 }
29870 CurrentVec = (Idx < (int)NumElts) ? OtherSVN->getOperand(Num: 0)
29871 : OtherSVN->getOperand(Num: 1);
29872 } else {
29873 // This shuffle index references an element within N1.
29874 CurrentVec = N1;
29875 }
29876
29877 // Simple case where 'CurrentVec' is UNDEF.
29878 if (CurrentVec.isUndef()) {
29879 Mask.push_back(Elt: -1);
29880 continue;
29881 }
29882
29883 // Canonicalize the shuffle index. We don't know yet if CurrentVec
29884 // will be the first or second operand of the combined shuffle.
29885 Idx = Idx % NumElts;
29886 if (!SV0.getNode() || SV0 == CurrentVec) {
29887 // Ok. CurrentVec is the left hand side.
29888 // Update the mask accordingly.
29889 SV0 = CurrentVec;
29890 Mask.push_back(Elt: Idx);
29891 continue;
29892 }
29893 if (!SV1.getNode() || SV1 == CurrentVec) {
29894 // Ok. CurrentVec is the right hand side.
29895 // Update the mask accordingly.
29896 SV1 = CurrentVec;
29897 Mask.push_back(Elt: Idx + NumElts);
29898 continue;
29899 }
29900
29901 // Last chance - see if the vector is another shuffle and if it
29902 // uses one of the existing candidate shuffle ops.
29903 if (auto *CurrentSVN = dyn_cast<ShuffleVectorSDNode>(Val&: CurrentVec)) {
29904 int InnerIdx = CurrentSVN->getMaskElt(Idx);
29905 if (InnerIdx < 0) {
29906 Mask.push_back(Elt: -1);
29907 continue;
29908 }
29909 SDValue InnerVec = (InnerIdx < (int)NumElts)
29910 ? CurrentSVN->getOperand(Num: 0)
29911 : CurrentSVN->getOperand(Num: 1);
29912 if (InnerVec.isUndef()) {
29913 Mask.push_back(Elt: -1);
29914 continue;
29915 }
29916 InnerIdx %= NumElts;
29917 if (InnerVec == SV0) {
29918 Mask.push_back(Elt: InnerIdx);
29919 continue;
29920 }
29921 if (InnerVec == SV1) {
29922 Mask.push_back(Elt: InnerIdx + NumElts);
29923 continue;
29924 }
29925 }
29926
29927 // Bail out if we cannot convert the shuffle pair into a single shuffle.
29928 return false;
29929 }
29930
29931 if (llvm::all_of(Range&: Mask, P: [](int M) { return M < 0; }))
29932 return true;
29933
29934 // Avoid introducing shuffles with illegal mask.
29935 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(A, B, M2)
29936 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(A, C, M2)
29937 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(B, C, M2)
29938 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(B, A, M2)
29939 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(C, A, M2)
29940 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(C, B, M2)
29941 if (TLI.isShuffleMaskLegal(Mask, VT))
29942 return true;
29943
29944 std::swap(a&: SV0, b&: SV1);
29945 ShuffleVectorSDNode::commuteMask(Mask);
29946 return TLI.isShuffleMaskLegal(Mask, VT);
29947 };
29948
29949 if (Level < AfterLegalizeDAG && TLI.isTypeLegal(VT)) {
29950 // Canonicalize shuffles according to rules:
29951 // shuffle(A, shuffle(A, B)) -> shuffle(shuffle(A,B), A)
29952 // shuffle(B, shuffle(A, B)) -> shuffle(shuffle(A,B), B)
29953 // shuffle(B, shuffle(A, Undef)) -> shuffle(shuffle(A, Undef), B)
29954 if (N1.getOpcode() == ISD::VECTOR_SHUFFLE &&
29955 N0.getOpcode() != ISD::VECTOR_SHUFFLE) {
29956 // The incoming shuffle must be of the same type as the result of the
29957 // current shuffle.
29958 assert(N1->getOperand(0).getValueType() == VT &&
29959 "Shuffle types don't match");
29960
29961 SDValue SV0 = N1->getOperand(Num: 0);
29962 SDValue SV1 = N1->getOperand(Num: 1);
29963 bool HasSameOp0 = N0 == SV0;
29964 bool IsSV1Undef = SV1.isUndef();
29965 if (HasSameOp0 || IsSV1Undef || N0 == SV1)
29966 // Commute the operands of this shuffle so merging below will trigger.
29967 return DAG.getCommutedVectorShuffle(SV: *SVN);
29968 }
29969
29970 // Canonicalize splat shuffles to the RHS to improve merging below.
29971 // shuffle(splat(A,u), shuffle(C,D)) -> shuffle'(shuffle(C,D), splat(A,u))
29972 if (N0.getOpcode() == ISD::VECTOR_SHUFFLE &&
29973 N1.getOpcode() == ISD::VECTOR_SHUFFLE &&
29974 cast<ShuffleVectorSDNode>(Val&: N0)->isSplat() &&
29975 !cast<ShuffleVectorSDNode>(Val&: N1)->isSplat()) {
29976 return DAG.getCommutedVectorShuffle(SV: *SVN);
29977 }
29978
29979 // Try to fold according to rules:
29980 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(A, B, M2)
29981 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(A, C, M2)
29982 // shuffle(shuffle(A, B, M0), C, M1) -> shuffle(B, C, M2)
29983 // Don't try to fold shuffles with illegal type.
29984 // Only fold if this shuffle is the only user of the other shuffle.
29985 // Try matching shuffle(C,shuffle(A,B)) commutted patterns as well.
29986 for (int i = 0; i != 2; ++i) {
29987 if (N->getOperand(Num: i).getOpcode() == ISD::VECTOR_SHUFFLE &&
29988 N->isOnlyUserOf(N: N->getOperand(Num: i).getNode())) {
29989 // The incoming shuffle must be of the same type as the result of the
29990 // current shuffle.
29991 auto *OtherSV = cast<ShuffleVectorSDNode>(Val: N->getOperand(Num: i));
29992 assert(OtherSV->getOperand(0).getValueType() == VT &&
29993 "Shuffle types don't match");
29994
29995 SDValue SV0, SV1;
29996 SmallVector<int, 4> Mask;
29997 if (MergeInnerShuffle(i != 0, SVN, OtherSV, N->getOperand(Num: 1 - i), TLI,
29998 SV0, SV1, Mask)) {
29999 // Check if all indices in Mask are poison. In case, propagate poison.
30000 if (llvm::all_of(Range&: Mask, P: [](int M) { return M < 0; }))
30001 return DAG.getPOISON(VT);
30002
30003 return DAG.getVectorShuffle(VT, dl: SDLoc(N),
30004 N1: SV0 ? SV0 : DAG.getPOISON(VT),
30005 N2: SV1 ? SV1 : DAG.getPOISON(VT), Mask);
30006 }
30007 }
30008 }
30009
30010 // Merge shuffles through binops if we are able to merge it with at least
30011 // one other shuffles.
30012 // shuffle(bop(shuffle(x,y),shuffle(z,w)),undef)
30013 // shuffle(bop(shuffle(x,y),shuffle(z,w)),bop(shuffle(a,b),shuffle(c,d)))
30014 unsigned SrcOpcode = N0.getOpcode();
30015 if (TLI.isBinOp(Opcode: SrcOpcode) && N->isOnlyUserOf(N: N0.getNode()) &&
30016 (N1.isUndef() ||
30017 (SrcOpcode == N1.getOpcode() && N->isOnlyUserOf(N: N1.getNode()) &&
30018 N0.getResNo() == N1.getResNo()))) {
30019 // Get binop source ops, or just pass on the undef.
30020 SDValue Op00 = N0.getOperand(i: 0);
30021 SDValue Op01 = N0.getOperand(i: 1);
30022 SDValue Op10 = N1.isUndef() ? N1 : N1.getOperand(i: 0);
30023 SDValue Op11 = N1.isUndef() ? N1 : N1.getOperand(i: 1);
30024 // TODO: We might be able to relax the VT check but we don't currently
30025 // have any isBinOp() that has different result/ops VTs so play safe until
30026 // we have test coverage.
30027 if (Op00.getValueType() == VT && Op10.getValueType() == VT &&
30028 Op01.getValueType() == VT && Op11.getValueType() == VT &&
30029 (Op00.getOpcode() == ISD::VECTOR_SHUFFLE ||
30030 Op10.getOpcode() == ISD::VECTOR_SHUFFLE ||
30031 Op01.getOpcode() == ISD::VECTOR_SHUFFLE ||
30032 Op11.getOpcode() == ISD::VECTOR_SHUFFLE)) {
30033 auto CanMergeInnerShuffle = [&](SDValue &SV0, SDValue &SV1,
30034 SmallVectorImpl<int> &Mask, bool LeftOp,
30035 bool Commute) {
30036 SDValue InnerN = Commute ? N1 : N0;
30037 SDValue Op0 = LeftOp ? Op00 : Op01;
30038 SDValue Op1 = LeftOp ? Op10 : Op11;
30039 if (Commute)
30040 std::swap(a&: Op0, b&: Op1);
30041 // Only accept the merged shuffle if we don't introduce undef elements,
30042 // or the inner shuffle already contained undef elements.
30043 auto *SVN0 = dyn_cast<ShuffleVectorSDNode>(Val&: Op0);
30044 return SVN0 && InnerN->isOnlyUserOf(N: SVN0) &&
30045 MergeInnerShuffle(Commute, SVN, SVN0, Op1, TLI, SV0, SV1,
30046 Mask) &&
30047 (llvm::any_of(Range: SVN0->getMask(), P: [](int M) { return M < 0; }) ||
30048 llvm::none_of(Range&: Mask, P: [](int M) { return M < 0; }));
30049 };
30050
30051 // Ensure we don't increase the number of shuffles - we must merge a
30052 // shuffle from at least one of the LHS and RHS ops.
30053 bool MergedLeft = false;
30054 SDValue LeftSV0, LeftSV1;
30055 SmallVector<int, 4> LeftMask;
30056 if (CanMergeInnerShuffle(LeftSV0, LeftSV1, LeftMask, true, false) ||
30057 CanMergeInnerShuffle(LeftSV0, LeftSV1, LeftMask, true, true)) {
30058 MergedLeft = true;
30059 } else {
30060 LeftMask.assign(in_start: SVN->getMask().begin(), in_end: SVN->getMask().end());
30061 LeftSV0 = Op00, LeftSV1 = Op10;
30062 }
30063
30064 bool MergedRight = false;
30065 SDValue RightSV0, RightSV1;
30066 SmallVector<int, 4> RightMask;
30067 if (CanMergeInnerShuffle(RightSV0, RightSV1, RightMask, false, false) ||
30068 CanMergeInnerShuffle(RightSV0, RightSV1, RightMask, false, true)) {
30069 MergedRight = true;
30070 } else {
30071 RightMask.assign(in_start: SVN->getMask().begin(), in_end: SVN->getMask().end());
30072 RightSV0 = Op01, RightSV1 = Op11;
30073 }
30074
30075 if (MergedLeft || MergedRight) {
30076 SDLoc DL(N);
30077 SDValue LHS = DAG.getVectorShuffle(
30078 VT, dl: DL, N1: LeftSV0 ? LeftSV0 : DAG.getPOISON(VT),
30079 N2: LeftSV1 ? LeftSV1 : DAG.getPOISON(VT), Mask: LeftMask);
30080 SDValue RHS = DAG.getVectorShuffle(
30081 VT, dl: DL, N1: RightSV0 ? RightSV0 : DAG.getPOISON(VT),
30082 N2: RightSV1 ? RightSV1 : DAG.getPOISON(VT), Mask: RightMask);
30083 return DAG.getNode(Opcode: SrcOpcode, DL, VTList: N0->getVTList(), N1: LHS, N2: RHS)
30084 .getValue(R: N0.getResNo());
30085 }
30086 }
30087 }
30088 }
30089
30090 if (SDValue V = foldShuffleOfConcatUndefs(Shuf: SVN, DAG))
30091 return V;
30092
30093 // Match shuffles that can be converted to ISD::ZERO_EXTEND_VECTOR_INREG.
30094 // Perform this really late, because it could eliminate knowledge
30095 // of undef elements created by this shuffle.
30096 if (Level < AfterLegalizeTypes)
30097 if (SDValue V = combineShuffleToZeroExtendVectorInReg(SVN, DAG, TLI,
30098 LegalOperations))
30099 return V;
30100
30101 return SDValue();
30102}
30103
30104SDValue DAGCombiner::visitSCALAR_TO_VECTOR(SDNode *N) {
30105 EVT VT = N->getValueType(ResNo: 0);
30106 if (!VT.isFixedLengthVector())
30107 return SDValue();
30108
30109 // Try to convert a scalar binop with an extracted vector element to a vector
30110 // binop. This is intended to reduce potentially expensive register moves.
30111 // TODO: Check if both operands are extracted.
30112 // TODO: How to prefer scalar/vector ops with multiple uses of the extact?
30113 // TODO: Generalize this, so it can be called from visitINSERT_VECTOR_ELT().
30114 SDValue Scalar = N->getOperand(Num: 0);
30115 unsigned Opcode = Scalar.getOpcode();
30116 EVT VecEltVT = VT.getScalarType();
30117 if (Scalar.hasOneUse() && Scalar->getNumValues() == 1 &&
30118 TLI.isBinOp(Opcode) && Scalar.getValueType() == VecEltVT &&
30119 Scalar.getOperand(i: 0).getValueType() == VecEltVT &&
30120 Scalar.getOperand(i: 1).getValueType() == VecEltVT &&
30121 Scalar->isOnlyUserOf(N: Scalar.getOperand(i: 0).getNode()) &&
30122 Scalar->isOnlyUserOf(N: Scalar.getOperand(i: 1).getNode()) &&
30123 DAG.isSafeToSpeculativelyExecute(Opcode) && hasOperation(Opcode, VT)) {
30124 // Match an extract element and get a shuffle mask equivalent.
30125 SmallVector<int, 8> ShufMask(VT.getVectorNumElements(), -1);
30126
30127 for (int i : {0, 1}) {
30128 // s2v (bo (extelt V, Idx), C) --> shuffle (bo V, C'), {Idx, -1, -1...}
30129 // s2v (bo C, (extelt V, Idx)) --> shuffle (bo C', V), {Idx, -1, -1...}
30130 SDValue EE = Scalar.getOperand(i);
30131 auto *C = dyn_cast<ConstantSDNode>(Val: Scalar.getOperand(i: i ? 0 : 1));
30132 if (C && EE.getOpcode() == ISD::EXTRACT_VECTOR_ELT &&
30133 EE.getOperand(i: 0).getValueType() == VT &&
30134 isa<ConstantSDNode>(Val: EE.getOperand(i: 1))) {
30135 // Mask = {ExtractIndex, undef, undef....}
30136 ShufMask[0] = EE.getConstantOperandVal(i: 1);
30137 // Make sure the shuffle is legal if we are crossing lanes.
30138 if (TLI.isShuffleMaskLegal(ShufMask, VT)) {
30139 SDLoc DL(N);
30140 SDValue V[] = {EE.getOperand(i: 0),
30141 DAG.getConstant(Val: C->getAPIntValue(), DL, VT)};
30142 SDValue VecBO = DAG.getNode(Opcode, DL, VT, N1: V[i], N2: V[1 - i]);
30143 return DAG.getVectorShuffle(VT, dl: DL, N1: VecBO, N2: DAG.getPOISON(VT),
30144 Mask: ShufMask);
30145 }
30146 }
30147 }
30148 }
30149
30150 // Replace a SCALAR_TO_VECTOR(EXTRACT_VECTOR_ELT(V,C0)) pattern
30151 // with a VECTOR_SHUFFLE and possible truncate.
30152 if (Opcode != ISD::EXTRACT_VECTOR_ELT ||
30153 !Scalar.getOperand(i: 0).getValueType().isFixedLengthVector())
30154 return SDValue();
30155
30156 // If we have an implicit truncate, truncate here if it is legal.
30157 if (VecEltVT != Scalar.getValueType() &&
30158 Scalar.getValueType().isScalarInteger() && isTypeLegal(VT: VecEltVT)) {
30159 SDValue Val = DAG.getNode(Opcode: ISD::TRUNCATE, DL: SDLoc(Scalar), VT: VecEltVT, Operand: Scalar);
30160 return DAG.getNode(Opcode: ISD::SCALAR_TO_VECTOR, DL: SDLoc(N), VT, Operand: Val);
30161 }
30162
30163 auto *ExtIndexC = dyn_cast<ConstantSDNode>(Val: Scalar.getOperand(i: 1));
30164 if (!ExtIndexC)
30165 return SDValue();
30166
30167 SDValue SrcVec = Scalar.getOperand(i: 0);
30168 EVT SrcVT = SrcVec.getValueType();
30169 unsigned SrcNumElts = SrcVT.getVectorNumElements();
30170 unsigned VTNumElts = VT.getVectorNumElements();
30171 if (VecEltVT == SrcVT.getScalarType() && VTNumElts <= SrcNumElts) {
30172 // Create a shuffle equivalent for scalar-to-vector: {ExtIndex, -1, -1, ...}
30173 SmallVector<int, 8> Mask(SrcNumElts, -1);
30174 Mask[0] = ExtIndexC->getZExtValue();
30175 SDValue LegalShuffle = TLI.buildLegalVectorShuffle(
30176 VT: SrcVT, DL: SDLoc(N), N0: SrcVec, N1: DAG.getPOISON(VT: SrcVT), Mask, DAG);
30177 if (!LegalShuffle)
30178 return SDValue();
30179
30180 // If the initial vector is the same size, the shuffle is the result.
30181 if (VT == SrcVT)
30182 return LegalShuffle;
30183
30184 // If not, shorten the shuffled vector.
30185 if (VTNumElts != SrcNumElts) {
30186 SDValue ZeroIdx = DAG.getVectorIdxConstant(Val: 0, DL: SDLoc(N));
30187 EVT SubVT = EVT::getVectorVT(Context&: *DAG.getContext(),
30188 VT: SrcVT.getVectorElementType(), NumElements: VTNumElts);
30189 return DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL: SDLoc(N), VT: SubVT, N1: LegalShuffle,
30190 N2: ZeroIdx);
30191 }
30192 }
30193
30194 return SDValue();
30195}
30196
30197SDValue DAGCombiner::visitINSERT_SUBVECTOR(SDNode *N) {
30198 EVT VT = N->getValueType(ResNo: 0);
30199 SDValue N0 = N->getOperand(Num: 0);
30200 SDValue N1 = N->getOperand(Num: 1);
30201 SDValue N2 = N->getOperand(Num: 2);
30202 uint64_t InsIdx = N->getConstantOperandVal(Num: 2);
30203
30204 // Remove insert of UNDEF/POISON.
30205 if (N1.isUndef()) {
30206 if (N1.getOpcode() == ISD::POISON || N0.getOpcode() == ISD::UNDEF)
30207 return N0;
30208 return DAG.getFreeze(V: N0);
30209 }
30210
30211 // If this is an insert of an extracted vector into an undef/poison vector, we
30212 // can just use the input to the extract if the types match, and can simplify
30213 // in some cases even if they don't.
30214 if (N0.isUndef() && N1.getOpcode() == ISD::EXTRACT_SUBVECTOR &&
30215 N1.getOperand(i: 1) == N2) {
30216 EVT N1VT = N1.getValueType();
30217 EVT SrcVT = N1.getOperand(i: 0).getValueType();
30218 if (SrcVT == VT) {
30219 // Need to ensure that result isn't more poisonous if skipping both the
30220 // extract+insert.
30221 if (N0.getOpcode() == ISD::POISON)
30222 return N1.getOperand(i: 0);
30223 if (VT.isFixedLengthVector() && N1VT.isFixedLengthVector()) {
30224 unsigned SubVecNumElts = N1VT.getVectorNumElements();
30225 APInt EltMask = APInt::getBitsSet(numBits: VT.getVectorNumElements(), loBit: InsIdx,
30226 hiBit: InsIdx + SubVecNumElts);
30227 if (DAG.isGuaranteedNotToBePoison(Op: N1.getOperand(i: 0), DemandedElts: ~EltMask))
30228 return N1.getOperand(i: 0);
30229 } else if (DAG.isGuaranteedNotToBePoison(Op: N1.getOperand(i: 0)))
30230 return N1.getOperand(i: 0);
30231 }
30232 // TODO: To remove the zero check, need to adjust the offset to
30233 // a multiple of the new src type.
30234 if (isNullConstant(V: N2)) {
30235 if (VT.knownBitsGE(VT: SrcVT) &&
30236 !(VT.isFixedLengthVector() && SrcVT.isScalableVector()))
30237 return DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL: SDLoc(N),
30238 VT, N1: N0, N2: N1.getOperand(i: 0), N3: N2);
30239 else if (VT.knownBitsLE(VT: SrcVT) &&
30240 !(VT.isScalableVector() && SrcVT.isFixedLengthVector()))
30241 return DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL: SDLoc(N),
30242 VT, N1: N1.getOperand(i: 0), N2);
30243 }
30244 }
30245
30246 // Handle case where we've ended up inserting back into the source vector
30247 // we extracted the subvector from.
30248 // insert_subvector(N0, extract_subvector(N0, N2), N2) --> N0
30249 if (N1.getOpcode() == ISD::EXTRACT_SUBVECTOR && N1.getOperand(i: 0) == N0 &&
30250 N1.getOperand(i: 1) == N2)
30251 return N0;
30252
30253 // Simplify scalar inserts into an undef vector:
30254 // insert_subvector undef, (splat X), N2 -> splat X
30255 if (N0.isUndef() && N1.getOpcode() == ISD::SPLAT_VECTOR)
30256 if (DAG.isConstantValueOfAnyType(N: N1.getOperand(i: 0)) || N1.hasOneUse())
30257 return DAG.getNode(Opcode: ISD::SPLAT_VECTOR, DL: SDLoc(N), VT, Operand: N1.getOperand(i: 0));
30258
30259 // insert_subvector (splat X), (splat X), N2 -> splat X
30260 if (N0.getOpcode() == ISD::SPLAT_VECTOR && N0.getOpcode() == N1.getOpcode() &&
30261 N0.getOperand(i: 0) == N1.getOperand(i: 0))
30262 return N0;
30263
30264 // If we are inserting a bitcast value into an undef, with the same
30265 // number of elements, just use the bitcast input of the extract.
30266 // i.e. INSERT_SUBVECTOR UNDEF (BITCAST N1) N2 ->
30267 // BITCAST (INSERT_SUBVECTOR UNDEF N1 N2)
30268 if (N0.isUndef() && N1.getOpcode() == ISD::BITCAST &&
30269 N1.getOperand(i: 0).getOpcode() == ISD::EXTRACT_SUBVECTOR &&
30270 N1.getOperand(i: 0).getOperand(i: 1) == N2 &&
30271 N1.getOperand(i: 0).getOperand(i: 0).getValueType().getVectorElementCount() ==
30272 VT.getVectorElementCount() &&
30273 N1.getOperand(i: 0).getOperand(i: 0).getValueType().getSizeInBits() ==
30274 VT.getSizeInBits()) {
30275 return DAG.getBitcast(VT, V: N1.getOperand(i: 0).getOperand(i: 0));
30276 }
30277
30278 // If both N1 and N2 are bitcast values on which insert_subvector
30279 // would makes sense, pull the bitcast through.
30280 // i.e. INSERT_SUBVECTOR (BITCAST N0) (BITCAST N1) N2 ->
30281 // BITCAST (INSERT_SUBVECTOR N0 N1 N2)
30282 if (N0.getOpcode() == ISD::BITCAST && N1.getOpcode() == ISD::BITCAST) {
30283 SDValue CN0 = N0.getOperand(i: 0);
30284 SDValue CN1 = N1.getOperand(i: 0);
30285 EVT CN0VT = CN0.getValueType();
30286 EVT CN1VT = CN1.getValueType();
30287 if (CN0VT.isVector() && CN1VT.isVector() &&
30288 CN0VT.getVectorElementType() == CN1VT.getVectorElementType() &&
30289 CN0VT.getVectorElementCount() == VT.getVectorElementCount()) {
30290 SDValue NewINSERT = DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL: SDLoc(N),
30291 VT: CN0.getValueType(), N1: CN0, N2: CN1, N3: N2);
30292 return DAG.getBitcast(VT, V: NewINSERT);
30293 }
30294 }
30295
30296 // Combine INSERT_SUBVECTORs where we are inserting to the same index.
30297 // INSERT_SUBVECTOR( INSERT_SUBVECTOR( Vec, SubOld, Idx ), SubNew, Idx )
30298 // --> INSERT_SUBVECTOR( Vec, SubNew, Idx )
30299 if (N0.getOpcode() == ISD::INSERT_SUBVECTOR &&
30300 N0.getOperand(i: 1).getValueType() == N1.getValueType() &&
30301 N0.getOperand(i: 2) == N2)
30302 return DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL: SDLoc(N), VT, N1: N0.getOperand(i: 0),
30303 N2: N1, N3: N2);
30304
30305 // Eliminate an intermediate insert into an undef vector:
30306 // insert_subvector undef, (insert_subvector undef, X, 0), 0 -->
30307 // insert_subvector undef, X, 0
30308 if (N0.isUndef() && N1.getOpcode() == ISD::INSERT_SUBVECTOR &&
30309 N1.getOperand(i: 0).isUndef() && isNullConstant(V: N1.getOperand(i: 2)) &&
30310 isNullConstant(V: N2))
30311 return DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL: SDLoc(N), VT, N1: N0,
30312 N2: N1.getOperand(i: 1), N3: N2);
30313
30314 // Push subvector bitcasts to the output, adjusting the index as we go.
30315 // insert_subvector(bitcast(v), bitcast(s), c1)
30316 // -> bitcast(insert_subvector(v, s, c2))
30317 if ((N0.isUndef() || N0.getOpcode() == ISD::BITCAST) &&
30318 N1.getOpcode() == ISD::BITCAST) {
30319 SDValue N0Src = peekThroughBitcasts(V: N0);
30320 SDValue N1Src = peekThroughBitcasts(V: N1);
30321 EVT N0SrcSVT = N0Src.getValueType().getScalarType();
30322 EVT N1SrcSVT = N1Src.getValueType().getScalarType();
30323 if ((N0.isUndef() || N0SrcSVT == N1SrcSVT) &&
30324 N0Src.getValueType().isVector() && N1Src.getValueType().isVector()) {
30325 EVT NewVT;
30326 SDLoc DL(N);
30327 SDValue NewIdx;
30328 LLVMContext &Ctx = *DAG.getContext();
30329 ElementCount NumElts = VT.getVectorElementCount();
30330 unsigned EltSizeInBits = VT.getScalarSizeInBits();
30331 if ((EltSizeInBits % N1SrcSVT.getSizeInBits()) == 0) {
30332 unsigned Scale = EltSizeInBits / N1SrcSVT.getSizeInBits();
30333 NewVT = EVT::getVectorVT(Context&: Ctx, VT: N1SrcSVT, EC: NumElts * Scale);
30334 NewIdx = DAG.getVectorIdxConstant(Val: InsIdx * Scale, DL);
30335 } else if ((N1SrcSVT.getSizeInBits() % EltSizeInBits) == 0) {
30336 unsigned Scale = N1SrcSVT.getSizeInBits() / EltSizeInBits;
30337 if (NumElts.isKnownMultipleOf(RHS: Scale) && (InsIdx % Scale) == 0) {
30338 NewVT = EVT::getVectorVT(Context&: Ctx, VT: N1SrcSVT,
30339 EC: NumElts.divideCoefficientBy(RHS: Scale));
30340 NewIdx = DAG.getVectorIdxConstant(Val: InsIdx / Scale, DL);
30341 }
30342 }
30343 if (NewIdx && hasOperation(Opcode: ISD::INSERT_SUBVECTOR, VT: NewVT)) {
30344 SDValue Res = DAG.getBitcast(VT: NewVT, V: N0Src);
30345 Res = DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL, VT: NewVT, N1: Res, N2: N1Src, N3: NewIdx);
30346 return DAG.getBitcast(VT, V: Res);
30347 }
30348 }
30349 }
30350
30351 // Canonicalize insert_subvector dag nodes.
30352 // Example:
30353 // (insert_subvector (insert_subvector A, Idx0), Idx1)
30354 // -> (insert_subvector (insert_subvector A, Idx1), Idx0)
30355 if (N0.getOpcode() == ISD::INSERT_SUBVECTOR && N0.hasOneUse() &&
30356 N1.getValueType() == N0.getOperand(i: 1).getValueType()) {
30357 unsigned OtherIdx = N0.getConstantOperandVal(i: 2);
30358 if (InsIdx < OtherIdx) {
30359 // Swap nodes.
30360 SDValue NewOp = DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL: SDLoc(N), VT,
30361 N1: N0.getOperand(i: 0), N2: N1, N3: N2);
30362 AddToWorklist(N: NewOp.getNode());
30363 return DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL: SDLoc(N0.getNode()),
30364 VT, N1: NewOp, N2: N0.getOperand(i: 1), N3: N0.getOperand(i: 2));
30365 }
30366 }
30367
30368 // If the input vector is a concatenation and the insert is wholly contained
30369 // in one of its operands, push the insertion into that operand.
30370 if (N0.getOpcode() == ISD::CONCAT_VECTORS && N0.hasOneUse()) {
30371 EVT ConcatOpVT = N0.getOperand(i: 0).getValueType();
30372 EVT InsVT = N1.getValueType();
30373 unsigned Factor = ConcatOpVT.getVectorMinNumElements();
30374 unsigned ConcatOpIdx = InsIdx / Factor;
30375 unsigned RelativeIdx = InsIdx - ConcatOpIdx * Factor;
30376 assert(ConcatOpIdx < N0.getNumOperands() && "subvector index mismatch");
30377
30378 // If the insert replaces a whole concat operand, optimize into a single
30379 // concat_vectors.
30380 if (RelativeIdx == 0 && ConcatOpVT == InsVT) {
30381 SmallVector<SDValue, 8> Ops(N0->ops());
30382 Ops[ConcatOpIdx] = N1;
30383 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT, Ops);
30384 }
30385
30386 if (VT.isFixedLengthVector() && ConcatOpVT.isFixedLengthVector() &&
30387 InsVT.isFixedLengthVector() &&
30388 hasOperation(Opcode: ISD::INSERT_SUBVECTOR, VT: ConcatOpVT)) {
30389 unsigned NumConcatOpElts = ConcatOpVT.getVectorNumElements();
30390 unsigned NumInsElts = InsVT.getVectorNumElements();
30391 if (RelativeIdx % NumInsElts == 0 &&
30392 RelativeIdx + NumInsElts <= NumConcatOpElts) {
30393 SmallVector<SDValue, 8> Ops(N0->ops());
30394 Ops[ConcatOpIdx] =
30395 DAG.getInsertSubvector(DL: SDLoc(N), Vec: Ops[ConcatOpIdx], SubVec: N1, Idx: RelativeIdx);
30396 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL: SDLoc(N), VT, Ops);
30397 }
30398 }
30399 }
30400
30401 // Simplify source operands based on insertion.
30402 if (SimplifyDemandedVectorElts(Op: SDValue(N, 0)))
30403 return SDValue(N, 0);
30404
30405 return SDValue();
30406}
30407
30408SDValue DAGCombiner::visitFP_TO_FP16(SDNode *N) {
30409 SDValue N0 = N->getOperand(Num: 0);
30410
30411 // fold (fp_to_fp16 (fp16_to_fp op)) -> op
30412 if (N0->getOpcode() == ISD::FP16_TO_FP)
30413 return N0->getOperand(Num: 0);
30414
30415 return SDValue();
30416}
30417
30418SDValue DAGCombiner::visitFP16_TO_FP(SDNode *N) {
30419 SelectionDAG::FlagInserter FlagsInserter(DAG, N);
30420 auto Op = N->getOpcode();
30421 assert((Op == ISD::FP16_TO_FP || Op == ISD::BF16_TO_FP) &&
30422 "opcode should be FP16_TO_FP or BF16_TO_FP.");
30423 SDValue N0 = N->getOperand(Num: 0);
30424
30425 // fold fp16_to_fp(op & 0xffff) -> fp16_to_fp(op) or
30426 // fold bf16_to_fp(op & 0xffff) -> bf16_to_fp(op)
30427 if (!TLI.shouldKeepZExtForFP16Conv() && N0->getOpcode() == ISD::AND) {
30428 ConstantSDNode *AndConst = getAsNonOpaqueConstant(N: N0.getOperand(i: 1));
30429 if (AndConst && AndConst->getAPIntValue() == 0xffff) {
30430 return DAG.getNode(Opcode: Op, DL: SDLoc(N), VT: N->getValueType(ResNo: 0), Operand: N0.getOperand(i: 0));
30431 }
30432 }
30433
30434 if (SDValue CastEliminated = eliminateFPCastPair(N))
30435 return CastEliminated;
30436
30437 // Sometimes constants manage to survive very late in the pipeline, e.g.,
30438 // because they are wrapped inside the <1 x f16> type. Try one last time to
30439 // get rid of them.
30440 SDValue Folded = DAG.FoldConstantArithmetic(Opcode: N->getOpcode(), DL: SDLoc(N),
30441 VT: N->getValueType(ResNo: 0), Ops: {N0});
30442 return Folded;
30443}
30444
30445SDValue DAGCombiner::visitFP_TO_BF16(SDNode *N) {
30446 SDValue N0 = N->getOperand(Num: 0);
30447
30448 // fold (fp_to_bf16 (bf16_to_fp op)) -> op
30449 if (N0->getOpcode() == ISD::BF16_TO_FP)
30450 return N0->getOperand(Num: 0);
30451
30452 return SDValue();
30453}
30454
30455SDValue DAGCombiner::visitBF16_TO_FP(SDNode *N) {
30456 // fold bf16_to_fp(op & 0xffff) -> bf16_to_fp(op)
30457 return visitFP16_TO_FP(N);
30458}
30459
30460SDValue DAGCombiner::visitVECREDUCE(SDNode *N) {
30461 SDValue N0 = N->getOperand(Num: 0);
30462 EVT VT = N0.getValueType();
30463 unsigned Opcode = N->getOpcode();
30464
30465 // VECREDUCE over 1-element vector is just an extract.
30466 if (VT.getVectorElementCount().isScalar()) {
30467 SDLoc dl(N);
30468 SDValue Res =
30469 DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL: dl, VT: VT.getVectorElementType(), N1: N0,
30470 N2: DAG.getVectorIdxConstant(Val: 0, DL: dl));
30471 if (Res.getValueType() != N->getValueType(ResNo: 0))
30472 Res = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL: dl, VT: N->getValueType(ResNo: 0), Operand: Res);
30473 return Res;
30474 }
30475
30476 // On an boolean vector an and/or reduction is the same as a umin/umax
30477 // reduction. Convert them if the latter is legal while the former isn't.
30478 if (Opcode == ISD::VECREDUCE_AND || Opcode == ISD::VECREDUCE_OR) {
30479 unsigned NewOpcode = Opcode == ISD::VECREDUCE_AND
30480 ? ISD::VECREDUCE_UMIN : ISD::VECREDUCE_UMAX;
30481 if (!TLI.isOperationLegalOrCustom(Op: Opcode, VT) &&
30482 TLI.isOperationLegalOrCustom(Op: NewOpcode, VT) &&
30483 DAG.ComputeNumSignBits(Op: N0) == VT.getScalarSizeInBits())
30484 return DAG.getNode(Opcode: NewOpcode, DL: SDLoc(N), VT: N->getValueType(ResNo: 0), Operand: N0);
30485 }
30486
30487 // vecreduce_or(insert_subvector(zero or undef, val)) -> vecreduce_or(val)
30488 // vecreduce_and(insert_subvector(ones or undef, val)) -> vecreduce_and(val)
30489 if (N0.getOpcode() == ISD::INSERT_SUBVECTOR &&
30490 TLI.isTypeLegal(VT: N0.getOperand(i: 1).getValueType())) {
30491 SDValue Vec = N0.getOperand(i: 0);
30492 SDValue Subvec = N0.getOperand(i: 1);
30493 if ((Opcode == ISD::VECREDUCE_OR &&
30494 (N0.getOperand(i: 0).isUndef() || isNullOrNullSplat(V: Vec))) ||
30495 (Opcode == ISD::VECREDUCE_AND &&
30496 (N0.getOperand(i: 0).isUndef() || isAllOnesOrAllOnesSplat(V: Vec))))
30497 return DAG.getNode(Opcode, DL: SDLoc(N), VT: N->getValueType(ResNo: 0), Operand: Subvec);
30498 }
30499
30500 // vecreduce_or(sext(x)) -> sext(vecreduce_or(x))
30501 // Same for zext and anyext, and for and/or/xor reductions.
30502 if ((Opcode == ISD::VECREDUCE_OR || Opcode == ISD::VECREDUCE_AND ||
30503 Opcode == ISD::VECREDUCE_XOR) &&
30504 (N0.getOpcode() == ISD::SIGN_EXTEND ||
30505 N0.getOpcode() == ISD::ZERO_EXTEND ||
30506 N0.getOpcode() == ISD::ANY_EXTEND) &&
30507 TLI.isOperationLegalOrCustom(Op: Opcode, VT: N0.getOperand(i: 0).getValueType())) {
30508 SDValue Red = DAG.getNode(Opcode, DL: SDLoc(N),
30509 VT: N0.getOperand(i: 0).getValueType().getScalarType(),
30510 Operand: N0.getOperand(i: 0));
30511 return DAG.getNode(Opcode: N0.getOpcode(), DL: SDLoc(N), VT: N->getValueType(ResNo: 0), Operand: Red);
30512 }
30513 return SDValue();
30514}
30515
30516SDValue DAGCombiner::visitVPOp(SDNode *N) {
30517
30518 if (N->getOpcode() == ISD::VP_GATHER)
30519 if (SDValue SD = visitVPGATHER(N))
30520 return SD;
30521
30522 if (N->getOpcode() == ISD::VP_SCATTER)
30523 if (SDValue SD = visitVPSCATTER(N))
30524 return SD;
30525
30526 if (N->getOpcode() == ISD::EXPERIMENTAL_VP_STRIDED_LOAD)
30527 if (SDValue SD = visitVP_STRIDED_LOAD(N))
30528 return SD;
30529
30530 if (N->getOpcode() == ISD::EXPERIMENTAL_VP_STRIDED_STORE)
30531 if (SDValue SD = visitVP_STRIDED_STORE(N))
30532 return SD;
30533
30534 // VP operations in which all vector elements are disabled - either by
30535 // determining that the mask is all false or that the EVL is 0 - can be
30536 // eliminated.
30537 bool AreAllEltsDisabled = false;
30538 if (auto EVLIdx = ISD::getVPExplicitVectorLengthIdx(Opcode: N->getOpcode()))
30539 AreAllEltsDisabled |= isNullConstant(V: N->getOperand(Num: *EVLIdx));
30540 if (auto MaskIdx = ISD::getVPMaskIdx(Opcode: N->getOpcode()))
30541 AreAllEltsDisabled |=
30542 ISD::isConstantSplatVectorAllZeros(N: N->getOperand(Num: *MaskIdx).getNode());
30543
30544 // This is the only generic VP combine we support for now.
30545 if (!AreAllEltsDisabled)
30546 return SDValue();
30547
30548 // Binary operations can be replaced by UNDEF.
30549 if (ISD::isVPBinaryOp(Opcode: N->getOpcode()))
30550 return DAG.getUNDEF(VT: N->getValueType(ResNo: 0));
30551
30552 // VP Memory operations can be replaced by either the chain (stores) or the
30553 // chain + undef (loads).
30554 if (const auto *MemSD = dyn_cast<MemSDNode>(Val: N)) {
30555 if (MemSD->writeMem())
30556 return MemSD->getChain();
30557 return CombineTo(N, Res0: DAG.getUNDEF(VT: N->getValueType(ResNo: 0)), Res1: MemSD->getChain());
30558 }
30559
30560 // Reduction operations return the start operand when no elements are active.
30561 if (ISD::isVPReduction(Opcode: N->getOpcode()))
30562 return N->getOperand(Num: 0);
30563
30564 return SDValue();
30565}
30566
30567SDValue DAGCombiner::visitGET_FPENV_MEM(SDNode *N) {
30568 SDValue Chain = N->getOperand(Num: 0);
30569 SDValue Ptr = N->getOperand(Num: 1);
30570 EVT MemVT = cast<FPStateAccessSDNode>(Val: N)->getMemoryVT();
30571
30572 // Check if the memory, where FP state is written to, is used only in a single
30573 // load operation.
30574 LoadSDNode *LdNode = nullptr;
30575 for (auto *U : Ptr->users()) {
30576 if (U == N)
30577 continue;
30578 if (auto *Ld = dyn_cast<LoadSDNode>(Val: U)) {
30579 if (LdNode && LdNode != Ld)
30580 return SDValue();
30581 LdNode = Ld;
30582 continue;
30583 }
30584 return SDValue();
30585 }
30586 if (!LdNode || !LdNode->isSimple() || LdNode->isIndexed() ||
30587 !LdNode->getOffset().isUndef() || LdNode->getMemoryVT() != MemVT ||
30588 !LdNode->getChain().reachesChainWithoutSideEffects(Dest: SDValue(N, 0)))
30589 return SDValue();
30590
30591 // Check if the loaded value is used only in a store operation.
30592 StoreSDNode *StNode = nullptr;
30593 for (SDUse &U : LdNode->uses()) {
30594 if (U.getResNo() == 0) {
30595 if (auto *St = dyn_cast<StoreSDNode>(Val: U.getUser())) {
30596 if (StNode)
30597 return SDValue();
30598 StNode = St;
30599 } else {
30600 return SDValue();
30601 }
30602 }
30603 }
30604 if (!StNode || !StNode->isSimple() || StNode->isIndexed() ||
30605 !StNode->getOffset().isUndef() || StNode->getMemoryVT() != MemVT ||
30606 !StNode->getChain().reachesChainWithoutSideEffects(Dest: SDValue(LdNode, 1)))
30607 return SDValue();
30608
30609 // The new node replaces N, so the store address must not depend on N (for
30610 // example through a CopyFromReg chained after the load), or the DAG would
30611 // become cyclic.
30612 if (StNode->getBasePtr()->hasPredecessor(N))
30613 return SDValue();
30614
30615 // Create new node GET_FPENV_MEM, which uses the store address to write FP
30616 // environment.
30617 SDValue Res = DAG.getGetFPEnv(Chain, dl: SDLoc(N), Ptr: StNode->getBasePtr(), MemVT,
30618 MMO: StNode->getMemOperand());
30619 CombineTo(N: StNode, Res, AddTo: false);
30620 return Res;
30621}
30622
30623SDValue DAGCombiner::visitSET_FPENV_MEM(SDNode *N) {
30624 SDValue Chain = N->getOperand(Num: 0);
30625 SDValue Ptr = N->getOperand(Num: 1);
30626 EVT MemVT = cast<FPStateAccessSDNode>(Val: N)->getMemoryVT();
30627
30628 // Check if the address of FP state is used also in a store operation only.
30629 StoreSDNode *StNode = nullptr;
30630 for (auto *U : Ptr->users()) {
30631 if (U == N)
30632 continue;
30633 if (auto *St = dyn_cast<StoreSDNode>(Val: U)) {
30634 if (StNode && StNode != St)
30635 return SDValue();
30636 StNode = St;
30637 continue;
30638 }
30639 return SDValue();
30640 }
30641 if (!StNode || !StNode->isSimple() || StNode->isIndexed() ||
30642 !StNode->getOffset().isUndef() || StNode->getMemoryVT() != MemVT ||
30643 !Chain.reachesChainWithoutSideEffects(Dest: SDValue(StNode, 0)))
30644 return SDValue();
30645
30646 // Check if the stored value is loaded from some location and the loaded
30647 // value is used only in the store operation.
30648 SDValue StValue = StNode->getValue();
30649 auto *LdNode = dyn_cast<LoadSDNode>(Val&: StValue);
30650 if (!LdNode || !LdNode->isSimple() || LdNode->isIndexed() ||
30651 !LdNode->getOffset().isUndef() || LdNode->getMemoryVT() != MemVT ||
30652 !StNode->getChain().reachesChainWithoutSideEffects(Dest: SDValue(LdNode, 1)))
30653 return SDValue();
30654
30655 // Create new node SET_FPENV_MEM, which uses the load address to read FP
30656 // environment.
30657 SDValue Res =
30658 DAG.getSetFPEnv(Chain: LdNode->getChain(), dl: SDLoc(N), Ptr: LdNode->getBasePtr(), MemVT,
30659 MMO: LdNode->getMemOperand());
30660 return Res;
30661}
30662
30663/// Returns a vector_shuffle if it able to transform an AND to a vector_shuffle
30664/// with the destination vector and a zero vector.
30665/// e.g. AND V, <0xffffffff, 0, 0xffffffff, 0>. ==>
30666/// vector_shuffle V, Zero, <0, 4, 2, 4>
30667SDValue DAGCombiner::XformToShuffleWithZero(SDNode *N) {
30668 assert(N->getOpcode() == ISD::AND && "Unexpected opcode!");
30669
30670 EVT VT = N->getValueType(ResNo: 0);
30671 SDValue LHS = N->getOperand(Num: 0);
30672 SDValue RHS = peekThroughBitcasts(V: N->getOperand(Num: 1));
30673 SDLoc DL(N);
30674
30675 // Make sure we're not running after operation legalization where it
30676 // may have custom lowered the vector shuffles.
30677 if (LegalOperations)
30678 return SDValue();
30679
30680 if (RHS.getOpcode() != ISD::BUILD_VECTOR)
30681 return SDValue();
30682
30683 EVT RVT = RHS.getValueType();
30684 unsigned NumElts = RHS.getNumOperands();
30685
30686 // Attempt to create a valid clear mask, splitting the mask into
30687 // sub elements and checking to see if each is
30688 // all zeros or all ones - suitable for shuffle masking.
30689 auto BuildClearMask = [&](int Split) {
30690 int NumSubElts = NumElts * Split;
30691 int NumSubBits = RVT.getScalarSizeInBits() / Split;
30692
30693 SmallVector<int, 8> Indices;
30694 for (int i = 0; i != NumSubElts; ++i) {
30695 int EltIdx = i / Split;
30696 int SubIdx = i % Split;
30697 SDValue Elt = RHS.getOperand(i: EltIdx);
30698 // X & undef --> 0 (not undef). So this lane must be converted to choose
30699 // from the zero constant vector (same as if the element had all 0-bits).
30700 if (Elt.isUndef()) {
30701 Indices.push_back(Elt: i + NumSubElts);
30702 continue;
30703 }
30704
30705 std::optional<APInt> Bits = Elt->bitcastToAPInt();
30706 if (!Bits)
30707 return SDValue();
30708
30709 // Extract the sub element from the constant bit mask.
30710 if (DAG.getDataLayout().isBigEndian())
30711 *Bits =
30712 Bits->extractBits(numBits: NumSubBits, bitPosition: (Split - SubIdx - 1) * NumSubBits);
30713 else
30714 *Bits = Bits->extractBits(numBits: NumSubBits, bitPosition: SubIdx * NumSubBits);
30715
30716 if (Bits->isAllOnes())
30717 Indices.push_back(Elt: i);
30718 else if (*Bits == 0)
30719 Indices.push_back(Elt: i + NumSubElts);
30720 else
30721 return SDValue();
30722 }
30723
30724 // Let's see if the target supports this vector_shuffle.
30725 EVT ClearSVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: NumSubBits);
30726 EVT ClearVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: ClearSVT, NumElements: NumSubElts);
30727 if (!TLI.isVectorClearMaskLegal(Indices, ClearVT))
30728 return SDValue();
30729
30730 SDValue Zero = DAG.getConstant(Val: 0, DL, VT: ClearVT);
30731 return DAG.getBitcast(VT, V: DAG.getVectorShuffle(VT: ClearVT, dl: DL,
30732 N1: DAG.getBitcast(VT: ClearVT, V: LHS),
30733 N2: Zero, Mask: Indices));
30734 };
30735
30736 // Determine maximum split level (byte level masking).
30737 int MaxSplit = 1;
30738 if (RVT.getScalarSizeInBits() % 8 == 0)
30739 MaxSplit = RVT.getScalarSizeInBits() / 8;
30740
30741 for (int Split = 1; Split <= MaxSplit; ++Split)
30742 if (RVT.getScalarSizeInBits() % Split == 0)
30743 if (SDValue S = BuildClearMask(Split))
30744 return S;
30745
30746 return SDValue();
30747}
30748
30749/// If a vector binop is performed on splat values, it may be profitable to
30750/// extract, scalarize, and insert/splat.
30751static SDValue scalarizeBinOpOfSplats(SDNode *N, SelectionDAG &DAG,
30752 const SDLoc &DL, bool LegalTypes) {
30753 SDValue N0 = N->getOperand(Num: 0);
30754 SDValue N1 = N->getOperand(Num: 1);
30755 unsigned Opcode = N->getOpcode();
30756 EVT VT = N->getValueType(ResNo: 0);
30757 EVT EltVT = VT.getVectorElementType();
30758 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
30759
30760 // TODO: Remove/replace the extract cost check? If the elements are available
30761 // as scalars, then there may be no extract cost. Should we ask if
30762 // inserting a scalar back into a vector is cheap instead?
30763 int Index0, Index1;
30764 SDValue Src0 = DAG.getSplatSourceVector(V: N0, SplatIndex&: Index0);
30765 SDValue Src1 = DAG.getSplatSourceVector(V: N1, SplatIndex&: Index1);
30766 // Extract element from splat_vector should be free.
30767 // TODO: use DAG.isSplatValue instead?
30768 bool IsBothSplatVector = N0.getOpcode() == ISD::SPLAT_VECTOR &&
30769 N1.getOpcode() == ISD::SPLAT_VECTOR;
30770 if (!Src0 || !Src1 || Index0 != Index1 ||
30771 Src0.getValueType().getVectorElementType() != EltVT ||
30772 Src1.getValueType().getVectorElementType() != EltVT ||
30773 !(IsBothSplatVector || TLI.isExtractVecEltCheap(VT, Index: Index0)) ||
30774 // If before type legalization, allow scalar types that will eventually be
30775 // made legal.
30776 !TLI.isOperationLegalOrCustom(
30777 Op: Opcode, VT: LegalTypes
30778 ? EltVT
30779 : TLI.getTypeToTransformTo(Context&: *DAG.getContext(), VT: EltVT)))
30780 return SDValue();
30781
30782 // FIXME: Type legalization can't handle illegal MULHS/MULHU.
30783 if ((Opcode == ISD::MULHS || Opcode == ISD::MULHU) && !TLI.isTypeLegal(VT: EltVT))
30784 return SDValue();
30785
30786 if (N0.getOpcode() == ISD::BUILD_VECTOR && N0.getOpcode() == N1.getOpcode()) {
30787 // All but one element should have an undef input, which will fold to a
30788 // constant or undef. Avoid splatting which would over-define potentially
30789 // undefined elements.
30790
30791 // bo (build_vec ..undef, X, undef...), (build_vec ..undef, Y, undef...) -->
30792 // build_vec ..undef, (bo X, Y), undef...
30793 SmallVector<SDValue, 16> EltsX, EltsY, EltsResult;
30794 DAG.ExtractVectorElements(Op: Src0, Args&: EltsX);
30795 DAG.ExtractVectorElements(Op: Src1, Args&: EltsY);
30796
30797 for (auto [X, Y] : zip(t&: EltsX, u&: EltsY))
30798 EltsResult.push_back(Elt: DAG.getNode(Opcode, DL, VT: EltVT, N1: X, N2: Y, Flags: N->getFlags()));
30799 return DAG.getBuildVector(VT, DL, Ops: EltsResult);
30800 }
30801
30802 SDValue IndexC = DAG.getVectorIdxConstant(Val: Index0, DL);
30803 SDValue X = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: Src0, N2: IndexC);
30804 SDValue Y = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: Src1, N2: IndexC);
30805 SDValue ScalarBO = DAG.getNode(Opcode, DL, VT: EltVT, N1: X, N2: Y, Flags: N->getFlags());
30806
30807 // bo (splat X, Index), (splat Y, Index) --> splat (bo X, Y), Index
30808 return DAG.getSplat(VT, DL, Op: ScalarBO);
30809}
30810
30811/// Visit a vector cast operation, like FP_EXTEND.
30812SDValue DAGCombiner::SimplifyVCastOp(SDNode *N, const SDLoc &DL) {
30813 EVT VT = N->getValueType(ResNo: 0);
30814 assert(VT.isVector() && "SimplifyVCastOp only works on vectors!");
30815 EVT EltVT = VT.getVectorElementType();
30816 unsigned Opcode = N->getOpcode();
30817
30818 SDValue N0 = N->getOperand(Num: 0);
30819 const TargetLowering &TLI = DAG.getTargetLoweringInfo();
30820
30821 // TODO: promote operation might be also good here?
30822 int Index0;
30823 SDValue Src0 = DAG.getSplatSourceVector(V: N0, SplatIndex&: Index0);
30824 if (Src0 &&
30825 (N0.getOpcode() == ISD::SPLAT_VECTOR ||
30826 TLI.isExtractVecEltCheap(VT, Index: Index0)) &&
30827 TLI.isOperationLegalOrCustom(Op: Opcode, VT: EltVT) &&
30828 TLI.preferScalarizeSplat(N)) {
30829 EVT SrcVT = N0.getValueType();
30830 EVT SrcEltVT = SrcVT.getVectorElementType();
30831 if (!LegalTypes || TLI.isTypeLegal(VT: SrcEltVT)) {
30832 SDValue IndexC = DAG.getVectorIdxConstant(Val: Index0, DL);
30833 SDValue Elt =
30834 DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: SrcEltVT, N1: Src0, N2: IndexC);
30835 SDValue ScalarBO = DAG.getNode(Opcode, DL, VT: EltVT, Operand: Elt, Flags: N->getFlags());
30836 if (VT.isScalableVector())
30837 return DAG.getSplatVector(VT, DL, Op: ScalarBO);
30838 SmallVector<SDValue, 8> Ops(VT.getVectorNumElements(), ScalarBO);
30839 return DAG.getBuildVector(VT, DL, Ops);
30840 }
30841 }
30842
30843 return SDValue();
30844}
30845
30846/// Visit a binary vector operation, like ADD.
30847SDValue DAGCombiner::SimplifyVBinOp(SDNode *N, const SDLoc &DL) {
30848 EVT VT = N->getValueType(ResNo: 0);
30849 assert(VT.isVector() && "SimplifyVBinOp only works on vectors!");
30850
30851 SDValue LHS = N->getOperand(Num: 0);
30852 SDValue RHS = N->getOperand(Num: 1);
30853 unsigned Opcode = N->getOpcode();
30854 SDNodeFlags Flags = N->getFlags();
30855
30856 // Move unary shuffles with identical masks after a vector binop:
30857 // VBinOp (shuffle A, Undef, Mask), (shuffle B, Undef, Mask))
30858 // --> shuffle (VBinOp A, B), Undef, Mask
30859 // This does not require type legality checks because we are creating the
30860 // same types of operations that are in the original sequence. We do have to
30861 // restrict ops like integer div that have immediate UB (eg, div-by-zero)
30862 // though. This code is adapted from the identical transform in instcombine.
30863 if (DAG.isSafeToSpeculativelyExecute(Opcode)) {
30864 auto *Shuf0 = dyn_cast<ShuffleVectorSDNode>(Val&: LHS);
30865 auto *Shuf1 = dyn_cast<ShuffleVectorSDNode>(Val&: RHS);
30866 if (Shuf0 && Shuf1 && Shuf0->getMask().equals(RHS: Shuf1->getMask()) &&
30867 LHS.getOperand(i: 1).isUndef() && RHS.getOperand(i: 1).isUndef() &&
30868 (LHS.hasOneUse() || RHS.hasOneUse() || LHS == RHS)) {
30869 SDValue NewBinOp = DAG.getNode(Opcode, DL, VT, N1: LHS.getOperand(i: 0),
30870 N2: RHS.getOperand(i: 0), Flags);
30871 SDValue UndefV = LHS.getOperand(i: 1);
30872 return DAG.getVectorShuffle(VT, dl: DL, N1: NewBinOp, N2: UndefV, Mask: Shuf0->getMask());
30873 }
30874
30875 // Try to sink a splat shuffle after a binop with a uniform constant.
30876 // This is limited to cases where neither the shuffle nor the constant have
30877 // undefined elements because that could be poison-unsafe or inhibit
30878 // demanded elements analysis. It is further limited to not change a splat
30879 // of an inserted scalar because that may be optimized better by
30880 // load-folding or other target-specific behaviors.
30881 if (isConstOrConstSplat(N: RHS) && Shuf0 && all_equal(Range: Shuf0->getMask()) &&
30882 Shuf0->hasOneUse() && Shuf0->getOperand(Num: 1).isUndef() &&
30883 Shuf0->getOperand(Num: 0).getOpcode() != ISD::INSERT_VECTOR_ELT) {
30884 // binop (splat X), (splat C) --> splat (binop X, C)
30885 SDValue X = Shuf0->getOperand(Num: 0);
30886 SDValue NewBinOp = DAG.getNode(Opcode, DL, VT, N1: X, N2: RHS, Flags);
30887 return DAG.getVectorShuffle(VT, dl: DL, N1: NewBinOp, N2: DAG.getPOISON(VT),
30888 Mask: Shuf0->getMask());
30889 }
30890 if (isConstOrConstSplat(N: LHS) && Shuf1 && all_equal(Range: Shuf1->getMask()) &&
30891 Shuf1->hasOneUse() && Shuf1->getOperand(Num: 1).isUndef() &&
30892 Shuf1->getOperand(Num: 0).getOpcode() != ISD::INSERT_VECTOR_ELT) {
30893 // binop (splat C), (splat X) --> splat (binop C, X)
30894 SDValue X = Shuf1->getOperand(Num: 0);
30895 SDValue NewBinOp = DAG.getNode(Opcode, DL, VT, N1: LHS, N2: X, Flags);
30896 return DAG.getVectorShuffle(VT, dl: DL, N1: NewBinOp, N2: DAG.getPOISON(VT),
30897 Mask: Shuf1->getMask());
30898 }
30899 }
30900
30901 // The following pattern is likely to emerge with vector reduction ops. Moving
30902 // the binary operation ahead of insertion may allow using a narrower vector
30903 // instruction that has better performance than the wide version of the op:
30904 // VBinOp (ins undef, X, Z), (ins undef, Y, Z) --> ins VecC, (VBinOp X, Y), Z
30905 if (LHS.getOpcode() == ISD::INSERT_SUBVECTOR && LHS.getOperand(i: 0).isUndef() &&
30906 RHS.getOpcode() == ISD::INSERT_SUBVECTOR && RHS.getOperand(i: 0).isUndef() &&
30907 LHS.getOperand(i: 2) == RHS.getOperand(i: 2) &&
30908 (LHS.hasOneUse() || RHS.hasOneUse())) {
30909 SDValue X = LHS.getOperand(i: 1);
30910 SDValue Y = RHS.getOperand(i: 1);
30911 SDValue Z = LHS.getOperand(i: 2);
30912 EVT NarrowVT = X.getValueType();
30913 if (NarrowVT == Y.getValueType() &&
30914 TLI.isOperationLegalOrCustomOrPromote(Op: Opcode, VT: NarrowVT,
30915 LegalOnly: LegalOperations)) {
30916 // (binop undef, undef) may not return undef, so compute that result.
30917 SDValue VecC =
30918 DAG.getNode(Opcode, DL, VT, N1: DAG.getUNDEF(VT), N2: DAG.getUNDEF(VT));
30919 SDValue NarrowBO = DAG.getNode(Opcode, DL, VT: NarrowVT, N1: X, N2: Y);
30920 return DAG.getNode(Opcode: ISD::INSERT_SUBVECTOR, DL, VT, N1: VecC, N2: NarrowBO, N3: Z);
30921 }
30922 }
30923
30924 // Make sure all but the first op are undef or constant.
30925 auto ConcatWithConstantOrUndef = [](SDValue Concat) {
30926 return Concat.getOpcode() == ISD::CONCAT_VECTORS &&
30927 all_of(Range: drop_begin(RangeOrContainer: Concat->ops()), P: [](const SDValue &Op) {
30928 return Op.isUndef() ||
30929 ISD::isBuildVectorOfConstantSDNodes(N: Op.getNode());
30930 });
30931 };
30932
30933 // The following pattern is likely to emerge with vector reduction ops. Moving
30934 // the binary operation ahead of the concat may allow using a narrower vector
30935 // instruction that has better performance than the wide version of the op:
30936 // VBinOp (concat X, undef/constant), (concat Y, undef/constant) -->
30937 // concat (VBinOp X, Y), VecC
30938 if (ConcatWithConstantOrUndef(LHS) && ConcatWithConstantOrUndef(RHS) &&
30939 (LHS.hasOneUse() || RHS.hasOneUse())) {
30940 EVT NarrowVT = LHS.getOperand(i: 0).getValueType();
30941 if (NarrowVT == RHS.getOperand(i: 0).getValueType() &&
30942 TLI.isOperationLegalOrCustomOrPromote(Op: Opcode, VT: NarrowVT)) {
30943 unsigned NumOperands = LHS.getNumOperands();
30944 SmallVector<SDValue, 4> ConcatOps;
30945 for (unsigned i = 0; i != NumOperands; ++i) {
30946 // This constant fold for operands 1 and up.
30947 ConcatOps.push_back(Elt: DAG.getNode(Opcode, DL, VT: NarrowVT, N1: LHS.getOperand(i),
30948 N2: RHS.getOperand(i)));
30949 }
30950
30951 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT, Ops: ConcatOps);
30952 }
30953 }
30954
30955 if (SDValue V = scalarizeBinOpOfSplats(N, DAG, DL, LegalTypes))
30956 return V;
30957
30958 return SDValue();
30959}
30960
30961SDValue DAGCombiner::SimplifySelect(const SDLoc &DL, SDValue N0, SDValue N1,
30962 SDValue N2) {
30963 assert(N0.getOpcode() == ISD::SETCC &&
30964 "First argument must be a SetCC node!");
30965
30966 SDValue SCC = SimplifySelectCC(DL, N0: N0.getOperand(i: 0), N1: N0.getOperand(i: 1), N2: N1, N3: N2,
30967 CC: cast<CondCodeSDNode>(Val: N0.getOperand(i: 2))->get());
30968
30969 // If we got a simplified select_cc node back from SimplifySelectCC, then
30970 // break it down into a new SETCC node, and a new SELECT node, and then return
30971 // the SELECT node, since we were called with a SELECT node.
30972 if (SCC.getNode()) {
30973 // Check to see if we got a select_cc back (to turn into setcc/select).
30974 // Otherwise, just return whatever node we got back, like fabs.
30975 if (SCC.getOpcode() == ISD::SELECT_CC) {
30976 const SDNodeFlags Flags = N0->getFlags();
30977 SDValue SETCC = DAG.getNode(Opcode: ISD::SETCC, DL: SDLoc(N0),
30978 VT: N0.getValueType(),
30979 N1: SCC.getOperand(i: 0), N2: SCC.getOperand(i: 1),
30980 N3: SCC.getOperand(i: 4), Flags);
30981 AddToWorklist(N: SETCC.getNode());
30982 return DAG.getSelect(DL: SDLoc(SCC), VT: SCC.getValueType(), Cond: SETCC,
30983 LHS: SCC.getOperand(i: 2), RHS: SCC.getOperand(i: 3), Flags);
30984 }
30985
30986 return SCC;
30987 }
30988 return SDValue();
30989}
30990
30991/// Given a SELECT or a SELECT_CC node, where LHS and RHS are the two values
30992/// being selected between, see if we can simplify the select. Callers of this
30993/// should assume that TheSelect is deleted if this returns true. As such, they
30994/// should return the appropriate thing (e.g. the node) back to the top-level of
30995/// the DAG combiner loop to avoid it being looked at.
30996bool DAGCombiner::SimplifySelectOps(SDNode *TheSelect, SDValue LHS,
30997 SDValue RHS) {
30998 // fold (select (setcc x, [+-]0.0, *lt), NaN, (fsqrt x))
30999 // The select + setcc is redundant, because fsqrt returns NaN for X < 0.
31000 if (const ConstantFPSDNode *NaN = isConstOrConstSplatFP(N: LHS)) {
31001 if (NaN->isNaN() && RHS.getOpcode() == ISD::FSQRT) {
31002 // We have: (select (setcc ?, ?, ?), NaN, (fsqrt ?))
31003 SDValue Sqrt = RHS;
31004 ISD::CondCode CC;
31005 SDValue CmpLHS;
31006 const ConstantFPSDNode *Zero = nullptr;
31007
31008 if (TheSelect->getOpcode() == ISD::SELECT_CC) {
31009 CC = cast<CondCodeSDNode>(Val: TheSelect->getOperand(Num: 4))->get();
31010 CmpLHS = TheSelect->getOperand(Num: 0);
31011 Zero = isConstOrConstSplatFP(N: TheSelect->getOperand(Num: 1));
31012 } else {
31013 // SELECT or VSELECT
31014 SDValue Cmp = TheSelect->getOperand(Num: 0);
31015 if (Cmp.getOpcode() == ISD::SETCC) {
31016 CC = cast<CondCodeSDNode>(Val: Cmp.getOperand(i: 2))->get();
31017 CmpLHS = Cmp.getOperand(i: 0);
31018 Zero = isConstOrConstSplatFP(N: Cmp.getOperand(i: 1));
31019 }
31020 }
31021 if (Zero && Zero->isZero() &&
31022 Sqrt.getOperand(i: 0) == CmpLHS && (CC == ISD::SETOLT ||
31023 CC == ISD::SETULT || CC == ISD::SETLT)) {
31024 // We have: (select (setcc x, [+-]0.0, *lt), NaN, (fsqrt x))
31025 CombineTo(N: TheSelect, Res: Sqrt);
31026 return true;
31027 }
31028 }
31029 }
31030 // Cannot simplify select with vector condition
31031 if (TheSelect->getOperand(Num: 0).getValueType().isVector()) return false;
31032
31033 // If this is a select from two identical things, try to pull the operation
31034 // through the select.
31035 if (LHS.getOpcode() != RHS.getOpcode() ||
31036 !LHS.hasOneUse() || !RHS.hasOneUse())
31037 return false;
31038
31039 // If this is a load and the token chain is identical, replace the select
31040 // of two loads with a load through a select of the address to load from.
31041 // This triggers in things like "select bool X, 10.0, 123.0" after the FP
31042 // constants have been dropped into the constant pool.
31043 if (LHS.getOpcode() == ISD::LOAD) {
31044 LoadSDNode *LLD = cast<LoadSDNode>(Val&: LHS);
31045 LoadSDNode *RLD = cast<LoadSDNode>(Val&: RHS);
31046
31047 // Token chains must be identical.
31048 if (LHS.getOperand(i: 0) != RHS.getOperand(i: 0) ||
31049 // Do not let this transformation reduce the number of volatile loads.
31050 // Be conservative for atomics for the moment
31051 // TODO: This does appear to be legal for unordered atomics (see D66309)
31052 !LLD->isSimple() || !RLD->isSimple() ||
31053 // FIXME: If either is a pre/post inc/dec load,
31054 // we'd need to split out the address adjustment.
31055 LLD->isIndexed() || RLD->isIndexed() ||
31056 // If this is an EXTLOAD, the VT's must match.
31057 LLD->getMemoryVT() != RLD->getMemoryVT() ||
31058 // If this is an EXTLOAD, the kind of extension must match.
31059 (LLD->getExtensionType() != RLD->getExtensionType() &&
31060 // The only exception is if one of the extensions is anyext.
31061 LLD->getExtensionType() != ISD::EXTLOAD &&
31062 RLD->getExtensionType() != ISD::EXTLOAD) ||
31063 // FIXME: this discards src value information. This is
31064 // over-conservative. It would be beneficial to be able to remember
31065 // both potential memory locations. Since we are discarding
31066 // src value info, don't do the transformation if the memory
31067 // locations are not in the same address space.
31068 LLD->getPointerInfo().getAddrSpace() !=
31069 RLD->getPointerInfo().getAddrSpace() ||
31070 // We can't produce a CMOV of a TargetFrameIndex since we won't
31071 // generate the address generation required.
31072 LLD->getBasePtr().getOpcode() == ISD::TargetFrameIndex ||
31073 RLD->getBasePtr().getOpcode() == ISD::TargetFrameIndex ||
31074 !TLI.isOperationLegalOrCustom(Op: TheSelect->getOpcode(),
31075 VT: LLD->getBasePtr().getValueType()))
31076 return false;
31077
31078 // The loads must not depend on one another.
31079 if (LLD->isPredecessorOf(N: RLD) || RLD->isPredecessorOf(N: LLD))
31080 return false;
31081
31082 // Check that the select condition doesn't reach either load. If so,
31083 // folding this will induce a cycle into the DAG. If not, this is safe to
31084 // xform, so create a select of the addresses.
31085
31086 SmallPtrSet<const SDNode *, 32> Visited;
31087 SmallVector<const SDNode *, 16> Worklist;
31088
31089 // Always fail if LLD and RLD are not independent. TheSelect is a
31090 // predecessor to all Nodes in question so we need not search past it.
31091
31092 Visited.insert(Ptr: TheSelect);
31093 Worklist.push_back(Elt: LLD);
31094 Worklist.push_back(Elt: RLD);
31095
31096 if (SDNode::hasPredecessorHelper(N: LLD, Visited, Worklist) ||
31097 SDNode::hasPredecessorHelper(N: RLD, Visited, Worklist))
31098 return false;
31099
31100 SDValue Addr;
31101 if (TheSelect->getOpcode() == ISD::SELECT) {
31102 // We cannot do this optimization if any pair of {RLD, LLD} is a
31103 // predecessor to {RLD, LLD, CondNode}. As we've already compared the
31104 // Loads, we only need to check if CondNode is a successor to one of the
31105 // loads. We can further avoid this if there's no use of their chain
31106 // value.
31107 SDNode *CondNode = TheSelect->getOperand(Num: 0).getNode();
31108 Worklist.push_back(Elt: CondNode);
31109
31110 if ((LLD->hasAnyUseOfValue(Value: 1) &&
31111 SDNode::hasPredecessorHelper(N: LLD, Visited, Worklist)) ||
31112 (RLD->hasAnyUseOfValue(Value: 1) &&
31113 SDNode::hasPredecessorHelper(N: RLD, Visited, Worklist)))
31114 return false;
31115
31116 // If the condition is poison, originally this would result in a poison
31117 // result. After the transform, this would result in a load of poison,
31118 // which is UB. Freeze the condition to prevent this.
31119 Addr = DAG.getSelect(DL: SDLoc(TheSelect), VT: LLD->getBasePtr().getValueType(),
31120 Cond: DAG.getFreeze(V: TheSelect->getOperand(Num: 0)),
31121 LHS: LLD->getBasePtr(), RHS: RLD->getBasePtr());
31122 } else { // Otherwise SELECT_CC
31123 // We cannot do this optimization if any pair of {RLD, LLD} is a
31124 // predecessor to {RLD, LLD, CondLHS, CondRHS}. As we've already compared
31125 // the Loads, we only need to check if CondLHS/CondRHS is a successor to
31126 // one of the loads. We can further avoid this if there's no use of their
31127 // chain value.
31128
31129 SDNode *CondLHS = TheSelect->getOperand(Num: 0).getNode();
31130 SDNode *CondRHS = TheSelect->getOperand(Num: 1).getNode();
31131 Worklist.push_back(Elt: CondLHS);
31132 Worklist.push_back(Elt: CondRHS);
31133
31134 if ((LLD->hasAnyUseOfValue(Value: 1) &&
31135 SDNode::hasPredecessorHelper(N: LLD, Visited, Worklist)) ||
31136 (RLD->hasAnyUseOfValue(Value: 1) &&
31137 SDNode::hasPredecessorHelper(N: RLD, Visited, Worklist)))
31138 return false;
31139
31140 SDValue FrozenOp0 = DAG.getFreeze(V: TheSelect->getOperand(Num: 0));
31141 SDValue FrozenOp1 = DAG.getFreeze(V: TheSelect->getOperand(Num: 1));
31142 Addr = DAG.getNode(Opcode: ISD::SELECT_CC, DL: SDLoc(TheSelect),
31143 VT: LLD->getBasePtr().getValueType(), N1: FrozenOp0, N2: FrozenOp1,
31144 N3: LLD->getBasePtr(), N4: RLD->getBasePtr(),
31145 N5: TheSelect->getOperand(Num: 4));
31146 }
31147
31148 SDValue Load;
31149 // It is safe to replace the two loads if they have different alignments,
31150 // but the new load must be the minimum (most restrictive) alignment of the
31151 // inputs.
31152 Align Alignment = std::min(a: LLD->getAlign(), b: RLD->getAlign());
31153 unsigned AddrSpace = LLD->getAddressSpace();
31154 assert(AddrSpace == RLD->getAddressSpace());
31155
31156 MachineMemOperand::Flags MMOFlags = LLD->getMemOperand()->getFlags();
31157 if (!RLD->isInvariant())
31158 MMOFlags &= ~MachineMemOperand::MOInvariant;
31159 if (!RLD->isDereferenceable())
31160 MMOFlags &= ~MachineMemOperand::MODereferenceable;
31161 if (LLD->getExtensionType() == ISD::NON_EXTLOAD) {
31162 // FIXME: Discards pointer and AA info.
31163 Load = DAG.getLoad(VT: TheSelect->getValueType(ResNo: 0), dl: SDLoc(TheSelect),
31164 Chain: LLD->getChain(), Ptr: Addr, PtrInfo: MachinePointerInfo(AddrSpace),
31165 Alignment, MMOFlags);
31166 } else {
31167 // FIXME: Discards pointer and AA info.
31168 Load = DAG.getExtLoad(
31169 ExtType: LLD->getExtensionType() == ISD::EXTLOAD ? RLD->getExtensionType()
31170 : LLD->getExtensionType(),
31171 dl: SDLoc(TheSelect), VT: TheSelect->getValueType(ResNo: 0), Chain: LLD->getChain(), Ptr: Addr,
31172 PtrInfo: MachinePointerInfo(AddrSpace), MemVT: LLD->getMemoryVT(), Alignment,
31173 MMOFlags);
31174 }
31175
31176 // Users of the select now use the result of the load.
31177 CombineTo(N: TheSelect, Res: Load);
31178
31179 // Users of the old loads now use the new load's chain. We know the
31180 // old-load value is dead now.
31181 CombineTo(N: LHS.getNode(), Res0: Load.getValue(R: 0), Res1: Load.getValue(R: 1));
31182 CombineTo(N: RHS.getNode(), Res0: Load.getValue(R: 0), Res1: Load.getValue(R: 1));
31183 return true;
31184 }
31185
31186 return false;
31187}
31188
31189/// Try to fold an expression of the form (N0 cond N1) ? N2 : N3 to a shift and
31190/// bitwise 'and'.
31191SDValue DAGCombiner::foldSelectCCToShiftAnd(const SDLoc &DL, SDValue N0,
31192 SDValue N1, SDValue N2, SDValue N3,
31193 ISD::CondCode CC) {
31194 // If this is a select where the false operand is zero and the compare is a
31195 // check of the sign bit, see if we can perform the "gzip trick":
31196 // select_cc setlt X, 0, A, 0 -> and (sra X, size(X)-1), A
31197 // select_cc setgt X, 0, A, 0 -> and (not (sra X, size(X)-1)), A
31198 EVT XType = N0.getValueType();
31199 EVT AType = N2.getValueType();
31200 if (!isNullConstant(V: N3) || !XType.bitsGE(VT: AType))
31201 return SDValue();
31202
31203 // If the comparison is testing for a positive value, we have to invert
31204 // the sign bit mask, so only do that transform if the target has a bitwise
31205 // 'and not' instruction (the invert is free).
31206 if (CC == ISD::SETGT && TLI.hasAndNot(X: N2)) {
31207 // (X > -1) ? A : 0
31208 // (X > 0) ? X : 0 <-- This is canonical signed max.
31209 if (!(isAllOnesConstant(V: N1) || (isNullConstant(V: N1) && N0 == N2)))
31210 return SDValue();
31211 } else if (CC == ISD::SETLT) {
31212 // (X < 0) ? A : 0
31213 // (X < 1) ? X : 0 <-- This is un-canonicalized signed min.
31214 if (!(isNullConstant(V: N1) || (isOneConstant(V: N1) && N0 == N2)))
31215 return SDValue();
31216 } else {
31217 return SDValue();
31218 }
31219
31220 // and (sra X, size(X)-1), A -> "and (srl X, C2), A" iff A is a single-bit
31221 // constant.
31222 auto *N2C = dyn_cast<ConstantSDNode>(Val: N2.getNode());
31223 if (N2C && ((N2C->getAPIntValue() & (N2C->getAPIntValue() - 1)) == 0)) {
31224 unsigned ShCt = XType.getSizeInBits() - N2C->getAPIntValue().logBase2() - 1;
31225 if (!TLI.shouldAvoidTransformToShift(VT: XType, Amount: ShCt)) {
31226 SDValue ShiftAmt = DAG.getShiftAmountConstant(Val: ShCt, VT: XType, DL);
31227 SDValue Shift = DAG.getNode(Opcode: ISD::SRL, DL, VT: XType, N1: N0, N2: ShiftAmt);
31228 AddToWorklist(N: Shift.getNode());
31229
31230 if (XType.bitsGT(VT: AType)) {
31231 Shift = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: AType, Operand: Shift);
31232 AddToWorklist(N: Shift.getNode());
31233 }
31234
31235 if (CC == ISD::SETGT)
31236 Shift = DAG.getNOT(DL, Val: Shift, VT: AType);
31237
31238 return DAG.getNode(Opcode: ISD::AND, DL, VT: AType, N1: Shift, N2);
31239 }
31240 }
31241
31242 unsigned ShCt = XType.getSizeInBits() - 1;
31243 if (TLI.shouldAvoidTransformToShift(VT: XType, Amount: ShCt))
31244 return SDValue();
31245
31246 SDValue ShiftAmt = DAG.getShiftAmountConstant(Val: ShCt, VT: XType, DL);
31247 SDValue Shift = DAG.getNode(Opcode: ISD::SRA, DL, VT: XType, N1: N0, N2: ShiftAmt);
31248 AddToWorklist(N: Shift.getNode());
31249
31250 if (XType.bitsGT(VT: AType)) {
31251 Shift = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: AType, Operand: Shift);
31252 AddToWorklist(N: Shift.getNode());
31253 }
31254
31255 if (CC == ISD::SETGT)
31256 Shift = DAG.getNOT(DL, Val: Shift, VT: AType);
31257
31258 return DAG.getNode(Opcode: ISD::AND, DL, VT: AType, N1: Shift, N2);
31259}
31260
31261// Fold select(cc, binop(), binop()) -> binop(select(), select()) etc.
31262SDValue DAGCombiner::foldSelectOfBinops(SDNode *N) {
31263 SDValue N0 = N->getOperand(Num: 0);
31264 SDValue N1 = N->getOperand(Num: 1);
31265 SDValue N2 = N->getOperand(Num: 2);
31266 SDLoc DL(N);
31267
31268 unsigned BinOpc = N1.getOpcode();
31269 if (!TLI.isBinOp(Opcode: BinOpc) || (N2.getOpcode() != BinOpc) ||
31270 (N1.getResNo() != N2.getResNo()))
31271 return SDValue();
31272
31273 // The use checks are intentionally on SDNode because we may be dealing
31274 // with opcodes that produce more than one SDValue.
31275 if (!N1->hasOneUse() || !N2->hasOneUse())
31276 return SDValue();
31277
31278 // Binops may include opcodes that return multiple values, so all values
31279 // must be created/propagated from the newly created binops below.
31280 SDVTList OpVTs = N1->getVTList();
31281
31282 // Fold select(cond, binop(x, y), binop(z, y))
31283 // --> binop(select(cond, x, z), y)
31284 if (N1.getOperand(i: 1) == N2.getOperand(i: 1)) {
31285 SDValue N10 = N1.getOperand(i: 0);
31286 SDValue N20 = N2.getOperand(i: 0);
31287 SDValue NewSel = DAG.getSelect(DL, VT: N10.getValueType(), Cond: N0, LHS: N10, RHS: N20);
31288 SDNodeFlags Flags = N1->getFlags() & N2->getFlags();
31289 SDValue NewBinOp =
31290 DAG.getNode(Opcode: BinOpc, DL, VTList: OpVTs, Ops: {NewSel, N1.getOperand(i: 1)}, Flags);
31291 return SDValue(NewBinOp.getNode(), N1.getResNo());
31292 }
31293
31294 // Fold select(cond, binop(x, y), binop(x, z))
31295 // --> binop(x, select(cond, y, z))
31296 if (N1.getOperand(i: 0) == N2.getOperand(i: 0)) {
31297 SDValue N11 = N1.getOperand(i: 1);
31298 SDValue N21 = N2.getOperand(i: 1);
31299 // Second op VT might be different (e.g. shift amount type)
31300 if (N11.getValueType() == N21.getValueType()) {
31301 SDValue NewSel = DAG.getSelect(DL, VT: N11.getValueType(), Cond: N0, LHS: N11, RHS: N21);
31302 SDNodeFlags Flags = N1->getFlags() & N2->getFlags();
31303 SDValue NewBinOp =
31304 DAG.getNode(Opcode: BinOpc, DL, VTList: OpVTs, Ops: {N1.getOperand(i: 0), NewSel}, Flags);
31305 return SDValue(NewBinOp.getNode(), N1.getResNo());
31306 }
31307 }
31308
31309 // TODO: Handle isCommutativeBinOp patterns as well?
31310 return SDValue();
31311}
31312
31313// Transform (fneg/fabs (bitconvert x)) to avoid loading constant pool values.
31314SDValue DAGCombiner::foldSignChangeInBitcast(SDNode *N) {
31315 SDValue N0 = N->getOperand(Num: 0);
31316 EVT VT = N->getValueType(ResNo: 0);
31317 bool IsFabs = N->getOpcode() == ISD::FABS;
31318 bool IsFree = IsFabs ? TLI.isFAbsFree(VT) : TLI.isFNegFree(VT);
31319
31320 if (IsFree || N0.getOpcode() != ISD::BITCAST || !N0.hasOneUse())
31321 return SDValue();
31322
31323 SDValue Int = N0.getOperand(i: 0);
31324 EVT IntVT = Int.getValueType();
31325
31326 // The operand to cast should be integer.
31327 if (!IntVT.isInteger() || IntVT.isVector())
31328 return SDValue();
31329
31330 // (fneg (bitconvert x)) -> (bitconvert (xor x sign))
31331 // (fabs (bitconvert x)) -> (bitconvert (and x ~sign))
31332 APInt SignMask;
31333 if (N0.getValueType().isVector()) {
31334 // For vector, create a sign mask (0x80...) or its inverse (for fabs,
31335 // 0x7f...) per element and splat it.
31336 SignMask = APInt::getSignMask(BitWidth: N0.getScalarValueSizeInBits());
31337 if (IsFabs)
31338 SignMask = ~SignMask;
31339 SignMask = APInt::getSplat(NewLen: IntVT.getSizeInBits(), V: SignMask);
31340 } else {
31341 // For scalar, just use the sign mask (0x80... or the inverse, 0x7f...)
31342 SignMask = APInt::getSignMask(BitWidth: IntVT.getSizeInBits());
31343 if (IsFabs)
31344 SignMask = ~SignMask;
31345 }
31346 SDLoc DL(N0);
31347 Int = DAG.getNode(Opcode: IsFabs ? ISD::AND : ISD::XOR, DL, VT: IntVT, N1: Int,
31348 N2: DAG.getConstant(Val: SignMask, DL, VT: IntVT));
31349 AddToWorklist(N: Int.getNode());
31350 return DAG.getBitcast(VT, V: Int);
31351}
31352
31353/// Turn "(a cond b) ? 1.0f : 2.0f" into "load (tmp + ((a cond b) ? 0 : 4)"
31354/// where "tmp" is a constant pool entry containing an array with 1.0 and 2.0
31355/// in it. This may be a win when the constant is not otherwise available
31356/// because it replaces two constant pool loads with one.
31357SDValue DAGCombiner::convertSelectOfFPConstantsToLoadOffset(
31358 const SDLoc &DL, SDValue N0, SDValue N1, SDValue N2, SDValue N3,
31359 ISD::CondCode CC) {
31360 if (!TLI.reduceSelectOfFPConstantLoads(CmpOpVT: N0.getValueType()))
31361 return SDValue();
31362
31363 // If we are before legalize types, we want the other legalization to happen
31364 // first (for example, to avoid messing with soft float).
31365 auto *TV = dyn_cast<ConstantFPSDNode>(Val&: N2);
31366 auto *FV = dyn_cast<ConstantFPSDNode>(Val&: N3);
31367 EVT VT = N2.getValueType();
31368 if (!TV || !FV || !TLI.isTypeLegal(VT))
31369 return SDValue();
31370
31371 // If a constant can be materialized without loads, this does not make sense.
31372 if (TLI.getOperationAction(Op: ISD::ConstantFP, VT) == TargetLowering::Legal ||
31373 TLI.isFPImmLegal(TV->getValueAPF(), TV->getValueType(ResNo: 0), ForCodeSize) ||
31374 TLI.isFPImmLegal(FV->getValueAPF(), FV->getValueType(ResNo: 0), ForCodeSize))
31375 return SDValue();
31376
31377 // If both constants have multiple uses, then we won't need to do an extra
31378 // load. The values are likely around in registers for other users.
31379 if (!TV->hasOneUse() && !FV->hasOneUse())
31380 return SDValue();
31381
31382 Constant *Elts[] = { const_cast<ConstantFP*>(FV->getConstantFPValue()),
31383 const_cast<ConstantFP*>(TV->getConstantFPValue()) };
31384 Type *FPTy = Elts[0]->getType();
31385 const DataLayout &TD = DAG.getDataLayout();
31386
31387 // Create a ConstantArray of the two constants.
31388 Constant *CA = ConstantArray::get(T: ArrayType::get(ElementType: FPTy, NumElements: 2), V: Elts);
31389 SDValue CPIdx = DAG.getConstantPool(C: CA, VT: TLI.getPointerTy(DL: DAG.getDataLayout()),
31390 Align: TD.getPrefTypeAlign(Ty: FPTy));
31391 Align Alignment = cast<ConstantPoolSDNode>(Val&: CPIdx)->getAlign();
31392
31393 // Get offsets to the 0 and 1 elements of the array, so we can select between
31394 // them.
31395 SDValue Zero = DAG.getIntPtrConstant(Val: 0, DL);
31396 unsigned EltSize = (unsigned)TD.getTypeAllocSize(Ty: Elts[0]->getType());
31397 SDValue One = DAG.getIntPtrConstant(Val: EltSize, DL: SDLoc(FV));
31398 SDValue Cond =
31399 DAG.getSetCC(DL, VT: getSetCCResultType(VT: N0.getValueType()), LHS: N0, RHS: N1, Cond: CC);
31400 AddToWorklist(N: Cond.getNode());
31401 SDValue CstOffset = DAG.getSelect(DL, VT: Zero.getValueType(), Cond, LHS: One, RHS: Zero);
31402 AddToWorklist(N: CstOffset.getNode());
31403 CPIdx = DAG.getNode(Opcode: ISD::ADD, DL, VT: CPIdx.getValueType(), N1: CPIdx, N2: CstOffset);
31404 AddToWorklist(N: CPIdx.getNode());
31405 return DAG.getLoad(VT: TV->getValueType(ResNo: 0), dl: DL, Chain: DAG.getEntryNode(), Ptr: CPIdx,
31406 PtrInfo: MachinePointerInfo::getConstantPool(
31407 MF&: DAG.getMachineFunction()), Alignment);
31408}
31409
31410/// Simplify an expression of the form (N0 cond N1) ? N2 : N3
31411/// where 'cond' is the comparison specified by CC.
31412SDValue DAGCombiner::SimplifySelectCC(const SDLoc &DL, SDValue N0, SDValue N1,
31413 SDValue N2, SDValue N3, ISD::CondCode CC,
31414 bool NotExtCompare) {
31415 // (x ? y : y) -> y.
31416 if (N2 == N3) return N2;
31417
31418 EVT CmpOpVT = N0.getValueType();
31419 EVT CmpResVT = getSetCCResultType(VT: CmpOpVT);
31420 EVT VT = N2.getValueType();
31421 auto *N1C = dyn_cast<ConstantSDNode>(Val: N1.getNode());
31422 auto *N2C = dyn_cast<ConstantSDNode>(Val: N2.getNode());
31423 auto *N3C = dyn_cast<ConstantSDNode>(Val: N3.getNode());
31424
31425 // Determine if the condition we're dealing with is constant.
31426 if (SDValue SCC = DAG.FoldSetCC(VT: CmpResVT, N1: N0, N2: N1, Cond: CC, dl: DL)) {
31427 AddToWorklist(N: SCC.getNode());
31428 if (auto *SCCC = dyn_cast<ConstantSDNode>(Val&: SCC)) {
31429 // fold select_cc true, x, y -> x
31430 // fold select_cc false, x, y -> y
31431 return !(SCCC->isZero()) ? N2 : N3;
31432 }
31433 }
31434
31435 if (SDValue V =
31436 convertSelectOfFPConstantsToLoadOffset(DL, N0, N1, N2, N3, CC))
31437 return V;
31438
31439 if (SDValue V = foldSelectCCToShiftAnd(DL, N0, N1, N2, N3, CC))
31440 return V;
31441
31442 // fold (select_cc seteq (and x, y), 0, 0, A) -> (and (sra (shl x)) A)
31443 // where y is has a single bit set.
31444 // A plaintext description would be, we can turn the SELECT_CC into an AND
31445 // when the condition can be materialized as an all-ones register. Any
31446 // single bit-test can be materialized as an all-ones register with
31447 // shift-left and shift-right-arith.
31448 if (CC == ISD::SETEQ && N0->getOpcode() == ISD::AND &&
31449 N0->getValueType(ResNo: 0) == VT && isNullConstant(V: N1) && isNullConstant(V: N2)) {
31450 SDValue AndLHS = N0->getOperand(Num: 0);
31451 auto *ConstAndRHS = dyn_cast<ConstantSDNode>(Val: N0->getOperand(Num: 1));
31452 if (ConstAndRHS && ConstAndRHS->getAPIntValue().isPowerOf2()) {
31453 // Shift the tested bit over the sign bit.
31454 const APInt &AndMask = ConstAndRHS->getAPIntValue();
31455 if (TLI.shouldFoldSelectWithSingleBitTest(VT, AndMask)) {
31456 unsigned ShCt = AndMask.getBitWidth() - 1;
31457 SDValue ShlAmt = DAG.getShiftAmountConstant(Val: AndMask.countl_zero(), VT,
31458 DL: SDLoc(AndLHS));
31459 SDValue Shl = DAG.getNode(Opcode: ISD::SHL, DL: SDLoc(N0), VT, N1: AndLHS, N2: ShlAmt);
31460
31461 // Now arithmetic right shift it all the way over, so the result is
31462 // either all-ones, or zero.
31463 SDValue ShrAmt = DAG.getShiftAmountConstant(Val: ShCt, VT, DL: SDLoc(Shl));
31464 SDValue Shr = DAG.getNode(Opcode: ISD::SRA, DL: SDLoc(N0), VT, N1: Shl, N2: ShrAmt);
31465
31466 return DAG.getNode(Opcode: ISD::AND, DL, VT, N1: Shr, N2: N3);
31467 }
31468 }
31469 }
31470
31471 // fold select C, 16, 0 -> shl C, 4
31472 bool Fold = N2C && isNullConstant(V: N3) && N2C->getAPIntValue().isPowerOf2();
31473 bool Swap = N3C && isNullConstant(V: N2) && N3C->getAPIntValue().isPowerOf2();
31474
31475 if ((Fold || Swap) &&
31476 TLI.getBooleanContents(Type: CmpOpVT) ==
31477 TargetLowering::ZeroOrOneBooleanContent &&
31478 (!LegalOperations || TLI.isOperationLegal(Op: ISD::SETCC, VT: CmpOpVT)) &&
31479 TLI.convertSelectOfConstantsToMath(VT)) {
31480
31481 if (Swap) {
31482 CC = ISD::getSetCCInverse(Operation: CC, Type: CmpOpVT);
31483 std::swap(a&: N2C, b&: N3C);
31484 }
31485
31486 // If the caller doesn't want us to simplify this into a zext of a compare,
31487 // don't do it.
31488 if (NotExtCompare && N2C->isOne())
31489 return SDValue();
31490
31491 SDValue Temp, SCC;
31492 // zext (setcc n0, n1)
31493 if (LegalTypes) {
31494 SCC = DAG.getSetCC(DL, VT: CmpResVT, LHS: N0, RHS: N1, Cond: CC);
31495 Temp = DAG.getZExtOrTrunc(Op: SCC, DL: SDLoc(N2), VT);
31496 } else {
31497 SCC = DAG.getSetCC(DL: SDLoc(N0), VT: MVT::i1, LHS: N0, RHS: N1, Cond: CC);
31498 Temp = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL: SDLoc(N2), VT, Operand: SCC);
31499 }
31500
31501 AddToWorklist(N: SCC.getNode());
31502 AddToWorklist(N: Temp.getNode());
31503
31504 if (N2C->isOne())
31505 return Temp;
31506
31507 unsigned ShCt = N2C->getAPIntValue().logBase2();
31508 if (TLI.shouldAvoidTransformToShift(VT, Amount: ShCt))
31509 return SDValue();
31510
31511 // shl setcc result by log2 n2c
31512 return DAG.getNode(
31513 Opcode: ISD::SHL, DL, VT: N2.getValueType(), N1: Temp,
31514 N2: DAG.getShiftAmountConstant(Val: ShCt, VT: N2.getValueType(), DL: SDLoc(Temp)));
31515 }
31516
31517 // select_cc seteq X, 0, sizeof(X), ctlz(X) -> ctlz(X)
31518 // select_cc seteq X, 0, sizeof(X), ctlz_zero_poison(X) -> ctlz(X)
31519 // select_cc seteq X, 0, sizeof(X), cttz(X) -> cttz(X)
31520 // select_cc seteq X, 0, sizeof(X), cttz_zero_poison(X) -> cttz(X)
31521 // select_cc setne X, 0, ctlz(X), sizeof(X) -> ctlz(X)
31522 // select_cc setne X, 0, ctlz_zero_poison(X), sizeof(X) -> ctlz(X)
31523 // select_cc setne X, 0, cttz(X), sizeof(X) -> cttz(X)
31524 // select_cc setne X, 0, cttz_zero_poison(X), sizeof(X) -> cttz(X)
31525 if (N1C && N1C->isZero() && (CC == ISD::SETEQ || CC == ISD::SETNE)) {
31526 SDValue ValueOnZero = N2;
31527 SDValue Count = N3;
31528 // If the condition is NE instead of E, swap the operands.
31529 if (CC == ISD::SETNE)
31530 std::swap(a&: ValueOnZero, b&: Count);
31531 // Check if the value on zero is a constant equal to the bits in the type.
31532 if (auto *ValueOnZeroC = dyn_cast<ConstantSDNode>(Val&: ValueOnZero)) {
31533 if (ValueOnZeroC->getAPIntValue() == VT.getSizeInBits()) {
31534 // If the other operand is cttz/cttz_zero_poison of N0, and cttz is
31535 // legal, combine to just cttz.
31536 if ((Count.getOpcode() == ISD::CTTZ ||
31537 Count.getOpcode() == ISD::CTTZ_ZERO_POISON) &&
31538 N0 == Count.getOperand(i: 0) &&
31539 (!LegalOperations || TLI.isOperationLegal(Op: ISD::CTTZ, VT)))
31540 return DAG.getNode(Opcode: ISD::CTTZ, DL, VT, Operand: N0);
31541 // If the other operand is ctlz/ctlz_zero_poison of N0, and ctlz is
31542 // legal, combine to just ctlz.
31543 if ((Count.getOpcode() == ISD::CTLZ ||
31544 Count.getOpcode() == ISD::CTLZ_ZERO_POISON) &&
31545 N0 == Count.getOperand(i: 0) &&
31546 (!LegalOperations || TLI.isOperationLegal(Op: ISD::CTLZ, VT)))
31547 return DAG.getNode(Opcode: ISD::CTLZ, DL, VT, Operand: N0);
31548 }
31549 }
31550 }
31551
31552 // Fold select_cc setgt X, -1, C, ~C -> xor (ashr X, BW-1), C
31553 // Fold select_cc setlt X, 0, C, ~C -> xor (ashr X, BW-1), ~C
31554 if (!NotExtCompare && N1C && N2C && N3C &&
31555 N2C->getAPIntValue() == ~N3C->getAPIntValue() &&
31556 ((N1C->isAllOnes() && CC == ISD::SETGT) ||
31557 (N1C->isZero() && CC == ISD::SETLT)) &&
31558 !TLI.shouldAvoidTransformToShift(VT, Amount: CmpOpVT.getScalarSizeInBits() - 1)) {
31559 SDValue ASHR =
31560 DAG.getNode(Opcode: ISD::SRA, DL, VT: CmpOpVT, N1: N0,
31561 N2: DAG.getShiftAmountConstant(
31562 Val: CmpOpVT.getScalarSizeInBits() - 1, VT: CmpOpVT, DL));
31563 return DAG.getNode(Opcode: ISD::XOR, DL, VT, N1: DAG.getSExtOrTrunc(Op: ASHR, DL, VT),
31564 N2: DAG.getSExtOrTrunc(Op: CC == ISD::SETLT ? N3 : N2, DL, VT));
31565 }
31566
31567 // Fold sign pattern select_cc setgt X, -1, 1, -1 -> or (ashr X, BW-1), 1
31568 if (CC == ISD::SETGT && N1C && N2C && N3C && N1C->isAllOnes() &&
31569 N2C->isOne() && N3C->isAllOnes() &&
31570 !TLI.shouldAvoidTransformToShift(VT: CmpOpVT,
31571 Amount: CmpOpVT.getScalarSizeInBits() - 1)) {
31572 SDValue ASHR =
31573 DAG.getNode(Opcode: ISD::SRA, DL, VT: CmpOpVT, N1: N0,
31574 N2: DAG.getShiftAmountConstant(
31575 Val: CmpOpVT.getScalarSizeInBits() - 1, VT: CmpOpVT, DL));
31576 return DAG.getNode(Opcode: ISD::OR, DL, VT, N1: DAG.getSExtOrTrunc(Op: ASHR, DL, VT),
31577 N2: DAG.getConstant(Val: 1, DL, VT));
31578 }
31579
31580 if (SDValue S = PerformMinMaxFpToSatCombine(N0, N1, N2, N3, CC, DAG))
31581 return S;
31582 if (SDValue S = PerformUMinFpToSatCombine(N0, N1, N2, N3, CC, DAG))
31583 return S;
31584 if (SDValue ABD = foldSelectToABD(LHS: N0, RHS: N1, True: N2, False: N3, CC, DL))
31585 return ABD;
31586
31587 return SDValue();
31588}
31589
31590static SDValue matchMergedBFX(SDValue Root, SelectionDAG &DAG,
31591 const TargetLowering &TLI) {
31592 // Match a pattern such as:
31593 // (X | (X >> C0) | (X >> C1) | ...) & Mask
31594 // This extracts contiguous parts of X and ORs them together before comparing.
31595 // We can optimize this so that we directly check (X & SomeMask) instead,
31596 // eliminating the shifts.
31597
31598 EVT VT = Root.getValueType();
31599
31600 // TODO: Support vectors?
31601 if (!VT.isScalarInteger() || Root.getOpcode() != ISD::AND)
31602 return SDValue();
31603
31604 SDValue N0 = Root.getOperand(i: 0);
31605 SDValue N1 = Root.getOperand(i: 1);
31606
31607 if (N0.getOpcode() != ISD::OR || !isa<ConstantSDNode>(Val: N1))
31608 return SDValue();
31609
31610 APInt RootMask = cast<ConstantSDNode>(Val&: N1)->getAsAPIntVal();
31611
31612 SDValue Src;
31613 const auto IsSrc = [&](SDValue V) {
31614 if (!Src) {
31615 Src = V;
31616 return true;
31617 }
31618
31619 return Src == V;
31620 };
31621
31622 SmallVector<SDValue> Worklist = {N0};
31623 APInt PartsMask(VT.getSizeInBits(), 0);
31624 while (!Worklist.empty()) {
31625 SDValue V = Worklist.pop_back_val();
31626 if (!V.hasOneUse() && (Src && Src != V))
31627 return SDValue();
31628
31629 if (V.getOpcode() == ISD::OR) {
31630 Worklist.push_back(Elt: V.getOperand(i: 0));
31631 Worklist.push_back(Elt: V.getOperand(i: 1));
31632 continue;
31633 }
31634
31635 if (V.getOpcode() == ISD::SRL) {
31636 SDValue ShiftSrc = V.getOperand(i: 0);
31637 SDValue ShiftAmt = V.getOperand(i: 1);
31638
31639 if (!IsSrc(ShiftSrc) || !isa<ConstantSDNode>(Val: ShiftAmt))
31640 return SDValue();
31641
31642 auto ShiftAmtVal = cast<ConstantSDNode>(Val&: ShiftAmt)->getAsZExtVal();
31643 if (ShiftAmtVal > RootMask.getBitWidth())
31644 return SDValue();
31645
31646 PartsMask |= (RootMask << ShiftAmtVal);
31647 continue;
31648 }
31649
31650 if (IsSrc(V)) {
31651 PartsMask |= RootMask;
31652 continue;
31653 }
31654
31655 return SDValue();
31656 }
31657
31658 if (!Src)
31659 return SDValue();
31660
31661 SDLoc DL(Root);
31662 return DAG.getNode(Opcode: ISD::AND, DL, VT,
31663 Ops: {Src, DAG.getConstant(Val: PartsMask, DL, VT)});
31664}
31665
31666/// This is a stub for TargetLowering::SimplifySetCC.
31667SDValue DAGCombiner::SimplifySetCC(EVT VT, SDValue N0, SDValue N1,
31668 ISD::CondCode Cond, const SDLoc &DL,
31669 bool foldBooleans) {
31670 TargetLowering::DAGCombinerInfo
31671 DagCombineInfo(DAG, Level, false, this);
31672 if (SDValue C =
31673 TLI.SimplifySetCC(VT, N0, N1, Cond, foldBooleans, DCI&: DagCombineInfo, dl: DL))
31674 return C;
31675
31676 if (ISD::isIntEqualitySetCC(Code: Cond) && N0.getOpcode() == ISD::AND &&
31677 isNullConstant(V: N1)) {
31678
31679 if (SDValue Res = matchMergedBFX(Root: N0, DAG, TLI))
31680 return DAG.getSetCC(DL, VT, LHS: Res, RHS: N1, Cond);
31681 }
31682
31683 return SDValue();
31684}
31685
31686/// Given an ISD::SDIV node expressing a divide by constant, return
31687/// a DAG expression to select that will generate the same value by multiplying
31688/// by a magic number.
31689/// Ref: "Hacker's Delight" or "The PowerPC Compiler Writer's Guide".
31690SDValue DAGCombiner::BuildSDIV(SDNode *N) {
31691 // when optimising for minimum size, we don't want to expand a div to a mul
31692 // and a shift.
31693 if (DAG.getMachineFunction().getFunction().hasMinSize())
31694 return SDValue();
31695
31696 SmallVector<SDNode *, 8> Built;
31697 if (SDValue S = TLI.BuildSDIV(N, DAG, IsAfterLegalization: LegalOperations, IsAfterLegalTypes: LegalTypes, Created&: Built)) {
31698 for (SDNode *N : Built)
31699 AddToWorklist(N);
31700 return S;
31701 }
31702
31703 return SDValue();
31704}
31705
31706/// Given an ISD::SDIV node expressing a divide by constant power of 2, return a
31707/// DAG expression that will generate the same value by right shifting.
31708SDValue DAGCombiner::BuildSDIVPow2(SDNode *N) {
31709 ConstantSDNode *C = isConstOrConstSplat(N: N->getOperand(Num: 1));
31710 if (!C)
31711 return SDValue();
31712
31713 // Avoid division by zero.
31714 if (C->isZero())
31715 return SDValue();
31716
31717 SmallVector<SDNode *, 8> Built;
31718 if (SDValue S = TLI.BuildSDIVPow2(N, Divisor: C->getAPIntValue(), DAG, Created&: Built)) {
31719 for (SDNode *N : Built)
31720 AddToWorklist(N);
31721 return S;
31722 }
31723
31724 return SDValue();
31725}
31726
31727/// Given an ISD::UDIV node expressing a divide by constant, return a DAG
31728/// expression that will generate the same value by multiplying by a magic
31729/// number.
31730/// Ref: "Hacker's Delight" or "The PowerPC Compiler Writer's Guide".
31731SDValue DAGCombiner::BuildUDIV(SDNode *N) {
31732 // when optimising for minimum size, we don't want to expand a div to a mul
31733 // and a shift.
31734 if (DAG.getMachineFunction().getFunction().hasMinSize())
31735 return SDValue();
31736
31737 SmallVector<SDNode *, 8> Built;
31738 if (SDValue S = TLI.BuildUDIV(N, DAG, IsAfterLegalization: LegalOperations, IsAfterLegalTypes: LegalTypes, Created&: Built)) {
31739 for (SDNode *N : Built)
31740 AddToWorklist(N);
31741 return S;
31742 }
31743
31744 return SDValue();
31745}
31746
31747/// Given an ISD::SREM node expressing a remainder by constant power of 2,
31748/// return a DAG expression that will generate the same value.
31749SDValue DAGCombiner::BuildSREMPow2(SDNode *N) {
31750 ConstantSDNode *C = isConstOrConstSplat(N: N->getOperand(Num: 1));
31751 if (!C)
31752 return SDValue();
31753
31754 // Avoid division by zero.
31755 if (C->isZero())
31756 return SDValue();
31757
31758 SmallVector<SDNode *, 8> Built;
31759 if (SDValue S = TLI.BuildSREMPow2(N, Divisor: C->getAPIntValue(), DAG, Created&: Built)) {
31760 for (SDNode *N : Built)
31761 AddToWorklist(N);
31762 return S;
31763 }
31764
31765 return SDValue();
31766}
31767
31768// This is basically just a port of takeLog2 from InstCombineMulDivRem.cpp
31769//
31770// Returns the node that represents `Log2(Op)`. This may create a new node. If
31771// we are unable to compute `Log2(Op)` its return `SDValue()`.
31772//
31773// All nodes will be created at `DL` and the output will be of type `VT`.
31774//
31775// This will only return `Log2(Op)` if we can prove `Op` is non-zero. Set
31776// `AssumeNonZero` if this function should simply assume (not require proving
31777// `Op` is non-zero).
31778static SDValue takeInexpensiveLog2(SelectionDAG &DAG, const SDLoc &DL, EVT VT,
31779 SDValue Op, unsigned Depth,
31780 bool AssumeNonZero) {
31781 assert(VT.isInteger() && "Only integer types are supported!");
31782
31783 auto PeekThroughCastsAndTrunc = [](SDValue V) {
31784 while (true) {
31785 switch (V.getOpcode()) {
31786 case ISD::TRUNCATE:
31787 case ISD::ZERO_EXTEND:
31788 V = V.getOperand(i: 0);
31789 break;
31790 default:
31791 return V;
31792 }
31793 }
31794 };
31795
31796 if (VT.isScalableVector())
31797 return SDValue();
31798
31799 Op = PeekThroughCastsAndTrunc(Op);
31800
31801 // Helper for determining whether a value is a power-2 constant scalar or a
31802 // vector of such elements.
31803 SmallVector<APInt> Pow2Constants;
31804 auto IsPowerOfTwo = [&Pow2Constants](ConstantSDNode *C) {
31805 if (C->isZero() || C->isOpaque())
31806 return false;
31807 // TODO: We may also be able to support negative powers of 2 here.
31808 if (C->getAPIntValue().isPowerOf2()) {
31809 Pow2Constants.emplace_back(Args: C->getAPIntValue());
31810 return true;
31811 }
31812 return false;
31813 };
31814
31815 if (ISD::matchUnaryPredicate(Op, Match: IsPowerOfTwo, /*AllowUndefs=*/false,
31816 /*AllowTruncation=*/true)) {
31817 if (!VT.isVector())
31818 return DAG.getConstant(Val: Pow2Constants.back().logBase2(), DL, VT);
31819 // We need to create a build vector
31820 if (Op.getOpcode() == ISD::SPLAT_VECTOR)
31821 return DAG.getSplat(VT, DL,
31822 Op: DAG.getConstant(Val: Pow2Constants.back().logBase2(), DL,
31823 VT: VT.getScalarType()));
31824 SmallVector<SDValue> Log2Ops;
31825 for (const APInt &Pow2 : Pow2Constants)
31826 Log2Ops.emplace_back(
31827 Args: DAG.getConstant(Val: Pow2.logBase2(), DL, VT: VT.getScalarType()));
31828 return DAG.getBuildVector(VT, DL, Ops: Log2Ops);
31829 }
31830
31831 if (Depth >= DAG.MaxRecursionDepth)
31832 return SDValue();
31833
31834 auto CastToVT = [&](EVT NewVT, SDValue ToCast) {
31835 // Peek through zero extend. We can't peek through truncates since this
31836 // function is called on a shift amount. We must ensure that all of the bits
31837 // above the original shift amount are zeroed by this function.
31838 while (ToCast.getOpcode() == ISD::ZERO_EXTEND)
31839 ToCast = ToCast.getOperand(i: 0);
31840 EVT CurVT = ToCast.getValueType();
31841 if (NewVT == CurVT)
31842 return ToCast;
31843
31844 if (NewVT.getSizeInBits() == CurVT.getSizeInBits())
31845 return DAG.getBitcast(VT: NewVT, V: ToCast);
31846
31847 return DAG.getZExtOrTrunc(Op: ToCast, DL, VT: NewVT);
31848 };
31849
31850 // log2(X << Y) -> log2(X) + Y
31851 if (Op.getOpcode() == ISD::SHL) {
31852 // 1 << Y and X nuw/nsw << Y are all non-zero.
31853 if (AssumeNonZero || Op->getFlags().hasNoUnsignedWrap() ||
31854 Op->getFlags().hasNoSignedWrap() || isOneConstant(V: Op.getOperand(i: 0)))
31855 if (SDValue LogX = takeInexpensiveLog2(DAG, DL, VT, Op: Op.getOperand(i: 0),
31856 Depth: Depth + 1, AssumeNonZero))
31857 return DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: LogX,
31858 N2: CastToVT(VT, Op.getOperand(i: 1)));
31859 }
31860
31861 // c ? X : Y -> c ? Log2(X) : Log2(Y)
31862 SDValue Cond, TVal, FVal;
31863 if (sd_match(N: Op, P: m_OneUse(P: m_SelectLike(Cond: m_Value(N&: Cond), T: m_Value(N&: TVal),
31864 F: m_Value(N&: FVal))))) {
31865 if (SDValue LogX =
31866 takeInexpensiveLog2(DAG, DL, VT, Op: TVal, Depth: Depth + 1, AssumeNonZero))
31867 if (SDValue LogY =
31868 takeInexpensiveLog2(DAG, DL, VT, Op: FVal, Depth: Depth + 1, AssumeNonZero))
31869 return DAG.getSelect(DL, VT, Cond, LHS: LogX, RHS: LogY);
31870 }
31871
31872 // log2(umin(X, Y)) -> umin(log2(X), log2(Y))
31873 // log2(umax(X, Y)) -> umax(log2(X), log2(Y))
31874 if ((Op.getOpcode() == ISD::UMIN || Op.getOpcode() == ISD::UMAX) &&
31875 Op.hasOneUse()) {
31876 // Use AssumeNonZero as false here. Otherwise we can hit case where
31877 // log2(umax(X, Y)) != umax(log2(X), log2(Y)) (because overflow).
31878 if (SDValue LogX =
31879 takeInexpensiveLog2(DAG, DL, VT, Op: Op.getOperand(i: 0), Depth: Depth + 1,
31880 /*AssumeNonZero*/ false))
31881 if (SDValue LogY =
31882 takeInexpensiveLog2(DAG, DL, VT, Op: Op.getOperand(i: 1), Depth: Depth + 1,
31883 /*AssumeNonZero*/ false))
31884 return DAG.getNode(Opcode: Op.getOpcode(), DL, VT, N1: LogX, N2: LogY);
31885 }
31886
31887 return SDValue();
31888}
31889
31890/// Determines the LogBase2 value for a non-null input value using the
31891/// transform: LogBase2(V) = (EltBits - 1) - ctlz(V).
31892SDValue DAGCombiner::BuildLogBase2(SDValue V, const SDLoc &DL,
31893 bool KnownNonZero, bool InexpensiveOnly,
31894 std::optional<EVT> OutVT) {
31895 EVT VT = OutVT ? *OutVT : V.getValueType();
31896 SDValue InexpensiveLogBase2 =
31897 takeInexpensiveLog2(DAG, DL, VT, Op: V, /*Depth*/ 0, AssumeNonZero: KnownNonZero);
31898 if (InexpensiveLogBase2 || InexpensiveOnly || !DAG.isKnownToBeAPowerOfTwo(Val: V))
31899 return InexpensiveLogBase2;
31900
31901 SDValue Ctlz = DAG.getNode(Opcode: ISD::CTLZ, DL, VT, Operand: V);
31902 SDValue Base = DAG.getConstant(Val: VT.getScalarSizeInBits() - 1, DL, VT);
31903 SDValue LogBase2 = DAG.getNode(Opcode: ISD::SUB, DL, VT, N1: Base, N2: Ctlz);
31904 return LogBase2;
31905}
31906
31907/// Newton iteration for a function: F(X) is X_{i+1} = X_i - F(X_i)/F'(X_i)
31908/// For the reciprocal, we need to find the zero of the function:
31909/// F(X) = 1/X - A [which has a zero at X = 1/A]
31910/// =>
31911/// X_{i+1} = X_i (2 - A X_i) = X_i + X_i (1 - A X_i) [this second form
31912/// does not require additional intermediate precision]
31913/// For the last iteration, put numerator N into it to gain more precision:
31914/// Result = N X_i + X_i (N - N A X_i)
31915SDValue DAGCombiner::BuildDivEstimate(SDValue N, SDValue Op,
31916 SDNodeFlags Flags) {
31917 if (LegalDAG)
31918 return SDValue();
31919
31920 // TODO: Handle extended types?
31921 EVT VT = Op.getValueType();
31922 if (VT.getScalarType() != MVT::f16 && VT.getScalarType() != MVT::f32 &&
31923 VT.getScalarType() != MVT::f64)
31924 return SDValue();
31925
31926 // If estimates are explicitly disabled for this function, we're done.
31927 const Function &F = DAG.getMachineFunction().getFunction();
31928 int Enabled = TLI.getRecipEstimateDivEnabled(VT, MF: F);
31929 if (Enabled == TLI.ReciprocalEstimate::Disabled)
31930 return SDValue();
31931
31932 // Estimates may be explicitly enabled for this type with a custom number of
31933 // refinement steps.
31934 int Iterations = TLI.getDivRefinementSteps(VT, MF: F);
31935 if (SDValue Est = TLI.getRecipEstimate(Operand: Op, DAG, Enabled, RefinementSteps&: Iterations)) {
31936 AddToWorklist(N: Est.getNode());
31937
31938 SDLoc DL(Op);
31939 if (Iterations) {
31940 SDValue FPOne = DAG.getConstantFP(Val: 1.0, DL, VT);
31941
31942 // Newton iterations: Est = Est + Est (N - Arg * Est)
31943 // If this is the last iteration, also multiply by the numerator.
31944 for (int i = 0; i < Iterations; ++i) {
31945 SDValue MulEst = Est;
31946
31947 if (i == Iterations - 1) {
31948 MulEst = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: N, N2: Est, Flags);
31949 AddToWorklist(N: MulEst.getNode());
31950 }
31951
31952 SDValue NewEst = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Op, N2: MulEst, Flags);
31953 AddToWorklist(N: NewEst.getNode());
31954
31955 NewEst = DAG.getNode(Opcode: ISD::FSUB, DL, VT,
31956 N1: (i == Iterations - 1 ? N : FPOne), N2: NewEst, Flags);
31957 AddToWorklist(N: NewEst.getNode());
31958
31959 NewEst = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Est, N2: NewEst, Flags);
31960 AddToWorklist(N: NewEst.getNode());
31961
31962 Est = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: MulEst, N2: NewEst, Flags);
31963 AddToWorklist(N: Est.getNode());
31964 }
31965 } else {
31966 // If no iterations are available, multiply with N.
31967 Est = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Est, N2: N, Flags);
31968 AddToWorklist(N: Est.getNode());
31969 }
31970
31971 return Est;
31972 }
31973
31974 return SDValue();
31975}
31976
31977/// Newton iteration for a function: F(X) is X_{i+1} = X_i - F(X_i)/F'(X_i)
31978/// For the reciprocal sqrt, we need to find the zero of the function:
31979/// F(X) = 1/X^2 - A [which has a zero at X = 1/sqrt(A)]
31980/// =>
31981/// X_{i+1} = X_i (1.5 - A X_i^2 / 2)
31982/// As a result, we precompute A/2 prior to the iteration loop.
31983SDValue DAGCombiner::buildSqrtNROneConst(SDValue Arg, SDValue Est,
31984 unsigned Iterations, bool Reciprocal) {
31985 EVT VT = Arg.getValueType();
31986 SDLoc DL(Arg);
31987 SDValue ThreeHalves = DAG.getConstantFP(Val: 1.5, DL, VT);
31988
31989 // We now need 0.5 * Arg which we can write as (1.5 * Arg - Arg) so that
31990 // this entire sequence requires only one FP constant.
31991 SDValue HalfArg = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: ThreeHalves, N2: Arg);
31992 HalfArg = DAG.getNode(Opcode: ISD::FSUB, DL, VT, N1: HalfArg, N2: Arg);
31993
31994 // Newton iterations: Est = Est * (1.5 - HalfArg * Est * Est)
31995 for (unsigned i = 0; i < Iterations; ++i) {
31996 SDValue NewEst = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Est, N2: Est);
31997 NewEst = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: HalfArg, N2: NewEst);
31998 NewEst = DAG.getNode(Opcode: ISD::FSUB, DL, VT, N1: ThreeHalves, N2: NewEst);
31999 Est = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Est, N2: NewEst);
32000 }
32001
32002 // If non-reciprocal square root is requested, multiply the result by Arg.
32003 if (!Reciprocal)
32004 Est = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Est, N2: Arg);
32005
32006 return Est;
32007}
32008
32009/// Newton iteration for a function: F(X) is X_{i+1} = X_i - F(X_i)/F'(X_i)
32010/// For the reciprocal sqrt, we need to find the zero of the function:
32011/// F(X) = 1/X^2 - A [which has a zero at X = 1/sqrt(A)]
32012/// =>
32013/// X_{i+1} = (-0.5 * X_i) * (A * X_i * X_i + (-3.0))
32014SDValue DAGCombiner::buildSqrtNRTwoConst(SDValue Arg, SDValue Est,
32015 unsigned Iterations, bool Reciprocal) {
32016 EVT VT = Arg.getValueType();
32017 SDLoc DL(Arg);
32018 SDValue MinusThree = DAG.getConstantFP(Val: -3.0, DL, VT);
32019 SDValue MinusHalf = DAG.getConstantFP(Val: -0.5, DL, VT);
32020
32021 // This routine must enter the loop below to work correctly
32022 // when (Reciprocal == false).
32023 assert(Iterations > 0);
32024
32025 // Newton iterations for reciprocal square root:
32026 // E = (E * -0.5) * ((A * E) * E + -3.0)
32027 for (unsigned i = 0; i < Iterations; ++i) {
32028 SDValue AE = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Arg, N2: Est);
32029 SDValue AEE = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: AE, N2: Est);
32030 SDValue RHS = DAG.getNode(Opcode: ISD::FADD, DL, VT, N1: AEE, N2: MinusThree);
32031
32032 // When calculating a square root at the last iteration build:
32033 // S = ((A * E) * -0.5) * ((A * E) * E + -3.0)
32034 // (notice a common subexpression)
32035 SDValue LHS;
32036 if (Reciprocal || (i + 1) < Iterations) {
32037 // RSQRT: LHS = (E * -0.5)
32038 LHS = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: Est, N2: MinusHalf);
32039 } else {
32040 // SQRT: LHS = (A * E) * -0.5
32041 LHS = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: AE, N2: MinusHalf);
32042 }
32043
32044 Est = DAG.getNode(Opcode: ISD::FMUL, DL, VT, N1: LHS, N2: RHS);
32045 }
32046
32047 return Est;
32048}
32049
32050/// Build code to calculate either rsqrt(Op) or sqrt(Op). In the latter case
32051/// Op*rsqrt(Op) is actually computed, so additional postprocessing is needed if
32052/// Op can be zero.
32053SDValue DAGCombiner::buildSqrtEstimateImpl(SDValue Op, bool Reciprocal,
32054 SDNodeFlags Flags) {
32055 if (LegalDAG)
32056 return SDValue();
32057
32058 // TODO: Handle extended types?
32059 EVT VT = Op.getValueType();
32060 if (VT.getScalarType() != MVT::f16 && VT.getScalarType() != MVT::f32 &&
32061 VT.getScalarType() != MVT::f64)
32062 return SDValue();
32063
32064 // If estimates are explicitly disabled for this function, we're done.
32065 const Function &F = DAG.getMachineFunction().getFunction();
32066 int Enabled = TLI.getRecipEstimateSqrtEnabled(VT, F);
32067 if (Enabled == TLI.ReciprocalEstimate::Disabled)
32068 return SDValue();
32069
32070 // Estimates may be explicitly enabled for this type with a custom number of
32071 // refinement steps.
32072 int Iterations = TLI.getSqrtRefinementSteps(VT, MF: F);
32073
32074 bool UseOneConstNR = false;
32075 if (SDValue Est =
32076 TLI.getSqrtEstimate(Operand: Op, DAG, Enabled, RefinementSteps&: Iterations, UseOneConstNR,
32077 Reciprocal)) {
32078 AddToWorklist(N: Est.getNode());
32079
32080 if (Iterations > 0)
32081 Est = UseOneConstNR
32082 ? buildSqrtNROneConst(Arg: Op, Est, Iterations, Reciprocal)
32083 : buildSqrtNRTwoConst(Arg: Op, Est, Iterations, Reciprocal);
32084 if (!Reciprocal) {
32085 SDLoc DL(Op);
32086 // Try the target specific test first.
32087 SDValue Test =
32088 TLI.getSqrtInputTest(Operand: Op, DAG, Mode: DAG.getDenormalMode(VT), Flags);
32089
32090 // The estimate is now completely wrong if the input was exactly 0.0 or
32091 // possibly a denormal. Force the answer to 0.0 or value provided by
32092 // target for those cases.
32093 Est = DAG.getSelect(DL, VT, Cond: Test,
32094 LHS: TLI.getSqrtResultForDenormInput(Operand: Op, DAG), RHS: Est);
32095 }
32096 return Est;
32097 }
32098
32099 return SDValue();
32100}
32101
32102SDValue DAGCombiner::buildRsqrtEstimate(SDValue Op, SDNodeFlags Flags) {
32103 return buildSqrtEstimateImpl(Op, Reciprocal: true, Flags);
32104}
32105
32106SDValue DAGCombiner::buildSqrtEstimate(SDValue Op, SDNodeFlags Flags) {
32107 return buildSqrtEstimateImpl(Op, Reciprocal: false, Flags);
32108}
32109
32110/// Return true if there is any possibility that the two addresses overlap.
32111bool DAGCombiner::mayAlias(SDNode *Op0, SDNode *Op1) const {
32112
32113 struct MemUseCharacteristics {
32114 bool IsVolatile;
32115 bool IsAtomic;
32116 SDValue BasePtr;
32117 int64_t Offset;
32118 LocationSize NumBytes;
32119 MachineMemOperand *MMO;
32120 };
32121
32122 auto getCharacteristics = [this](SDNode *N) -> MemUseCharacteristics {
32123 if (const auto *LSN = dyn_cast<LSBaseSDNode>(Val: N)) {
32124 int64_t Offset = 0;
32125 if (auto *C = dyn_cast<ConstantSDNode>(Val: LSN->getOffset()))
32126 Offset = (LSN->getAddressingMode() == ISD::PRE_INC) ? C->getSExtValue()
32127 : (LSN->getAddressingMode() == ISD::PRE_DEC)
32128 ? -1 * C->getSExtValue()
32129 : 0;
32130 TypeSize Size = LSN->getMemoryVT().getStoreSize();
32131 return {.IsVolatile: LSN->isVolatile(), .IsAtomic: LSN->isAtomic(),
32132 .BasePtr: LSN->getBasePtr(), .Offset: Offset /*base offset*/,
32133 .NumBytes: LocationSize::precise(Value: Size), .MMO: LSN->getMemOperand()};
32134 }
32135 if (const auto *LN = cast<LifetimeSDNode>(Val: N)) {
32136 MachineFrameInfo &MFI = DAG.getMachineFunction().getFrameInfo();
32137 return {.IsVolatile: false /*isVolatile*/,
32138 /*isAtomic*/ .IsAtomic: false,
32139 .BasePtr: LN->getOperand(Num: 1),
32140 .Offset: 0,
32141 .NumBytes: LocationSize::precise(Value: MFI.getObjectSize(ObjectIdx: LN->getFrameIndex())),
32142 .MMO: (MachineMemOperand *)nullptr};
32143 }
32144 // Default.
32145 return {.IsVolatile: false /*isvolatile*/,
32146 /*isAtomic*/ .IsAtomic: false,
32147 .BasePtr: SDValue(),
32148 .Offset: (int64_t)0 /*offset*/,
32149 .NumBytes: LocationSize::beforeOrAfterPointer() /*size*/,
32150 .MMO: (MachineMemOperand *)nullptr};
32151 };
32152
32153 MemUseCharacteristics MUC0 = getCharacteristics(Op0),
32154 MUC1 = getCharacteristics(Op1);
32155
32156 // If they are to the same address, then they must be aliases.
32157 if (MUC0.BasePtr.getNode() && MUC0.BasePtr == MUC1.BasePtr &&
32158 MUC0.Offset == MUC1.Offset)
32159 return true;
32160
32161 // If they are both volatile then they cannot be reordered.
32162 if (MUC0.IsVolatile && MUC1.IsVolatile)
32163 return true;
32164
32165 // Be conservative about atomics for the moment
32166 // TODO: This is way overconservative for unordered atomics (see D66309)
32167 if (MUC0.IsAtomic && MUC1.IsAtomic)
32168 return true;
32169
32170 if (MUC0.MMO && MUC1.MMO) {
32171 if ((MUC0.MMO->isInvariant() && MUC1.MMO->isStore()) ||
32172 (MUC1.MMO->isInvariant() && MUC0.MMO->isStore()))
32173 return false;
32174 }
32175
32176 // If NumBytes is scalable and offset is not 0, conservatively return may
32177 // alias
32178 if ((MUC0.NumBytes.hasValue() && MUC0.NumBytes.isScalable() &&
32179 MUC0.Offset != 0) ||
32180 (MUC1.NumBytes.hasValue() && MUC1.NumBytes.isScalable() &&
32181 MUC1.Offset != 0))
32182 return true;
32183 // Try to prove that there is aliasing, or that there is no aliasing. Either
32184 // way, we can return now. If nothing can be proved, proceed with more tests.
32185 bool IsAlias;
32186 if (BaseIndexOffset::computeAliasing(Op0, NumBytes0: MUC0.NumBytes, Op1, NumBytes1: MUC1.NumBytes,
32187 DAG, IsAlias))
32188 return IsAlias;
32189
32190 // The following all rely on MMO0 and MMO1 being valid. Fail conservatively if
32191 // either are not known.
32192 if (!MUC0.MMO || !MUC1.MMO)
32193 return true;
32194
32195 // If one operation reads from invariant memory, and the other may store, they
32196 // cannot alias. These should really be checking the equivalent of mayWrite,
32197 // but it only matters for memory nodes other than load /store.
32198 if ((MUC0.MMO->isInvariant() && MUC1.MMO->isStore()) ||
32199 (MUC1.MMO->isInvariant() && MUC0.MMO->isStore()))
32200 return false;
32201
32202 // If we know required SrcValue1 and SrcValue2 have relatively large
32203 // alignment compared to the size and offset of the access, we may be able
32204 // to prove they do not alias. This check is conservative for now to catch
32205 // cases created by splitting vector types, it only works when the offsets are
32206 // multiples of the size of the data.
32207 int64_t SrcValOffset0 = MUC0.MMO->getOffset();
32208 int64_t SrcValOffset1 = MUC1.MMO->getOffset();
32209 Align OrigAlignment0 = MUC0.MMO->getBaseAlign();
32210 Align OrigAlignment1 = MUC1.MMO->getBaseAlign();
32211 LocationSize Size0 = MUC0.NumBytes;
32212 LocationSize Size1 = MUC1.NumBytes;
32213
32214 if (OrigAlignment0 == OrigAlignment1 && SrcValOffset0 != SrcValOffset1 &&
32215 Size0.hasValue() && Size1.hasValue() && !Size0.isScalable() &&
32216 !Size1.isScalable() && Size0 == Size1 &&
32217 OrigAlignment0 > Size0.getValue().getKnownMinValue() &&
32218 SrcValOffset0 % Size0.getValue().getKnownMinValue() == 0 &&
32219 SrcValOffset1 % Size1.getValue().getKnownMinValue() == 0) {
32220 int64_t OffAlign0 = SrcValOffset0 % OrigAlignment0.value();
32221 int64_t OffAlign1 = SrcValOffset1 % OrigAlignment1.value();
32222
32223 // There is no overlap between these relatively aligned accesses of
32224 // similar size. Return no alias.
32225 if ((OffAlign0 + static_cast<int64_t>(
32226 Size0.getValue().getKnownMinValue())) <= OffAlign1 ||
32227 (OffAlign1 + static_cast<int64_t>(
32228 Size1.getValue().getKnownMinValue())) <= OffAlign0)
32229 return false;
32230 }
32231
32232 bool UseAA = CombinerGlobalAA.getNumOccurrences() > 0
32233 ? CombinerGlobalAA
32234 : DAG.getSubtarget().useAA();
32235#ifndef NDEBUG
32236 if (CombinerAAOnlyFunc.getNumOccurrences() &&
32237 CombinerAAOnlyFunc != DAG.getMachineFunction().getName())
32238 UseAA = false;
32239#endif
32240
32241 if (UseAA && BatchAA && MUC0.MMO->getValue() && MUC1.MMO->getValue() &&
32242 Size0.hasValue() && Size1.hasValue() &&
32243 // Can't represent a scalable size + fixed offset in LocationSize
32244 (!Size0.isScalable() || SrcValOffset0 == 0) &&
32245 (!Size1.isScalable() || SrcValOffset1 == 0)) {
32246 // Use alias analysis information.
32247 int64_t MinOffset = std::min(a: SrcValOffset0, b: SrcValOffset1);
32248 int64_t Overlap0 =
32249 Size0.getValue().getKnownMinValue() + SrcValOffset0 - MinOffset;
32250 int64_t Overlap1 =
32251 Size1.getValue().getKnownMinValue() + SrcValOffset1 - MinOffset;
32252 LocationSize Loc0 =
32253 Size0.isScalable() ? Size0 : LocationSize::precise(Value: Overlap0);
32254 LocationSize Loc1 =
32255 Size1.isScalable() ? Size1 : LocationSize::precise(Value: Overlap1);
32256 if (BatchAA->isNoAlias(
32257 LocA: MemoryLocation(MUC0.MMO->getValue(), Loc0,
32258 UseTBAA ? MUC0.MMO->getAAInfo() : AAMDNodes()),
32259 LocB: MemoryLocation(MUC1.MMO->getValue(), Loc1,
32260 UseTBAA ? MUC1.MMO->getAAInfo() : AAMDNodes())))
32261 return false;
32262 }
32263
32264 // Otherwise we have to assume they alias.
32265 return true;
32266}
32267
32268/// Walk up chain skipping non-aliasing memory nodes,
32269/// looking for aliasing nodes and adding them to the Aliases vector.
32270void DAGCombiner::GatherAllAliases(SDNode *N, SDValue OriginalChain,
32271 SmallVectorImpl<SDValue> &Aliases) {
32272 SmallVector<SDValue, 8> Chains; // List of chains to visit.
32273 SmallPtrSet<SDNode *, 16> Visited; // Visited node set.
32274
32275 // Get alias information for node.
32276 // TODO: relax aliasing for unordered atomics (see D66309)
32277 const bool IsLoad = isa<LoadSDNode>(Val: N) && cast<LoadSDNode>(Val: N)->isSimple();
32278
32279 // Starting off.
32280 Chains.push_back(Elt: OriginalChain);
32281 unsigned Depth = 0;
32282
32283 // Attempt to improve chain by a single step
32284 auto ImproveChain = [&](SDValue &C) -> bool {
32285 switch (C.getOpcode()) {
32286 case ISD::EntryToken:
32287 // No need to mark EntryToken.
32288 C = SDValue();
32289 return true;
32290 case ISD::LOAD:
32291 case ISD::STORE: {
32292 // Get alias information for C.
32293 // TODO: Relax aliasing for unordered atomics (see D66309)
32294 bool IsOpLoad = isa<LoadSDNode>(Val: C.getNode()) &&
32295 cast<LSBaseSDNode>(Val: C.getNode())->isSimple();
32296 if ((IsLoad && IsOpLoad) || !mayAlias(Op0: N, Op1: C.getNode())) {
32297 // Look further up the chain.
32298 C = C.getOperand(i: 0);
32299 return true;
32300 }
32301 // Alias, so stop here.
32302 return false;
32303 }
32304
32305 case ISD::CopyFromReg:
32306 // Always forward past CopyFromReg.
32307 C = C.getOperand(i: 0);
32308 return true;
32309
32310 case ISD::LIFETIME_START:
32311 case ISD::LIFETIME_END: {
32312 // We can forward past any lifetime start/end that can be proven not to
32313 // alias the memory access.
32314 if (!mayAlias(Op0: N, Op1: C.getNode())) {
32315 // Look further up the chain.
32316 C = C.getOperand(i: 0);
32317 return true;
32318 }
32319 return false;
32320 }
32321 default:
32322 return false;
32323 }
32324 };
32325
32326 // Look at each chain and determine if it is an alias. If so, add it to the
32327 // aliases list. If not, then continue up the chain looking for the next
32328 // candidate.
32329 while (!Chains.empty()) {
32330 SDValue Chain = Chains.pop_back_val();
32331
32332 // Don't bother if we've seen Chain before.
32333 if (!Visited.insert(Ptr: Chain.getNode()).second)
32334 continue;
32335
32336 // For TokenFactor nodes, look at each operand and only continue up the
32337 // chain until we reach the depth limit.
32338 //
32339 // FIXME: The depth check could be made to return the last non-aliasing
32340 // chain we found before we hit a tokenfactor rather than the original
32341 // chain.
32342 if (Depth > TLI.getGatherAllAliasesMaxDepth()) {
32343 Aliases.clear();
32344 Aliases.push_back(Elt: OriginalChain);
32345 return;
32346 }
32347
32348 if (Chain.getOpcode() == ISD::TokenFactor) {
32349 // We have to check each of the operands of the token factor for "small"
32350 // token factors, so we queue them up. Adding the operands to the queue
32351 // (stack) in reverse order maintains the original order and increases the
32352 // likelihood that getNode will find a matching token factor (CSE.)
32353 if (Chain.getNumOperands() > 16) {
32354 Aliases.push_back(Elt: Chain);
32355 continue;
32356 }
32357 for (unsigned n = Chain.getNumOperands(); n;)
32358 Chains.push_back(Elt: Chain.getOperand(i: --n));
32359 ++Depth;
32360 continue;
32361 }
32362 // Everything else
32363 if (ImproveChain(Chain)) {
32364 // Updated Chain Found, Consider new chain if one exists.
32365 if (Chain.getNode())
32366 Chains.push_back(Elt: Chain);
32367 ++Depth;
32368 continue;
32369 }
32370 // No Improved Chain Possible, treat as Alias.
32371 Aliases.push_back(Elt: Chain);
32372 }
32373}
32374
32375/// Walk up chain skipping non-aliasing memory nodes, looking for a better chain
32376/// (aliasing node.)
32377SDValue DAGCombiner::FindBetterChain(SDNode *N, SDValue OldChain) {
32378 if (OptLevel == CodeGenOptLevel::None)
32379 return OldChain;
32380
32381 // Ops for replacing token factor.
32382 SmallVector<SDValue, 8> Aliases;
32383
32384 // Accumulate all the aliases to this node.
32385 GatherAllAliases(N, OriginalChain: OldChain, Aliases);
32386
32387 // If no operands then chain to entry token.
32388 if (Aliases.empty())
32389 return DAG.getEntryNode();
32390
32391 // If a single operand then chain to it. We don't need to revisit it.
32392 if (Aliases.size() == 1)
32393 return Aliases[0];
32394
32395 // Construct a custom tailored token factor.
32396 return DAG.getTokenFactor(DL: SDLoc(N), Vals&: Aliases);
32397}
32398
32399// This function tries to collect a bunch of potentially interesting
32400// nodes to improve the chains of, all at once. This might seem
32401// redundant, as this function gets called when visiting every store
32402// node, so why not let the work be done on each store as it's visited?
32403//
32404// I believe this is mainly important because mergeConsecutiveStores
32405// is unable to deal with merging stores of different sizes, so unless
32406// we improve the chains of all the potential candidates up-front
32407// before running mergeConsecutiveStores, it might only see some of
32408// the nodes that will eventually be candidates, and then not be able
32409// to go from a partially-merged state to the desired final
32410// fully-merged state.
32411
32412bool DAGCombiner::parallelizeChainedStores(StoreSDNode *St) {
32413 SmallVector<StoreSDNode *, 8> ChainedStores;
32414 StoreSDNode *STChain = St;
32415 // Intervals records which offsets from BaseIndex have been covered. In
32416 // the common case, every store writes to the immediately previous address
32417 // space and thus merged with the previous interval at insertion time.
32418
32419 using IMap = llvm::IntervalMap<int64_t, std::monostate, 8,
32420 IntervalMapHalfOpenInfo<int64_t>>;
32421 IMap::Allocator A;
32422 IMap Intervals(A);
32423
32424 // This holds the base pointer, index, and the offset in bytes from the base
32425 // pointer.
32426 const BaseIndexOffset BasePtr = BaseIndexOffset::match(N: St, DAG);
32427
32428 // We must have a base and an offset.
32429 if (!BasePtr.getBase().getNode())
32430 return false;
32431
32432 // Do not handle stores to undef base pointers.
32433 if (BasePtr.getBase().isUndef())
32434 return false;
32435
32436 // Do not handle stores to opaque types
32437 if (St->getMemoryVT().isZeroSized())
32438 return false;
32439
32440 // BaseIndexOffset assumes that offsets are fixed-size, which
32441 // is not valid for scalable vectors where the offsets are
32442 // scaled by `vscale`, so bail out early.
32443 if (St->getMemoryVT().isScalableVT())
32444 return false;
32445
32446 // Add ST's interval.
32447 Intervals.insert(a: 0, b: (St->getMemoryVT().getSizeInBits() + 7) / 8,
32448 y: std::monostate{});
32449
32450 while (StoreSDNode *Chain = dyn_cast<StoreSDNode>(Val: STChain->getChain())) {
32451 if (Chain->getMemoryVT().isScalableVector())
32452 return false;
32453
32454 // If the chain has more than one use, then we can't reorder the mem ops.
32455 if (!SDValue(Chain, 0)->hasOneUse())
32456 break;
32457 // TODO: Relax for unordered atomics (see D66309)
32458 if (!Chain->isSimple() || Chain->isIndexed())
32459 break;
32460
32461 // Find the base pointer and offset for this memory node.
32462 const BaseIndexOffset Ptr = BaseIndexOffset::match(N: Chain, DAG);
32463 // Check that the base pointer is the same as the original one.
32464 int64_t Offset;
32465 if (!BasePtr.equalBaseIndex(Other: Ptr, DAG, Off&: Offset))
32466 break;
32467 int64_t Length = (Chain->getMemoryVT().getSizeInBits() + 7) / 8;
32468 // Make sure we don't overlap with other intervals by checking the ones to
32469 // the left or right before inserting.
32470 auto I = Intervals.find(x: Offset);
32471 // If there's a next interval, we should end before it.
32472 if (I != Intervals.end() && I.start() < (Offset + Length))
32473 break;
32474 // If there's a previous interval, we should start after it.
32475 if (I != Intervals.begin() && (--I).stop() <= Offset)
32476 break;
32477 Intervals.insert(a: Offset, b: Offset + Length, y: std::monostate{});
32478
32479 ChainedStores.push_back(Elt: Chain);
32480 STChain = Chain;
32481 }
32482
32483 // If we didn't find a chained store, exit.
32484 if (ChainedStores.empty())
32485 return false;
32486
32487 // Improve all chained stores (St and ChainedStores members) starting from
32488 // where the store chain ended and return single TokenFactor.
32489 SDValue NewChain = STChain->getChain();
32490 SmallVector<SDValue, 8> TFOps;
32491 for (unsigned I = ChainedStores.size(); I;) {
32492 StoreSDNode *S = ChainedStores[--I];
32493 SDValue BetterChain = FindBetterChain(N: S, OldChain: NewChain);
32494 S = cast<StoreSDNode>(Val: DAG.UpdateNodeOperands(
32495 N: S, Op1: BetterChain, Op2: S->getOperand(Num: 1), Op3: S->getOperand(Num: 2), Op4: S->getOperand(Num: 3)));
32496 TFOps.push_back(Elt: SDValue(S, 0));
32497 ChainedStores[I] = S;
32498 }
32499
32500 // Improve St's chain. Use a new node to avoid creating a loop from CombineTo.
32501 SDValue BetterChain = FindBetterChain(N: St, OldChain: NewChain);
32502 SDValue NewST;
32503 if (St->isTruncatingStore())
32504 NewST = DAG.getTruncStore(Chain: BetterChain, dl: SDLoc(St), Val: St->getValue(),
32505 Ptr: St->getBasePtr(), SVT: St->getMemoryVT(),
32506 MMO: St->getMemOperand());
32507 else
32508 NewST = DAG.getStore(Chain: BetterChain, dl: SDLoc(St), Val: St->getValue(),
32509 Ptr: St->getBasePtr(), MMO: St->getMemOperand());
32510
32511 TFOps.push_back(Elt: NewST);
32512
32513 // If we improved every element of TFOps, then we've lost the dependence on
32514 // NewChain to successors of St and we need to add it back to TFOps. Do so at
32515 // the beginning to keep relative order consistent with FindBetterChains.
32516 auto hasImprovedChain = [&](SDValue ST) -> bool {
32517 return ST->getOperand(Num: 0) != NewChain;
32518 };
32519 bool AddNewChain = llvm::all_of(Range&: TFOps, P: hasImprovedChain);
32520 if (AddNewChain)
32521 TFOps.insert(I: TFOps.begin(), Elt: NewChain);
32522
32523 SDValue TF = DAG.getTokenFactor(DL: SDLoc(STChain), Vals&: TFOps);
32524 CombineTo(N: St, Res: TF);
32525
32526 // Add TF and its operands to the worklist.
32527 AddToWorklist(N: TF.getNode());
32528 for (const SDValue &Op : TF->ops())
32529 AddToWorklist(N: Op.getNode());
32530 AddToWorklist(N: STChain);
32531 return true;
32532}
32533
32534bool DAGCombiner::findBetterNeighborChains(StoreSDNode *St) {
32535 if (OptLevel == CodeGenOptLevel::None)
32536 return false;
32537
32538 const BaseIndexOffset BasePtr = BaseIndexOffset::match(N: St, DAG);
32539
32540 // We must have a base and an offset.
32541 if (!BasePtr.getBase().getNode())
32542 return false;
32543
32544 // Do not handle stores to undef base pointers.
32545 if (BasePtr.getBase().isUndef())
32546 return false;
32547
32548 // Directly improve a chain of disjoint stores starting at St.
32549 if (parallelizeChainedStores(St))
32550 return true;
32551
32552 // Improve St's Chain..
32553 SDValue BetterChain = FindBetterChain(N: St, OldChain: St->getChain());
32554 if (St->getChain() != BetterChain) {
32555 replaceStoreChain(ST: St, BetterChain);
32556 return true;
32557 }
32558 return false;
32559}
32560
32561/// This is the entry point for the file.
32562void SelectionDAG::Combine(CombineLevel Level, BatchAAResults *BatchAA,
32563 CodeGenOptLevel OptLevel) {
32564 /// This is the main entry point to this class.
32565 DAGCombiner(*this, BatchAA, OptLevel).Run(AtLevel: Level);
32566}
32567