Line data Source code
1 : /**
2 : * Copyright (c) 2025 Huawei Technologies Co., Ltd.
3 : * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4 : * CANN Open Software License Agreement Version 2.0 (the "License").
5 : * Please refer to the License for details. You may not use this file except in compliance with the License.
6 : * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7 : * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8 : * See LICENSE in the root of the software repository for the full text of the License.
9 : */
10 : #include "fp16_t.h"
11 : #include "securec.h"
12 : #include "log/hdc_log.h"
13 :
14 : namespace Adx {
15 : namespace {
16 37 : inline float uint32ToFloat(uint32_t val)
17 : {
18 37 : float result = 0.0f;
19 37 : auto ret = memcpy_s(&result, sizeof(result), &val, sizeof(val));
20 37 : if (ret != EOK) {
21 0 : IDE_LOGE("memcpy_s from uint32 to float failed, ret=%d.", ret);
22 0 : return 0.0f;
23 : }
24 37 : return result;
25 : }
26 : }
27 :
28 : /**
29 : * @ingroup fp16_t global filed
30 : * @brief round mode of last valid digital
31 : */
32 : const fp16RoundMode_t g_RoundMode = ROUND_TO_NEAREST;
33 :
34 40 : void ExtractFP16(const uint16_t &val, uint16_t *s, int16_t *e, uint16_t *m)
35 : {
36 : // 1.Extract
37 40 : *s = FP16_EXTRAC_SIGN(val);
38 40 : *e = FP16_EXTRAC_EXP(val);
39 40 : *m = FP16_EXTRAC_MAN(val);
40 :
41 : // Denormal
42 40 : if (0 == (*e)) {
43 4 : *e = 1;
44 : }
45 40 : }
46 : /**
47 : * @ingroup fp16_t static method
48 : * @param [in] man truncated mantissa
49 : * @param [in] shiftOut left shift bits based on ten bits
50 : * @brief judge whether to add one to the result while converting fp16_t to other datatype
51 : * @return Return true if add one, otherwise false
52 : */
53 38 : static bool IsRoundOne(uint64_t man, uint16_t truncLen)
54 : {
55 38 : uint16_t shiftOut = truncLen - DIM_2;
56 38 : uint64_t mask0 = 0x4;
57 38 : mask0 = mask0 << shiftOut;
58 38 : uint64_t mask1 = 0x2;
59 38 : mask1 = mask1 << shiftOut;
60 : uint64_t mask2;
61 38 : mask2 = mask1 - 1;
62 :
63 38 : bool lastBit = ((man & mask0) > 0);
64 38 : bool truncHigh = false;
65 38 : bool truncLeft = false;
66 : if (ROUND_TO_NEAREST == g_RoundMode) {
67 38 : truncHigh = ((man & mask1) > 0);
68 38 : truncLeft = ((man & mask2) > 0);
69 : }
70 38 : return (truncHigh && (truncLeft || lastBit));
71 : }
72 : /**
73 : * @ingroup fp16_t public method
74 : * @param [in] exp exponent of fp16_t value
75 : * @param [in] man exponent of fp16_t value
76 : * @brief normalize fp16_t value
77 : * @return
78 : */
79 48 : static void Fp16Normalize(int16_t &exp, uint16_t &man)
80 : {
81 48 : if (exp >= FP16_MAX_EXP) {
82 0 : exp = FP16_MAX_EXP - 1;
83 0 : man = FP16_MAX_MAN;
84 48 : } else if (exp == 0 && man == FP16_MAN_HIDE_BIT) {
85 0 : exp++;
86 0 : man = 0;
87 : }
88 48 : }
89 :
90 : /**
91 : * @ingroup fp16_t math conversion static method
92 : * @param [in] fpVal uint16_t value of fp16_t object
93 : * @brief Convert fp16_t to float/fp32
94 : * @return Return float/fp32 value of fpVal which is the value of fp16_t object
95 : */
96 37 : static float fp16ToFloat(const uint16_t &fpVal)
97 : {
98 : uint16_t hfSign;
99 : uint16_t hfMan;
100 : int16_t hfExp;
101 37 : ExtractFP16(fpVal, &hfSign, &hfExp, &hfMan);
102 :
103 37 : if (hfExp == FP16_MAX_EXP) {
104 7 : if (hfMan == FP16_MAN_HIDE_BIT) {
105 : // Infinity
106 5 : uint32_t fVal = (hfSign << FP32_SIGN_INDEX) | FP32_EXP_MASK;
107 5 : return uint32ToFloat(fVal);
108 : } else {
109 : // NaN
110 2 : uint32_t mRet = (hfMan & FP16_MAN_MASK) << (FP32_MAN_LEN - FP16_MAN_LEN);
111 2 : uint32_t fVal = (hfSign << FP32_SIGN_INDEX) | FP32_EXP_MASK | mRet;
112 2 : return uint32ToFloat(fVal);
113 : }
114 : }
115 :
116 40 : while (hfMan && !(hfMan & FP16_MAN_HIDE_BIT)) {
117 10 : hfMan <<= 1;
118 10 : hfExp--;
119 : }
120 :
121 : uint32_t sRet;
122 : uint32_t eRet;
123 : uint32_t mRet;
124 : uint32_t fVal;
125 :
126 30 : sRet = hfSign;
127 30 : if (!hfMan) {
128 2 : eRet = 0;
129 2 : mRet = 0;
130 : } else {
131 28 : eRet = static_cast<uint32_t>(hfExp - FP16_EXP_BIAS + FP32_EXP_BIAS);
132 28 : mRet = hfMan & FP16_MAN_MASK;
133 28 : mRet = mRet << (FP32_MAN_LEN - FP16_MAN_LEN);
134 : }
135 30 : fVal = FP32_CONSTRUCTOR(sRet, eRet, mRet);
136 30 : return uint32ToFloat(fVal);
137 : }
138 :
139 : // evaluation
140 22 : fp16_t &fp16_t::operator=(const fp16_t &fp)
141 : {
142 22 : if (this == &fp) {
143 1 : return *this;
144 : }
145 21 : val = fp.val;
146 21 : return *this;
147 : }
148 17 : fp16_t &fp16_t::operator=(const float &fVal)
149 : {
150 : uint16_t sRet;
151 : uint16_t mRet;
152 : int16_t eRet;
153 : uint32_t eF;
154 : uint32_t mF;
155 17 : uint32_t ui32V = *(reinterpret_cast<const uint32_t *>(&fVal)); // 1:8:23bit sign:exp:man
156 : uint32_t mLenDelta;
157 :
158 17 : sRet = static_cast<uint16_t>((ui32V & FP32_SIGN_MASK) >> FP32_SIGN_INDEX); // 4Byte->2Byte
159 17 : eF = (ui32V & FP32_EXP_MASK) >> FP32_MAN_LEN; // 8 bit exponent
160 17 : mF = (ui32V & FP32_MAN_MASK); // 23 bit mantissa dont't need to care about denormal
161 17 : mLenDelta = FP32_MAN_LEN - FP16_MAN_LEN;
162 :
163 17 : bool needRound = false;
164 : // Exponent overflow/NaN converts to signed inf/NaN
165 17 : if (eF > 0x8Fu) { // 0x8Fu:142=127+15
166 3 : eRet = FP16_MAX_EXP - 1;
167 3 : mRet = FP16_MAX_MAN;
168 14 : } else if (eF <= 0x70u) { // 0x70u:112=127-15 Exponent underflow converts to denormalized half or signed zero
169 4 : eRet = 0;
170 4 : if (eF >= 0x67) { // 0x67:103=127-24 Denormal
171 1 : mF = (mF | FP32_MAN_HIDE_BIT);
172 1 : uint16_t shiftOut = FP32_MAN_LEN;
173 1 : uint64_t mTmp = ((uint64_t) mF) << (eF - 0x67);
174 :
175 1 : needRound = IsRoundOne(mTmp, shiftOut);
176 1 : mRet = static_cast<uint16_t>(mTmp >> shiftOut);
177 1 : if (needRound) {
178 1 : mRet++;
179 : }
180 3 : } else if (eF == 0x66 && mF > 0) { // 0x66:102 Denormal 0<f_v<min(Denormal)
181 1 : mRet = 1;
182 : } else {
183 2 : mRet = 0;
184 : }
185 : } else { // Regular case with no overflow or underflow
186 10 : eRet = (int16_t) (eF - 0x70u);
187 :
188 10 : needRound = IsRoundOne(mF, mLenDelta);
189 10 : mRet = static_cast<uint16_t>(mF >> mLenDelta);
190 10 : if (needRound) {
191 2 : mRet++;
192 : }
193 10 : if (mRet & FP16_MAN_HIDE_BIT) {
194 0 : eRet++;
195 : }
196 : }
197 :
198 17 : Fp16Normalize(eRet, mRet);
199 17 : val = FP16_CONSTRUCTOR(sRet, (uint16_t) eRet, mRet);
200 17 : return *this;
201 : }
202 8 : fp16_t &fp16_t::operator=(const int8_t &iVal)
203 : {
204 : uint16_t sRet;
205 : uint16_t eRet;
206 : uint16_t mRet;
207 :
208 8 : sRet = (uint16_t) ((((uint8_t) iVal) & 0x80) >> DIM_7);
209 8 : mRet = (uint16_t) ((((uint8_t) iVal) & INT8_T_MAX));
210 :
211 8 : if (mRet == 0) {
212 2 : eRet = 0;
213 : } else {
214 6 : if (sRet) { // negative number(<0)
215 3 : mRet = (uint16_t) std::abs(iVal); // complement
216 : }
217 :
218 6 : eRet = FP16_MAN_LEN;
219 44 : while ((mRet & FP16_MAN_HIDE_BIT) == 0) {
220 38 : mRet = mRet << DIM_1;
221 38 : eRet = eRet - DIM_1;
222 : }
223 6 : eRet = eRet + FP16_EXP_BIAS;
224 : }
225 :
226 8 : val = FP16_CONSTRUCTOR(sRet, eRet, mRet);
227 8 : return *this;
228 : }
229 6 : fp16_t &fp16_t::operator=(const uint8_t &uiVal)
230 : {
231 : uint16_t sRet;
232 : uint16_t eRet;
233 : uint16_t mRet;
234 6 : sRet = 0;
235 6 : eRet = 0;
236 6 : mRet = uiVal;
237 6 : if (mRet) {
238 5 : eRet = FP16_MAN_LEN;
239 30 : while ((mRet & FP16_MAN_HIDE_BIT) == 0) {
240 25 : mRet = mRet << DIM_1;
241 25 : eRet = eRet - DIM_1;
242 : }
243 5 : eRet = eRet + FP16_EXP_BIAS;
244 : }
245 :
246 6 : val = FP16_CONSTRUCTOR(sRet, eRet, mRet);
247 6 : return *this;
248 : }
249 9 : fp16_t &fp16_t::operator=(const int16_t &iVal)
250 : {
251 9 : if (iVal == 0) {
252 1 : val = 0;
253 : } else {
254 : uint16_t sRet;
255 8 : uint16_t uiVal = *(reinterpret_cast<const uint16_t *>(&iVal));
256 8 : sRet = (uint16_t) (uiVal >> BitShift_15);
257 8 : if (sRet) {
258 4 : int16_t iValM = -iVal;
259 4 : uiVal = *(reinterpret_cast<uint16_t *>(&iValM));
260 : }
261 8 : uint32_t mTmp = (uiVal & FP32_ABS_MAX);
262 8 : uint16_t mMin = FP16_MAN_HIDE_BIT;
263 8 : uint16_t mMax = mMin << 1;
264 8 : uint16_t len = (uint16_t) GetManBitLength(mTmp);
265 8 : if (mTmp) {
266 : int16_t eRet;
267 8 : if (len > DIM_11) {
268 4 : eRet = FP16_EXP_BIAS + FP16_MAN_LEN;
269 4 : uint16_t eTmp = len - DIM_11;
270 4 : uint32_t truncMask = 1;
271 16 : for (int i = 1; i < eTmp; i++) {
272 12 : truncMask = (truncMask << 1) + 1;
273 : }
274 4 : uint32_t mTrunc = (mTmp & truncMask) << (BitShift_32 - eTmp);
275 20 : for (int i = 0; i < eTmp; i++) {
276 16 : mTmp = (mTmp >> 1);
277 16 : eRet = eRet + 1;
278 : }
279 4 : bool bLastBit = ((mTmp & 1) > 0);
280 4 : bool bTruncHigh = false;
281 4 : bool bTruncLeft = false;
282 : if (ROUND_TO_NEAREST == g_RoundMode) { // trunc
283 4 : bTruncHigh = ((mTrunc & FP32_SIGN_MASK) > 0);
284 4 : bTruncLeft = ((mTrunc & FP32_ABS_MAX) > 0);
285 : }
286 4 : mTmp = ManRoundToNearest(bLastBit, bTruncHigh, bTruncLeft, mTmp);
287 5 : while (mTmp >= mMax || eRet < 0) {
288 1 : mTmp = mTmp >> 1;
289 1 : eRet = eRet + 1;
290 : }
291 : } else {
292 4 : eRet = FP16_EXP_BIAS;
293 4 : mTmp = mTmp << (11 - len);
294 4 : eRet = eRet + (len - 1);
295 : }
296 8 : uint16_t mRet = (uint16_t) mTmp;
297 8 : val = FP16_CONSTRUCTOR(sRet, (uint16_t) eRet, mRet);
298 : }
299 : }
300 9 : return *this;
301 : }
302 6 : fp16_t &fp16_t::operator=(const uint16_t &uiVal)
303 : {
304 6 : if (uiVal == 0) {
305 1 : val = 0;
306 : } else {
307 : int16_t eRet;
308 5 : uint16_t mRet = uiVal;
309 :
310 5 : uint16_t mMin = FP16_MAN_HIDE_BIT;
311 5 : uint16_t mMax = mMin << 1;
312 5 : uint16_t len = (uint16_t) GetManBitLength(mRet);
313 :
314 5 : if (len > 11) {
315 3 : eRet = FP16_EXP_BIAS + FP16_MAN_LEN;
316 : uint32_t mTrunc;
317 3 : uint32_t truncMask = 1;
318 3 : uint16_t eTmp = len - DIM_11;
319 15 : for (int i = 1; i < eTmp; i++) {
320 12 : truncMask = (truncMask << 1) + 1;
321 : }
322 3 : mTrunc = (mRet & truncMask) << (32 - eTmp);
323 18 : for (int i = 0; i < eTmp; i++) {
324 15 : mRet = (mRet >> 1);
325 15 : eRet = eRet + 1;
326 : }
327 3 : bool bLastBit = ((mRet & 1) > 0);
328 3 : bool bTruncHigh = false;
329 3 : bool bTruncLeft = false;
330 : if (ROUND_TO_NEAREST == g_RoundMode) { // trunc
331 3 : bTruncHigh = ((mTrunc & FP32_SIGN_MASK) > 0);
332 3 : bTruncLeft = ((mTrunc & FP32_ABS_MAX) > 0);
333 : }
334 3 : mRet = ManRoundToNearest(bLastBit, bTruncHigh, bTruncLeft, mRet);
335 5 : while (mRet >= mMax || eRet < 0) {
336 2 : mRet = mRet >> 1;
337 2 : eRet = eRet + 1;
338 : }
339 3 : if (FP16_IS_INVALID(val)) {
340 0 : val = FP16_MAX;
341 : }
342 : } else {
343 2 : eRet = FP16_EXP_BIAS;
344 2 : mRet = mRet << (DIM_11 - len);
345 2 : eRet = eRet + (len - 1);
346 : }
347 5 : val = FP16_CONSTRUCTOR(0u, (uint16_t) eRet, mRet);
348 : }
349 6 : return *this;
350 : }
351 13 : fp16_t &fp16_t::operator=(const int32_t &iVal)
352 : {
353 13 : if (iVal == 0) {
354 1 : val = 0;
355 : } else {
356 12 : uint32_t uiVal = *(reinterpret_cast<const uint32_t *>(&iVal));
357 12 : uint16_t sRet = (uint16_t) (uiVal >> BitShift_31);
358 12 : if (sRet) {
359 1 : int32_t iValM = -iVal;
360 1 : uiVal = *(reinterpret_cast<uint32_t *>(&iValM));
361 : }
362 : int16_t eRet;
363 12 : uint32_t mTmp = (uiVal & FP32_ABS_MAX);
364 12 : uint32_t mMin = FP16_MAN_HIDE_BIT;
365 12 : uint32_t mMax = mMin << 1;
366 12 : uint16_t len = (uint16_t) GetManBitLength(mTmp);
367 12 : if (len > DIM_11) {
368 6 : eRet = FP16_EXP_BIAS + FP16_MAN_LEN;
369 6 : uint32_t mTrunc = 0;
370 6 : uint32_t truncMask = 1;
371 6 : uint16_t eTmp = len - DIM_11;
372 46 : for (int i = 1; i < eTmp; i++) {
373 40 : truncMask = (truncMask << 1) + 1;
374 : }
375 6 : mTrunc = (mTmp & truncMask) << (BitShift_32 - eTmp);
376 52 : for (int i = 0; i < eTmp; i++) {
377 46 : mTmp = (mTmp >> 1);
378 46 : eRet = eRet + 1;
379 : }
380 6 : bool bLastBit = ((mTmp & 1) > 0);
381 6 : bool bTruncHigh = false;
382 6 : bool bTruncLeft = false;
383 : if (ROUND_TO_NEAREST == g_RoundMode) { // trunc
384 6 : bTruncHigh = ((mTrunc & FP32_SIGN_MASK) > 0);
385 6 : bTruncLeft = ((mTrunc & FP32_ABS_MAX) > 0);
386 : }
387 6 : mTmp = ManRoundToNearest(bLastBit, bTruncHigh, bTruncLeft, mTmp);
388 7 : while (mTmp >= mMax || eRet < 0) {
389 1 : mTmp = mTmp >> 1;
390 1 : eRet = eRet + 1;
391 : }
392 6 : if (eRet >= FP16_MAX_EXP) {
393 2 : eRet = FP16_MAX_EXP - 1;
394 2 : mTmp = FP16_MAX_MAN;
395 : }
396 : } else {
397 6 : eRet = FP16_EXP_BIAS;
398 6 : mTmp = mTmp << (DIM_11 - len);
399 6 : eRet = eRet + (len - 1);
400 : }
401 12 : uint16_t mRet = (uint16_t) mTmp;
402 12 : val = FP16_CONSTRUCTOR(sRet, (uint16_t) eRet, mRet);
403 : }
404 13 : return *this;
405 : }
406 11 : fp16_t &fp16_t::operator=(const uint32_t &uiVal)
407 : {
408 11 : if (uiVal == 0) {
409 1 : val = 0;
410 : } else {
411 : int16_t eRet;
412 10 : uint32_t mTmp = uiVal;
413 :
414 10 : uint32_t mMin = FP16_MAN_HIDE_BIT;
415 10 : uint32_t mMax = mMin << 1;
416 10 : uint16_t len = (uint16_t) GetManBitLength(mTmp);
417 :
418 10 : if (len > DIM_11) {
419 5 : eRet = FP16_EXP_BIAS + FP16_MAN_LEN;
420 5 : uint32_t mTrunc = 0;
421 5 : uint32_t truncMask = 1;
422 5 : uint16_t eTmp = len - DIM_11;
423 27 : for (int i = 1; i < eTmp; i++) {
424 22 : truncMask = (truncMask << 1) + 1;
425 : }
426 5 : mTrunc = (mTmp & truncMask) << (BitShift_32 - eTmp);
427 32 : for (int i = 0; i < eTmp; i++) {
428 27 : mTmp = (mTmp >> 1);
429 27 : eRet = eRet + 1;
430 : }
431 5 : bool bLastBit = ((mTmp & 1) > 0);
432 5 : bool bTruncHigh = false;
433 5 : bool bTruncLeft = false;
434 : if (ROUND_TO_NEAREST == g_RoundMode) { // trunc
435 5 : bTruncHigh = ((mTrunc & FP32_SIGN_MASK) > 0);
436 5 : bTruncLeft = ((mTrunc & FP32_ABS_MAX) > 0);
437 : }
438 5 : mTmp = ManRoundToNearest(bLastBit, bTruncHigh, bTruncLeft, mTmp);
439 5 : while (mTmp >= mMax || eRet < 0) {
440 0 : mTmp = mTmp >> 1;
441 0 : eRet = eRet + 1;
442 : }
443 5 : if (eRet >= FP16_MAX_EXP) {
444 1 : eRet = FP16_MAX_EXP - 1;
445 1 : mTmp = FP16_MAX_MAN;
446 : }
447 : } else {
448 5 : eRet = FP16_EXP_BIAS;
449 5 : mTmp = mTmp << (DIM_11 - len);
450 5 : eRet = eRet + (len - 1);
451 : }
452 10 : uint16_t mRet = (uint16_t) mTmp;
453 10 : val = FP16_CONSTRUCTOR(0u, (uint16_t) eRet, mRet);
454 : }
455 11 : return *this;
456 : }
457 31 : fp16_t &fp16_t::operator=(const double &dVal)
458 : {
459 : uint16_t sRet;
460 : uint16_t mRet;
461 : int16_t eRet;
462 : uint64_t eD;
463 : uint64_t mD;
464 31 : uint64_t ui64V = *((uint64_t *) &dVal); // 1:11:52bit sign:exp:man
465 : uint32_t mLenDelta;
466 :
467 31 : sRet = static_cast<uint16_t>((ui64V & FP64_SIGN_MASK) >> FP64_SIGN_INDEX); // 4Byte
468 31 : eD = (ui64V & FP64_EXP_MASK) >> FP64_MAN_LEN; // 10 bit exponent
469 31 : mD = (ui64V & FP64_MAN_MASK); // 52 bit mantissa
470 31 : mLenDelta = FP64_MAN_LEN - FP16_MAN_LEN;
471 :
472 31 : bool needRound = false;
473 : // Exponent overflow/NaN converts to signed inf/NaN
474 31 : if (eD >= 0x410u) { // 0x410:1040=1023+16
475 1 : eRet = FP16_MAX_EXP - 1;
476 1 : mRet = FP16_MAX_MAN;
477 1 : val = FP16_CONSTRUCTOR(sRet, (uint16_t) eRet, mRet);
478 30 : } else if (eD <= 0x3F0u) { // Exponent underflow converts to denormalized half or signed zero
479 : // 0x3F0:1008=1023-15
480 : /**
481 : * Signed zeros, denormalized floats, and floats with small
482 : * exponents all convert to signed zero half precision.
483 : */
484 4 : eRet = 0;
485 4 : if (eD >= 0x3E7u) { // 0x3E7u:999=1023-24 Denormal
486 : // Underflows to a denormalized value
487 1 : mD = (FP64_MAN_HIDE_BIT | mD);
488 1 : uint16_t shiftOut = FP64_MAN_LEN;
489 1 : uint64_t mTmp = ((uint64_t) mD) << (eD - 0x3E7u);
490 :
491 1 : needRound = IsRoundOne(mTmp, shiftOut);
492 1 : mRet = static_cast<uint16_t>(mTmp >> shiftOut);
493 1 : if (needRound) {
494 1 : mRet++;
495 : }
496 3 : } else if (eD == 0x3E6u && mD > 0) {
497 1 : mRet = 1;
498 : } else {
499 2 : mRet = 0;
500 : }
501 : } else { // Regular case with no overflow or underflow
502 26 : eRet = (int16_t) (eD - 0x3F0u);
503 :
504 26 : needRound = IsRoundOne(mD, mLenDelta);
505 26 : mRet = static_cast<uint16_t>(mD >> mLenDelta);
506 26 : if (needRound) {
507 11 : mRet++;
508 : }
509 26 : if (mRet & FP16_MAN_HIDE_BIT) {
510 0 : eRet++;
511 : }
512 : }
513 :
514 31 : Fp16Normalize(eRet, mRet);
515 31 : val = FP16_CONSTRUCTOR(sRet, (uint16_t) eRet, mRet);
516 31 : return *this;
517 : }
518 :
519 4 : tagFp16 &fp16_t::operator=(const int64_t &iVal)
520 : {
521 4 : return *this = (static_cast<int32_t>(iVal));
522 : }
523 4 : tagFp16 &fp16_t::operator=(const uint64_t &uiVal)
524 : {
525 4 : return *this = (static_cast<uint32_t>(uiVal));
526 : }
527 :
528 37 : float fp16_t::toFloat()
529 : {
530 37 : return fp16ToFloat(val);
531 : }
532 : } // namespace op
|