target/arm: Use FloatParts64 in f16_dotadd
Use softfloat-parts.h so that we can more naturally perform the required operations witha single rounding step. This happens to also simplify the NaN detection step. Signed-off-by: Richard Henderson <richard.henderson@linaro.org> Reviewed-by: Philippe Mathieu-Daudé <philmd@linaro.org> Message-id: 20260609192110.752384-4-richard.henderson@linaro.org Message-Id: <20260517002550.321291-11-richard.henderson@linaro.org> Signed-off-by: Peter Maydell <peter.maydell@linaro.org>
Richard Henderson committed
Jun 9, 2026 at 12:20 UTC
091b7cdeae8e6d95baa58ea9a078f3c66ac925f6
1 file changed
+40
-57
target/arm/tcg/sme_helper.c
+40
-57
@@ -27,6 +27,7 @@
27
#include "accel/tcg/helper-retaddr.h"
28
#include "qemu/int128.h"
29
#include "fpu/softfloat.h"
30
+#include "fpu/softfloat-parts.h"
31
#include "vec_internal.h"
32
#include "sve_ldst_internal.h"
33
@@ -1221,18 +1222,15 @@ static inline uint32_t bf16mop_ah_neg_adj_pair(uint32_t pair, uint32_t pg)
1222
}
1223
1224
static float32 f16_dotadd(float32 sum, uint32_t e1, uint32_t e2,
1224
- float_status *s_f16, float_status *s_std,
1225
- float_status *s_odd)
1225
+ float_status *s_f16, float_status *s_std)
1226
{
1227
/*
1228
- * We need three different float_status for different parts of this
1228
+ * We need two different float_status for different parts of this
1229
* operation:
1230
* - the input conversion of the float16 values must use the
1231
* f16-specific float_status, so that the FPCR.FZ16 control is applied
1232
* - operations on float32 including the final accumulation must use
1233
* the normal float_status, so that FPCR.FZ is applied
1234
- * - we have pre-set-up copy of s_std which is set to round-to-odd,
1235
- * for the multiply (see below)
1234
*/
1235
float16 h1r = e1 & 0xffff;
1236
float16 h1c = e1 >> 16;
@@ -1240,48 +1238,48 @@ static float32 f16_dotadd(float32 sum, uint32_t e1, uint32_t e2,
1238
float16 h2c = e2 >> 16;
1239
float32 t32;
1240
1241
+ FloatParts64 p1r = float16_unpack_canonical(h1r, s_f16);
1242
+ FloatParts64 p1c = float16_unpack_canonical(h1c, s_f16);
1243
+ FloatParts64 p2r = float16_unpack_canonical(h2r, s_f16);
1244
+ FloatParts64 p2c = float16_unpack_canonical(h2c, s_f16);
1245
+
1246
+ int all_mask = (float_cmask(p1r.cls) | float_cmask(p1c.cls) |
1247
+ float_cmask(p2r.cls) | float_cmask(p2c.cls));
1248
+
1249
/* C.f. FPProcessNaNs4 */
1244
- if (float16_is_any_nan(h1r) || float16_is_any_nan(h1c) ||
1245
- float16_is_any_nan(h2r) || float16_is_any_nan(h2c)) {
1250
+ if (unlikely(all_mask & float_cmask_anynan)) {
1251
float16 t16;
1252
1248
- if (float16_is_signaling_nan(h1r, s_f16)) {
1249
- t16 = h1r;
1250
- } else if (float16_is_signaling_nan(h1c, s_f16)) {
1251
- t16 = h1c;
1252
- } else if (float16_is_signaling_nan(h2r, s_f16)) {
1253
- t16 = h2r;
1254
- } else if (float16_is_signaling_nan(h2c, s_f16)) {
1255
- t16 = h2c;
1256
- } else if (float16_is_any_nan(h1r)) {
1257
- t16 = h1r;
1258
- } else if (float16_is_any_nan(h1c)) {
1259
- t16 = h1c;
1260
- } else if (float16_is_any_nan(h2r)) {
1261
- t16 = h2r;
1253
+ if (unlikely(all_mask & float_cmask_snan)) {
1254
+ if (p1r.cls == float_class_snan) {
1255
+ t16 = h1r;
1256
+ } else if (p1c.cls == float_class_snan) {
1257
+ t16 = h1c;
1258
+ } else if (p2r.cls == float_class_snan) {
1259
+ t16 = h2r;
1260
+ } else {
1261
+ t16 = h2c;
1262
+ }
1263
} else {
1263
- t16 = h2c;
1264
+ if (p1r.cls == float_class_qnan) {
1265
+ t16 = h1r;
1266
+ } else if (p1c.cls == float_class_qnan) {
1267
+ t16 = h1c;
1268
+ } else if (p2r.cls == float_class_qnan) {
1269
+ t16 = h2r;
1270
+ } else {
1271
+ t16 = h2c;
1272
+ }
1273
}
1274
t32 = float16_to_float32(t16, true, s_f16);
1275
} else {
1267
- float64 e1r = float16_to_float64(h1r, true, s_f16);
1268
- float64 e1c = float16_to_float64(h1c, true, s_f16);
1269
- float64 e2r = float16_to_float64(h2r, true, s_f16);
1270
- float64 e2c = float16_to_float64(h2c, true, s_f16);
1271
- float64 t64;
1272
-
1276
/*
1277
* The ARM pseudocode function FPDot performs both multiplies
1275
- * and the add with a single rounding operation. Emulate this
1276
- * by performing the first multiply in round-to-odd, then doing
1277
- * the second multiply as fused multiply-add, and rounding to
1278
- * float32 all in one step.
1278
+ * and the add with a single rounding operation.
1279
*/
1280
- t64 = float64_mul(e1r, e2r, s_odd);
1281
- t64 = float64r32_muladd(e1c, e2c, t64, 0, s_std);
1282
-
1283
- /* This conversion is exact, because we've already rounded. */
1284
- t32 = float64_to_float32(t64, s_std);
1280
+ FloatParts64 tmp = parts64_mul(&p1r, &p2r, s_f16);
1281
+ tmp = parts64_muladd(&p1c, &p2c, &tmp, 0, s_f16);
1282
+ t32 = float32_round_pack_canonical(&tmp, s_f16);
1283
}
1284
1285
/* The final accumulation step is not fused. */
@@ -1293,9 +1291,6 @@ static void do_fmopa_w_h(void *vza, void *vzn, void *vzm, uint16_t *pn,
1291
uint32_t negx, bool ah_neg)
1292
{
1293
intptr_t row, col, oprsz = simd_maxsz(desc);
1296
- float_status fpst_odd = env->vfp.fp_status[FPST_ZA];
1297
-
1298
- set_float_rounding_mode(float_round_to_odd, &fpst_odd);
1294
1295
for (row = 0; row < oprsz; ) {
1296
uint16_t prow = pn[H2(row >> 4)];
@@ -1319,8 +1314,7 @@ static void do_fmopa_w_h(void *vza, void *vzn, void *vzm, uint16_t *pn,
1314
m = f16mop_adj_pair(m, pcol, 0);
1315
*a = f16_dotadd(*a, n, m,
1316
&env->vfp.fp_status[FPST_ZA_F16],
1322
- &env->vfp.fp_status[FPST_ZA],
1323
- &fpst_odd);
1317
+ &env->vfp.fp_status[FPST_ZA]);
1318
}
1319
col += 4;
1320
pcol >>= 4;
@@ -1357,15 +1351,12 @@ void HELPER(sme2_fdot_h)(void *vd, void *vn, void *vm, void *va,
1351
bool za = extract32(desc, SIMD_DATA_SHIFT, 1);
1352
float_status *fpst_std = &env->vfp.fp_status[za ? FPST_ZA : FPST_A64];
1353
float_status *fpst_f16 = &env->vfp.fp_status[za ? FPST_ZA_F16 : FPST_A64_F16];
1360
- float_status fpst_odd = *fpst_std;
1354
float32 *d = vd, *a = va;
1355
uint32_t *n = vn, *m = vm;
1356
1364
- set_float_rounding_mode(float_round_to_odd, &fpst_odd);
1365
-
1357
for (i = 0; i < oprsz / sizeof(float32); ++i) {
1358
d[H4(i)] = f16_dotadd(a[H4(i)], n[H4(i)], m[H4(i)],
1368
- fpst_f16, fpst_std, &fpst_odd);
1359
+ fpst_f16, fpst_std);
1360
}
1361
}
1362
@@ -1379,17 +1370,14 @@ void HELPER(sme2_fdot_idx_h)(void *vd, void *vn, void *vm, void *va,
1370
bool za = extract32(desc, SIMD_DATA_SHIFT + 2, 1);
1371
float_status *fpst_std = &env->vfp.fp_status[za ? FPST_ZA : FPST_A64];
1372
float_status *fpst_f16 = &env->vfp.fp_status[za ? FPST_ZA_F16 : FPST_A64_F16];
1382
- float_status fpst_odd = *fpst_std;
1373
float32 *d = vd, *a = va;
1374
uint32_t *n = vn, *m = (uint32_t *)vm + H4(idx);
1375
1386
- set_float_rounding_mode(float_round_to_odd, &fpst_odd);
1387
-
1376
for (i = 0; i < elements; i += eltspersegment) {
1377
uint32_t mm = m[i];
1378
for (j = 0; j < eltspersegment; ++j) {
1379
d[H4(i + j)] = f16_dotadd(a[H4(i + j)], n[H4(i + j)], mm,
1392
- fpst_f16, fpst_std, &fpst_odd);
1380
+ fpst_f16, fpst_std);
1381
}
1382
}
1383
}
@@ -1402,24 +1390,19 @@ void HELPER(sme2_fvdot_idx_h)(void *vd, void *vn, void *vm, void *va,
1390
intptr_t eltspersegment = MIN(4, elements);
1391
int idx = extract32(desc, SIMD_DATA_SHIFT, 2);
1392
int sel = extract32(desc, SIMD_DATA_SHIFT + 2, 1);
1405
- float_status fpst_odd, *fpst_std, *fpst_f16;
1393
float32 *d = vd, *a = va;
1394
uint16_t *n0 = vn;
1395
uint16_t *n1 = vn + sizeof(ARMVectorReg);
1396
uint32_t *m = (uint32_t *)vm + H4(idx);
1397
1411
- fpst_std = &env->vfp.fp_status[FPST_ZA];
1412
- fpst_f16 = &env->vfp.fp_status[FPST_ZA_F16];
1413
- fpst_odd = *fpst_std;
1414
- set_float_rounding_mode(float_round_to_odd, &fpst_odd);
1415
-
1398
for (i = 0; i < elements; i += eltspersegment) {
1399
uint32_t mm = m[i];
1400
for (j = 0; j < eltspersegment; ++j) {
1401
uint32_t nn = (n0[H2(2 * (i + j) + sel)])
1402
| (n1[H2(2 * (i + j) + sel)] << 16);
1403
d[i + H4(j)] = f16_dotadd(a[i + H4(j)], nn, mm,
1422
- fpst_f16, fpst_std, &fpst_odd);
1404
+ &env->vfp.fp_status[FPST_ZA_F16],
1405
+ &env->vfp.fp_status[FPST_ZA]);
1406
}
1407
}
1408
}