master
c 3,184 lines 115 KB
Raw
1 /*
2 * ARM SME Operations
3 *
4 * Copyright (c) 2022 Linaro, Ltd.
5 *
6 * This library is free software; you can redistribute it and/or
7 * modify it under the terms of the GNU Lesser General Public
8 * License as published by the Free Software Foundation; either
9 * version 2.1 of the License, or (at your option) any later version.
10 *
11 * This library is distributed in the hope that it will be useful,
12 * but WITHOUT ANY WARRANTY; without even the implied warranty of
13 * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU
14 * Lesser General Public License for more details.
15 *
16 * You should have received a copy of the GNU Lesser General Public
17 * License along with this library; if not, see <http://www.gnu.org/licenses/>.
18 */
19
20 #include "qemu/osdep.h"
21 #include "cpu.h"
22 #include "internals.h"
23 #include "tcg/tcg-gvec-desc.h"
24 #include "helper.h"
25 #include "helper-sme.h"
26 #include "accel/tcg/cpu-ldst.h"
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
34 #define HELPER_H "tcg/helper-sme-defs.h"
35 #include "exec/helper-info.c.inc"
36
37 void helper_set_svcr(CPUARMState *env, uint32_t val, uint32_t mask)
38 {
39 aarch64_set_svcr(env, val, mask);
40 }
41
42 void helper_sme_zero(CPUARMState *env, uint32_t imm, uint32_t svl)
43 {
44 uint32_t i;
45
46 /*
47 * Special case clearing the entire ZArray.
48 * This falls into the CONSTRAINED UNPREDICTABLE zeroing of any
49 * parts of the ZA storage outside of SVL.
50 */
51 if (imm == 0xff) {
52 memset(env->za_state.za, 0, sizeof(env->za_state.za));
53 return;
54 }
55
56 /*
57 * Recall that ZAnH.D[m] is spread across ZA[n+8*m],
58 * so each row is discontiguous within ZA[].
59 */
60 for (i = 0; i < svl; i++) {
61 if (imm & (1 << (i % 8))) {
62 memset(&env->za_state.za[i], 0, svl);
63 }
64 }
65 }
66
67 /*
68 * Move Zreg vector to ZArray column.
69 */
70 #define DO_MOVA_C(NAME, TYPE, H) \
71 void HELPER(NAME)(void *za, void *vn, void *vg, uint32_t desc) \
72 { \
73 int i, oprsz = simd_oprsz(desc); \
74 for (i = 0; i < oprsz; ) { \
75 uint16_t pg = *(uint16_t *)(vg + H1_2(i >> 3)); \
76 do { \
77 if (pg & 1) { \
78 *(TYPE *)(za + tile_vslice_offset(i)) = *(TYPE *)(vn + H(i)); \
79 } \
80 i += sizeof(TYPE); \
81 pg >>= sizeof(TYPE); \
82 } while (i & 15); \
83 } \
84 }
85
86 DO_MOVA_C(sme_mova_cz_b, uint8_t, H1)
87 DO_MOVA_C(sme_mova_cz_h, uint16_t, H1_2)
88 DO_MOVA_C(sme_mova_cz_s, uint32_t, H1_4)
89
90 void HELPER(sme_mova_cz_d)(void *za, void *vn, void *vg, uint32_t desc)
91 {
92 int i, oprsz = simd_oprsz(desc) / 8;
93 uint8_t *pg = vg;
94 uint64_t *n = vn;
95 uint64_t *a = za;
96
97 for (i = 0; i < oprsz; i++) {
98 if (pg[H1(i)] & 1) {
99 a[tile_vslice_index(i)] = n[i];
100 }
101 }
102 }
103
104 void HELPER(sme_mova_cz_q)(void *za, void *vn, void *vg, uint32_t desc)
105 {
106 int i, oprsz = simd_oprsz(desc) / 16;
107 uint16_t *pg = vg;
108 Int128 *n = vn;
109 Int128 *a = za;
110
111 /*
112 * Int128 is used here simply to copy 16 bytes, and to simplify
113 * the address arithmetic.
114 */
115 for (i = 0; i < oprsz; i++) {
116 if (pg[H2(i)] & 1) {
117 a[tile_vslice_index(i)] = n[i];
118 }
119 }
120 }
121
122 #undef DO_MOVA_C
123
124 /*
125 * Move ZArray column to Zreg vector.
126 */
127 #define DO_MOVA_Z(NAME, TYPE, H) \
128 void HELPER(NAME)(void *vd, void *za, void *vg, uint32_t desc) \
129 { \
130 int i, oprsz = simd_oprsz(desc); \
131 for (i = 0; i < oprsz; ) { \
132 uint16_t pg = *(uint16_t *)(vg + H1_2(i >> 3)); \
133 do { \
134 if (pg & 1) { \
135 *(TYPE *)(vd + H(i)) = *(TYPE *)(za + tile_vslice_offset(i)); \
136 } \
137 i += sizeof(TYPE); \
138 pg >>= sizeof(TYPE); \
139 } while (i & 15); \
140 } \
141 }
142
143 DO_MOVA_Z(sme_mova_zc_b, uint8_t, H1)
144 DO_MOVA_Z(sme_mova_zc_h, uint16_t, H1_2)
145 DO_MOVA_Z(sme_mova_zc_s, uint32_t, H1_4)
146
147 void HELPER(sme_mova_zc_d)(void *vd, void *za, void *vg, uint32_t desc)
148 {
149 int i, oprsz = simd_oprsz(desc) / 8;
150 uint8_t *pg = vg;
151 uint64_t *d = vd;
152 uint64_t *a = za;
153
154 for (i = 0; i < oprsz; i++) {
155 if (pg[H1(i)] & 1) {
156 d[i] = a[tile_vslice_index(i)];
157 }
158 }
159 }
160
161 void HELPER(sme_mova_zc_q)(void *vd, void *za, void *vg, uint32_t desc)
162 {
163 int i, oprsz = simd_oprsz(desc) / 16;
164 uint16_t *pg = vg;
165 Int128 *d = vd;
166 Int128 *a = za;
167
168 /*
169 * Int128 is used here simply to copy 16 bytes, and to simplify
170 * the address arithmetic.
171 */
172 for (i = 0; i < oprsz; i++, za += sizeof(ARMVectorReg)) {
173 if (pg[H2(i)] & 1) {
174 d[i] = a[tile_vslice_index(i)];
175 }
176 }
177 }
178
179 #undef DO_MOVA_Z
180
181 void HELPER(sme2_mova_zc_b)(void *vdst, void *vsrc, uint32_t desc)
182 {
183 const uint8_t *src = vsrc;
184 uint8_t *dst = vdst;
185 size_t i, n = simd_oprsz(desc);
186
187 for (i = 0; i < n; ++i) {
188 dst[i] = src[tile_vslice_index(i)];
189 }
190 }
191
192 void HELPER(sme2_mova_zc_h)(void *vdst, void *vsrc, uint32_t desc)
193 {
194 const uint16_t *src = vsrc;
195 uint16_t *dst = vdst;
196 size_t i, n = simd_oprsz(desc) / 2;
197
198 for (i = 0; i < n; ++i) {
199 dst[i] = src[tile_vslice_index(i)];
200 }
201 }
202
203 void HELPER(sme2_mova_zc_s)(void *vdst, void *vsrc, uint32_t desc)
204 {
205 const uint32_t *src = vsrc;
206 uint32_t *dst = vdst;
207 size_t i, n = simd_oprsz(desc) / 4;
208
209 for (i = 0; i < n; ++i) {
210 dst[i] = src[tile_vslice_index(i)];
211 }
212 }
213
214 void HELPER(sme2_mova_zc_d)(void *vdst, void *vsrc, uint32_t desc)
215 {
216 const uint64_t *src = vsrc;
217 uint64_t *dst = vdst;
218 size_t i, n = simd_oprsz(desc) / 8;
219
220 for (i = 0; i < n; ++i) {
221 dst[i] = src[tile_vslice_index(i)];
222 }
223 }
224
225 void HELPER(sme2p1_movaz_zc_b)(void *vdst, void *vsrc, uint32_t desc)
226 {
227 uint8_t *src = vsrc;
228 uint8_t *dst = vdst;
229 size_t i, n = simd_oprsz(desc);
230
231 for (i = 0; i < n; ++i) {
232 dst[i] = src[tile_vslice_index(i)];
233 src[tile_vslice_index(i)] = 0;
234 }
235 }
236
237 void HELPER(sme2p1_movaz_zc_h)(void *vdst, void *vsrc, uint32_t desc)
238 {
239 uint16_t *src = vsrc;
240 uint16_t *dst = vdst;
241 size_t i, n = simd_oprsz(desc) / 2;
242
243 for (i = 0; i < n; ++i) {
244 dst[i] = src[tile_vslice_index(i)];
245 src[tile_vslice_index(i)] = 0;
246 }
247 }
248
249 void HELPER(sme2p1_movaz_zc_s)(void *vdst, void *vsrc, uint32_t desc)
250 {
251 uint32_t *src = vsrc;
252 uint32_t *dst = vdst;
253 size_t i, n = simd_oprsz(desc) / 4;
254
255 for (i = 0; i < n; ++i) {
256 dst[i] = src[tile_vslice_index(i)];
257 src[tile_vslice_index(i)] = 0;
258 }
259 }
260
261 void HELPER(sme2p1_movaz_zc_d)(void *vdst, void *vsrc, uint32_t desc)
262 {
263 uint64_t *src = vsrc;
264 uint64_t *dst = vdst;
265 size_t i, n = simd_oprsz(desc) / 8;
266
267 for (i = 0; i < n; ++i) {
268 dst[i] = src[tile_vslice_index(i)];
269 src[tile_vslice_index(i)] = 0;
270 }
271 }
272
273 void HELPER(sme2p1_movaz_zc_q)(void *vdst, void *vsrc, uint32_t desc)
274 {
275 Int128 *src = vsrc;
276 Int128 *dst = vdst;
277 size_t i, n = simd_oprsz(desc) / 16;
278
279 for (i = 0; i < n; ++i) {
280 dst[i] = src[tile_vslice_index(i)];
281 memset(&src[tile_vslice_index(i)], 0, 16);
282 }
283 }
284
285 /*
286 * Clear elements in a tile slice comprising len bytes.
287 */
288
289 typedef void ClearFn(void *ptr, size_t off, size_t len);
290
291 static void clear_horizontal(void *ptr, size_t off, size_t len)
292 {
293 memset(ptr + off, 0, len);
294 }
295
296 static void clear_vertical_b(void *vptr, size_t off, size_t len)
297 {
298 for (size_t i = 0; i < len; ++i) {
299 *(uint8_t *)(vptr + tile_vslice_offset(i + off)) = 0;
300 }
301 }
302
303 static void clear_vertical_h(void *vptr, size_t off, size_t len)
304 {
305 for (size_t i = 0; i < len; i += 2) {
306 *(uint16_t *)(vptr + tile_vslice_offset(i + off)) = 0;
307 }
308 }
309
310 static void clear_vertical_s(void *vptr, size_t off, size_t len)
311 {
312 for (size_t i = 0; i < len; i += 4) {
313 *(uint32_t *)(vptr + tile_vslice_offset(i + off)) = 0;
314 }
315 }
316
317 static void clear_vertical_d(void *vptr, size_t off, size_t len)
318 {
319 for (size_t i = 0; i < len; i += 8) {
320 *(uint64_t *)(vptr + tile_vslice_offset(i + off)) = 0;
321 }
322 }
323
324 static void clear_vertical_q(void *vptr, size_t off, size_t len)
325 {
326 for (size_t i = 0; i < len; i += 16) {
327 memset(vptr + tile_vslice_offset(i + off), 0, 16);
328 }
329 }
330
331 /*
332 * Copy elements from an array into a tile slice comprising len bytes.
333 */
334
335 typedef void CopyFn(void *dst, const void *src, size_t len);
336
337 static void copy_horizontal(void *dst, const void *src, size_t len)
338 {
339 memcpy(dst, src, len);
340 }
341
342 static void copy_vertical_b(void *vdst, const void *vsrc, size_t len)
343 {
344 const uint8_t *src = vsrc;
345 uint8_t *dst = vdst;
346 size_t i;
347
348 for (i = 0; i < len; ++i) {
349 dst[tile_vslice_index(i)] = src[i];
350 }
351 }
352
353 static void copy_vertical_h(void *vdst, const void *vsrc, size_t len)
354 {
355 const uint16_t *src = vsrc;
356 uint16_t *dst = vdst;
357 size_t i;
358
359 for (i = 0; i < len / 2; ++i) {
360 dst[tile_vslice_index(i)] = src[i];
361 }
362 }
363
364 static void copy_vertical_s(void *vdst, const void *vsrc, size_t len)
365 {
366 const uint32_t *src = vsrc;
367 uint32_t *dst = vdst;
368 size_t i;
369
370 for (i = 0; i < len / 4; ++i) {
371 dst[tile_vslice_index(i)] = src[i];
372 }
373 }
374
375 static void copy_vertical_d(void *vdst, const void *vsrc, size_t len)
376 {
377 const uint64_t *src = vsrc;
378 uint64_t *dst = vdst;
379 size_t i;
380
381 for (i = 0; i < len / 8; ++i) {
382 dst[tile_vslice_index(i)] = src[i];
383 }
384 }
385
386 static void copy_vertical_q(void *vdst, const void *vsrc, size_t len)
387 {
388 for (size_t i = 0; i < len; i += 16) {
389 memcpy(vdst + tile_vslice_offset(i), vsrc + i, 16);
390 }
391 }
392
393 void HELPER(sme2_mova_cz_b)(void *vdst, void *vsrc, uint32_t desc)
394 {
395 copy_vertical_b(vdst, vsrc, simd_oprsz(desc));
396 }
397
398 void HELPER(sme2_mova_cz_h)(void *vdst, void *vsrc, uint32_t desc)
399 {
400 copy_vertical_h(vdst, vsrc, simd_oprsz(desc));
401 }
402
403 void HELPER(sme2_mova_cz_s)(void *vdst, void *vsrc, uint32_t desc)
404 {
405 copy_vertical_s(vdst, vsrc, simd_oprsz(desc));
406 }
407
408 void HELPER(sme2_mova_cz_d)(void *vdst, void *vsrc, uint32_t desc)
409 {
410 copy_vertical_d(vdst, vsrc, simd_oprsz(desc));
411 }
412
413 /*
414 * Host and TLB primitives for vertical tile slice addressing.
415 */
416
417 #define DO_LD(NAME, TYPE, HOST, TLB) \
418 static inline void sme_##NAME##_v_host(void *za, intptr_t off, void *host) \
419 { \
420 TYPE val = HOST(host); \
421 *(TYPE *)(za + tile_vslice_offset(off)) = val; \
422 } \
423 static inline void sme_##NAME##_v_tlb(CPUARMState *env, void *za, \
424 intptr_t off, target_ulong addr, uintptr_t ra) \
425 { \
426 TYPE val = TLB(env, useronly_clean_ptr(addr), ra); \
427 *(TYPE *)(za + tile_vslice_offset(off)) = val; \
428 }
429
430 #define DO_ST(NAME, TYPE, HOST, TLB) \
431 static inline void sme_##NAME##_v_host(void *za, intptr_t off, void *host) \
432 { \
433 TYPE val = *(TYPE *)(za + tile_vslice_offset(off)); \
434 HOST(host, val); \
435 } \
436 static inline void sme_##NAME##_v_tlb(CPUARMState *env, void *za, \
437 intptr_t off, target_ulong addr, uintptr_t ra) \
438 { \
439 TYPE val = *(TYPE *)(za + tile_vslice_offset(off)); \
440 TLB(env, useronly_clean_ptr(addr), val, ra); \
441 }
442
443 #define DO_LDQ(HNAME, VNAME) \
444 static inline void VNAME##_v_host(void *za, intptr_t off, void *host) \
445 { \
446 HNAME##_host(za, tile_vslice_offset(off), host); \
447 } \
448 static inline void VNAME##_v_tlb(CPUARMState *env, void *za, intptr_t off, \
449 target_ulong addr, uintptr_t ra) \
450 { \
451 HNAME##_tlb(env, za, tile_vslice_offset(off), addr, ra); \
452 }
453
454 #define DO_STQ(HNAME, VNAME) \
455 static inline void VNAME##_v_host(void *za, intptr_t off, void *host) \
456 { \
457 HNAME##_host(za, tile_vslice_offset(off), host); \
458 } \
459 static inline void VNAME##_v_tlb(CPUARMState *env, void *za, intptr_t off, \
460 target_ulong addr, uintptr_t ra) \
461 { \
462 HNAME##_tlb(env, za, tile_vslice_offset(off), addr, ra); \
463 }
464
465 DO_LD(ld1b, uint8_t, ldub_p, cpu_ldub_data_ra)
466 DO_LD(ld1h_be, uint16_t, lduw_be_p, cpu_lduw_be_data_ra)
467 DO_LD(ld1h_le, uint16_t, lduw_le_p, cpu_lduw_le_data_ra)
468 DO_LD(ld1s_be, uint32_t, ldl_be_p, cpu_ldl_be_data_ra)
469 DO_LD(ld1s_le, uint32_t, ldl_le_p, cpu_ldl_le_data_ra)
470 DO_LD(ld1d_be, uint64_t, ldq_be_p, cpu_ldq_be_data_ra)
471 DO_LD(ld1d_le, uint64_t, ldq_le_p, cpu_ldq_le_data_ra)
472
473 DO_LDQ(sve_ld1qq_be, sme_ld1q_be)
474 DO_LDQ(sve_ld1qq_le, sme_ld1q_le)
475
476 DO_ST(st1b, uint8_t, stb_p, cpu_stb_data_ra)
477 DO_ST(st1h_be, uint16_t, stw_be_p, cpu_stw_be_data_ra)
478 DO_ST(st1h_le, uint16_t, stw_le_p, cpu_stw_le_data_ra)
479 DO_ST(st1s_be, uint32_t, stl_be_p, cpu_stl_be_data_ra)
480 DO_ST(st1s_le, uint32_t, stl_le_p, cpu_stl_le_data_ra)
481 DO_ST(st1d_be, uint64_t, stq_be_p, cpu_stq_be_data_ra)
482 DO_ST(st1d_le, uint64_t, stq_le_p, cpu_stq_le_data_ra)
483
484 DO_STQ(sve_st1qq_be, sme_st1q_be)
485 DO_STQ(sve_st1qq_le, sme_st1q_le)
486
487 #undef DO_LD
488 #undef DO_ST
489 #undef DO_LDQ
490 #undef DO_STQ
491
492 /*
493 * Common helper for all contiguous predicated loads.
494 */
495
496 static inline QEMU_ALWAYS_INLINE
497 void sme_ld1(CPUARMState *env, void *za, uint64_t *vg,
498 const target_ulong addr, uint32_t desc, const uintptr_t ra,
499 const int esz, uint32_t mtedesc, bool vertical,
500 sve_ldst1_host_fn *host_fn,
501 sve_ldst1_tlb_fn *tlb_fn,
502 ClearFn *clr_fn,
503 CopyFn *cpy_fn)
504 {
505 const intptr_t reg_max = simd_oprsz(desc);
506 const intptr_t esize = 1 << esz;
507 intptr_t reg_off, reg_last;
508 SVEContLdSt info;
509 void *host;
510 int flags;
511
512 /* Find the active elements. */
513 if (!sve_cont_ldst_elements(&info, addr, vg, reg_max, esz, esize)) {
514 /* The entire predicate was false; no load occurs. */
515 clr_fn(za, 0, reg_max);
516 return;
517 }
518
519 /* Probe the page(s). Exit with exception for any invalid page. */
520 sve_cont_ldst_pages(&info, FAULT_ALL, env, addr, MMU_DATA_LOAD, ra);
521
522 /* Handle watchpoints for all active elements. */
523 sve_cont_ldst_watchpoints(&info, env, vg, addr, esize, esize,
524 BP_MEM_READ, ra);
525
526 /*
527 * Handle mte checks for all active elements.
528 * Since TBI must be set for MTE, !mtedesc => !mte_active.
529 */
530 if (mtedesc) {
531 sve_cont_ldst_mte_check(&info, env, vg, addr, esize, esize,
532 mtedesc, ra);
533 }
534
535 flags = info.page[0].flags | info.page[1].flags;
536 if (unlikely(flags != 0)) {
537 #ifdef CONFIG_USER_ONLY
538 g_assert_not_reached();
539 #else
540 /*
541 * At least one page includes MMIO.
542 * Any bus operation can fail with cpu_transaction_failed,
543 * which for ARM will raise SyncExternal. Perform the load
544 * into scratch memory to preserve register state until the end.
545 */
546 ARMVectorReg scratch = { };
547
548 reg_off = info.reg_off_first[0];
549 reg_last = info.reg_off_last[1];
550 if (reg_last < 0) {
551 reg_last = info.reg_off_split;
552 if (reg_last < 0) {
553 reg_last = info.reg_off_last[0];
554 }
555 }
556
557 do {
558 uint64_t pg = vg[reg_off >> 6];
559 do {
560 if ((pg >> (reg_off & 63)) & 1) {
561 tlb_fn(env, &scratch, reg_off, addr + reg_off, ra);
562 }
563 reg_off += esize;
564 } while (reg_off & 63);
565 } while (reg_off <= reg_last);
566
567 cpy_fn(za, &scratch, reg_max);
568 return;
569 #endif
570 }
571
572 /* The entire operation is in RAM, on valid pages. */
573
574 reg_off = info.reg_off_first[0];
575 reg_last = info.reg_off_last[0];
576 host = info.page[0].host;
577
578 if (!vertical) {
579 memset(za, 0, reg_max);
580 } else if (reg_off) {
581 clr_fn(za, 0, reg_off);
582 }
583
584 set_helper_retaddr(ra);
585
586 while (reg_off <= reg_last) {
587 uint64_t pg = vg[reg_off >> 6];
588 do {
589 if ((pg >> (reg_off & 63)) & 1) {
590 host_fn(za, reg_off, host + reg_off);
591 } else if (vertical) {
592 clr_fn(za, reg_off, esize);
593 }
594 reg_off += esize;
595 } while (reg_off <= reg_last && (reg_off & 63));
596 }
597
598 clear_helper_retaddr();
599
600 /*
601 * Use the slow path to manage the cross-page misalignment.
602 * But we know this is RAM and cannot trap.
603 */
604 reg_off = info.reg_off_split;
605 if (unlikely(reg_off >= 0)) {
606 tlb_fn(env, za, reg_off, addr + reg_off, ra);
607 }
608
609 reg_off = info.reg_off_first[1];
610 if (unlikely(reg_off >= 0)) {
611 reg_last = info.reg_off_last[1];
612 host = info.page[1].host;
613
614 set_helper_retaddr(ra);
615
616 do {
617 uint64_t pg = vg[reg_off >> 6];
618 do {
619 if ((pg >> (reg_off & 63)) & 1) {
620 host_fn(za, reg_off, host + reg_off);
621 } else if (vertical) {
622 clr_fn(za, reg_off, esize);
623 }
624 reg_off += esize;
625 } while (reg_off & 63);
626 } while (reg_off <= reg_last);
627
628 clear_helper_retaddr();
629 }
630 }
631
632 static inline QEMU_ALWAYS_INLINE
633 void sme_ld1_mte(CPUARMState *env, void *za, uint64_t *vg,
634 target_ulong addr, uint64_t desc, uintptr_t ra,
635 const int esz, bool vertical,
636 sve_ldst1_host_fn *host_fn,
637 sve_ldst1_tlb_fn *tlb_fn,
638 ClearFn *clr_fn,
639 CopyFn *cpy_fn)
640 {
641 uint32_t mtedesc = desc >> 32;
642 int bit55 = extract64(addr, 55, 1);
643
644 /* Perform gross MTE suppression early. */
645 if (!tbi_or_mtx_check(mtedesc, bit55) ||
646 tcma_check(mtedesc, bit55, allocation_tag_from_addr(addr))) {
647 mtedesc = 0;
648 }
649
650 sme_ld1(env, za, vg, addr, desc, ra, esz, mtedesc, vertical,
651 host_fn, tlb_fn, clr_fn, cpy_fn);
652 }
653
654 #define DO_LD(L, END, ESZ) \
655 void HELPER(sme_ld1##L##END##_h)(CPUARMState *env, void *za, void *vg, \
656 target_ulong addr, uint64_t desc) \
657 { \
658 sme_ld1(env, za, vg, addr, desc, GETPC(), ESZ, 0, false, \
659 sve_ld1##L##L##END##_host, sve_ld1##L##L##END##_tlb, \
660 clear_horizontal, copy_horizontal); \
661 } \
662 void HELPER(sme_ld1##L##END##_v)(CPUARMState *env, void *za, void *vg, \
663 target_ulong addr, uint64_t desc) \
664 { \
665 sme_ld1(env, za, vg, addr, desc, GETPC(), ESZ, 0, true, \
666 sme_ld1##L##END##_v_host, sme_ld1##L##END##_v_tlb, \
667 clear_vertical_##L, copy_vertical_##L); \
668 } \
669 void HELPER(sme_ld1##L##END##_h_mte)(CPUARMState *env, void *za, void *vg, \
670 target_ulong addr, uint64_t desc) \
671 { \
672 sme_ld1_mte(env, za, vg, addr, desc, GETPC(), ESZ, false, \
673 sve_ld1##L##L##END##_host, sve_ld1##L##L##END##_tlb, \
674 clear_horizontal, copy_horizontal); \
675 } \
676 void HELPER(sme_ld1##L##END##_v_mte)(CPUARMState *env, void *za, void *vg, \
677 target_ulong addr, uint64_t desc) \
678 { \
679 sme_ld1_mte(env, za, vg, addr, desc, GETPC(), ESZ, true, \
680 sme_ld1##L##END##_v_host, sme_ld1##L##END##_v_tlb, \
681 clear_vertical_##L, copy_vertical_##L); \
682 }
683
684 DO_LD(b, , MO_8)
685 DO_LD(h, _be, MO_16)
686 DO_LD(h, _le, MO_16)
687 DO_LD(s, _be, MO_32)
688 DO_LD(s, _le, MO_32)
689 DO_LD(d, _be, MO_64)
690 DO_LD(d, _le, MO_64)
691 DO_LD(q, _be, MO_128)
692 DO_LD(q, _le, MO_128)
693
694 #undef DO_LD
695
696 /*
697 * Common helper for all contiguous predicated stores.
698 */
699
700 static inline QEMU_ALWAYS_INLINE
701 void sme_st1(CPUARMState *env, void *za, uint64_t *vg,
702 const target_ulong addr, uint32_t desc, const uintptr_t ra,
703 const int esz, uint32_t mtedesc, bool vertical,
704 sve_ldst1_host_fn *host_fn,
705 sve_ldst1_tlb_fn *tlb_fn)
706 {
707 const intptr_t reg_max = simd_oprsz(desc);
708 const intptr_t esize = 1 << esz;
709 intptr_t reg_off, reg_last;
710 SVEContLdSt info;
711 void *host;
712 int flags;
713
714 /* Find the active elements. */
715 if (!sve_cont_ldst_elements(&info, addr, vg, reg_max, esz, esize)) {
716 /* The entire predicate was false; no store occurs. */
717 return;
718 }
719
720 /* Probe the page(s). Exit with exception for any invalid page. */
721 sve_cont_ldst_pages(&info, FAULT_ALL, env, addr, MMU_DATA_STORE, ra);
722
723 /* Handle watchpoints for all active elements. */
724 sve_cont_ldst_watchpoints(&info, env, vg, addr, esize, esize,
725 BP_MEM_WRITE, ra);
726
727 /*
728 * Handle mte checks for all active elements.
729 * Since TBI must be set for MTE, !mtedesc => !mte_active.
730 */
731 if (mtedesc) {
732 sve_cont_ldst_mte_check(&info, env, vg, addr, esize, esize,
733 mtedesc, ra);
734 }
735
736 flags = info.page[0].flags | info.page[1].flags;
737 if (unlikely(flags != 0)) {
738 #ifdef CONFIG_USER_ONLY
739 g_assert_not_reached();
740 #else
741 /*
742 * At least one page includes MMIO.
743 * Any bus operation can fail with cpu_transaction_failed,
744 * which for ARM will raise SyncExternal. We cannot avoid
745 * this fault and will leave with the store incomplete.
746 */
747 reg_off = info.reg_off_first[0];
748 reg_last = info.reg_off_last[1];
749 if (reg_last < 0) {
750 reg_last = info.reg_off_split;
751 if (reg_last < 0) {
752 reg_last = info.reg_off_last[0];
753 }
754 }
755
756 do {
757 uint64_t pg = vg[reg_off >> 6];
758 do {
759 if ((pg >> (reg_off & 63)) & 1) {
760 tlb_fn(env, za, reg_off, addr + reg_off, ra);
761 }
762 reg_off += esize;
763 } while (reg_off & 63);
764 } while (reg_off <= reg_last);
765 return;
766 #endif
767 }
768
769 reg_off = info.reg_off_first[0];
770 reg_last = info.reg_off_last[0];
771 host = info.page[0].host;
772
773 set_helper_retaddr(ra);
774
775 while (reg_off <= reg_last) {
776 uint64_t pg = vg[reg_off >> 6];
777 do {
778 if ((pg >> (reg_off & 63)) & 1) {
779 host_fn(za, reg_off, host + reg_off);
780 }
781 reg_off += 1 << esz;
782 } while (reg_off <= reg_last && (reg_off & 63));
783 }
784
785 clear_helper_retaddr();
786
787 /*
788 * Use the slow path to manage the cross-page misalignment.
789 * But we know this is RAM and cannot trap.
790 */
791 reg_off = info.reg_off_split;
792 if (unlikely(reg_off >= 0)) {
793 tlb_fn(env, za, reg_off, addr + reg_off, ra);
794 }
795
796 reg_off = info.reg_off_first[1];
797 if (unlikely(reg_off >= 0)) {
798 reg_last = info.reg_off_last[1];
799 host = info.page[1].host;
800
801 set_helper_retaddr(ra);
802
803 do {
804 uint64_t pg = vg[reg_off >> 6];
805 do {
806 if ((pg >> (reg_off & 63)) & 1) {
807 host_fn(za, reg_off, host + reg_off);
808 }
809 reg_off += 1 << esz;
810 } while (reg_off & 63);
811 } while (reg_off <= reg_last);
812
813 clear_helper_retaddr();
814 }
815 }
816
817 static inline QEMU_ALWAYS_INLINE
818 void sme_st1_mte(CPUARMState *env, void *za, uint64_t *vg, target_ulong addr,
819 uint64_t desc, uintptr_t ra, int esz, bool vertical,
820 sve_ldst1_host_fn *host_fn,
821 sve_ldst1_tlb_fn *tlb_fn)
822 {
823 uint32_t mtedesc = desc >> 32;
824 int bit55 = extract64(addr, 55, 1);
825
826 /* Perform gross MTE suppression early. */
827 if (!tbi_or_mtx_check(mtedesc, bit55) ||
828 tcma_check(mtedesc, bit55, allocation_tag_from_addr(addr))) {
829 mtedesc = 0;
830 }
831
832 sme_st1(env, za, vg, addr, desc, ra, esz, mtedesc,
833 vertical, host_fn, tlb_fn);
834 }
835
836 #define DO_ST(L, END, ESZ) \
837 void HELPER(sme_st1##L##END##_h)(CPUARMState *env, void *za, void *vg, \
838 target_ulong addr, uint64_t desc) \
839 { \
840 sme_st1(env, za, vg, addr, desc, GETPC(), ESZ, 0, false, \
841 sve_st1##L##L##END##_host, sve_st1##L##L##END##_tlb); \
842 } \
843 void HELPER(sme_st1##L##END##_v)(CPUARMState *env, void *za, void *vg, \
844 target_ulong addr, uint64_t desc) \
845 { \
846 sme_st1(env, za, vg, addr, desc, GETPC(), ESZ, 0, true, \
847 sme_st1##L##END##_v_host, sme_st1##L##END##_v_tlb); \
848 } \
849 void HELPER(sme_st1##L##END##_h_mte)(CPUARMState *env, void *za, void *vg, \
850 target_ulong addr, uint64_t desc) \
851 { \
852 sme_st1_mte(env, za, vg, addr, desc, GETPC(), ESZ, false, \
853 sve_st1##L##L##END##_host, sve_st1##L##L##END##_tlb); \
854 } \
855 void HELPER(sme_st1##L##END##_v_mte)(CPUARMState *env, void *za, void *vg, \
856 target_ulong addr, uint64_t desc) \
857 { \
858 sme_st1_mte(env, za, vg, addr, desc, GETPC(), ESZ, true, \
859 sme_st1##L##END##_v_host, sme_st1##L##END##_v_tlb); \
860 }
861
862 DO_ST(b, , MO_8)
863 DO_ST(h, _be, MO_16)
864 DO_ST(h, _le, MO_16)
865 DO_ST(s, _be, MO_32)
866 DO_ST(s, _le, MO_32)
867 DO_ST(d, _be, MO_64)
868 DO_ST(d, _le, MO_64)
869 DO_ST(q, _be, MO_128)
870 DO_ST(q, _le, MO_128)
871
872 #undef DO_ST
873
874 void HELPER(sme_addha_s)(void *vzda, void *vzn, void *vpn,
875 void *vpm, uint32_t desc)
876 {
877 intptr_t row, col, oprsz = simd_oprsz(desc) / 4;
878 uint64_t *pn = vpn, *pm = vpm;
879 uint32_t *zda = vzda, *zn = vzn;
880
881 for (row = 0; row < oprsz; ) {
882 uint64_t pa = pn[row >> 4];
883 do {
884 if (pa & 1) {
885 for (col = 0; col < oprsz; ) {
886 uint64_t pb = pm[col >> 4];
887 do {
888 if (pb & 1) {
889 zda[tile_vslice_index(row) + H4(col)] += zn[H4(col)];
890 }
891 pb >>= 4;
892 } while (++col & 15);
893 }
894 }
895 pa >>= 4;
896 } while (++row & 15);
897 }
898 }
899
900 void HELPER(sme_addha_d)(void *vzda, void *vzn, void *vpn,
901 void *vpm, uint32_t desc)
902 {
903 intptr_t row, col, oprsz = simd_oprsz(desc) / 8;
904 uint8_t *pn = vpn, *pm = vpm;
905 uint64_t *zda = vzda, *zn = vzn;
906
907 for (row = 0; row < oprsz; ++row) {
908 if (pn[H1(row)] & 1) {
909 for (col = 0; col < oprsz; ++col) {
910 if (pm[H1(col)] & 1) {
911 zda[tile_vslice_index(row) + col] += zn[col];
912 }
913 }
914 }
915 }
916 }
917
918 void HELPER(sme_addva_s)(void *vzda, void *vzn, void *vpn,
919 void *vpm, uint32_t desc)
920 {
921 intptr_t row, col, oprsz = simd_oprsz(desc) / 4;
922 uint64_t *pn = vpn, *pm = vpm;
923 uint32_t *zda = vzda, *zn = vzn;
924
925 for (row = 0; row < oprsz; ) {
926 uint64_t pa = pn[row >> 4];
927 do {
928 if (pa & 1) {
929 uint32_t zn_row = zn[H4(row)];
930 for (col = 0; col < oprsz; ) {
931 uint64_t pb = pm[col >> 4];
932 do {
933 if (pb & 1) {
934 zda[tile_vslice_index(row) + H4(col)] += zn_row;
935 }
936 pb >>= 4;
937 } while (++col & 15);
938 }
939 }
940 pa >>= 4;
941 } while (++row & 15);
942 }
943 }
944
945 void HELPER(sme_addva_d)(void *vzda, void *vzn, void *vpn,
946 void *vpm, uint32_t desc)
947 {
948 intptr_t row, col, oprsz = simd_oprsz(desc) / 8;
949 uint8_t *pn = vpn, *pm = vpm;
950 uint64_t *zda = vzda, *zn = vzn;
951
952 for (row = 0; row < oprsz; ++row) {
953 if (pn[H1(row)] & 1) {
954 uint64_t zn_row = zn[row];
955 for (col = 0; col < oprsz; ++col) {
956 if (pm[H1(col)] & 1) {
957 zda[tile_vslice_index(row) + col] += zn_row;
958 }
959 }
960 }
961 }
962 }
963
964 static void do_fmopa_h(void *vza, void *vzn, void *vzm, uint16_t *pn,
965 uint16_t *pm, float_status *fpst, uint32_t desc,
966 uint16_t negx, int negf)
967 {
968 intptr_t row, col, oprsz = simd_maxsz(desc);
969
970 for (row = 0; row < oprsz; ) {
971 uint16_t pa = pn[H2(row >> 4)];
972 do {
973 if (pa & 1) {
974 void *vza_row = vza + tile_vslice_offset(row);
975 uint16_t n = *(uint32_t *)(vzn + H1_2(row)) ^ negx;
976
977 for (col = 0; col < oprsz; ) {
978 uint16_t pb = pm[H2(col >> 4)];
979 do {
980 if (pb & 1) {
981 uint16_t *a = vza_row + H1_2(col);
982 uint16_t *m = vzm + H1_2(col);
983 *a = float16_muladd(n, *m, *a, negf, fpst);
984 }
985 col += 2;
986 pb >>= 2;
987 } while (col & 15);
988 }
989 }
990 row += 2;
991 pa >>= 2;
992 } while (row & 15);
993 }
994 }
995
996 void HELPER(sme_fmopa_h)(void *vza, void *vzn, void *vzm, void *vpn,
997 void *vpm, float_status *fpst, uint32_t desc)
998 {
999 do_fmopa_h(vza, vzn, vzm, vpn, vpm, fpst, desc, 0, 0);
1000 }
1001
1002 void HELPER(sme_fmops_h)(void *vza, void *vzn, void *vzm, void *vpn,
1003 void *vpm, float_status *fpst, uint32_t desc)
1004 {
1005 do_fmopa_h(vza, vzn, vzm, vpn, vpm, fpst, desc, 1u << 15, 0);
1006 }
1007
1008 void HELPER(sme_ah_fmops_h)(void *vza, void *vzn, void *vzm, void *vpn,
1009 void *vpm, float_status *fpst, uint32_t desc)
1010 {
1011 do_fmopa_h(vza, vzn, vzm, vpn, vpm, fpst, desc, 0,
1012 float_muladd_negate_product);
1013 }
1014
1015 static void do_fmopa_s(void *vza, void *vzn, void *vzm, uint16_t *pn,
1016 uint16_t *pm, float_status *fpst, uint32_t desc,
1017 uint32_t negx, int negf)
1018 {
1019 intptr_t row, col, oprsz = simd_maxsz(desc);
1020
1021 for (row = 0; row < oprsz; ) {
1022 uint16_t pa = pn[H2(row >> 4)];
1023 do {
1024 if (pa & 1) {
1025 void *vza_row = vza + tile_vslice_offset(row);
1026 uint32_t n = *(uint32_t *)(vzn + H1_4(row)) ^ negx;
1027
1028 for (col = 0; col < oprsz; ) {
1029 uint16_t pb = pm[H2(col >> 4)];
1030 do {
1031 if (pb & 1) {
1032 uint32_t *a = vza_row + H1_4(col);
1033 uint32_t *m = vzm + H1_4(col);
1034 *a = float32_muladd(n, *m, *a, negf, fpst);
1035 }
1036 col += 4;
1037 pb >>= 4;
1038 } while (col & 15);
1039 }
1040 }
1041 row += 4;
1042 pa >>= 4;
1043 } while (row & 15);
1044 }
1045 }
1046
1047 void HELPER(sme_fmopa_s)(void *vza, void *vzn, void *vzm, void *vpn,
1048 void *vpm, float_status *fpst, uint32_t desc)
1049 {
1050 do_fmopa_s(vza, vzn, vzm, vpn, vpm, fpst, desc, 0, 0);
1051 }
1052
1053 void HELPER(sme_fmops_s)(void *vza, void *vzn, void *vzm, void *vpn,
1054 void *vpm, float_status *fpst, uint32_t desc)
1055 {
1056 do_fmopa_s(vza, vzn, vzm, vpn, vpm, fpst, desc, 1u << 31, 0);
1057 }
1058
1059 void HELPER(sme_ah_fmops_s)(void *vza, void *vzn, void *vzm, void *vpn,
1060 void *vpm, float_status *fpst, uint32_t desc)
1061 {
1062 do_fmopa_s(vza, vzn, vzm, vpn, vpm, fpst, desc, 0,
1063 float_muladd_negate_product);
1064 }
1065
1066 static void do_fmopa_d(uint64_t *za, uint64_t *zn, uint64_t *zm, uint8_t *pn,
1067 uint8_t *pm, float_status *fpst, uint32_t desc,
1068 uint64_t negx, int negf)
1069 {
1070 intptr_t row, col, oprsz = simd_oprsz(desc) / 8;
1071
1072 for (row = 0; row < oprsz; ++row) {
1073 if (pn[H1(row)] & 1) {
1074 uint64_t *za_row = &za[tile_vslice_index(row)];
1075 uint64_t n = zn[row] ^ negx;
1076
1077 for (col = 0; col < oprsz; ++col) {
1078 if (pm[H1(col)] & 1) {
1079 uint64_t *a = &za_row[col];
1080 *a = float64_muladd(n, zm[col], *a, negf, fpst);
1081 }
1082 }
1083 }
1084 }
1085 }
1086
1087 void HELPER(sme_fmopa_d)(void *vza, void *vzn, void *vzm, void *vpn,
1088 void *vpm, float_status *fpst, uint32_t desc)
1089 {
1090 do_fmopa_d(vza, vzn, vzm, vpn, vpm, fpst, desc, 0, 0);
1091 }
1092
1093 void HELPER(sme_fmops_d)(void *vza, void *vzn, void *vzm, void *vpn,
1094 void *vpm, float_status *fpst, uint32_t desc)
1095 {
1096 do_fmopa_d(vza, vzn, vzm, vpn, vpm, fpst, desc, 1ull << 63, 0);
1097 }
1098
1099 void HELPER(sme_ah_fmops_d)(void *vza, void *vzn, void *vzm, void *vpn,
1100 void *vpm, float_status *fpst, uint32_t desc)
1101 {
1102 do_fmopa_d(vza, vzn, vzm, vpn, vpm, fpst, desc, 0,
1103 float_muladd_negate_product);
1104 }
1105
1106 static void do_bfmopa(void *vza, void *vzn, void *vzm, uint16_t *pn,
1107 uint16_t *pm, float_status *fpst, uint32_t desc,
1108 uint16_t negx, int negf)
1109 {
1110 intptr_t row, col, oprsz = simd_maxsz(desc);
1111
1112 for (row = 0; row < oprsz; ) {
1113 uint16_t pa = pn[H2(row >> 4)];
1114 do {
1115 if (pa & 1) {
1116 void *vza_row = vza + tile_vslice_offset(row);
1117 uint16_t n = *(uint32_t *)(vzn + H1_2(row)) ^ negx;
1118
1119 for (col = 0; col < oprsz; ) {
1120 uint16_t pb = pm[H2(col >> 4)];
1121 do {
1122 if (pb & 1) {
1123 uint16_t *a = vza_row + H1_2(col);
1124 uint16_t *m = vzm + H1_2(col);
1125 *a = bfloat16_muladd(n, *m, *a, negf, fpst);
1126 }
1127 col += 2;
1128 pb >>= 2;
1129 } while (col & 15);
1130 }
1131 }
1132 row += 2;
1133 pa >>= 2;
1134 } while (row & 15);
1135 }
1136 }
1137
1138 void HELPER(sme_bfmopa)(void *vza, void *vzn, void *vzm, void *vpn,
1139 void *vpm, float_status *fpst, uint32_t desc)
1140 {
1141 do_bfmopa(vza, vzn, vzm, vpn, vpm, fpst, desc, 0, 0);
1142 }
1143
1144 void HELPER(sme_bfmops)(void *vza, void *vzn, void *vzm, void *vpn,
1145 void *vpm, float_status *fpst, uint32_t desc)
1146 {
1147 do_bfmopa(vza, vzn, vzm, vpn, vpm, fpst, desc, 1u << 15, 0);
1148 }
1149
1150 void HELPER(sme_ah_bfmops)(void *vza, void *vzn, void *vzm, void *vpn,
1151 void *vpm, float_status *fpst, uint32_t desc)
1152 {
1153 do_bfmopa(vza, vzn, vzm, vpn, vpm, fpst, desc, 0,
1154 float_muladd_negate_product);
1155 }
1156
1157 /*
1158 * Alter PAIR as needed for controlling predicates being false,
1159 * and for NEG on an enabled row element.
1160 */
1161 static inline uint32_t f16mop_adj_pair(uint32_t pair, uint32_t pg, uint32_t neg)
1162 {
1163 /*
1164 * The pseudocode uses a conditional negate after the conditional zero.
1165 * It is simpler here to unconditionally negate before conditional zero.
1166 */
1167 pair ^= neg;
1168 if (!(pg & 1)) {
1169 pair &= 0xffff0000u;
1170 }
1171 if (!(pg & 4)) {
1172 pair &= 0x0000ffffu;
1173 }
1174 return pair;
1175 }
1176
1177 static inline uint32_t f16mop_ah_neg_adj_pair(uint32_t pair, uint32_t pg)
1178 {
1179 uint32_t l = pg & 1 ? float16_ah_chs(pair) : 0;
1180 uint32_t h = pg & 4 ? float16_ah_chs(pair >> 16) : 0;
1181 return l | (h << 16);
1182 }
1183
1184 static inline uint32_t bf16mop_ah_neg_adj_pair(uint32_t pair, uint32_t pg)
1185 {
1186 uint32_t l = pg & 1 ? bfloat16_ah_chs(pair) : 0;
1187 uint32_t h = pg & 4 ? bfloat16_ah_chs(pair >> 16) : 0;
1188 return l | (h << 16);
1189 }
1190
1191 static float32 f16_dotadd(float32 sum, uint32_t e1, uint32_t e2,
1192 float_status *s_f16, float_status *s_std)
1193 {
1194 /*
1195 * We need two different float_status for different parts of this
1196 * operation:
1197 * - the input conversion of the float16 values must use the
1198 * f16-specific float_status, so that the FPCR.FZ16 control is applied
1199 * - operations on float32 including the final accumulation must use
1200 * the normal float_status, so that FPCR.FZ is applied
1201 */
1202 float16 h1r = e1 & 0xffff;
1203 float16 h1c = e1 >> 16;
1204 float16 h2r = e2 & 0xffff;
1205 float16 h2c = e2 >> 16;
1206 float32 t32;
1207
1208 FloatParts64 p1r = float16_unpack_canonical(h1r, s_f16);
1209 FloatParts64 p1c = float16_unpack_canonical(h1c, s_f16);
1210 FloatParts64 p2r = float16_unpack_canonical(h2r, s_f16);
1211 FloatParts64 p2c = float16_unpack_canonical(h2c, s_f16);
1212
1213 int all_mask = (float_cmask(p1r.cls) | float_cmask(p1c.cls) |
1214 float_cmask(p2r.cls) | float_cmask(p2c.cls));
1215
1216 /* C.f. FPProcessNaNs4 */
1217 if (unlikely(all_mask & float_cmask_anynan)) {
1218 float16 t16;
1219
1220 if (unlikely(all_mask & float_cmask_snan)) {
1221 if (p1r.cls == float_class_snan) {
1222 t16 = h1r;
1223 } else if (p1c.cls == float_class_snan) {
1224 t16 = h1c;
1225 } else if (p2r.cls == float_class_snan) {
1226 t16 = h2r;
1227 } else {
1228 t16 = h2c;
1229 }
1230 } else {
1231 if (p1r.cls == float_class_qnan) {
1232 t16 = h1r;
1233 } else if (p1c.cls == float_class_qnan) {
1234 t16 = h1c;
1235 } else if (p2r.cls == float_class_qnan) {
1236 t16 = h2r;
1237 } else {
1238 t16 = h2c;
1239 }
1240 }
1241 t32 = float16_to_float32(t16, true, s_f16);
1242 } else {
1243 /*
1244 * The ARM pseudocode function FPDot performs both multiplies
1245 * and the add with a single rounding operation.
1246 */
1247 FloatParts64 tmp = parts64_mul(&p1r, &p2r, s_f16);
1248 tmp = parts64_muladd(&p1c, &p2c, &tmp, 0, s_f16);
1249 t32 = float32_round_pack_canonical(&tmp, s_f16);
1250 }
1251
1252 /* The final accumulation step is not fused. */
1253 return float32_add(sum, t32, s_std);
1254 }
1255
1256 static void do_fmopa_w_h(void *vza, void *vzn, void *vzm, uint16_t *pn,
1257 uint16_t *pm, CPUARMState *env, uint32_t desc,
1258 uint32_t negx, bool ah_neg)
1259 {
1260 intptr_t row, col, oprsz = simd_maxsz(desc);
1261
1262 for (row = 0; row < oprsz; ) {
1263 uint16_t prow = pn[H2(row >> 4)];
1264 do {
1265 void *vza_row = vza + tile_vslice_offset(row);
1266 uint32_t n = *(uint32_t *)(vzn + H1_4(row));
1267
1268 if (ah_neg) {
1269 n = f16mop_ah_neg_adj_pair(n, prow);
1270 } else {
1271 n = f16mop_adj_pair(n, prow, negx);
1272 }
1273
1274 for (col = 0; col < oprsz; ) {
1275 uint16_t pcol = pm[H2(col >> 4)];
1276 do {
1277 if (prow & pcol & 0b0101) {
1278 uint32_t *a = vza_row + H1_4(col);
1279 uint32_t m = *(uint32_t *)(vzm + H1_4(col));
1280
1281 m = f16mop_adj_pair(m, pcol, 0);
1282 *a = f16_dotadd(*a, n, m,
1283 &env->vfp.fp_status[FPST_ZA_F16],
1284 &env->vfp.fp_status[FPST_ZA]);
1285 }
1286 col += 4;
1287 pcol >>= 4;
1288 } while (col & 15);
1289 }
1290 row += 4;
1291 prow >>= 4;
1292 } while (row & 15);
1293 }
1294 }
1295
1296 void HELPER(sme_fmopa_w_h)(void *vza, void *vzn, void *vzm, void *vpn,
1297 void *vpm, CPUARMState *env, uint32_t desc)
1298 {
1299 do_fmopa_w_h(vza, vzn, vzm, vpn, vpm, env, desc, 0, false);
1300 }
1301
1302 void HELPER(sme_fmops_w_h)(void *vza, void *vzn, void *vzm, void *vpn,
1303 void *vpm, CPUARMState *env, uint32_t desc)
1304 {
1305 do_fmopa_w_h(vza, vzn, vzm, vpn, vpm, env, desc, 0x80008000u, false);
1306 }
1307
1308 void HELPER(sme_ah_fmops_w_h)(void *vza, void *vzn, void *vzm, void *vpn,
1309 void *vpm, CPUARMState *env, uint32_t desc)
1310 {
1311 do_fmopa_w_h(vza, vzn, vzm, vpn, vpm, env, desc, 0, true);
1312 }
1313
1314 void HELPER(sme2_fdot_h)(void *vd, void *vn, void *vm, void *va,
1315 CPUARMState *env, uint32_t desc)
1316 {
1317 intptr_t i, oprsz = simd_maxsz(desc);
1318 bool za = extract32(desc, SIMD_DATA_SHIFT, 1);
1319 float_status *fpst_std = &env->vfp.fp_status[za ? FPST_ZA : FPST_A64];
1320 float_status *fpst_f16 = &env->vfp.fp_status[za ? FPST_ZA_F16 : FPST_A64_F16];
1321 float32 *d = vd, *a = va;
1322 uint32_t *n = vn, *m = vm;
1323
1324 for (i = 0; i < oprsz / sizeof(float32); ++i) {
1325 d[H4(i)] = f16_dotadd(a[H4(i)], n[H4(i)], m[H4(i)],
1326 fpst_f16, fpst_std);
1327 }
1328 }
1329
1330 void HELPER(sme2_fdot_idx_h)(void *vd, void *vn, void *vm, void *va,
1331 CPUARMState *env, uint32_t desc)
1332 {
1333 intptr_t i, j, oprsz = simd_maxsz(desc);
1334 intptr_t elements = oprsz / sizeof(float32);
1335 intptr_t eltspersegment = MIN(4, elements);
1336 int idx = extract32(desc, SIMD_DATA_SHIFT, 2);
1337 bool za = extract32(desc, SIMD_DATA_SHIFT + 2, 1);
1338 float_status *fpst_std = &env->vfp.fp_status[za ? FPST_ZA : FPST_A64];
1339 float_status *fpst_f16 = &env->vfp.fp_status[za ? FPST_ZA_F16 : FPST_A64_F16];
1340 float32 *d = vd, *a = va;
1341 uint32_t *n = vn, *m = (uint32_t *)vm + H4(idx);
1342
1343 for (i = 0; i < elements; i += eltspersegment) {
1344 uint32_t mm = m[i];
1345 for (j = 0; j < eltspersegment; ++j) {
1346 d[H4(i + j)] = f16_dotadd(a[H4(i + j)], n[H4(i + j)], mm,
1347 fpst_f16, fpst_std);
1348 }
1349 }
1350 }
1351
1352 void HELPER(sme2_fvdot_idx_h)(void *vd, void *vn, void *vm, void *va,
1353 CPUARMState *env, uint32_t desc)
1354 {
1355 intptr_t i, j, oprsz = simd_maxsz(desc);
1356 intptr_t elements = oprsz / sizeof(float32);
1357 intptr_t eltspersegment = MIN(4, elements);
1358 int idx = extract32(desc, SIMD_DATA_SHIFT, 2);
1359 int sel = extract32(desc, SIMD_DATA_SHIFT + 2, 1);
1360 float32 *d = vd, *a = va;
1361 uint16_t *n0 = vn;
1362 uint16_t *n1 = vn + sizeof(ARMVectorReg);
1363 uint32_t *m = (uint32_t *)vm + H4(idx);
1364
1365 for (i = 0; i < elements; i += eltspersegment) {
1366 uint32_t mm = m[i];
1367 for (j = 0; j < eltspersegment; ++j) {
1368 uint32_t nn = (n0[H2(2 * (i + j) + sel)])
1369 | (n1[H2(2 * (i + j) + sel)] << 16);
1370 d[i + H4(j)] = f16_dotadd(a[i + H4(j)], nn, mm,
1371 &env->vfp.fp_status[FPST_ZA_F16],
1372 &env->vfp.fp_status[FPST_ZA]);
1373 }
1374 }
1375 }
1376
1377 static void do_bfmopa_w(void *vza, void *vzn, void *vzm,
1378 uint16_t *pn, uint16_t *pm, CPUARMState *env,
1379 uint32_t desc, uint32_t negx, bool ah_neg)
1380 {
1381 intptr_t row, col, oprsz = simd_maxsz(desc);
1382 float_status fpst;
1383
1384 if (is_ebf(env, &fpst)) {
1385 for (row = 0; row < oprsz; ) {
1386 uint16_t prow = pn[H2(row >> 4)];
1387 do {
1388 void *vza_row = vza + tile_vslice_offset(row);
1389 uint32_t n = *(uint32_t *)(vzn + H1_4(row));
1390
1391 if (ah_neg) {
1392 n = bf16mop_ah_neg_adj_pair(n, prow);
1393 } else {
1394 n = f16mop_adj_pair(n, prow, negx);
1395 }
1396
1397 for (col = 0; col < oprsz; ) {
1398 uint16_t pcol = pm[H2(col >> 4)];
1399 do {
1400 if (prow & pcol & 0b0101) {
1401 uint32_t *a = vza_row + H1_4(col);
1402 uint32_t m = *(uint32_t *)(vzm + H1_4(col));
1403
1404 m = f16mop_adj_pair(m, pcol, 0);
1405 *a = bfdotadd_ebf(*a, n, m, &fpst);
1406 }
1407 col += 4;
1408 pcol >>= 4;
1409 } while (col & 15);
1410 }
1411 row += 4;
1412 prow >>= 4;
1413 } while (row & 15);
1414 }
1415 } else {
1416 for (row = 0; row < oprsz; ) {
1417 uint16_t prow = pn[H2(row >> 4)];
1418 do {
1419 void *vza_row = vza + tile_vslice_offset(row);
1420 uint32_t n = *(uint32_t *)(vzn + H1_4(row));
1421
1422 if (ah_neg) {
1423 n = bf16mop_ah_neg_adj_pair(n, prow);
1424 } else {
1425 n = f16mop_adj_pair(n, prow, negx);
1426 }
1427
1428 for (col = 0; col < oprsz; ) {
1429 uint16_t pcol = pm[H2(col >> 4)];
1430 do {
1431 if (prow & pcol & 0b0101) {
1432 uint32_t *a = vza_row + H1_4(col);
1433 uint32_t m = *(uint32_t *)(vzm + H1_4(col));
1434
1435 m = f16mop_adj_pair(m, pcol, 0);
1436 *a = bfdotadd(*a, n, m, &fpst);
1437 }
1438 col += 4;
1439 pcol >>= 4;
1440 } while (col & 15);
1441 }
1442 row += 4;
1443 prow >>= 4;
1444 } while (row & 15);
1445 }
1446 }
1447 }
1448
1449 void HELPER(sme_bfmopa_w)(void *vza, void *vzn, void *vzm, void *vpn,
1450 void *vpm, CPUARMState *env, uint32_t desc)
1451 {
1452 do_bfmopa_w(vza, vzn, vzm, vpn, vpm, env, desc, 0, false);
1453 }
1454
1455 void HELPER(sme_bfmops_w)(void *vza, void *vzn, void *vzm, void *vpn,
1456 void *vpm, CPUARMState *env, uint32_t desc)
1457 {
1458 do_bfmopa_w(vza, vzn, vzm, vpn, vpm, env, desc, 0x80008000u, false);
1459 }
1460
1461 void HELPER(sme_ah_bfmops_w)(void *vza, void *vzn, void *vzm, void *vpn,
1462 void *vpm, CPUARMState *env, uint32_t desc)
1463 {
1464 do_bfmopa_w(vza, vzn, vzm, vpn, vpm, env, desc, 0, true);
1465 }
1466
1467 typedef uint32_t IMOPFn32(uint32_t, uint32_t, uint32_t, uint8_t, bool);
1468 static inline void do_imopa_s(uint32_t *za, uint32_t *zn, uint32_t *zm,
1469 uint8_t *pn, uint8_t *pm,
1470 uint32_t desc, IMOPFn32 *fn)
1471 {
1472 intptr_t row, col, oprsz = simd_oprsz(desc) / 4;
1473 bool neg = simd_data(desc);
1474
1475 for (row = 0; row < oprsz; ++row) {
1476 uint8_t pa = (pn[H1(row >> 1)] >> ((row & 1) * 4)) & 0xf;
1477 uint32_t *za_row = &za[tile_vslice_index(row)];
1478 uint32_t n = zn[H4(row)];
1479
1480 for (col = 0; col < oprsz; ++col) {
1481 uint8_t pb = pm[H1(col >> 1)] >> ((col & 1) * 4);
1482 uint32_t *a = &za_row[H4(col)];
1483
1484 *a = fn(n, zm[H4(col)], *a, pa & pb, neg);
1485 }
1486 }
1487 }
1488
1489 typedef uint64_t IMOPFn64(uint64_t, uint64_t, uint64_t, uint8_t, bool);
1490 static inline void do_imopa_d(uint64_t *za, uint64_t *zn, uint64_t *zm,
1491 uint8_t *pn, uint8_t *pm,
1492 uint32_t desc, IMOPFn64 *fn)
1493 {
1494 intptr_t row, col, oprsz = simd_oprsz(desc) / 8;
1495 bool neg = simd_data(desc);
1496
1497 for (row = 0; row < oprsz; ++row) {
1498 uint8_t pa = pn[H1(row)];
1499 uint64_t *za_row = &za[tile_vslice_index(row)];
1500 uint64_t n = zn[row];
1501
1502 for (col = 0; col < oprsz; ++col) {
1503 uint8_t pb = pm[H1(col)];
1504 uint64_t *a = &za_row[col];
1505
1506 *a = fn(n, zm[col], *a, pa & pb, neg);
1507 }
1508 }
1509 }
1510
1511 #define DEF_IMOP_8x4_32(NAME, NTYPE, MTYPE) \
1512 static uint32_t NAME(uint32_t n, uint32_t m, uint32_t a, uint8_t p, bool neg) \
1513 { \
1514 uint32_t sum = 0; \
1515 /* Apply P to N as a mask, making the inactive elements 0. */ \
1516 n &= expand_pred_b(p); \
1517 sum += (NTYPE)(n >> 0) * (MTYPE)(m >> 0); \
1518 sum += (NTYPE)(n >> 8) * (MTYPE)(m >> 8); \
1519 sum += (NTYPE)(n >> 16) * (MTYPE)(m >> 16); \
1520 sum += (NTYPE)(n >> 24) * (MTYPE)(m >> 24); \
1521 return neg ? a - sum : a + sum; \
1522 }
1523
1524 #define DEF_IMOP_16x4_64(NAME, NTYPE, MTYPE) \
1525 static uint64_t NAME(uint64_t n, uint64_t m, uint64_t a, uint8_t p, bool neg) \
1526 { \
1527 uint64_t sum = 0; \
1528 /* Apply P to N as a mask, making the inactive elements 0. */ \
1529 n &= expand_pred_h(p); \
1530 sum += (int64_t)(NTYPE)(n >> 0) * (MTYPE)(m >> 0); \
1531 sum += (int64_t)(NTYPE)(n >> 16) * (MTYPE)(m >> 16); \
1532 sum += (int64_t)(NTYPE)(n >> 32) * (MTYPE)(m >> 32); \
1533 sum += (int64_t)(NTYPE)(n >> 48) * (MTYPE)(m >> 48); \
1534 return neg ? a - sum : a + sum; \
1535 }
1536
1537 DEF_IMOP_8x4_32(smopa_s, int8_t, int8_t)
1538 DEF_IMOP_8x4_32(umopa_s, uint8_t, uint8_t)
1539 DEF_IMOP_8x4_32(sumopa_s, int8_t, uint8_t)
1540 DEF_IMOP_8x4_32(usmopa_s, uint8_t, int8_t)
1541
1542 DEF_IMOP_16x4_64(smopa_d, int16_t, int16_t)
1543 DEF_IMOP_16x4_64(umopa_d, uint16_t, uint16_t)
1544 DEF_IMOP_16x4_64(sumopa_d, int16_t, uint16_t)
1545 DEF_IMOP_16x4_64(usmopa_d, uint16_t, int16_t)
1546
1547 #define DEF_IMOPH(P, NAME, S) \
1548 void HELPER(P##_##NAME##_##S)(void *vza, void *vzn, void *vzm, \
1549 void *vpn, void *vpm, uint32_t desc) \
1550 { do_imopa_##S(vza, vzn, vzm, vpn, vpm, desc, NAME##_##S); }
1551
1552 DEF_IMOPH(sme, smopa, s)
1553 DEF_IMOPH(sme, umopa, s)
1554 DEF_IMOPH(sme, sumopa, s)
1555 DEF_IMOPH(sme, usmopa, s)
1556
1557 DEF_IMOPH(sme, smopa, d)
1558 DEF_IMOPH(sme, umopa, d)
1559 DEF_IMOPH(sme, sumopa, d)
1560 DEF_IMOPH(sme, usmopa, d)
1561
1562 static uint32_t bmopa_s(uint32_t n, uint32_t m, uint32_t a, uint8_t p, bool neg)
1563 {
1564 uint32_t sum = ctpop32(~(n ^ m));
1565 if (neg) {
1566 sum = -sum;
1567 }
1568 if (!(p & 1)) {
1569 sum = 0;
1570 }
1571 return a + sum;
1572 }
1573
1574 DEF_IMOPH(sme2, bmopa, s)
1575
1576 #define DEF_IMOP_16x2_32(NAME, NTYPE, MTYPE) \
1577 static uint32_t NAME(uint32_t n, uint32_t m, uint32_t a, uint8_t p, bool neg) \
1578 { \
1579 uint32_t sum = 0; \
1580 /* Apply P to N as a mask, making the inactive elements 0. */ \
1581 n &= expand_pred_h(p); \
1582 sum += (NTYPE)(n >> 0) * (MTYPE)(m >> 0); \
1583 sum += (NTYPE)(n >> 16) * (MTYPE)(m >> 16); \
1584 return neg ? a - sum : a + sum; \
1585 }
1586
1587 DEF_IMOP_16x2_32(smopa2_s, int16_t, int16_t)
1588 DEF_IMOP_16x2_32(umopa2_s, uint16_t, uint16_t)
1589
1590 DEF_IMOPH(sme2, smopa2, s)
1591 DEF_IMOPH(sme2, umopa2, s)
1592
1593 #define DO_VDOT_IDX(NAME, TYPED, TYPEN, TYPEM, HD, HN) \
1594 void HELPER(NAME)(void *vd, void *vn, void *vm, uint32_t desc) \
1595 { \
1596 intptr_t svl = simd_oprsz(desc); \
1597 intptr_t elements = svl / sizeof(TYPED); \
1598 intptr_t eltperseg = 16 / sizeof(TYPED); \
1599 intptr_t nreg = sizeof(TYPED) / sizeof(TYPEN); \
1600 intptr_t vstride = (svl / nreg) * sizeof(ARMVectorReg); \
1601 intptr_t zstride = sizeof(ARMVectorReg) / sizeof(TYPEN); \
1602 intptr_t idx = extract32(desc, SIMD_DATA_SHIFT, 2); \
1603 TYPEN *n = vn; \
1604 TYPEM *m = vm; \
1605 for (intptr_t r = 0; r < nreg; r++) { \
1606 TYPED *d = vd + r * vstride; \
1607 for (intptr_t seg = 0; seg < elements; seg += eltperseg) { \
1608 intptr_t s = seg + idx; \
1609 for (intptr_t e = seg; e < seg + eltperseg; e++) { \
1610 TYPED sum = d[HD(e)]; \
1611 for (intptr_t i = 0; i < nreg; i++) { \
1612 TYPED nn = n[i * zstride + HN(nreg * e + r)]; \
1613 TYPED mm = m[HN(nreg * s + i)]; \
1614 sum += nn * mm; \
1615 } \
1616 d[HD(e)] = sum; \
1617 } \
1618 } \
1619 } \
1620 }
1621
1622 DO_VDOT_IDX(sme2_svdot_idx_4b, int32_t, int8_t, int8_t, H4, H1)
1623 DO_VDOT_IDX(sme2_uvdot_idx_4b, uint32_t, uint8_t, uint8_t, H4, H1)
1624 DO_VDOT_IDX(sme2_suvdot_idx_4b, int32_t, int8_t, uint8_t, H4, H1)
1625 DO_VDOT_IDX(sme2_usvdot_idx_4b, int32_t, uint8_t, int8_t, H4, H1)
1626
1627 DO_VDOT_IDX(sme2_svdot_idx_4h, int64_t, int16_t, int16_t, H8, H2)
1628 DO_VDOT_IDX(sme2_uvdot_idx_4h, uint64_t, uint16_t, uint16_t, H8, H2)
1629
1630 DO_VDOT_IDX(sme2_svdot_idx_2h, int32_t, int16_t, int16_t, H4, H2)
1631 DO_VDOT_IDX(sme2_uvdot_idx_2h, uint32_t, uint16_t, uint16_t, H4, H2)
1632
1633 #undef DO_VDOT_IDX
1634
1635 #define DO_MLALL(NAME, TYPEW, TYPEN, TYPEM, HW, HN, OP) \
1636 void HELPER(NAME)(void *vd, void *vn, void *vm, void *va, uint32_t desc) \
1637 { \
1638 intptr_t elements = simd_oprsz(desc) / sizeof(TYPEW); \
1639 intptr_t sel = extract32(desc, SIMD_DATA_SHIFT, 2); \
1640 TYPEW *d = vd, *a = va; TYPEN *n = vn; TYPEM *m = vm; \
1641 for (intptr_t i = 0; i < elements; ++i) { \
1642 TYPEW nn = n[HN(i * 4 + sel)]; \
1643 TYPEM mm = m[HN(i * 4 + sel)]; \
1644 d[HW(i)] = a[HW(i)] OP (nn * mm); \
1645 } \
1646 }
1647
1648 DO_MLALL(sme2_smlall_s, int32_t, int8_t, int8_t, H4, H1, +)
1649 DO_MLALL(sme2_smlall_d, int64_t, int16_t, int16_t, H8, H2, +)
1650 DO_MLALL(sme2_smlsll_s, int32_t, int8_t, int8_t, H4, H1, -)
1651 DO_MLALL(sme2_smlsll_d, int64_t, int16_t, int16_t, H8, H2, -)
1652
1653 DO_MLALL(sme2_umlall_s, uint32_t, uint8_t, uint8_t, H4, H1, +)
1654 DO_MLALL(sme2_umlall_d, uint64_t, uint16_t, uint16_t, H8, H2, +)
1655 DO_MLALL(sme2_umlsll_s, uint32_t, uint8_t, uint8_t, H4, H1, -)
1656 DO_MLALL(sme2_umlsll_d, uint64_t, uint16_t, uint16_t, H8, H2, -)
1657
1658 DO_MLALL(sme2_usmlall_s, uint32_t, uint8_t, int8_t, H4, H1, +)
1659
1660 #undef DO_MLALL
1661
1662 #define DO_MLALL_IDX(NAME, TYPEW, TYPEN, TYPEM, HW, HN, OP) \
1663 void HELPER(NAME)(void *vd, void *vn, void *vm, void *va, uint32_t desc) \
1664 { \
1665 intptr_t elements = simd_oprsz(desc) / sizeof(TYPEW); \
1666 intptr_t eltspersegment = 16 / sizeof(TYPEW); \
1667 intptr_t sel = extract32(desc, SIMD_DATA_SHIFT, 2); \
1668 intptr_t idx = extract32(desc, SIMD_DATA_SHIFT + 2, 4); \
1669 TYPEW *d = vd, *a = va; TYPEN *n = vn; TYPEM *m = vm; \
1670 for (intptr_t i = 0; i < elements; i += eltspersegment) { \
1671 TYPEW mm = m[HN(i * 4 + idx)]; \
1672 for (intptr_t j = 0; j < eltspersegment; ++j) { \
1673 TYPEN nn = n[HN((i + j) * 4 + sel)]; \
1674 d[HW(i + j)] = a[HW(i + j)] OP (nn * mm); \
1675 } \
1676 } \
1677 }
1678
1679 DO_MLALL_IDX(sme2_smlall_idx_s, int32_t, int8_t, int8_t, H4, H1, +)
1680 DO_MLALL_IDX(sme2_smlall_idx_d, int64_t, int16_t, int16_t, H8, H2, +)
1681 DO_MLALL_IDX(sme2_smlsll_idx_s, int32_t, int8_t, int8_t, H4, H1, -)
1682 DO_MLALL_IDX(sme2_smlsll_idx_d, int64_t, int16_t, int16_t, H8, H2, -)
1683
1684 DO_MLALL_IDX(sme2_umlall_idx_s, uint32_t, uint8_t, uint8_t, H4, H1, +)
1685 DO_MLALL_IDX(sme2_umlall_idx_d, uint64_t, uint16_t, uint16_t, H8, H2, +)
1686 DO_MLALL_IDX(sme2_umlsll_idx_s, uint32_t, uint8_t, uint8_t, H4, H1, -)
1687 DO_MLALL_IDX(sme2_umlsll_idx_d, uint64_t, uint16_t, uint16_t, H8, H2, -)
1688
1689 DO_MLALL_IDX(sme2_usmlall_idx_s, uint32_t, uint8_t, int8_t, H4, H1, +)
1690 DO_MLALL_IDX(sme2_sumlall_idx_s, uint32_t, int8_t, uint8_t, H4, H1, +)
1691
1692 #undef DO_MLALL_IDX
1693
1694 /* Convert and compress */
1695 void HELPER(sme2_bfcvt_hs)(void *vd, void *vs, float_status *fpst, uint32_t desc)
1696 {
1697 ARMVectorReg scratch;
1698 size_t oprsz = simd_oprsz(desc);
1699 size_t i, n = oprsz / 4;
1700 float32 *s0 = vs;
1701 float32 *s1 = vs + sizeof(ARMVectorReg);
1702 bfloat16 *d = vd;
1703
1704 if (vd == s1) {
1705 s1 = memcpy(&scratch, s1, oprsz);
1706 }
1707
1708 for (i = 0; i < n; ++i) {
1709 d[H2(i)] = float32_to_bfloat16(s0[H4(i)], fpst);
1710 }
1711 for (i = 0; i < n; ++i) {
1712 d[H2(i) + n] = float32_to_bfloat16(s1[H4(i)], fpst);
1713 }
1714 }
1715
1716 void HELPER(sme2_fcvt_n)(void *vd, void *vs, float_status *fpst, uint32_t desc)
1717 {
1718 ARMVectorReg scratch;
1719 size_t oprsz = simd_oprsz(desc);
1720 size_t i, n = oprsz / 4;
1721 float32 *s0 = vs;
1722 float32 *s1 = vs + sizeof(ARMVectorReg);
1723 float16 *d = vd;
1724
1725 if (vd == s1) {
1726 s1 = memcpy(&scratch, s1, oprsz);
1727 }
1728
1729 for (i = 0; i < n; ++i) {
1730 d[H2(i)] = sve_f32_to_f16(s0[H4(i)], fpst);
1731 }
1732 for (i = 0; i < n; ++i) {
1733 d[H2(i) + n] = sve_f32_to_f16(s1[H4(i)], fpst);
1734 }
1735 }
1736
1737 #define SQCVT2(NAME, TW, TN, HW, HN, SAT) \
1738 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1739 { \
1740 ARMVectorReg scratch; \
1741 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1742 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1743 TN *d = vd; \
1744 if (vectors_overlap(vd, 1, vs, 2)) { \
1745 d = (TN *)&scratch; \
1746 } \
1747 for (size_t i = 0; i < n; ++i) { \
1748 d[HN(i)] = SAT(s0[HW(i)]); \
1749 d[HN(i + n)] = SAT(s1[HW(i)]); \
1750 } \
1751 if (d != vd) { \
1752 memcpy(vd, d, oprsz); \
1753 } \
1754 }
1755
1756 SQCVT2(sme2_sqcvt_sh, int32_t, int16_t, H4, H2, do_ssat_h)
1757 SQCVT2(sme2_uqcvt_sh, uint32_t, uint16_t, H4, H2, do_usat_h)
1758 SQCVT2(sme2_sqcvtu_sh, int32_t, uint16_t, H4, H2, do_usat_h)
1759
1760 #undef SQCVT2
1761
1762 #define SQCVT4(NAME, TW, TN, HW, HN, SAT) \
1763 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1764 { \
1765 ARMVectorReg scratch; \
1766 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1767 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1768 TW *s2 = vs + 2 * sizeof(ARMVectorReg); \
1769 TW *s3 = vs + 3 * sizeof(ARMVectorReg); \
1770 TN *d = vd; \
1771 if (vectors_overlap(vd, 1, vs, 4)) { \
1772 d = (TN *)&scratch; \
1773 } \
1774 for (size_t i = 0; i < n; ++i) { \
1775 d[HN(i)] = SAT(s0[HW(i)]); \
1776 d[HN(i + n)] = SAT(s1[HW(i)]); \
1777 d[HN(i + 2 * n)] = SAT(s2[HW(i)]); \
1778 d[HN(i + 3 * n)] = SAT(s3[HW(i)]); \
1779 } \
1780 if (d != vd) { \
1781 memcpy(vd, d, oprsz); \
1782 } \
1783 }
1784
1785 SQCVT4(sme2_sqcvt_sb, int32_t, int8_t, H4, H2, do_ssat_b)
1786 SQCVT4(sme2_uqcvt_sb, uint32_t, uint8_t, H4, H2, do_usat_b)
1787 SQCVT4(sme2_sqcvtu_sb, int32_t, uint8_t, H4, H2, do_usat_b)
1788
1789 SQCVT4(sme2_sqcvt_dh, int64_t, int16_t, H8, H2, do_ssat_h)
1790 SQCVT4(sme2_uqcvt_dh, uint64_t, uint16_t, H8, H2, do_usat_h)
1791 SQCVT4(sme2_sqcvtu_dh, int64_t, uint16_t, H8, H2, do_usat_h)
1792
1793 #undef SQCVT4
1794
1795 #define SQRSHR2(NAME, TW, TN, HW, HN, RSHR, SAT) \
1796 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1797 { \
1798 ARMVectorReg scratch; \
1799 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1800 int shift = simd_data(desc); \
1801 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1802 TN *d = vd; \
1803 if (vectors_overlap(vd, 1, vs, 2)) { \
1804 d = (TN *)&scratch; \
1805 } \
1806 for (size_t i = 0; i < n; ++i) { \
1807 d[HN(i)] = SAT(RSHR(s0[HW(i)], shift)); \
1808 d[HN(i + n)] = SAT(RSHR(s1[HW(i)], shift)); \
1809 } \
1810 if (d != vd) { \
1811 memcpy(vd, d, oprsz); \
1812 } \
1813 }
1814
1815 SQRSHR2(sme2_sqrshr_sh, int32_t, int16_t, H4, H2, do_srshr, do_ssat_h)
1816 SQRSHR2(sme2_uqrshr_sh, uint32_t, uint16_t, H4, H2, do_urshr, do_usat_h)
1817 SQRSHR2(sme2_sqrshru_sh, int32_t, uint16_t, H4, H2, do_srshr, do_usat_h)
1818
1819 #undef SQRSHR2
1820
1821 #define SQRSHR4(NAME, TW, TN, HW, HN, RSHR, SAT) \
1822 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1823 { \
1824 ARMVectorReg scratch; \
1825 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1826 int shift = simd_data(desc); \
1827 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1828 TW *s2 = vs + 2 * sizeof(ARMVectorReg); \
1829 TW *s3 = vs + 3 * sizeof(ARMVectorReg); \
1830 TN *d = vd; \
1831 if (vectors_overlap(vd, 1, vs, 4)) { \
1832 d = (TN *)&scratch; \
1833 } \
1834 for (size_t i = 0; i < n; ++i) { \
1835 d[HN(i)] = SAT(RSHR(s0[HW(i)], shift)); \
1836 d[HN(i + n)] = SAT(RSHR(s1[HW(i)], shift)); \
1837 d[HN(i + 2 * n)] = SAT(RSHR(s2[HW(i)], shift)); \
1838 d[HN(i + 3 * n)] = SAT(RSHR(s3[HW(i)], shift)); \
1839 } \
1840 if (d != vd) { \
1841 memcpy(vd, d, oprsz); \
1842 } \
1843 }
1844
1845 SQRSHR4(sme2_sqrshr_sb, int32_t, int8_t, H4, H2, do_srshr, do_ssat_b)
1846 SQRSHR4(sme2_uqrshr_sb, uint32_t, uint8_t, H4, H2, do_urshr, do_usat_b)
1847 SQRSHR4(sme2_sqrshru_sb, int32_t, uint8_t, H4, H2, do_srshr, do_usat_b)
1848
1849 SQRSHR4(sme2_sqrshr_dh, int64_t, int16_t, H8, H2, do_srshr, do_ssat_h)
1850 SQRSHR4(sme2_uqrshr_dh, uint64_t, uint16_t, H8, H2, do_urshr, do_usat_h)
1851 SQRSHR4(sme2_sqrshru_dh, int64_t, uint16_t, H8, H2, do_srshr, do_usat_h)
1852
1853 #undef SQRSHR4
1854
1855 /* Convert and interleave */
1856 void HELPER(sme2_bfcvtn)(void *vd, void *vs, float_status *fpst, uint32_t desc)
1857 {
1858 size_t i, n = simd_oprsz(desc) / 4;
1859 float32 *s0 = vs;
1860 float32 *s1 = vs + sizeof(ARMVectorReg);
1861 bfloat16 *d = vd;
1862
1863 for (i = 0; i < n; ++i) {
1864 bfloat16 d0 = float32_to_bfloat16(s0[H4(i)], fpst);
1865 bfloat16 d1 = float32_to_bfloat16(s1[H4(i)], fpst);
1866 d[H2(i * 2 + 0)] = d0;
1867 d[H2(i * 2 + 1)] = d1;
1868 }
1869 }
1870
1871 void HELPER(sme2_fcvtn)(void *vd, void *vs, float_status *fpst, uint32_t desc)
1872 {
1873 size_t i, n = simd_oprsz(desc) / 4;
1874 float32 *s0 = vs;
1875 float32 *s1 = vs + sizeof(ARMVectorReg);
1876 bfloat16 *d = vd;
1877
1878 for (i = 0; i < n; ++i) {
1879 bfloat16 d0 = sve_f32_to_f16(s0[H4(i)], fpst);
1880 bfloat16 d1 = sve_f32_to_f16(s1[H4(i)], fpst);
1881 d[H2(i * 2 + 0)] = d0;
1882 d[H2(i * 2 + 1)] = d1;
1883 }
1884 }
1885
1886 #define SQCVTN2(NAME, TW, TN, HW, HN, SAT) \
1887 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1888 { \
1889 ARMVectorReg scratch; \
1890 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1891 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1892 TN *d = vd; \
1893 if (vectors_overlap(vd, 1, vs, 2)) { \
1894 d = (TN *)&scratch; \
1895 } \
1896 for (size_t i = 0; i < n; ++i) { \
1897 d[HN(2 * i + 0)] = SAT(s0[HW(i)]); \
1898 d[HN(2 * i + 1)] = SAT(s1[HW(i)]); \
1899 } \
1900 if (d != vd) { \
1901 memcpy(vd, d, oprsz); \
1902 } \
1903 }
1904
1905 SQCVTN2(sme2_sqcvtn_sh, int32_t, int16_t, H4, H2, do_ssat_h)
1906 SQCVTN2(sme2_uqcvtn_sh, uint32_t, uint16_t, H4, H2, do_usat_h)
1907 SQCVTN2(sme2_sqcvtun_sh, int32_t, uint16_t, H4, H2, do_usat_h)
1908
1909 #undef SQCVTN2
1910
1911 #define SQCVTN4(NAME, TW, TN, HW, HN, SAT) \
1912 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1913 { \
1914 ARMVectorReg scratch; \
1915 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1916 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1917 TW *s2 = vs + 2 * sizeof(ARMVectorReg); \
1918 TW *s3 = vs + 3 * sizeof(ARMVectorReg); \
1919 TN *d = vd; \
1920 if (vectors_overlap(vd, 1, vs, 4)) { \
1921 d = (TN *)&scratch; \
1922 } \
1923 for (size_t i = 0; i < n; ++i) { \
1924 d[HN(4 * i + 0)] = SAT(s0[HW(i)]); \
1925 d[HN(4 * i + 1)] = SAT(s1[HW(i)]); \
1926 d[HN(4 * i + 2)] = SAT(s2[HW(i)]); \
1927 d[HN(4 * i + 3)] = SAT(s3[HW(i)]); \
1928 } \
1929 if (d != vd) { \
1930 memcpy(vd, d, oprsz); \
1931 } \
1932 }
1933
1934 SQCVTN4(sme2_sqcvtn_sb, int32_t, int8_t, H4, H1, do_ssat_b)
1935 SQCVTN4(sme2_uqcvtn_sb, uint32_t, uint8_t, H4, H1, do_usat_b)
1936 SQCVTN4(sme2_sqcvtun_sb, int32_t, uint8_t, H4, H1, do_usat_b)
1937
1938 SQCVTN4(sme2_sqcvtn_dh, int64_t, int16_t, H8, H2, do_ssat_h)
1939 SQCVTN4(sme2_uqcvtn_dh, uint64_t, uint16_t, H8, H2, do_usat_h)
1940 SQCVTN4(sme2_sqcvtun_dh, int64_t, uint16_t, H8, H2, do_usat_h)
1941
1942 #undef SQCVTN4
1943
1944 #define SQRSHRN2(NAME, TW, TN, HW, HN, RSHR, SAT) \
1945 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1946 { \
1947 ARMVectorReg scratch; \
1948 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1949 int shift = simd_data(desc); \
1950 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1951 TN *d = vd; \
1952 if (vectors_overlap(vd, 1, vs, 2)) { \
1953 d = (TN *)&scratch; \
1954 } \
1955 for (size_t i = 0; i < n; ++i) { \
1956 d[HN(2 * i + 0)] = SAT(RSHR(s0[HW(i)], shift)); \
1957 d[HN(2 * i + 1)] = SAT(RSHR(s1[HW(i)], shift)); \
1958 } \
1959 if (d != vd) { \
1960 memcpy(vd, d, oprsz); \
1961 } \
1962 }
1963
1964 SQRSHRN2(sme2_sqrshrn_sh, int32_t, int16_t, H4, H2, do_srshr, do_ssat_h)
1965 SQRSHRN2(sme2_uqrshrn_sh, uint32_t, uint16_t, H4, H2, do_urshr, do_usat_h)
1966 SQRSHRN2(sme2_sqrshrun_sh, int32_t, uint16_t, H4, H2, do_srshr, do_usat_h)
1967
1968 #undef SQRSHRN2
1969
1970 #define SQRSHRN4(NAME, TW, TN, HW, HN, RSHR, SAT) \
1971 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
1972 { \
1973 ARMVectorReg scratch; \
1974 size_t oprsz = simd_oprsz(desc), n = oprsz / sizeof(TW); \
1975 int shift = simd_data(desc); \
1976 TW *s0 = vs, *s1 = vs + sizeof(ARMVectorReg); \
1977 TW *s2 = vs + 2 * sizeof(ARMVectorReg); \
1978 TW *s3 = vs + 3 * sizeof(ARMVectorReg); \
1979 TN *d = vd; \
1980 if (vectors_overlap(vd, 1, vs, 4)) { \
1981 d = (TN *)&scratch; \
1982 } \
1983 for (size_t i = 0; i < n; ++i) { \
1984 d[HN(4 * i + 0)] = SAT(RSHR(s0[HW(i)], shift)); \
1985 d[HN(4 * i + 1)] = SAT(RSHR(s1[HW(i)], shift)); \
1986 d[HN(4 * i + 2)] = SAT(RSHR(s2[HW(i)], shift)); \
1987 d[HN(4 * i + 3)] = SAT(RSHR(s3[HW(i)], shift)); \
1988 } \
1989 if (d != vd) { \
1990 memcpy(vd, d, oprsz); \
1991 } \
1992 }
1993
1994 SQRSHRN4(sme2_sqrshrn_sb, int32_t, int8_t, H4, H1, do_srshr, do_ssat_b)
1995 SQRSHRN4(sme2_uqrshrn_sb, uint32_t, uint8_t, H4, H1, do_urshr, do_usat_b)
1996 SQRSHRN4(sme2_sqrshrun_sb, int32_t, uint8_t, H4, H1, do_srshr, do_usat_b)
1997
1998 SQRSHRN4(sme2_sqrshrn_dh, int64_t, int16_t, H8, H2, do_srshr, do_ssat_h)
1999 SQRSHRN4(sme2_uqrshrn_dh, uint64_t, uint16_t, H8, H2, do_urshr, do_usat_h)
2000 SQRSHRN4(sme2_sqrshrun_dh, int64_t, uint16_t, H8, H2, do_srshr, do_usat_h)
2001
2002 #undef SQRSHRN4
2003
2004 /* Expand and convert */
2005 void HELPER(sme2_fcvt_w)(void *vd, void *vs, float_status *fpst, uint32_t desc)
2006 {
2007 ARMVectorReg scratch;
2008 size_t oprsz = simd_oprsz(desc);
2009 size_t i, n = oprsz / 4;
2010 float16 *s = vs;
2011 float32 *d0 = vd;
2012 float32 *d1 = vd + sizeof(ARMVectorReg);
2013
2014 if (vectors_overlap(vd, 1, vs, 2)) {
2015 s = memcpy(&scratch, s, oprsz);
2016 }
2017
2018 for (i = 0; i < n; ++i) {
2019 d0[H4(i)] = sve_f16_to_f32(s[H2(i)], fpst);
2020 }
2021 for (i = 0; i < n; ++i) {
2022 d1[H4(i)] = sve_f16_to_f32(s[H2(n + i)], fpst);
2023 }
2024 }
2025
2026 #define UNPK(NAME, SREG, TW, TN, HW, HN) \
2027 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
2028 { \
2029 ARMVectorReg scratch[SREG]; \
2030 size_t oprsz = simd_oprsz(desc); \
2031 size_t n = oprsz / sizeof(TW); \
2032 if (vectors_overlap(vd, 2 * SREG, vs, SREG)) { \
2033 vs = memcpy(scratch, vs, sizeof(scratch)); \
2034 } \
2035 for (size_t r = 0; r < SREG; ++r) { \
2036 TN *s = vs + r * sizeof(ARMVectorReg); \
2037 for (size_t i = 0; i < 2; ++i) { \
2038 TW *d = vd + (2 * r + i) * sizeof(ARMVectorReg); \
2039 for (size_t e = 0; e < n; ++e) { \
2040 d[HW(e)] = s[HN(i * n + e)]; \
2041 } \
2042 } \
2043 } \
2044 }
2045
2046 UNPK(sme2_sunpk2_bh, 1, int16_t, int8_t, H2, H1)
2047 UNPK(sme2_sunpk2_hs, 1, int32_t, int16_t, H4, H2)
2048 UNPK(sme2_sunpk2_sd, 1, int64_t, int32_t, H8, H4)
2049
2050 UNPK(sme2_sunpk4_bh, 2, int16_t, int8_t, H2, H1)
2051 UNPK(sme2_sunpk4_hs, 2, int32_t, int16_t, H4, H2)
2052 UNPK(sme2_sunpk4_sd, 2, int64_t, int32_t, H8, H4)
2053
2054 UNPK(sme2_uunpk2_bh, 1, uint16_t, uint8_t, H2, H1)
2055 UNPK(sme2_uunpk2_hs, 1, uint32_t, uint16_t, H4, H2)
2056 UNPK(sme2_uunpk2_sd, 1, uint64_t, uint32_t, H8, H4)
2057
2058 UNPK(sme2_uunpk4_bh, 2, uint16_t, uint8_t, H2, H1)
2059 UNPK(sme2_uunpk4_hs, 2, uint32_t, uint16_t, H4, H2)
2060 UNPK(sme2_uunpk4_sd, 2, uint64_t, uint32_t, H8, H4)
2061
2062 #undef UNPK
2063
2064 /* Deinterleave and convert. */
2065 void HELPER(sme2_fcvtl)(void *vd, void *vs, float_status *fpst, uint32_t desc)
2066 {
2067 size_t i, n = simd_oprsz(desc) / 4;
2068 float16 *s = vs;
2069 float32 *d0 = vd;
2070 float32 *d1 = vd + sizeof(ARMVectorReg);
2071
2072 for (i = 0; i < n; ++i) {
2073 float32 v0 = sve_f16_to_f32(s[H2(i * 2 + 0)], fpst);
2074 float32 v1 = sve_f16_to_f32(s[H2(i * 2 + 1)], fpst);
2075 d0[H4(i)] = v0;
2076 d1[H4(i)] = v1;
2077 }
2078 }
2079
2080 void HELPER(sme2_scvtf)(void *vd, void *vs, float_status *fpst, uint32_t desc)
2081 {
2082 size_t i, n = simd_oprsz(desc) / 4;
2083 int32_t *d = vd;
2084 float32 *s = vs;
2085
2086 for (i = 0; i < n; ++i) {
2087 d[i] = int32_to_float32(s[i], fpst);
2088 }
2089 }
2090
2091 void HELPER(sme2_ucvtf)(void *vd, void *vs, float_status *fpst, uint32_t desc)
2092 {
2093 size_t i, n = simd_oprsz(desc) / 4;
2094 uint32_t *d = vd;
2095 float32 *s = vs;
2096
2097 for (i = 0; i < n; ++i) {
2098 d[i] = uint32_to_float32(s[i], fpst);
2099 }
2100 }
2101
2102 #define ZIP2(NAME, TYPE, H) \
2103 void HELPER(NAME)(void *vd, void *vn, void *vm, uint32_t desc) \
2104 { \
2105 ARMVectorReg scratch[2]; \
2106 size_t oprsz = simd_oprsz(desc); \
2107 size_t pairs = oprsz / (sizeof(TYPE) * 2); \
2108 TYPE *n = vn, *m = vm; \
2109 if (vectors_overlap(vd, 2, vn, 1)) { \
2110 n = memcpy(&scratch[0], vn, oprsz); \
2111 } \
2112 if (vectors_overlap(vd, 2, vm, 1)) { \
2113 m = memcpy(&scratch[1], vm, oprsz); \
2114 } \
2115 for (size_t r = 0; r < 2; ++r) { \
2116 TYPE *d = vd + r * sizeof(ARMVectorReg); \
2117 size_t base = r * pairs; \
2118 for (size_t p = 0; p < pairs; ++p) { \
2119 d[H(2 * p + 0)] = n[base + H(p)]; \
2120 d[H(2 * p + 1)] = m[base + H(p)]; \
2121 } \
2122 } \
2123 }
2124
2125 ZIP2(sme2_zip2_b, uint8_t, H1)
2126 ZIP2(sme2_zip2_h, uint16_t, H2)
2127 ZIP2(sme2_zip2_s, uint32_t, H4)
2128 ZIP2(sme2_zip2_d, uint64_t, )
2129 ZIP2(sme2_zip2_q, Int128, )
2130
2131 #undef ZIP2
2132
2133 #define ZIP4(NAME, TYPE, H) \
2134 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
2135 { \
2136 ARMVectorReg scratch[4]; \
2137 size_t oprsz = simd_oprsz(desc); \
2138 size_t quads = oprsz / (sizeof(TYPE) * 4); \
2139 TYPE *s0, *s1, *s2, *s3; \
2140 if (vs == vd) { \
2141 vs = memcpy(scratch, vs, sizeof(scratch)); \
2142 } \
2143 s0 = vs; \
2144 s1 = vs + sizeof(ARMVectorReg); \
2145 s2 = vs + 2 * sizeof(ARMVectorReg); \
2146 s3 = vs + 3 * sizeof(ARMVectorReg); \
2147 for (size_t r = 0; r < 4; ++r) { \
2148 TYPE *d = vd + r * sizeof(ARMVectorReg); \
2149 size_t base = r * quads; \
2150 for (size_t q = 0; q < quads; ++q) { \
2151 d[H(4 * q + 0)] = s0[base + H(q)]; \
2152 d[H(4 * q + 1)] = s1[base + H(q)]; \
2153 d[H(4 * q + 2)] = s2[base + H(q)]; \
2154 d[H(4 * q + 3)] = s3[base + H(q)]; \
2155 } \
2156 } \
2157 }
2158
2159 ZIP4(sme2_zip4_b, uint8_t, H1)
2160 ZIP4(sme2_zip4_h, uint16_t, H2)
2161 ZIP4(sme2_zip4_s, uint32_t, H4)
2162 ZIP4(sme2_zip4_d, uint64_t, )
2163 ZIP4(sme2_zip4_q, Int128, )
2164
2165 #undef ZIP4
2166
2167 #define UZP2(NAME, TYPE, H) \
2168 void HELPER(NAME)(void *vd, void *vn, void *vm, uint32_t desc) \
2169 { \
2170 ARMVectorReg scratch[2]; \
2171 size_t oprsz = simd_oprsz(desc); \
2172 size_t pairs = oprsz / (sizeof(TYPE) * 2); \
2173 TYPE *d0 = vd, *d1 = vd + sizeof(ARMVectorReg); \
2174 if (vectors_overlap(vd, 2, vn, 1)) { \
2175 vn = memcpy(&scratch[0], vn, oprsz); \
2176 } \
2177 if (vectors_overlap(vd, 2, vm, 1)) { \
2178 vm = memcpy(&scratch[1], vm, oprsz); \
2179 } \
2180 for (size_t r = 0; r < 2; ++r) { \
2181 TYPE *s = r ? vm : vn; \
2182 size_t base = r * pairs; \
2183 for (size_t p = 0; p < pairs; ++p) { \
2184 d0[base + H(p)] = s[H(2 * p + 0)]; \
2185 d1[base + H(p)] = s[H(2 * p + 1)]; \
2186 } \
2187 } \
2188 }
2189
2190 UZP2(sme2_uzp2_b, uint8_t, H1)
2191 UZP2(sme2_uzp2_h, uint16_t, H2)
2192 UZP2(sme2_uzp2_s, uint32_t, H4)
2193 UZP2(sme2_uzp2_d, uint64_t, )
2194 UZP2(sme2_uzp2_q, Int128, )
2195
2196 #undef UZP2
2197
2198 #define UZP4(NAME, TYPE, H) \
2199 void HELPER(NAME)(void *vd, void *vs, uint32_t desc) \
2200 { \
2201 ARMVectorReg scratch[4]; \
2202 size_t oprsz = simd_oprsz(desc); \
2203 size_t quads = oprsz / (sizeof(TYPE) * 4); \
2204 TYPE *d0, *d1, *d2, *d3; \
2205 if (vs == vd) { \
2206 vs = memcpy(scratch, vs, sizeof(scratch)); \
2207 } \
2208 d0 = vd; \
2209 d1 = vd + sizeof(ARMVectorReg); \
2210 d2 = vd + 2 * sizeof(ARMVectorReg); \
2211 d3 = vd + 3 * sizeof(ARMVectorReg); \
2212 for (size_t r = 0; r < 4; ++r) { \
2213 TYPE *s = vs + r * sizeof(ARMVectorReg); \
2214 size_t base = r * quads; \
2215 for (size_t q = 0; q < quads; ++q) { \
2216 d0[base + H(q)] = s[H(4 * q + 0)]; \
2217 d1[base + H(q)] = s[H(4 * q + 1)]; \
2218 d2[base + H(q)] = s[H(4 * q + 2)]; \
2219 d3[base + H(q)] = s[H(4 * q + 3)]; \
2220 } \
2221 } \
2222 }
2223
2224 UZP4(sme2_uzp4_b, uint8_t, H1)
2225 UZP4(sme2_uzp4_h, uint16_t, H2)
2226 UZP4(sme2_uzp4_s, uint32_t, H4)
2227 UZP4(sme2_uzp4_d, uint64_t, )
2228 UZP4(sme2_uzp4_q, Int128, )
2229
2230 #undef UZP4
2231
2232 #define ICLAMP(NAME, TYPE, H) \
2233 void HELPER(NAME)(void *vd, void *vn, void *vm, uint32_t desc) \
2234 { \
2235 size_t stride = sizeof(ARMVectorReg) / sizeof(TYPE); \
2236 size_t elements = simd_oprsz(desc) / sizeof(TYPE); \
2237 size_t nreg = simd_data(desc); \
2238 TYPE *d = vd, *n = vn, *m = vm; \
2239 for (size_t e = 0; e < elements; e++) { \
2240 TYPE nn = n[H(e)], mm = m[H(e)]; \
2241 for (size_t r = 0; r < nreg; r++) { \
2242 TYPE *dd = &d[r * stride + H(e)]; \
2243 *dd = MIN(MAX(*dd, nn), mm); \
2244 } \
2245 } \
2246 }
2247
2248 ICLAMP(sme2_sclamp_b, int8_t, H1)
2249 ICLAMP(sme2_sclamp_h, int16_t, H2)
2250 ICLAMP(sme2_sclamp_s, int32_t, H4)
2251 ICLAMP(sme2_sclamp_d, int64_t, H8)
2252
2253 ICLAMP(sme2_uclamp_b, uint8_t, H1)
2254 ICLAMP(sme2_uclamp_h, uint16_t, H2)
2255 ICLAMP(sme2_uclamp_s, uint32_t, H4)
2256 ICLAMP(sme2_uclamp_d, uint64_t, H8)
2257
2258 #undef ICLAMP
2259
2260 /*
2261 * Note the argument ordering to minnum and maxnum must match
2262 * the ARM pseudocode so that NaNs are propagated properly.
2263 */
2264 #define FCLAMP(NAME, TYPE, H) \
2265 void HELPER(NAME)(void *vd, void *vn, void *vm, \
2266 float_status *fpst, uint32_t desc) \
2267 { \
2268 size_t stride = sizeof(ARMVectorReg) / sizeof(TYPE); \
2269 size_t elements = simd_oprsz(desc) / sizeof(TYPE); \
2270 size_t nreg = simd_data(desc); \
2271 TYPE *d = vd, *n = vn, *m = vm; \
2272 for (size_t e = 0; e < elements; e++) { \
2273 TYPE nn = n[H(e)], mm = m[H(e)]; \
2274 for (size_t r = 0; r < nreg; r++) { \
2275 TYPE *dd = &d[r * stride + H(e)]; \
2276 *dd = TYPE##_minnum(TYPE##_maxnum(nn, *dd, fpst), mm, fpst); \
2277 } \
2278 } \
2279 }
2280
2281 FCLAMP(sme2_fclamp_h, float16, H2)
2282 FCLAMP(sme2_fclamp_s, float32, H4)
2283 FCLAMP(sme2_fclamp_d, float64, H8)
2284 FCLAMP(sme2_bfclamp, bfloat16, H2)
2285
2286 #undef FCLAMP
2287
2288 void HELPER(sme2_sel_b)(void *vd, void *vn, void *vm,
2289 uint32_t png, uint32_t desc)
2290 {
2291 int vl = simd_oprsz(desc);
2292 int nreg = simd_data(desc);
2293 int elements = vl / sizeof(uint8_t);
2294 DecodeCounter p = decode_counter(png, vl, MO_8);
2295
2296 if (p.lg2_stride == 0) {
2297 if (p.invert) {
2298 for (int r = 0; r < nreg; r++) {
2299 uint8_t *d = vd + r * sizeof(ARMVectorReg);
2300 uint8_t *n = vn + r * sizeof(ARMVectorReg);
2301 uint8_t *m = vm + r * sizeof(ARMVectorReg);
2302 int split = p.count - r * elements;
2303
2304 if (split <= 0) {
2305 memcpy(d, n, vl); /* all true */
2306 } else if (elements <= split) {
2307 memcpy(d, m, vl); /* all false */
2308 } else {
2309 for (int e = 0; e < split; e++) {
2310 d[H1(e)] = m[H1(e)];
2311 }
2312 for (int e = split; e < elements; e++) {
2313 d[H1(e)] = n[H1(e)];
2314 }
2315 }
2316 }
2317 } else {
2318 for (int r = 0; r < nreg; r++) {
2319 uint8_t *d = vd + r * sizeof(ARMVectorReg);
2320 uint8_t *n = vn + r * sizeof(ARMVectorReg);
2321 uint8_t *m = vm + r * sizeof(ARMVectorReg);
2322 int split = p.count - r * elements;
2323
2324 if (split <= 0) {
2325 memcpy(d, m, vl); /* all false */
2326 } else if (elements <= split) {
2327 memcpy(d, n, vl); /* all true */
2328 } else {
2329 for (int e = 0; e < split; e++) {
2330 d[H1(e)] = n[H1(e)];
2331 }
2332 for (int e = split; e < elements; e++) {
2333 d[H1(e)] = m[H1(e)];
2334 }
2335 }
2336 }
2337 }
2338 } else {
2339 int estride = 1 << p.lg2_stride;
2340 if (p.invert) {
2341 for (int r = 0; r < nreg; r++) {
2342 uint8_t *d = vd + r * sizeof(ARMVectorReg);
2343 uint8_t *n = vn + r * sizeof(ARMVectorReg);
2344 uint8_t *m = vm + r * sizeof(ARMVectorReg);
2345 int split = p.count - r * elements;
2346 int e = 0;
2347
2348 for (; e < MIN(split, elements); e++) {
2349 d[H1(e)] = m[H1(e)];
2350 }
2351 for (; e < elements; e += estride) {
2352 d[H1(e)] = n[H1(e)];
2353 for (int i = 1; i < estride; i++) {
2354 d[H1(e + i)] = m[H1(e + i)];
2355 }
2356 }
2357 }
2358 } else {
2359 for (int r = 0; r < nreg; r++) {
2360 uint8_t *d = vd + r * sizeof(ARMVectorReg);
2361 uint8_t *n = vn + r * sizeof(ARMVectorReg);
2362 uint8_t *m = vm + r * sizeof(ARMVectorReg);
2363 int split = p.count - r * elements;
2364 int e = 0;
2365
2366 for (; e < MIN(split, elements); e += estride) {
2367 d[H1(e)] = n[H1(e)];
2368 for (int i = 1; i < estride; i++) {
2369 d[H1(e + i)] = m[H1(e + i)];
2370 }
2371 }
2372 for (; e < elements; e++) {
2373 d[H1(e)] = m[H1(e)];
2374 }
2375 }
2376 }
2377 }
2378 }
2379
2380 void HELPER(sme2_sel_h)(void *vd, void *vn, void *vm,
2381 uint32_t png, uint32_t desc)
2382 {
2383 int vl = simd_oprsz(desc);
2384 int nreg = simd_data(desc);
2385 int elements = vl / sizeof(uint16_t);
2386 DecodeCounter p = decode_counter(png, vl, MO_16);
2387
2388 if (p.lg2_stride == 0) {
2389 if (p.invert) {
2390 for (int r = 0; r < nreg; r++) {
2391 uint16_t *d = vd + r * sizeof(ARMVectorReg);
2392 uint16_t *n = vn + r * sizeof(ARMVectorReg);
2393 uint16_t *m = vm + r * sizeof(ARMVectorReg);
2394 int split = p.count - r * elements;
2395
2396 if (split <= 0) {
2397 memcpy(d, n, vl); /* all true */
2398 } else if (elements <= split) {
2399 memcpy(d, m, vl); /* all false */
2400 } else {
2401 for (int e = 0; e < split; e++) {
2402 d[H2(e)] = m[H2(e)];
2403 }
2404 for (int e = split; e < elements; e++) {
2405 d[H2(e)] = n[H2(e)];
2406 }
2407 }
2408 }
2409 } else {
2410 for (int r = 0; r < nreg; r++) {
2411 uint16_t *d = vd + r * sizeof(ARMVectorReg);
2412 uint16_t *n = vn + r * sizeof(ARMVectorReg);
2413 uint16_t *m = vm + r * sizeof(ARMVectorReg);
2414 int split = p.count - r * elements;
2415
2416 if (split <= 0) {
2417 memcpy(d, m, vl); /* all false */
2418 } else if (elements <= split) {
2419 memcpy(d, n, vl); /* all true */
2420 } else {
2421 for (int e = 0; e < split; e++) {
2422 d[H2(e)] = n[H2(e)];
2423 }
2424 for (int e = split; e < elements; e++) {
2425 d[H2(e)] = m[H2(e)];
2426 }
2427 }
2428 }
2429 }
2430 } else {
2431 int estride = 1 << p.lg2_stride;
2432 if (p.invert) {
2433 for (int r = 0; r < nreg; r++) {
2434 uint16_t *d = vd + r * sizeof(ARMVectorReg);
2435 uint16_t *n = vn + r * sizeof(ARMVectorReg);
2436 uint16_t *m = vm + r * sizeof(ARMVectorReg);
2437 int split = p.count - r * elements;
2438 int e = 0;
2439
2440 for (; e < MIN(split, elements); e++) {
2441 d[H2(e)] = m[H2(e)];
2442 }
2443 for (; e < elements; e += estride) {
2444 d[H2(e)] = n[H2(e)];
2445 for (int i = 1; i < estride; i++) {
2446 d[H2(e + i)] = m[H2(e + i)];
2447 }
2448 }
2449 }
2450 } else {
2451 for (int r = 0; r < nreg; r++) {
2452 uint16_t *d = vd + r * sizeof(ARMVectorReg);
2453 uint16_t *n = vn + r * sizeof(ARMVectorReg);
2454 uint16_t *m = vm + r * sizeof(ARMVectorReg);
2455 int split = p.count - r * elements;
2456 int e = 0;
2457
2458 for (; e < MIN(split, elements); e += estride) {
2459 d[H2(e)] = n[H2(e)];
2460 for (int i = 1; i < estride; i++) {
2461 d[H2(e + i)] = m[H2(e + i)];
2462 }
2463 }
2464 for (; e < elements; e++) {
2465 d[H2(e)] = m[H2(e)];
2466 }
2467 }
2468 }
2469 }
2470 }
2471
2472 void HELPER(sme2_sel_s)(void *vd, void *vn, void *vm,
2473 uint32_t png, uint32_t desc)
2474 {
2475 int vl = simd_oprsz(desc);
2476 int nreg = simd_data(desc);
2477 int elements = vl / sizeof(uint32_t);
2478 DecodeCounter p = decode_counter(png, vl, MO_32);
2479
2480 if (p.lg2_stride == 0) {
2481 if (p.invert) {
2482 for (int r = 0; r < nreg; r++) {
2483 uint32_t *d = vd + r * sizeof(ARMVectorReg);
2484 uint32_t *n = vn + r * sizeof(ARMVectorReg);
2485 uint32_t *m = vm + r * sizeof(ARMVectorReg);
2486 int split = p.count - r * elements;
2487
2488 if (split <= 0) {
2489 memcpy(d, n, vl); /* all true */
2490 } else if (elements <= split) {
2491 memcpy(d, m, vl); /* all false */
2492 } else {
2493 for (int e = 0; e < split; e++) {
2494 d[H4(e)] = m[H4(e)];
2495 }
2496 for (int e = split; e < elements; e++) {
2497 d[H4(e)] = n[H4(e)];
2498 }
2499 }
2500 }
2501 } else {
2502 for (int r = 0; r < nreg; r++) {
2503 uint32_t *d = vd + r * sizeof(ARMVectorReg);
2504 uint32_t *n = vn + r * sizeof(ARMVectorReg);
2505 uint32_t *m = vm + r * sizeof(ARMVectorReg);
2506 int split = p.count - r * elements;
2507
2508 if (split <= 0) {
2509 memcpy(d, m, vl); /* all false */
2510 } else if (elements <= split) {
2511 memcpy(d, n, vl); /* all true */
2512 } else {
2513 for (int e = 0; e < split; e++) {
2514 d[H4(e)] = n[H4(e)];
2515 }
2516 for (int e = split; e < elements; e++) {
2517 d[H4(e)] = m[H4(e)];
2518 }
2519 }
2520 }
2521 }
2522 } else {
2523 /* p.esz must be MO_64, so stride must be 2. */
2524 if (p.invert) {
2525 for (int r = 0; r < nreg; r++) {
2526 uint32_t *d = vd + r * sizeof(ARMVectorReg);
2527 uint32_t *n = vn + r * sizeof(ARMVectorReg);
2528 uint32_t *m = vm + r * sizeof(ARMVectorReg);
2529 int split = p.count - r * elements;
2530 int e = 0;
2531
2532 for (; e < MIN(split, elements); e++) {
2533 d[H4(e)] = m[H4(e)];
2534 }
2535 for (; e < elements; e += 2) {
2536 d[H4(e)] = n[H4(e)];
2537 d[H4(e + 1)] = m[H4(e + 1)];
2538 }
2539 }
2540 } else {
2541 for (int r = 0; r < nreg; r++) {
2542 uint32_t *d = vd + r * sizeof(ARMVectorReg);
2543 uint32_t *n = vn + r * sizeof(ARMVectorReg);
2544 uint32_t *m = vm + r * sizeof(ARMVectorReg);
2545 int split = p.count - r * elements;
2546 int e = 0;
2547
2548 for (; e < MIN(split, elements); e += 2) {
2549 d[H4(e)] = n[H4(e)];
2550 d[H4(e + 1)] = m[H4(e + 1)];
2551 }
2552 for (; e < elements; e++) {
2553 d[H4(e)] = m[H4(e)];
2554 }
2555 }
2556 }
2557 }
2558 }
2559
2560 void HELPER(sme2_sel_d)(void *vd, void *vn, void *vm,
2561 uint32_t png, uint32_t desc)
2562 {
2563 int vl = simd_oprsz(desc);
2564 int nreg = simd_data(desc);
2565 int elements = vl / sizeof(uint64_t);
2566 DecodeCounter p = decode_counter(png, vl, MO_64);
2567
2568 if (p.invert) {
2569 for (int r = 0; r < nreg; r++) {
2570 uint64_t *d = vd + r * sizeof(ARMVectorReg);
2571 uint64_t *n = vn + r * sizeof(ARMVectorReg);
2572 uint64_t *m = vm + r * sizeof(ARMVectorReg);
2573 int split = p.count - r * elements;
2574
2575 if (split <= 0) {
2576 memcpy(d, n, vl); /* all true */
2577 } else if (elements <= split) {
2578 memcpy(d, m, vl); /* all false */
2579 } else {
2580 memcpy(d, m, split * sizeof(uint64_t));
2581 memcpy(d + split, n + split,
2582 (elements - split) * sizeof(uint64_t));
2583 }
2584 }
2585 } else {
2586 for (int r = 0; r < nreg; r++) {
2587 uint64_t *d = vd + r * sizeof(ARMVectorReg);
2588 uint64_t *n = vn + r * sizeof(ARMVectorReg);
2589 uint64_t *m = vm + r * sizeof(ARMVectorReg);
2590 int split = p.count - r * elements;
2591
2592 if (split <= 0) {
2593 memcpy(d, m, vl); /* all false */
2594 } else if (elements <= split) {
2595 memcpy(d, n, vl); /* all true */
2596 } else {
2597 memcpy(d, n, split * sizeof(uint64_t));
2598 memcpy(d + split, m + split,
2599 (elements - split) * sizeof(uint64_t));
2600 }
2601 }
2602 }
2603 }
2604
2605 void sme_mop4(void *vza, void *vzn, void *vzm, void *fn_opaque,
2606 uint32_t desc, size_t esize,
2607 void (*fn)(void *, void *, void *, void *))
2608 {
2609 intptr_t oprsz = simd_maxsz(desc);
2610 intptr_t dim = oprsz / 2; /* in bytes */
2611 bool nreg_m1 = extract32(desc, SIMD_DATA_SHIFT + 0, 1);
2612 bool mreg_m1 = extract32(desc, SIMD_DATA_SHIFT + 1, 1);
2613 intptr_t host_adj = HOST_BIG_ENDIAN ? 8 - esize : 0;
2614
2615 for (int outprod = 0; outprod < 4; outprod++) {
2616 bool row_hv = outprod & 2;
2617 bool col_hv = outprod & 1;
2618 intptr_t row_base = row_hv ? dim : 0;
2619 intptr_t col_base = col_hv ? dim : 0;
2620 void *op1 = vzn + (col_hv && nreg_m1 ? sizeof(ARMVectorReg) : 0);
2621 void *op2 = vzm + (row_hv && mreg_m1 ? sizeof(ARMVectorReg) : 0);
2622
2623 for (intptr_t row = 0; row < dim; row += esize) {
2624 intptr_t row_idx = row_base + row;
2625 void *vza_row = vza + tile_vslice_offset(row_idx);
2626 void *e1 = op1 + (row_idx ^ host_adj);
2627
2628 for (intptr_t col = 0; col < dim; col += esize) {
2629 intptr_t col_idx = col_base + col;
2630 void *e2 = op2 + (col_idx ^ host_adj);
2631 void *e3 = vza_row + (col_idx ^ host_adj);
2632
2633 fn(e3, e1, e2, fn_opaque);
2634 }
2635 }
2636 }
2637 }
2638
2639 /*
2640 * Sparse outer product, non-widening. ESZ in {16, 32}.
2641 */
2642 static void sme_tmop(void *vza, void *vzn, void *vzm, uint64_t *zk,
2643 void *fn_opaque, uint32_t desc, MemOp esz,
2644 void (*fn)(void *, void *, void *, void *))
2645 {
2646 intptr_t oprsz = simd_maxsz(desc);
2647 intptr_t index = simd_data(desc);
2648 intptr_t esize = 1 << esz;
2649 intptr_t host_adj = HOST_BIG_ENDIAN ? 8 - esize : 0;
2650 /* Base in bits for op3[index*:csize], csize = (VL * 2) / esize. */
2651 intptr_t ctrl_base = index * oprsz * 2;
2652 /* Create a zero for use with the largest esz. */
2653 uint32_t zero = 0;
2654
2655 for (intptr_t row = 0; row < oprsz; row += esize) {
2656 void *vza_row = vza + tile_vslice_offset(row);
2657
2658 for (intptr_t col = 0; col < oprsz; col += esize) {
2659 void *e2 = vzm + (col ^ host_adj);
2660 void *e3 = vza_row + (col ^ host_adj);
2661
2662 /*
2663 * Two control bits select one element:
2664 * Zn[row], if [0] is set,
2665 * Zn+1[row], if [1] is set,
2666 * 0, otherwise.
2667 * Compute the address of that element.
2668 */
2669 void *e1 = &zero;
2670 uint64_t this_ctrl = extractn(zk, (ctrl_base + 2 * col) >> esz, 2);
2671 if (this_ctrl) {
2672 e1 = vzn + (row ^ host_adj);
2673 if (!(this_ctrl & 1)) {
2674 e1 += sizeof(ARMVectorReg);
2675 }
2676 }
2677 fn(e3, e1, e2, fn_opaque);
2678 }
2679 }
2680 }
2681
2682 /*
2683 * Sparse outer product, widening 2-way, 16 to 32-bit.
2684 */
2685 static void sme_tmop_2way_sh(uint32_t *za, uint16_t *zn0, uint32_t *zm,
2686 uint64_t *zk, void *fn_opaque, uint32_t desc,
2687 void (*fn)(void *, void *, void *, void *))
2688 {
2689 intptr_t oprsz = simd_maxsz(desc);
2690 intptr_t dim = oprsz >> MO_32;
2691 intptr_t index = simd_data(desc);
2692 intptr_t ctrl_base = (index * oprsz) >> 1;
2693 uint16_t *zn1 = zn0 + sizeof(ARMVectorReg) / 2;
2694
2695 for (intptr_t row = 0; row < dim; row++) {
2696 uint32_t *za_row = za + tile_vslice_offset(row);
2697
2698 for (intptr_t col = 0; col < dim; col++) {
2699 uint32_t *e2 = zm + H4(col);
2700 uint32_t *e3 = za_row + H4(col);
2701 uint32_t e1 = 0;
2702
2703 /*
2704 * Four control bits select two elements. The two elements
2705 * may be non-contiguous, so assemble them locally into e1.
2706 * Pseudo-code has a double loop running forward, with a
2707 * test for (i < 2) to limit construction to 2 elements.
2708 * Easier to run a single loop backward, shifting extra
2709 * elements off the top of our uint32_t.
2710 */
2711 uint64_t this_ctrl = extractn(zk, ctrl_base + col * 4, 4);
2712 for (int i = 3; i >= 0; i--) {
2713 if (this_ctrl & (1 << i)) {
2714 bool e = i & 1;
2715 bool r = i & 2;
2716 uint16_t *p = (r ? zn1 : zn0) + H2(2 * row + e);
2717 e1 = (e1 << 16) | *p;
2718 }
2719 }
2720
2721 fn(e3, &e1, e2, fn_opaque);
2722 }
2723 }
2724 }
2725
2726 void sme_tmop_4way_sb(uint32_t *za, uint8_t *zn0, uint32_t *zm,
2727 uint64_t *zk, void *fn_opaque, uint32_t desc,
2728 void (*fn)(void *, void *, void *, void *))
2729 {
2730 intptr_t oprsz = simd_maxsz(desc);
2731 intptr_t dim = oprsz >> MO_32;
2732 intptr_t index = simd_data(desc);
2733 intptr_t ctrl_base = (index * oprsz) >> 1;
2734 uint8_t *zn1 = zn0 + sizeof(ARMVectorReg);
2735
2736 for (intptr_t row = 0; row < dim; row++) {
2737 uint32_t *za_row = za + tile_vslice_offset(row);
2738
2739 for (intptr_t col = 0; col < dim; col++) {
2740 uint32_t *e2 = zm + H4(col);
2741 uint32_t *e3 = za_row + H4(col);
2742 uint16_t e1l = 0, e1h = 0;
2743 uint32_t e1;
2744
2745 /*
2746 * Eight control bits select two elements from each row.
2747 * The elements may be non-contiguous, so assemble them
2748 * locally into e1.
2749 * Pseudo-code has a triple loop running forward, with a
2750 * test for (i < 2) to limit construction to 2 elements.
2751 * Easier to run a single loop backward, shifting extra
2752 * elements off the top.
2753 */
2754 uint64_t this_ctrl = extractn(zk, ctrl_base + col * 8, 8);
2755 for (int e = 3; e >= 0; e--) {
2756 if (this_ctrl & (0x01 << e)) {
2757 e1l = (e1l << 8) | zn0[H1(4 * row + e)];
2758 }
2759 if (this_ctrl & (0x10 << e)) {
2760 e1h = (e1h << 8) | zn1[H1(4 * row + e)];
2761 }
2762 }
2763 e1 = (e1h << 16) | e1l;
2764
2765 fn(e3, &e1, e2, fn_opaque);
2766 }
2767 }
2768 }
2769
2770 static void inner_fmop4a_hh(void *vd, void *vn, void *vm, void *vinfo)
2771 {
2772 float16 *d = vd, *n = vn, *m = vm;
2773 float_status *fpst = vinfo;
2774
2775 *d = float16_muladd(*n, *m, *d, 0, fpst);
2776 }
2777
2778 void HELPER(sme_fmop4a_hh)(void *vza, void *vzn, void *vzm,
2779 float_status *fpst, uint32_t desc)
2780 {
2781 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float16), inner_fmop4a_hh);
2782 }
2783
2784 void HELPER(sme_ftmopa_hh)(void *vza, void *vzn, void *vzm, void *vzk,
2785 float_status *fpst, uint32_t desc)
2786 {
2787 sme_tmop(vza, vzn, vzm, vzk, fpst, desc, MO_16, inner_fmop4a_hh);
2788 }
2789
2790 static void inner_fmop4s_hh(void *vd, void *vn, void *vm, void *vinfo)
2791 {
2792 float16 *d = vd, *n = vn, *m = vm;
2793 float_status *fpst = vinfo;
2794
2795 *d = float16_muladd(float16_chs(*n), *m, *d, 0, fpst);
2796 }
2797
2798 void HELPER(sme_fmop4s_hh)(void *vza, void *vzn, void *vzm,
2799 float_status *fpst, uint32_t desc)
2800 {
2801 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float16), inner_fmop4s_hh);
2802 }
2803
2804 static void inner_ah_fmop4s_hh(void *vd, void *vn, void *vm, void *vinfo)
2805 {
2806 float16 *d = vd, *n = vn, *m = vm;
2807 float_status *fpst = vinfo;
2808
2809 *d = float16_muladd(*n, *m, *d, float_muladd_negate_product, fpst);
2810 }
2811
2812 void HELPER(sme_ah_fmop4s_hh)(void *vza, void *vzn, void *vzm,
2813 float_status *fpst, uint32_t desc)
2814 {
2815 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float16), inner_ah_fmop4s_hh);
2816 }
2817
2818 static void inner_fmop4a_ss(void *vd, void *vn, void *vm, void *vinfo)
2819 {
2820 float32 *d = vd, *n = vn, *m = vm;
2821 float_status *fpst = vinfo;
2822
2823 *d = float32_muladd(*n, *m, *d, 0, fpst);
2824 }
2825
2826 void HELPER(sme_fmop4a_ss)(void *vza, void *vzn, void *vzm,
2827 float_status *fpst, uint32_t desc)
2828 {
2829 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float32), inner_fmop4a_ss);
2830 }
2831
2832 void HELPER(sme_ftmopa_ss)(void *vza, void *vzn, void *vzm, void *vzk,
2833 float_status *fpst, uint32_t desc)
2834 {
2835 sme_tmop(vza, vzn, vzm, vzk, fpst, desc, MO_32, inner_fmop4a_ss);
2836 }
2837
2838 static void inner_fmop4s_ss(void *vd, void *vn, void *vm, void *vinfo)
2839 {
2840 float32 *d = vd, *n = vn, *m = vm;
2841 float_status *fpst = vinfo;
2842
2843 *d = float32_muladd(float32_chs(*n), *m, *d, 0, fpst);
2844 }
2845
2846 void HELPER(sme_fmop4s_ss)(void *vza, void *vzn, void *vzm,
2847 float_status *fpst, uint32_t desc)
2848 {
2849 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float32), inner_fmop4s_ss);
2850 }
2851
2852 static void inner_ah_fmop4s_ss(void *vd, void *vn, void *vm, void *vinfo)
2853 {
2854 float32 *d = vd, *n = vn, *m = vm;
2855 float_status *fpst = vinfo;
2856
2857 *d = float32_muladd(*n, *m, *d, float_muladd_negate_product, fpst);
2858 }
2859
2860 void HELPER(sme_ah_fmop4s_ss)(void *vza, void *vzn, void *vzm,
2861 float_status *fpst, uint32_t desc)
2862 {
2863 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float32), inner_ah_fmop4s_ss);
2864 }
2865
2866 static void inner_fmop4a_dd(void *vd, void *vn, void *vm, void *vinfo)
2867 {
2868 float64 *d = vd, *n = vn, *m = vm;
2869 float_status *fpst = vinfo;
2870
2871 *d = float64_muladd(*n, *m, *d, 0, fpst);
2872 }
2873
2874 void HELPER(sme_fmop4a_dd)(void *vza, void *vzn, void *vzm,
2875 float_status *fpst, uint32_t desc)
2876 {
2877 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float64), inner_fmop4a_dd);
2878 }
2879
2880 static void inner_fmop4s_dd(void *vd, void *vn, void *vm, void *vinfo)
2881 {
2882 float64 *d = vd, *n = vn, *m = vm;
2883 float_status *fpst = vinfo;
2884
2885 *d = float64_muladd(float64_chs(*n), *m, *d, 0, fpst);
2886 }
2887
2888 void HELPER(sme_fmop4s_dd)(void *vza, void *vzn, void *vzm,
2889 float_status *fpst, uint32_t desc)
2890 {
2891 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float64), inner_fmop4s_dd);
2892 }
2893
2894 static void inner_ah_fmop4s_dd(void *vd, void *vn, void *vm, void *vinfo)
2895 {
2896 float64 *d = vd, *n = vn, *m = vm;
2897 float_status *fpst = vinfo;
2898
2899 *d = float64_muladd(*n, *m, *d, float_muladd_negate_product, fpst);
2900 }
2901
2902 void HELPER(sme_ah_fmop4s_dd)(void *vza, void *vzn, void *vzm,
2903 float_status *fpst, uint32_t desc)
2904 {
2905 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(float64), inner_ah_fmop4s_dd);
2906 }
2907
2908 static void inner_bfmop4a_hh(void *vd, void *vn, void *vm, void *vinfo)
2909 {
2910 bfloat16 *d = vd, *n = vn, *m = vm;
2911 float_status *fpst = vinfo;
2912
2913 *d = bfloat16_muladd(*n, *m, *d, 0, fpst);
2914 }
2915
2916 void HELPER(sme_bfmop4a_hh)(void *vza, void *vzn, void *vzm,
2917 float_status *fpst, uint32_t desc)
2918 {
2919 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(bfloat16), inner_bfmop4a_hh);
2920 }
2921
2922 void HELPER(sme_bftmopa_hh)(void *vza, void *vzn, void *vzm, void *vzk,
2923 float_status *fpst, uint32_t desc)
2924 {
2925 sme_tmop(vza, vzn, vzm, vzk, fpst, desc, MO_16, inner_bfmop4a_hh);
2926 }
2927
2928 static void inner_bfmop4s_hh(void *vd, void *vn, void *vm, void *vinfo)
2929 {
2930 bfloat16 *d = vd, *n = vn, *m = vm;
2931 float_status *fpst = vinfo;
2932
2933 *d = bfloat16_muladd(bfloat16_chs(*n), *m, *d, 0, fpst);
2934 }
2935
2936 void HELPER(sme_bfmop4s_hh)(void *vza, void *vzn, void *vzm,
2937 float_status *fpst, uint32_t desc)
2938 {
2939 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(bfloat16), inner_bfmop4s_hh);
2940 }
2941
2942 static void inner_ah_bfmop4s_hh(void *vd, void *vn, void *vm, void *vinfo)
2943 {
2944 bfloat16 *d = vd, *n = vn, *m = vm;
2945 float_status *fpst = vinfo;
2946
2947 *d = bfloat16_muladd(*n, *m, *d, float_muladd_negate_product, fpst);
2948 }
2949
2950 void HELPER(sme_ah_bfmop4s_hh)(void *vza, void *vzn, void *vzm,
2951 float_status *fpst, uint32_t desc)
2952 {
2953 sme_mop4(vza, vzn, vzm, fpst, desc, sizeof(bfloat16), inner_ah_bfmop4s_hh);
2954 }
2955
2956 static void inner_bfmop4a_sh(void *vd, void *vn, void *vm, void *vinfo)
2957 {
2958 float32 *d = vd;
2959 uint32_t *n = vn, *m = vm;
2960 float_status *fpst = vinfo;
2961
2962 *d = bfdotadd(*d, *n, *m, fpst);
2963 }
2964
2965 static void inner_ebf_bfmop4a_sh(void *vd, void *vn, void *vm, void *vinfo)
2966 {
2967 float32 *d = vd;
2968 uint32_t *n = vn, *m = vm;
2969 float_status *fpst = vinfo;
2970
2971 *d = bfdotadd_ebf(*d, *n, *m, fpst);
2972 }
2973
2974 void HELPER(sme_bfmop4a_sh)(void *vza, void *vzn, void *vzm,
2975 CPUArchState *env, uint32_t desc)
2976 {
2977 float_status fpst;
2978
2979 sme_mop4(vza, vzn, vzm, &fpst, desc, sizeof(float32),
2980 is_ebf(env, &fpst) ? inner_ebf_bfmop4a_sh
2981 : inner_bfmop4a_sh);
2982 }
2983
2984 void HELPER(sme_bftmopa_sh)(void *vza, void *vzn, void *vzm, void *vzk,
2985 CPUArchState *env, uint32_t desc)
2986 {
2987 float_status fpst;
2988
2989 sme_tmop_2way_sh(vza, vzn, vzm, vzk, &fpst, desc,
2990 is_ebf(env, &fpst) ? inner_ebf_bfmop4a_sh
2991 : inner_bfmop4a_sh);
2992 }
2993
2994 static void inner_bfmop4s_sh(void *vd, void *vn, void *vm, void *vinfo)
2995 {
2996 float32 *d = vd;
2997 uint32_t *n = vn, *m = vm;
2998 float_status *fpst = vinfo;
2999
3000 *d = bfdotadd(*d, *n ^ 0x80008000u, *m, fpst);
3001 }
3002
3003 static void inner_ebf_bfmop4s_sh(void *vd, void *vn, void *vm, void *vinfo)
3004 {
3005 float32 *d = vd;
3006 uint32_t *n = vn, *m = vm;
3007 float_status *fpst = vinfo;
3008
3009 *d = bfdotadd_ebf(*d, *n ^ 0x80008000u, *m, fpst);
3010 }
3011
3012 void HELPER(sme_bfmop4s_sh)(void *vza, void *vzn, void *vzm,
3013 CPUArchState *env, uint32_t desc)
3014 {
3015 float_status fpst;
3016
3017 sme_mop4(vza, vzn, vzm, &fpst, desc, sizeof(float32),
3018 is_ebf(env, &fpst) ? inner_ebf_bfmop4s_sh
3019 : inner_bfmop4s_sh);
3020 }
3021
3022 static void inner_ah_bfmop4s_sh(void *vd, void *vn, void *vm, void *vinfo)
3023 {
3024 float32 *d = vd;
3025 uint32_t *n = vn, *m = vm;
3026 float_status *fpst = vinfo;
3027
3028 *d = bfdotadd(*d, bf16mop_ah_neg_adj_pair(*n, -1), *m, fpst);
3029 }
3030
3031 static void inner_ebf_ah_bfmop4s_sh(void *vd, void *vn, void *vm, void *vinfo)
3032 {
3033 float32 *d = vd;
3034 uint32_t *n = vn, *m = vm;
3035 float_status *fpst = vinfo;
3036
3037 *d = bfdotadd_ebf(*d, bf16mop_ah_neg_adj_pair(*n, -1), *m, fpst);
3038 }
3039
3040 void HELPER(sme_ah_bfmop4s_sh)(void *vza, void *vzn, void *vzm,
3041 CPUArchState *env, uint32_t desc)
3042 {
3043 float_status fpst;
3044
3045 sme_mop4(vza, vzn, vzm, &fpst, desc, sizeof(float32),
3046 is_ebf(env, &fpst) ? inner_ebf_ah_bfmop4s_sh
3047 : inner_ah_bfmop4s_sh);
3048 }
3049
3050 static void inner_fmop4a_sh(void *vd, void *vn, void *vm, void *vinfo)
3051 {
3052 float32 *d = vd;
3053 uint32_t *n = vn, *m = vm;
3054 CPUArchState *env = vinfo;
3055
3056 *d = f16_dotadd(*d, *n, *m,
3057 &env->vfp.fp_status[FPST_ZA_F16],
3058 &env->vfp.fp_status[FPST_ZA]);
3059 }
3060
3061 void HELPER(sme_fmop4a_sh)(void *vza, void *vzn, void *vzm,
3062 CPUArchState *env, uint32_t desc)
3063 {
3064 sme_mop4(vza, vzn, vzm, env, desc, sizeof(float32), inner_fmop4a_sh);
3065 }
3066
3067 void HELPER(sme_ftmopa_sh)(void *vza, void *vzn, void *vzm, void *vzk,
3068 CPUArchState *env, uint32_t desc)
3069 {
3070 sme_tmop_2way_sh(vza, vzn, vzm, vzk, env, desc, inner_fmop4a_sh);
3071 }
3072
3073 static void inner_fmop4s_sh(void *vd, void *vn, void *vm, void *vinfo)
3074 {
3075 float32 *d = vd;
3076 uint32_t *n = vn, *m = vm;
3077 CPUArchState *env = vinfo;
3078
3079 *d = f16_dotadd(*d, *n ^ 0x80008000u, *m,
3080 &env->vfp.fp_status[FPST_ZA_F16],
3081 &env->vfp.fp_status[FPST_ZA]);
3082 }
3083
3084 void HELPER(sme_fmop4s_sh)(void *vza, void *vzn, void *vzm,
3085 CPUArchState *env, uint32_t desc)
3086 {
3087 sme_mop4(vza, vzn, vzm, env, desc, sizeof(float32), inner_fmop4s_sh);
3088 }
3089
3090 static void inner_ah_fmop4s_sh(void *vd, void *vn, void *vm, void *vinfo)
3091 {
3092 float32 *d = vd;
3093 uint32_t *n = vn, *m = vm;
3094 CPUArchState *env = vinfo;
3095
3096 *d = f16_dotadd(*d, f16mop_ah_neg_adj_pair(*n, -1), *m,
3097 &env->vfp.fp_status[FPST_ZA_F16],
3098 &env->vfp.fp_status[FPST_ZA]);
3099 }
3100
3101 void HELPER(sme_ah_fmop4s_sh)(void *vza, void *vzn, void *vzm,
3102 CPUArchState *env, uint32_t desc)
3103 {
3104 sme_mop4(vza, vzn, vzm, env, desc, sizeof(float32), inner_ah_fmop4s_sh);
3105 }
3106
3107 #define IMOP4_2WAY(NAME, OP, TYPED, TYPEN, TYPEM) \
3108 static void inner_##NAME(void *vd, void *vn, void *vm, void *vinfo) \
3109 { \
3110 TYPEN *n = vn; TYPEM *m = vm; TYPED *d = vd; \
3111 *d OP##= (TYPED)n[0] * m[0] + (TYPED)n[1] * m[1]; \
3112 } \
3113 void HELPER(sme_##NAME)(void *vza, void *vzn, void *vzm, uint32_t desc) \
3114 { \
3115 sme_mop4(vza, vzn, vzm, NULL, desc, sizeof(TYPED), inner_##NAME); \
3116 }
3117
3118 IMOP4_2WAY(smop4a_sh, +, int32_t, int16_t, int16_t)
3119 IMOP4_2WAY(smop4s_sh, -, int32_t, int16_t, int16_t)
3120
3121 IMOP4_2WAY(umop4a_sh, +, int32_t, uint16_t, uint16_t)
3122 IMOP4_2WAY(umop4s_sh, -, int32_t, uint16_t, uint16_t)
3123
3124 #undef IMOP4_2WAY
3125
3126 #define ITMOP_2WAY(TNAME, MNAME) \
3127 void HELPER(sme_##TNAME)(void *vza, void *vzn, void *vzm, \
3128 void *vzk, uint32_t desc) \
3129 { \
3130 sme_tmop_2way_sh(vza, vzn, vzm, vzk, NULL, desc, inner_##MNAME); \
3131 }
3132
3133 ITMOP_2WAY(stmopa_sh, smop4a_sh)
3134 ITMOP_2WAY(utmopa_sh, umop4a_sh)
3135
3136 #undef ITMOP_2WAY
3137
3138 #define IMOP4_4WAY(NAME, OP, TYPED, TYPEN, TYPEM) \
3139 static void inner_##NAME(void *vd, void *vn, void *vm, void *vinfo) \
3140 { \
3141 TYPEN *n = vn; TYPEM *m = vm; TYPED *d = vd; \
3142 *d OP##= (TYPED)n[0] * m[0] + (TYPED)n[1] * m[1] + \
3143 (TYPED)n[2] * m[2] + (TYPED)n[3] * m[3]; \
3144 } \
3145 void HELPER(sme_##NAME)(void *vza, void *vzn, void *vzm, uint32_t desc) \
3146 { \
3147 sme_mop4(vza, vzn, vzm, NULL, desc, sizeof(TYPED), inner_##NAME); \
3148 }
3149
3150 IMOP4_4WAY(smop4a_sb, +, int32_t, int8_t, int8_t)
3151 IMOP4_4WAY(smop4s_sb, -, int32_t, int8_t, int8_t)
3152 IMOP4_4WAY(smop4a_dh, +, int64_t, int16_t, int16_t)
3153 IMOP4_4WAY(smop4s_dh, -, int64_t, int16_t, int16_t)
3154
3155 IMOP4_4WAY(sumop4a_sb, +, int32_t, int8_t, uint8_t)
3156 IMOP4_4WAY(sumop4s_sb, -, int32_t, int8_t, uint8_t)
3157 IMOP4_4WAY(sumop4a_dh, +, int64_t, int16_t, uint16_t)
3158 IMOP4_4WAY(sumop4s_dh, -, int64_t, int16_t, uint16_t)
3159
3160 IMOP4_4WAY(umop4a_sb, +, int32_t, uint8_t, uint8_t)
3161 IMOP4_4WAY(umop4s_sb, -, int32_t, uint8_t, uint8_t)
3162 IMOP4_4WAY(umop4a_dh, +, int64_t, uint16_t, uint16_t)
3163 IMOP4_4WAY(umop4s_dh, -, int64_t, uint16_t, uint16_t)
3164
3165 IMOP4_4WAY(usmop4a_sb, +, int32_t, uint8_t, int8_t)
3166 IMOP4_4WAY(usmop4s_sb, -, int32_t, uint8_t, int8_t)
3167 IMOP4_4WAY(usmop4a_dh, +, int64_t, uint16_t, int16_t)
3168 IMOP4_4WAY(usmop4s_dh, -, int64_t, uint16_t, int16_t)
3169
3170 #undef IMOP4_4WAY
3171
3172 #define ITMOP_4WAY(TNAME, MNAME) \
3173 void HELPER(sme_##TNAME)(void *vza, void *vzn, void *vzm, \
3174 void *vzk, uint32_t desc) \
3175 { \
3176 sme_tmop_4way_sb(vza, vzn, vzm, vzk, NULL, desc, inner_##MNAME); \
3177 }
3178
3179 ITMOP_4WAY(stmopa_sb, smop4a_sb)
3180 ITMOP_4WAY(utmopa_sb, umop4a_sb)
3181 ITMOP_4WAY(sutmopa_sb, sumop4a_sb)
3182 ITMOP_4WAY(ustmopa_sb, usmop4a_sb)
3183
3184 #undef ITMOP_4WAY