1//===----------------------------------------------------------------------===//
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/// \file
10/// This file contains NVVM-specific IR verification logic. These checks are
11/// always compiled and linked as part of LLVMCore.
12///
13//===----------------------------------------------------------------------===//
14
15#include "VerifierInternal.h"
16#include "llvm/IR/Constants.h"
17#include "llvm/IR/DerivedTypes.h"
18#include "llvm/IR/IntrinsicsNVPTX.h"
19#include "llvm/Support/MathExtras.h"
20#include <optional>
21
22using namespace llvm;
23
24#define Check(C, ...) \
25 do { \
26 if (!(C)) { \
27 VS.CheckFailed(__VA_ARGS__); \
28 return; \
29 } \
30 } while (false)
31
32namespace {
33
34struct SPVectorInfo {
35 unsigned ElemSize;
36 unsigned NumElements;
37 unsigned NumRegisters;
38};
39
40// Register layout of the sparse intrinsic operands, used by the IR verifier.
41struct SPOperandLayout {
42 unsigned MetadataSize;
43 unsigned CompressedDataSize;
44 unsigned DataSize;
45};
46
47} // namespace
48
49// PTX limits the combined vector size of the mdata, cdata, and data operands
50// of spcompress and spdecompress to 253 32-bit registers.
51constexpr unsigned MaxSPOperandRegisters = 253;
52
53static bool isValidSPElemSize(unsigned ElemSize) {
54 return ElemSize == 8 || ElemSize == 16;
55}
56
57static bool isValidSPIdxSize(unsigned IdxSize) {
58 return IdxSize == 2 || IdxSize == 4;
59}
60
61static bool isValidSPRepeatFactor(unsigned RepeatFactor) {
62 return isPowerOf2_32(Value: RepeatFactor) && RepeatFactor <= 64;
63}
64
65static bool isValidSPDecompressFactor(unsigned NumSrc, unsigned NumTgt) {
66 switch (NumSrc) {
67 case 1:
68 return NumTgt == 2 || NumTgt == 4 || NumTgt == 8 || NumTgt == 16;
69 case 2:
70 return NumTgt == 4 || NumTgt == 8 || NumTgt == 16;
71 case 4:
72 return NumTgt == 8 || NumTgt == 16;
73 default:
74 return false;
75 }
76}
77
78static std::optional<SPOperandLayout>
79getSPCompressLayout(unsigned ElemSize, unsigned IdxSize,
80 unsigned RepeatFactor) {
81 if (!isValidSPElemSize(ElemSize) || !isValidSPIdxSize(IdxSize) ||
82 !isValidSPRepeatFactor(RepeatFactor))
83 return std::nullopt;
84
85 SPOperandLayout Layout = {.MetadataSize: divideCeil(Numerator: RepeatFactor * IdxSize, Denominator: ElemSize),
86 .CompressedDataSize: RepeatFactor, .DataSize: RepeatFactor * 2};
87 if (Layout.MetadataSize + Layout.CompressedDataSize + Layout.DataSize >
88 MaxSPOperandRegisters)
89 return std::nullopt;
90 return Layout;
91}
92
93static std::optional<SPOperandLayout>
94getSPDecompressLayout(unsigned NumSrc, unsigned NumTgt, unsigned ElemSize,
95 unsigned IdxSize, unsigned RepeatFactor) {
96 if (!isValidSPDecompressFactor(NumSrc, NumTgt) ||
97 !isValidSPElemSize(ElemSize) || !isValidSPIdxSize(IdxSize) ||
98 !isValidSPRepeatFactor(RepeatFactor) || NumSrc * ElemSize > 32 ||
99 (IdxSize == 2 && NumTgt > 4))
100 return std::nullopt;
101
102 unsigned DataBits = NumTgt * ElemSize * RepeatFactor;
103 if (DataBits < 32 || DataBits > 4096)
104 return std::nullopt;
105
106 SPOperandLayout Layout = {.MetadataSize: divideCeil(Numerator: NumSrc * IdxSize * RepeatFactor, Denominator: 32),
107 .CompressedDataSize: divideCeil(Numerator: NumSrc * ElemSize * RepeatFactor, Denominator: 32),
108 .DataSize: divideCeil(Numerator: DataBits, Denominator: 32)};
109 if (Layout.MetadataSize + Layout.CompressedDataSize + Layout.DataSize >
110 MaxSPOperandRegisters)
111 return std::nullopt;
112 return Layout;
113}
114
115static std::optional<SPVectorInfo> getSPVectorInfo(Type *Ty) {
116 auto *VT = dyn_cast<FixedVectorType>(Val: Ty);
117 if (!VT)
118 return std::nullopt;
119 auto *ElemTy = dyn_cast<IntegerType>(Val: VT->getElementType());
120 if (!ElemTy)
121 return std::nullopt;
122 unsigned ElemSize = ElemTy->getBitWidth();
123 if (ElemSize != 8 && ElemSize != 16)
124 return std::nullopt;
125 unsigned NumElements = VT->getNumElements();
126 return SPVectorInfo{.ElemSize: ElemSize, .NumElements: NumElements,
127 .NumRegisters: divideCeil(Numerator: NumElements, Denominator: 32 / ElemSize)};
128}
129
130static std::optional<unsigned> getSPMetadataRegisters(Type *Ty) {
131 if (Ty->isIntegerTy(BitWidth: 32))
132 return 1;
133 auto *VT = dyn_cast<FixedVectorType>(Val: Ty);
134 if (!VT || !VT->getElementType()->isIntegerTy(BitWidth: 32) || VT->getNumElements() < 2)
135 return std::nullopt;
136 return VT->getNumElements();
137}
138
139static void verifySPCompress(VerifierSupport &VS, CallBase &Call) {
140 auto *ResultTy = dyn_cast<StructType>(Val: Call.getType());
141 Check(ResultTy && ResultTy->getNumElements() == 2,
142 "invalid llvm.nvvm.spcompress result type", &Call);
143
144 auto MDataRegs = getSPMetadataRegisters(Ty: ResultTy->getElementType(N: 0));
145 auto CData = getSPVectorInfo(Ty: ResultTy->getElementType(N: 1));
146 auto Data = getSPVectorInfo(Ty: Call.getArgOperand(i: 0)->getType());
147 Check(MDataRegs && CData && Data && CData->ElemSize == Data->ElemSize,
148 "invalid llvm.nvvm.spcompress operand or result type", &Call);
149
150 unsigned IdxSize = cast<ConstantInt>(Val: Call.getArgOperand(i: 2))->getZExtValue();
151 unsigned NumTgt = cast<ConstantInt>(Val: Call.getArgOperand(i: 3))->getZExtValue();
152 // spcompress only implements the 2:4 pattern: each group of num_tgt = 4
153 // data elements keeps 2 of them in cdata. The repeat factor counts pairs of
154 // data registers, so data must occupy an even number of them.
155 Check(NumTgt == 4 && Data->NumElements % NumTgt == 0 &&
156 CData->NumElements == 2 * (Data->NumElements / NumTgt) &&
157 Data->NumRegisters % 2 == 0,
158 "invalid llvm.nvvm.spcompress layout", &Call);
159
160 unsigned RepeatFactor = Data->NumRegisters / 2;
161 auto Layout = getSPCompressLayout(ElemSize: Data->ElemSize, IdxSize, RepeatFactor);
162 // The declared types must match the register layout PTX gives these
163 // qualifiers.
164 Check(Layout && *MDataRegs == Layout->MetadataSize &&
165 CData->NumRegisters == Layout->CompressedDataSize &&
166 Data->NumRegisters == Layout->DataSize,
167 "invalid llvm.nvvm.spcompress layout", &Call);
168}
169
170static void verifySPDecompress(VerifierSupport &VS, CallBase &Call) {
171 auto Data = getSPVectorInfo(Ty: Call.getType());
172 auto MDataRegs = getSPMetadataRegisters(Ty: Call.getArgOperand(i: 0)->getType());
173 auto CData = getSPVectorInfo(Ty: Call.getArgOperand(i: 1)->getType());
174 Check(Data && MDataRegs && CData && Data->ElemSize == CData->ElemSize,
175 "invalid llvm.nvvm.spdecompress operand or result type", &Call);
176
177 unsigned IdxSize = cast<ConstantInt>(Val: Call.getArgOperand(i: 2))->getZExtValue();
178 unsigned NumTgt = cast<ConstantInt>(Val: Call.getArgOperand(i: 3))->getZExtValue();
179 // data holds repeat_factor groups of num_tgt elements.
180 Check(NumTgt != 0 && Data->NumElements % NumTgt == 0,
181 "invalid llvm.nvvm.spdecompress layout", &Call);
182
183 unsigned RepeatFactor = Data->NumElements / NumTgt;
184 // cdata holds num_src elements for each of those groups.
185 Check(RepeatFactor != 0 && CData->NumElements % RepeatFactor == 0,
186 "invalid llvm.nvvm.spdecompress layout", &Call);
187
188 unsigned NumSrc = CData->NumElements / RepeatFactor;
189 auto Layout = getSPDecompressLayout(NumSrc, NumTgt, ElemSize: Data->ElemSize, IdxSize,
190 RepeatFactor);
191 // The declared types must match the register layout PTX gives these
192 // qualifiers.
193 Check(Layout && *MDataRegs == Layout->MetadataSize &&
194 CData->NumRegisters == Layout->CompressedDataSize &&
195 Data->NumRegisters == Layout->DataSize,
196 "invalid llvm.nvvm.spdecompress layout", &Call);
197}
198
199void llvm::verifyNVVMIntrinsicCall(VerifierSupport &VS, Intrinsic::ID ID,
200 CallBase &Call) {
201 switch (ID) {
202 default:
203 return;
204 case Intrinsic::nvvm_spcompress:
205 verifySPCompress(VS, Call);
206 return;
207 case Intrinsic::nvvm_spdecompress:
208 verifySPDecompress(VS, Call);
209 return;
210 }
211}
212
213#undef Check
214