@samitouri / QOSamiQemu / commits / 091b7cdeae

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 }