1//===-- AArch64SelectionDAGInfo.cpp - AArch64 SelectionDAG Info -----------===//
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 file implements the AArch64SelectionDAGInfo class.
10//
11//===----------------------------------------------------------------------===//
12
13#include "AArch64SelectionDAGInfo.h"
14#include "AArch64MachineFunctionInfo.h"
15
16#define GET_SDNODE_DESC
17#include "AArch64GenSDNodeInfo.inc"
18#undef GET_SDNODE_DESC
19
20using namespace llvm;
21
22#define DEBUG_TYPE "aarch64-selectiondag-info"
23
24static cl::opt<bool>
25 LowerToSMERoutines("aarch64-lower-to-sme-routines", cl::Hidden,
26 cl::desc("Enable AArch64 SME memory operations "
27 "to lower to librt functions"),
28 cl::init(Val: true));
29
30static cl::opt<bool> UseMOPS("aarch64-use-mops", cl::Hidden,
31 cl::desc("Enable AArch64 MOPS instructions "
32 "for memcpy/memset/memmove"),
33 cl::init(Val: true));
34
35AArch64SelectionDAGInfo::AArch64SelectionDAGInfo()
36 : SelectionDAGGenTargetInfo(AArch64GenSDNodeInfo) {}
37
38void AArch64SelectionDAGInfo::verifyTargetNode(const SelectionDAG &DAG,
39 const SDNode *N) const {
40 switch (N->getOpcode()) {
41 case AArch64ISD::WrapperLarge:
42 // operand #0 must have type i32, but has type i64
43 return;
44 }
45
46 SelectionDAGGenTargetInfo::verifyTargetNode(DAG, N);
47
48#ifndef NDEBUG
49 // Some additional checks not yet implemented by verifyTargetNode.
50 switch (N->getOpcode()) {
51 case AArch64ISD::CTTZ_ELTS:
52 case AArch64ISD::CTTZ_ELTS_ZERO_POISON:
53 assert(N->getOperand(0).getValueType() == N->getOperand(1).getValueType() &&
54 "Expected the general-predicate and mask to have matching types");
55 break;
56 case AArch64ISD::SUNPKLO:
57 case AArch64ISD::SUNPKHI:
58 case AArch64ISD::UUNPKLO:
59 case AArch64ISD::UUNPKHI: {
60 EVT VT = N->getValueType(0);
61 EVT OpVT = N->getOperand(0).getValueType();
62 assert(OpVT.isVector() && VT.isVector() && OpVT.isInteger() &&
63 VT.isInteger() && "Expected integer vectors!");
64 assert(OpVT.getSizeInBits() == VT.getSizeInBits() &&
65 "Expected vectors of equal size!");
66 assert(OpVT.getVectorElementCount() == VT.getVectorElementCount() * 2 &&
67 "Expected result vector with half the lanes of its input!");
68 break;
69 }
70 case AArch64ISD::TRN1:
71 case AArch64ISD::TRN2:
72 case AArch64ISD::UZP1:
73 case AArch64ISD::UZP2:
74 case AArch64ISD::ZIP1:
75 case AArch64ISD::ZIP2: {
76 EVT VT = N->getValueType(0);
77 EVT Op0VT = N->getOperand(0).getValueType();
78 EVT Op1VT = N->getOperand(1).getValueType();
79 assert(VT.isVector() && Op0VT.isVector() && Op1VT.isVector() &&
80 "Expected vectors!");
81 assert(VT == Op0VT && VT == Op1VT && "Expected matching vectors!");
82 break;
83 }
84 case AArch64ISD::RSHRNB_I: {
85 EVT VT = N->getValueType(0);
86 EVT Op0VT = N->getOperand(0).getValueType();
87 assert(VT.isVector() && VT.isInteger() &&
88 "Expected integer vector result type!");
89 assert(Op0VT.isVector() && Op0VT.isInteger() &&
90 "Expected first operand to be an integer vector!");
91 assert(VT.getSizeInBits() == Op0VT.getSizeInBits() &&
92 "Expected vectors of equal size!");
93 assert(VT.getVectorElementCount() == Op0VT.getVectorElementCount() * 2 &&
94 "Expected input vector with half the lanes of its result!");
95 assert(isa<ConstantSDNode>(N->getOperand(1)) &&
96 "Expected second operand to be a constant!");
97 break;
98 }
99 }
100#endif
101}
102
103SDValue AArch64SelectionDAGInfo::EmitMOPS(unsigned Opcode, SelectionDAG &DAG,
104 const SDLoc &DL, SDValue Chain,
105 SDValue Dst, SDValue SrcOrValue,
106 SDValue Size, Align DstAlign,
107 Align SrcAlign, bool isVolatile,
108 MachinePointerInfo DstPtrInfo,
109 MachinePointerInfo SrcPtrInfo) const {
110
111 // Get the constant size of the copy/set.
112 LocationSize MemSize = LocationSize::afterPointer();
113 if (auto *C = dyn_cast<ConstantSDNode>(Val&: Size))
114 MemSize = LocationSize::precise(Value: C->getZExtValue());
115
116 const bool IsSet = Opcode == AArch64::MOPSMemorySetPseudo ||
117 Opcode == AArch64::MOPSMemorySetTaggingPseudo;
118
119 MachineFunction &MF = DAG.getMachineFunction();
120
121 auto Vol =
122 isVolatile ? MachineMemOperand::MOVolatile : MachineMemOperand::MONone;
123 auto DstFlags = MachineMemOperand::MOStore | Vol;
124 auto *DstOp =
125 MF.getMachineMemOperand(PtrInfo: DstPtrInfo, F: DstFlags, Size: MemSize, BaseAlignment: DstAlign);
126
127 if (IsSet) {
128 // Extend value to i64, if required.
129 if (SrcOrValue.getValueType() != MVT::i64)
130 SrcOrValue = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT: MVT::i64, Operand: SrcOrValue);
131 SDValue Ops[] = {Dst, Size, SrcOrValue, Chain};
132 const EVT ResultTys[] = {MVT::i64, MVT::i64, MVT::Other};
133 MachineSDNode *Node = DAG.getMachineNode(Opcode, dl: DL, ResultTys, Ops);
134 DAG.setNodeMemRefs(N: Node, NewMemRefs: {DstOp});
135 return SDValue(Node, 2);
136 } else {
137 SDValue Ops[] = {Dst, SrcOrValue, Size, Chain};
138 const EVT ResultTys[] = {MVT::i64, MVT::i64, MVT::i64, MVT::Other};
139 MachineSDNode *Node = DAG.getMachineNode(Opcode, dl: DL, ResultTys, Ops);
140
141 auto SrcFlags = MachineMemOperand::MOLoad | Vol;
142 auto *SrcOp =
143 MF.getMachineMemOperand(PtrInfo: SrcPtrInfo, F: SrcFlags, Size: MemSize, BaseAlignment: SrcAlign);
144 DAG.setNodeMemRefs(N: Node, NewMemRefs: {DstOp, SrcOp});
145 return SDValue(Node, 3);
146 }
147}
148
149SDValue AArch64SelectionDAGInfo::EmitStreamingCompatibleMemLibCall(
150 SelectionDAG &DAG, const SDLoc &DL, SDValue Chain, SDValue Op0, SDValue Op1,
151 SDValue Size, RTLIB::Libcall LC) const {
152 const AArch64Subtarget &STI =
153 DAG.getMachineFunction().getSubtarget<AArch64Subtarget>();
154 const AArch64TargetLowering *TLI = STI.getTargetLowering();
155 TargetLowering::ArgListTy Args;
156 Args.emplace_back(args&: Op0, args: PointerType::getUnqual(C&: *DAG.getContext()));
157
158 bool UsesResult = false;
159 RTLIB::Libcall NewLC;
160 switch (LC) {
161 case RTLIB::MEMCPY: {
162 NewLC = RTLIB::SC_MEMCPY;
163 Args.emplace_back(args&: Op1, args: PointerType::getUnqual(C&: *DAG.getContext()));
164 break;
165 }
166 case RTLIB::MEMMOVE: {
167 NewLC = RTLIB::SC_MEMMOVE;
168 Args.emplace_back(args&: Op1, args: PointerType::getUnqual(C&: *DAG.getContext()));
169 break;
170 }
171 case RTLIB::MEMSET: {
172 NewLC = RTLIB::SC_MEMSET;
173 Args.emplace_back(args: DAG.getZExtOrTrunc(Op: Op1, DL, VT: MVT::i32),
174 args: Type::getInt32Ty(C&: *DAG.getContext()));
175 break;
176 }
177 case RTLIB::MEMCHR: {
178 UsesResult = true;
179 NewLC = RTLIB::SC_MEMCHR;
180 Args.emplace_back(args: DAG.getZExtOrTrunc(Op: Op1, DL, VT: MVT::i32),
181 args: Type::getInt32Ty(C&: *DAG.getContext()));
182 break;
183 }
184 default:
185 return SDValue();
186 }
187
188 RTLIB::LibcallImpl NewLCImpl = DAG.getLibcalls().getLibcallImpl(Call: NewLC);
189 if (NewLCImpl == RTLIB::Unsupported)
190 return SDValue();
191
192 EVT PointerVT = TLI->getPointerTy(DL: DAG.getDataLayout());
193 SDValue Symbol = DAG.getExternalSymbol(LCImpl: NewLCImpl, VT: PointerVT);
194 Args.emplace_back(args&: Size, args: DAG.getDataLayout().getIntPtrType(C&: *DAG.getContext()));
195
196 TargetLowering::CallLoweringInfo CLI(DAG);
197 PointerType *RetTy = PointerType::getUnqual(C&: *DAG.getContext());
198 CLI.setDebugLoc(DL).setChain(Chain).setLibCallee(
199 CC: DAG.getLibcalls().getLibcallImplCallingConv(Call: NewLCImpl), ResultType: RetTy, Target: Symbol,
200 ArgsList: std::move(Args));
201
202 auto [Result, ChainOut] = TLI->LowerCallTo(CLI);
203 return UsesResult ? DAG.getMergeValues(Ops: {Result, ChainOut}, dl: DL) : ChainOut;
204}
205
206SDValue AArch64SelectionDAGInfo::EmitTargetCodeForMemcpy(
207 SelectionDAG &DAG, const SDLoc &DL, SDValue Chain, SDValue Dst, SDValue Src,
208 SDValue Size, Align DstAlign, Align SrcAlign, bool isVolatile,
209 bool AlwaysInline, MachinePointerInfo DstPtrInfo,
210 MachinePointerInfo SrcPtrInfo) const {
211 const AArch64Subtarget &STI =
212 DAG.getMachineFunction().getSubtarget<AArch64Subtarget>();
213
214 if (UseMOPS && STI.hasMOPS())
215 return EmitMOPS(Opcode: AArch64::MOPSMemoryCopyPseudo, DAG, DL, Chain, Dst, SrcOrValue: Src,
216 Size, DstAlign, SrcAlign, isVolatile, DstPtrInfo,
217 SrcPtrInfo);
218
219 auto *AFI = DAG.getMachineFunction().getInfo<AArch64FunctionInfo>();
220 SMEAttrs Attrs = AFI->getSMEFnAttrs();
221 if (LowerToSMERoutines && !Attrs.hasNonStreamingInterfaceAndBody())
222 return EmitStreamingCompatibleMemLibCall(DAG, DL, Chain, Op0: Dst, Op1: Src, Size,
223 LC: RTLIB::MEMCPY);
224 return SDValue();
225}
226
227SDValue AArch64SelectionDAGInfo::EmitTargetCodeForMemset(
228 SelectionDAG &DAG, const SDLoc &dl, SDValue Chain, SDValue Dst, SDValue Src,
229 SDValue Size, Align Alignment, bool isVolatile, bool AlwaysInline,
230 MachinePointerInfo DstPtrInfo) const {
231 const AArch64Subtarget &STI =
232 DAG.getMachineFunction().getSubtarget<AArch64Subtarget>();
233
234 if (UseMOPS && STI.hasMOPS())
235 return EmitMOPS(Opcode: AArch64::MOPSMemorySetPseudo, DAG, DL: dl, Chain, Dst, SrcOrValue: Src,
236 Size, DstAlign: Alignment, SrcAlign: Alignment, isVolatile, DstPtrInfo,
237 SrcPtrInfo: MachinePointerInfo{});
238
239 auto *AFI = DAG.getMachineFunction().getInfo<AArch64FunctionInfo>();
240 SMEAttrs Attrs = AFI->getSMEFnAttrs();
241 if (LowerToSMERoutines && !Attrs.hasNonStreamingInterfaceAndBody())
242 return EmitStreamingCompatibleMemLibCall(DAG, DL: dl, Chain, Op0: Dst, Op1: Src, Size,
243 LC: RTLIB::MEMSET);
244 return SDValue();
245}
246
247SDValue AArch64SelectionDAGInfo::EmitTargetCodeForMemmove(
248 SelectionDAG &DAG, const SDLoc &dl, SDValue Chain, SDValue Dst, SDValue Src,
249 SDValue Size, Align DstAlign, Align SrcAlign, bool isVolatile,
250 MachinePointerInfo DstPtrInfo, MachinePointerInfo SrcPtrInfo) const {
251 const AArch64Subtarget &STI =
252 DAG.getMachineFunction().getSubtarget<AArch64Subtarget>();
253
254 if (UseMOPS && STI.hasMOPS())
255 return EmitMOPS(Opcode: AArch64::MOPSMemoryMovePseudo, DAG, DL: dl, Chain, Dst, SrcOrValue: Src,
256 Size, DstAlign, SrcAlign, isVolatile, DstPtrInfo,
257 SrcPtrInfo);
258
259 auto *AFI = DAG.getMachineFunction().getInfo<AArch64FunctionInfo>();
260 SMEAttrs Attrs = AFI->getSMEFnAttrs();
261 if (LowerToSMERoutines && !Attrs.hasNonStreamingInterfaceAndBody())
262 return EmitStreamingCompatibleMemLibCall(DAG, DL: dl, Chain, Op0: Dst, Op1: Src, Size,
263 LC: RTLIB::MEMMOVE);
264 return SDValue();
265}
266
267std::pair<SDValue, SDValue> AArch64SelectionDAGInfo::EmitTargetCodeForMemchr(
268 SelectionDAG &DAG, const SDLoc &dl, SDValue Chain, SDValue Src,
269 SDValue Char, SDValue Length, MachinePointerInfo SrcPtrInfo) const {
270 auto *AFI = DAG.getMachineFunction().getInfo<AArch64FunctionInfo>();
271 SMEAttrs Attrs = AFI->getSMEFnAttrs();
272 if (LowerToSMERoutines && !Attrs.hasNonStreamingInterfaceAndBody()) {
273 SDValue Result = EmitStreamingCompatibleMemLibCall(
274 DAG, DL: dl, Chain, Op0: Src, Op1: Char, Size: Length, LC: RTLIB::MEMCHR);
275 return std::make_pair(x: Result.getValue(R: 0), y: Result.getValue(R: 1));
276 }
277 return std::make_pair(x: SDValue(), y: SDValue());
278}
279
280static const int kSetTagLoopThreshold = 176;
281
282static SDValue EmitUnrolledSetTag(SelectionDAG &DAG, const SDLoc &dl,
283 SDValue Chain, SDValue Ptr, uint64_t ObjSize,
284 const MachineMemOperand *BaseMemOperand,
285 bool ZeroData) {
286 MachineFunction &MF = DAG.getMachineFunction();
287 unsigned ObjSizeScaled = ObjSize / 16;
288
289 SDValue TagSrc = Ptr;
290 if (Ptr.getOpcode() == ISD::FrameIndex) {
291 int FI = cast<FrameIndexSDNode>(Val&: Ptr)->getIndex();
292 Ptr = DAG.getTargetFrameIndex(FI, VT: MVT::i64);
293 // A frame index operand may end up as [SP + offset] => it is fine to use SP
294 // register as the tag source.
295 TagSrc = DAG.getRegister(Reg: AArch64::SP, VT: MVT::i64);
296 }
297
298 const unsigned OpCode1 = ZeroData ? AArch64ISD::STZG : AArch64ISD::STG;
299 const unsigned OpCode2 = ZeroData ? AArch64ISD::STZ2G : AArch64ISD::ST2G;
300
301 SmallVector<SDValue, 8> OutChains;
302 unsigned OffsetScaled = 0;
303 while (OffsetScaled < ObjSizeScaled) {
304 if (ObjSizeScaled - OffsetScaled >= 2) {
305 SDValue AddrNode = DAG.getMemBasePlusOffset(
306 Base: Ptr, Offset: TypeSize::getFixed(ExactSize: OffsetScaled * 16), DL: dl);
307 SDValue St = DAG.getMemIntrinsicNode(
308 Opcode: OpCode2, dl, VTList: DAG.getVTList(VT: MVT::Other),
309 Ops: {Chain, TagSrc, AddrNode},
310 MemVT: MVT::v4i64,
311 MMO: MF.getMachineMemOperand(MMO: BaseMemOperand, Offset: OffsetScaled * 16, Size: 16 * 2));
312 OffsetScaled += 2;
313 OutChains.push_back(Elt: St);
314 continue;
315 }
316
317 if (ObjSizeScaled - OffsetScaled > 0) {
318 SDValue AddrNode = DAG.getMemBasePlusOffset(
319 Base: Ptr, Offset: TypeSize::getFixed(ExactSize: OffsetScaled * 16), DL: dl);
320 SDValue St = DAG.getMemIntrinsicNode(
321 Opcode: OpCode1, dl, VTList: DAG.getVTList(VT: MVT::Other),
322 Ops: {Chain, TagSrc, AddrNode},
323 MemVT: MVT::v2i64,
324 MMO: MF.getMachineMemOperand(MMO: BaseMemOperand, Offset: OffsetScaled * 16, Size: 16));
325 OffsetScaled += 1;
326 OutChains.push_back(Elt: St);
327 }
328 }
329
330 SDValue Res = DAG.getNode(Opcode: ISD::TokenFactor, DL: dl, VT: MVT::Other, Ops: OutChains);
331 return Res;
332}
333
334SDValue AArch64SelectionDAGInfo::EmitTargetCodeForSetTag(
335 SelectionDAG &DAG, const SDLoc &dl, SDValue Chain, SDValue Addr,
336 SDValue Size, MachinePointerInfo DstPtrInfo, bool ZeroData) const {
337 uint64_t ObjSize = Size->getAsZExtVal();
338 assert(ObjSize % 16 == 0);
339
340 MachineFunction &MF = DAG.getMachineFunction();
341 MachineMemOperand *BaseMemOperand = MF.getMachineMemOperand(
342 PtrInfo: DstPtrInfo, F: MachineMemOperand::MOStore, Size: ObjSize, BaseAlignment: Align(16));
343
344 bool UseSetTagRangeLoop =
345 kSetTagLoopThreshold >= 0 && (int)ObjSize >= kSetTagLoopThreshold;
346 if (!UseSetTagRangeLoop)
347 return EmitUnrolledSetTag(DAG, dl, Chain, Ptr: Addr, ObjSize, BaseMemOperand,
348 ZeroData);
349
350 const EVT ResTys[] = {MVT::i64, MVT::i64, MVT::Other};
351
352 unsigned Opcode;
353 if (Addr.getOpcode() == ISD::FrameIndex) {
354 int FI = cast<FrameIndexSDNode>(Val&: Addr)->getIndex();
355 Addr = DAG.getTargetFrameIndex(FI, VT: MVT::i64);
356 Opcode = ZeroData ? AArch64::STZGloop : AArch64::STGloop;
357 } else {
358 Opcode = ZeroData ? AArch64::STZGloop_wback : AArch64::STGloop_wback;
359 }
360 SDValue Ops[] = {DAG.getTargetConstant(Val: ObjSize, DL: dl, VT: MVT::i64), Addr, Chain};
361 SDNode *St = DAG.getMachineNode(Opcode, dl, ResultTys: ResTys, Ops);
362
363 DAG.setNodeMemRefs(N: cast<MachineSDNode>(Val: St), NewMemRefs: {BaseMemOperand});
364 return SDValue(St, 2);
365}
366