/src/spirv-tools/source/util/parse_number.cpp
Line | Count | Source |
1 | | // Copyright (c) 2016 Google Inc. |
2 | | // |
3 | | // Licensed under the Apache License, Version 2.0 (the "License"); |
4 | | // you may not use this file except in compliance with the License. |
5 | | // You may obtain a copy of the License at |
6 | | // |
7 | | // http://www.apache.org/licenses/LICENSE-2.0 |
8 | | // |
9 | | // Unless required by applicable law or agreed to in writing, software |
10 | | // distributed under the License is distributed on an "AS IS" BASIS, |
11 | | // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. |
12 | | // See the License for the specific language governing permissions and |
13 | | // limitations under the License. |
14 | | |
15 | | #include "source/util/parse_number.h" |
16 | | |
17 | | #include <functional> |
18 | | #include <iomanip> |
19 | | #include <memory> |
20 | | #include <sstream> |
21 | | #include <string> |
22 | | #include <tuple> |
23 | | |
24 | | #include "source/util/hex_float.h" |
25 | | #include "source/util/make_unique.h" |
26 | | |
27 | | namespace spvtools { |
28 | | namespace utils { |
29 | | namespace { |
30 | | |
31 | | // A helper class that temporarily stores error messages and dump the messages |
32 | | // to a string which given as as pointer when it is destructed. If the given |
33 | | // pointer is a nullptr, this class does not store error message. |
34 | | class ErrorMsgStream { |
35 | | public: |
36 | | explicit ErrorMsgStream(std::string* error_msg_sink) |
37 | 220k | : error_msg_sink_(error_msg_sink) { |
38 | 220k | if (error_msg_sink_) stream_ = MakeUnique<std::ostringstream>(); |
39 | 220k | } |
40 | 220k | ~ErrorMsgStream() { |
41 | 220k | if (error_msg_sink_ && stream_) *error_msg_sink_ = stream_->str(); |
42 | 220k | } |
43 | | template <typename T> |
44 | 447k | ErrorMsgStream& operator<<(T val) { |
45 | 447k | if (stream_) *stream_ << val; |
46 | 447k | return *this; |
47 | 447k | } parse_number.cpp:spvtools::utils::(anonymous namespace)::ErrorMsgStream& spvtools::utils::(anonymous namespace)::ErrorMsgStream::operator<< <char const*>(char const*) Line | Count | Source | 44 | 443k | ErrorMsgStream& operator<<(T val) { | 45 | 443k | if (stream_) *stream_ << val; | 46 | 443k | return *this; | 47 | 443k | } |
parse_number.cpp:spvtools::utils::(anonymous namespace)::ErrorMsgStream& spvtools::utils::(anonymous namespace)::ErrorMsgStream::operator<< <unsigned int>(unsigned int) Line | Count | Source | 44 | 827 | ErrorMsgStream& operator<<(T val) { | 45 | 827 | if (stream_) *stream_ << val; | 46 | 827 | return *this; | 47 | 827 | } |
parse_number.cpp:spvtools::utils::(anonymous namespace)::ErrorMsgStream& spvtools::utils::(anonymous namespace)::ErrorMsgStream::operator<< <std::__1::ios_base& (*)(std::__1::ios_base&)>(std::__1::ios_base& (*)(std::__1::ios_base&)) Line | Count | Source | 44 | 2.25k | ErrorMsgStream& operator<<(T val) { | 45 | 2.25k | if (stream_) *stream_ << val; | 46 | 2.25k | return *this; | 47 | 2.25k | } |
parse_number.cpp:spvtools::utils::(anonymous namespace)::ErrorMsgStream& spvtools::utils::(anonymous namespace)::ErrorMsgStream::operator<< <long>(long) Line | Count | Source | 44 | 319 | ErrorMsgStream& operator<<(T val) { | 45 | 319 | if (stream_) *stream_ << val; | 46 | 319 | return *this; | 47 | 319 | } |
parse_number.cpp:spvtools::utils::(anonymous namespace)::ErrorMsgStream& spvtools::utils::(anonymous namespace)::ErrorMsgStream::operator<< <unsigned long>(unsigned long) Line | Count | Source | 44 | 433 | ErrorMsgStream& operator<<(T val) { | 45 | 433 | if (stream_) *stream_ << val; | 46 | 433 | return *this; | 47 | 433 | } |
parse_number.cpp:spvtools::utils::(anonymous namespace)::ErrorMsgStream& spvtools::utils::(anonymous namespace)::ErrorMsgStream::operator<< <int>(int) Line | Count | Source | 44 | 10 | ErrorMsgStream& operator<<(T val) { | 45 | 10 | if (stream_) *stream_ << val; | 46 | 10 | return *this; | 47 | 10 | } |
|
48 | | |
49 | | private: |
50 | | std::unique_ptr<std::ostringstream> stream_; |
51 | | // The destination string to which this class dump the error message when |
52 | | // destructor is called. |
53 | | std::string* error_msg_sink_; |
54 | | }; |
55 | | } // namespace |
56 | | |
57 | | EncodeNumberStatus ParseAndEncodeIntegerNumber( |
58 | | const char* text, const NumberType& type, |
59 | 9.00M | std::function<void(uint32_t)> emit, std::string* error_msg) { |
60 | 9.00M | if (!text) { |
61 | 0 | ErrorMsgStream(error_msg) << "The given text is a nullptr"; |
62 | 0 | return EncodeNumberStatus::kInvalidText; |
63 | 0 | } |
64 | | |
65 | 9.00M | if (!IsIntegral(type)) { |
66 | 0 | ErrorMsgStream(error_msg) << "The expected type is not a integer type"; |
67 | 0 | return EncodeNumberStatus::kInvalidUsage; |
68 | 0 | } |
69 | | |
70 | 9.00M | const uint32_t bit_width = AssumedBitWidth(type); |
71 | | |
72 | 9.00M | if (bit_width > 64) { |
73 | 75 | ErrorMsgStream(error_msg) |
74 | 75 | << "Unsupported " << bit_width << "-bit integer literals"; |
75 | 75 | return EncodeNumberStatus::kUnsupported; |
76 | 75 | } |
77 | | |
78 | | // Either we are expecting anything or integer. |
79 | 9.00M | bool is_negative = text[0] == '-'; |
80 | 9.00M | bool can_be_signed = IsSigned(type); |
81 | | |
82 | 9.00M | if (is_negative && !can_be_signed) { |
83 | 9 | ErrorMsgStream(error_msg) |
84 | 9 | << "Cannot put a negative number in an unsigned literal"; |
85 | 9 | return EncodeNumberStatus::kInvalidUsage; |
86 | 9 | } |
87 | | |
88 | 9.00M | const bool is_hex = text[0] == '0' && (text[1] == 'x' || text[1] == 'X'); |
89 | | |
90 | 9.00M | uint64_t decoded_bits; |
91 | 9.00M | if (is_negative) { |
92 | 9.96k | int64_t decoded_signed = 0; |
93 | | |
94 | 9.96k | if (!ParseNumber(text, &decoded_signed)) { |
95 | 162 | ErrorMsgStream(error_msg) << "Invalid signed integer literal: " << text; |
96 | 162 | return EncodeNumberStatus::kInvalidText; |
97 | 162 | } |
98 | | |
99 | 9.80k | if (!CheckRangeAndIfHexThenSignExtend(decoded_signed, type, is_hex, |
100 | 9.80k | &decoded_signed)) { |
101 | 319 | ErrorMsgStream(error_msg) |
102 | 319 | << "Integer " << (is_hex ? std::hex : std::dec) << std::showbase |
103 | 319 | << decoded_signed << " does not fit in a " << std::dec << bit_width |
104 | 319 | << "-bit " << (IsSigned(type) ? "signed" : "unsigned") << " integer"; |
105 | 319 | return EncodeNumberStatus::kInvalidText; |
106 | 319 | } |
107 | 9.48k | decoded_bits = decoded_signed; |
108 | 8.99M | } else { |
109 | | // There's no leading minus sign, so parse it as an unsigned integer. |
110 | 8.99M | if (!ParseNumber(text, &decoded_bits)) { |
111 | 208k | ErrorMsgStream(error_msg) << "Invalid unsigned integer literal: " << text; |
112 | 208k | return EncodeNumberStatus::kInvalidText; |
113 | 208k | } |
114 | 8.78M | if (!CheckRangeAndIfHexThenSignExtend(decoded_bits, type, is_hex, |
115 | 8.78M | &decoded_bits)) { |
116 | 433 | ErrorMsgStream(error_msg) |
117 | 433 | << "Integer " << (is_hex ? std::hex : std::dec) << std::showbase |
118 | 433 | << decoded_bits << " does not fit in a " << std::dec << bit_width |
119 | 433 | << "-bit " << (IsSigned(type) ? "signed" : "unsigned") << " integer"; |
120 | 433 | return EncodeNumberStatus::kInvalidText; |
121 | 433 | } |
122 | 8.78M | } |
123 | 8.79M | if (bit_width > 32) { |
124 | 1.69k | uint32_t low = uint32_t(0x00000000ffffffff & decoded_bits); |
125 | 1.69k | uint32_t high = uint32_t((0xffffffff00000000 & decoded_bits) >> 32); |
126 | 1.69k | emit(low); |
127 | 1.69k | emit(high); |
128 | 8.79M | } else { |
129 | 8.79M | emit(uint32_t(decoded_bits)); |
130 | 8.79M | } |
131 | 8.79M | return EncodeNumberStatus::kSuccess; |
132 | 9.00M | } |
133 | | |
134 | 288k | spv_fp_encoding_t DeduceEncoding(const NumberType& type) { |
135 | 288k | if (type.encoding != SPV_FP_ENCODING_UNKNOWN) return type.encoding; |
136 | 85.1k | switch (type.bitwidth) { |
137 | 32.2k | case 16: |
138 | 32.2k | return SPV_FP_ENCODING_IEEE754_BINARY16; |
139 | 29.4k | case 32: |
140 | 29.4k | return SPV_FP_ENCODING_IEEE754_BINARY32; |
141 | 23.4k | case 64: |
142 | 23.4k | return SPV_FP_ENCODING_IEEE754_BINARY64; |
143 | 10 | default: |
144 | 10 | return SPV_FP_ENCODING_UNKNOWN; |
145 | 85.1k | } |
146 | 85.1k | } |
147 | | EncodeNumberStatus ParseAndEncodeFloatingPointNumber( |
148 | | const char* text, const NumberType& type, |
149 | 288k | std::function<void(uint32_t)> emit, std::string* error_msg) { |
150 | 288k | if (!text) { |
151 | 0 | ErrorMsgStream(error_msg) << "The given text is a nullptr"; |
152 | 0 | return EncodeNumberStatus::kInvalidText; |
153 | 0 | } |
154 | | |
155 | 288k | if (!IsFloating(type)) { |
156 | 0 | ErrorMsgStream(error_msg) << "The expected type is not a float type"; |
157 | 0 | return EncodeNumberStatus::kInvalidUsage; |
158 | 0 | } |
159 | | |
160 | 288k | const auto bit_width = AssumedBitWidth(type); |
161 | 288k | switch (DeduceEncoding(type)) { |
162 | 34.3k | case SPV_FP_ENCODING_FLOAT4_E2M1: { |
163 | 34.3k | HexFloat<FloatProxy<Float4_E2M1>> hVal(0); |
164 | 34.3k | if (!ParseNumber(text, &hVal)) { |
165 | 370 | ErrorMsgStream(error_msg) << "Invalid E2M1 float literal: " << text; |
166 | 370 | return EncodeNumberStatus::kInvalidText; |
167 | 370 | } |
168 | | // getAsFloat will return the Float16 value, and get_value |
169 | | // will return a uint8_t representing the bits of the float. |
170 | | // The encoding is therefore correct from the perspective of the SPIR-V |
171 | | // spec since the top 28 bits will be 0. |
172 | 33.9k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
173 | 33.9k | return EncodeNumberStatus::kSuccess; |
174 | 34.3k | } break; |
175 | 32.0k | case SPV_FP_ENCODING_FLOAT6_E2M3: { |
176 | 32.0k | HexFloat<FloatProxy<Float6_E2M3>> hVal(0); |
177 | 32.0k | if (!ParseNumber(text, &hVal)) { |
178 | 398 | ErrorMsgStream(error_msg) << "Invalid E2M3 float literal: " << text; |
179 | 398 | return EncodeNumberStatus::kInvalidText; |
180 | 398 | } |
181 | | // getAsFloat will return the Float16 value, and get_value |
182 | | // will return a uint8_t representing the bits of the float. |
183 | | // The encoding is therefore correct from the perspective of the SPIR-V |
184 | | // spec since the top 26 bits will be 0. |
185 | 31.6k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
186 | 31.6k | return EncodeNumberStatus::kSuccess; |
187 | 32.0k | } break; |
188 | 31.9k | case SPV_FP_ENCODING_FLOAT6_E3M2: { |
189 | 31.9k | HexFloat<FloatProxy<Float6_E3M2>> hVal(0); |
190 | 31.9k | if (!ParseNumber(text, &hVal)) { |
191 | 440 | ErrorMsgStream(error_msg) << "Invalid E3M2 float literal: " << text; |
192 | 440 | return EncodeNumberStatus::kInvalidText; |
193 | 440 | } |
194 | | // getAsFloat will return the Float16 value, and get_value |
195 | | // will return a uint8_t representing the bits of the float. |
196 | | // The encoding is therefore correct from the perspective of the SPIR-V |
197 | | // spec since the top 26 bits will be 0. |
198 | 31.5k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
199 | 31.5k | return EncodeNumberStatus::kSuccess; |
200 | 31.9k | } break; |
201 | 30.7k | case SPV_FP_ENCODING_FLOAT8_E4M3: { |
202 | 30.7k | HexFloat<FloatProxy<Float8_E4M3>> hVal(0); |
203 | 30.7k | if (!ParseNumber(text, &hVal)) { |
204 | 392 | ErrorMsgStream(error_msg) << "Invalid E4M3 float literal: " << text; |
205 | 392 | return EncodeNumberStatus::kInvalidText; |
206 | 392 | } |
207 | | // getAsFloat will return the Float16 value, and get_value |
208 | | // will return a uint16_t representing the bits of the float. |
209 | | // The encoding is therefore correct from the perspective of the SPIR-V |
210 | | // spec since the top 16 bits will be 0. |
211 | 30.3k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
212 | 30.3k | return EncodeNumberStatus::kSuccess; |
213 | 30.7k | } break; |
214 | 30.3k | case SPV_FP_ENCODING_FLOAT8_E5M2: { |
215 | 30.3k | HexFloat<FloatProxy<Float8_E5M2>> hVal(0); |
216 | 30.3k | if (!ParseNumber(text, &hVal)) { |
217 | 444 | ErrorMsgStream(error_msg) << "Invalid E5M2 float literal: " << text; |
218 | 444 | return EncodeNumberStatus::kInvalidText; |
219 | 444 | } |
220 | | // getAsFloat will return the Float16 value, and get_value |
221 | | // will return a uint16_t representing the bits of the float. |
222 | | // The encoding is therefore correct from the perspective of the SPIR-V |
223 | | // spec since the top 16 bits will be 0. |
224 | 29.9k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
225 | 29.9k | return EncodeNumberStatus::kSuccess; |
226 | 30.3k | } break; |
227 | 4.35k | case SPV_FP_ENCODING_FLOAT8_UNSIGNED_E8M0: { |
228 | 4.35k | HexFloat<FloatProxy<Float8_E8M0>> hVal(0); |
229 | 4.35k | if (!ParseNumber(text, &hVal)) { |
230 | 240 | ErrorMsgStream(error_msg) << "Invalid E8M0 float literal: " << text; |
231 | 240 | return EncodeNumberStatus::kInvalidText; |
232 | 240 | } |
233 | | // getAsFloat will return the Float16 value, and get_value |
234 | | // will return a uint16_t representing the bits of the float. |
235 | | // The encoding is therefore correct from the perspective of the SPIR-V |
236 | | // spec since the top 16 bits will be 0. |
237 | 4.11k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
238 | 4.11k | return EncodeNumberStatus::kSuccess; |
239 | 4.35k | } break; |
240 | 9.98k | case SPV_FP_ENCODING_MXINT8: { |
241 | 9.98k | HexFixedPoint<MXInt8> hVal; |
242 | 9.98k | if (!ParseNumber(text, &hVal)) { |
243 | 488 | ErrorMsgStream(error_msg) << "Invalid MXInt8 float literal: " << text; |
244 | 488 | return EncodeNumberStatus::kInvalidText; |
245 | 488 | } |
246 | 9.50k | emit(static_cast<uint32_t>(hVal.value())); |
247 | 9.50k | return EncodeNumberStatus::kSuccess; |
248 | 9.98k | } break; |
249 | 29.6k | case SPV_FP_ENCODING_BFLOAT16: { |
250 | 29.6k | HexFloat<FloatProxy<BFloat16>> hVal(0); |
251 | 29.6k | if (!ParseNumber(text, &hVal)) { |
252 | 498 | ErrorMsgStream(error_msg) << "Invalid bfloat16 literal: " << text; |
253 | 498 | return EncodeNumberStatus::kInvalidText; |
254 | 498 | } |
255 | 29.1k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
256 | 29.1k | return EncodeNumberStatus::kSuccess; |
257 | 29.6k | } break; |
258 | 32.2k | case SPV_FP_ENCODING_IEEE754_BINARY16: { |
259 | 32.2k | HexFloat<FloatProxy<Float16>> hVal(0); |
260 | 32.2k | if (!ParseNumber(text, &hVal)) { |
261 | 470 | ErrorMsgStream(error_msg) << "Invalid 16-bit float literal: " << text; |
262 | 470 | return EncodeNumberStatus::kInvalidText; |
263 | 470 | } |
264 | | // getAsFloat will return the Float16 value, and get_value |
265 | | // will return a uint16_t representing the bits of the float. |
266 | | // The encoding is therefore correct from the perspective of the SPIR-V |
267 | | // spec since the top 16 bits will be 0. |
268 | 31.7k | emit(static_cast<uint32_t>(hVal.value().getAsFloat().get_value())); |
269 | 31.7k | return EncodeNumberStatus::kSuccess; |
270 | 32.2k | } break; |
271 | 29.4k | case SPV_FP_ENCODING_IEEE754_BINARY32: { |
272 | 29.4k | HexFloat<FloatProxy<float>> fVal(0.0f); |
273 | 29.4k | if (!ParseNumber(text, &fVal)) { |
274 | 6.37k | ErrorMsgStream(error_msg) << "Invalid 32-bit float literal: " << text; |
275 | 6.37k | return EncodeNumberStatus::kInvalidText; |
276 | 6.37k | } |
277 | 23.0k | emit(BitwiseCast<uint32_t>(fVal)); |
278 | 23.0k | return EncodeNumberStatus::kSuccess; |
279 | 29.4k | } break; |
280 | 23.4k | case SPV_FP_ENCODING_IEEE754_BINARY64: { |
281 | 23.4k | HexFloat<FloatProxy<double>> dVal(0.0); |
282 | 23.4k | if (!ParseNumber(text, &dVal)) { |
283 | 768 | ErrorMsgStream(error_msg) << "Invalid 64-bit float literal: " << text; |
284 | 768 | return EncodeNumberStatus::kInvalidText; |
285 | 768 | } |
286 | 22.6k | uint64_t decoded_val = BitwiseCast<uint64_t>(dVal); |
287 | 22.6k | uint32_t low = uint32_t(0x00000000ffffffff & decoded_val); |
288 | 22.6k | uint32_t high = uint32_t((0xffffffff00000000 & decoded_val) >> 32); |
289 | 22.6k | emit(low); |
290 | 22.6k | emit(high); |
291 | 22.6k | return EncodeNumberStatus::kSuccess; |
292 | 23.4k | } break; |
293 | 10 | default: |
294 | 10 | break; |
295 | 288k | } |
296 | 10 | ErrorMsgStream(error_msg) |
297 | 10 | << "Unsupported " << bit_width << "-bit float literals"; |
298 | 10 | return EncodeNumberStatus::kUnsupported; |
299 | 288k | } |
300 | | |
301 | | EncodeNumberStatus ParseAndEncodeNumber(const char* text, |
302 | | const NumberType& type, |
303 | | std::function<void(uint32_t)> emit, |
304 | 9.29M | std::string* error_msg) { |
305 | 9.29M | if (!text) { |
306 | 0 | ErrorMsgStream(error_msg) << "The given text is a nullptr"; |
307 | 0 | return EncodeNumberStatus::kInvalidText; |
308 | 0 | } |
309 | | |
310 | 9.29M | if (IsUnknown(type)) { |
311 | 0 | ErrorMsgStream(error_msg) |
312 | 0 | << "The expected type is not a integer or float type"; |
313 | 0 | return EncodeNumberStatus::kInvalidUsage; |
314 | 0 | } |
315 | | |
316 | | // If we explicitly expect a floating-point number, we should handle that |
317 | | // first. |
318 | 9.29M | if (IsFloating(type)) { |
319 | 288k | return ParseAndEncodeFloatingPointNumber(text, type, emit, error_msg); |
320 | 288k | } |
321 | | |
322 | 9.00M | return ParseAndEncodeIntegerNumber(text, type, emit, error_msg); |
323 | 9.29M | } |
324 | | |
325 | | } // namespace utils |
326 | | } // namespace spvtools |