1
2
3
4
5 package main
6
7
8
9
10 import (
11 "bufio"
12 "bytes"
13 "flag"
14 "fmt"
15 "go/format"
16 "io"
17 "os"
18 "simd/archsimd/_gen/sgutil"
19 "strings"
20 "text/template"
21 )
22
23 type resultTypeFunc func(t string, w, c int) (ot string, ow int, oc int)
24
25
26 type shapes struct {
27 vecs []int
28 ints []int
29 uints []int
30 floats []int
31 output resultTypeFunc
32 }
33
34
35 type shapeAndTemplate struct {
36 s *shapes
37 t *template.Template
38 }
39
40 func Map[T, U any](f func(T) U, in []T) []U {
41 x := make([]U, len(in))
42 for i, v := range in {
43 x[i] = f(v)
44 }
45 return x
46 }
47
48
49
50 type shapeFilter int
51
52 const (
53 filterAll shapeFilter = iota
54 filterSmallOnly
55 filterLarge
56 )
57
58 func (sat shapeAndTemplate) target(outType string, width int) shapeAndTemplate {
59 newSat := sat
60 newShape := *sat.s
61 newShape.output = func(t string, w, c int) (ot string, ow int, oc int) {
62 oc = c
63 if width*c > 512 {
64 oc = 512 / width
65 } else if width*c < 128 {
66 oc = 128 / width
67 }
68 return outType, width, oc
69 }
70 newSat.s = &newShape
71 return newSat
72 }
73
74
75 func (sat shapeAndTemplate) arm64Target(outType string, width int) shapeAndTemplate {
76 newSat := sat
77 newShape := *sat.s
78 newShape.output = func(t string, w, c int) (ot string, ow int, oc int) {
79 oc = c
80 if width*c > 128 {
81 oc = 128 / width
82 } else if width*c < 128 {
83 oc = 128 / width
84 }
85 return outType, width, oc
86 }
87 newSat.s = &newShape
88 return newSat
89 }
90
91 func (sat shapeAndTemplate) targetFixed(outType string, width, count int) shapeAndTemplate {
92 newSat := sat
93 newShape := *sat.s
94 newShape.output = func(t string, w, c int) (ot string, ow int, oc int) {
95 return outType, width, count
96 }
97 newSat.s = &newShape
98 return newSat
99 }
100
101 func (s *shapes) forAllShapes(f func(seq int, t, upperT string, w, c int, out io.Writer), out io.Writer) {
102 vecs := s.vecs
103 ints := s.ints
104 uints := s.uints
105 floats := s.floats
106 seq := 0
107 for _, v := range vecs {
108 for _, w := range ints {
109 c := v / w
110 f(seq, "int", "Int", w, c, out)
111 seq++
112 }
113 for _, w := range uints {
114 c := v / w
115 f(seq, "uint", "Uint", w, c, out)
116 seq++
117 }
118 for _, w := range floats {
119 c := v / w
120 f(seq, "float", "Float", w, c, out)
121 seq++
122 }
123 }
124 }
125
126 var allShapes = &shapes{
127 vecs: []int{128, 256, 512},
128 ints: []int{8, 16, 32, 64},
129 uints: []int{8, 16, 32, 64},
130 floats: []int{32, 64},
131 }
132
133 var intShapes = &shapes{
134 vecs: []int{128, 256, 512},
135 ints: []int{8, 16, 32, 64},
136 }
137
138 var uintShapes = &shapes{
139 vecs: []int{128, 256, 512},
140 uints: []int{8, 16, 32, 64},
141 }
142
143 var floatShapes = &shapes{
144 vecs: []int{128, 256, 512},
145 floats: []int{32, 64},
146 }
147
148 var integerShapes = &shapes{
149 vecs: []int{128, 256, 512},
150 ints: []int{8, 16, 32, 64},
151 uints: []int{8, 16, 32, 64},
152 }
153
154 var avx512Shapes = &shapes{
155 vecs: []int{512},
156 ints: []int{8, 16, 32, 64},
157 uints: []int{8, 16, 32, 64},
158 floats: []int{32, 64},
159 }
160
161 var avx2Shapes = &shapes{
162 vecs: []int{128, 256},
163 ints: []int{8, 16, 32, 64},
164 uints: []int{8, 16, 32, 64},
165 floats: []int{32, 64},
166 }
167
168 var avx2MaskedLoadShapes = &shapes{
169 vecs: []int{128, 256},
170 ints: []int{32, 64},
171 uints: []int{32, 64},
172 floats: []int{32, 64},
173 }
174
175
176 var arm64Shapes = &shapes{
177 vecs: []int{128},
178 ints: []int{8, 16, 32, 64},
179 uints: []int{8, 16, 32, 64},
180 floats: []int{32, 64},
181 }
182
183
184 var arm64IntegerShapes = &shapes{
185 vecs: []int{128},
186 ints: []int{8, 16, 32, 64},
187 uints: []int{8, 16, 32, 64},
188 }
189
190 var arm64IntShapes = &shapes{
191 vecs: []int{128},
192 ints: []int{8, 16, 32, 64},
193 }
194
195
196 var arm64ReduceIntegerShapes = &shapes{
197 vecs: []int{128},
198 ints: []int{8, 16, 32},
199 uints: []int{8, 16, 32},
200 }
201
202
203 var arm64ReduceAllShapes = &shapes{
204 vecs: []int{128},
205 ints: []int{8, 16, 32},
206 uints: []int{8, 16, 32},
207 floats: []int{32},
208 }
209
210
211
212 var arm64UintToIntShapes = &shapes{
213 vecs: []int{128},
214 uints: []int{8, 16, 32, 64},
215 output: func(t string, w, c int) (string, int, int) {
216 return "int", w, c
217 },
218 }
219
220 var avx2SmallLoadPunShapes = &shapes{
221
222
223 vecs: []int{256},
224 uints: []int{8, 16},
225 }
226
227 var unaryFlaky = &shapes{
228 vecs: []int{128, 256, 512},
229 floats: []int{32, 64},
230 }
231
232 var ternaryFlaky = &shapes{
233 vecs: []int{128, 256, 512},
234 floats: []int{32, 64},
235 }
236
237 var avx2SignedComparisons = &shapes{
238 vecs: []int{128, 256},
239 ints: []int{8, 16, 32, 64},
240 }
241
242 var avx2UnsignedComparisons = &shapes{
243 vecs: []int{128, 256},
244 uints: []int{8, 16, 32, 64},
245 }
246
247
248 var amdIntShiftAllShapes = &shapes{
249 vecs: []int{128, 256, 512},
250 ints: []int{16, 32, 64},
251 }
252
253 var amdUintShiftAllShapes = &shapes{
254 vecs: []int{128, 256, 512},
255 uints: []int{16, 32, 64},
256 }
257
258 var neonIntShiftAllShapes = &shapes{
259 vecs: []int{128},
260 ints: []int{8, 16, 32, 64},
261 }
262
263 var neonUintShiftAllShapes = &shapes{
264 vecs: []int{128},
265 uints: []int{8, 16, 32, 64},
266 }
267
268 type templateData struct {
269 VType string
270 AOrAn string
271 EWidth int
272 Vwidth int
273 Count int
274 WxC string
275 BxC string
276 Base string
277 Etype string
278 OxFF string
279
280 OVType string
281 OEtype string
282 OEType string
283 OCount int
284 }
285
286 func (t templateData) As128BitVec() string {
287 return fmt.Sprintf("%s%dx%d", t.Base, t.EWidth, 128/t.EWidth)
288 }
289
290 func oneTemplate(t *template.Template, baseType string, width, count int, out io.Writer, rtf resultTypeFunc, filter shapeFilter) {
291 b := width * count
292 if b < 128 || b > 512 {
293 return
294 }
295
296 ot, ow, oc := baseType, width, count
297 if rtf != nil {
298 ot, ow, oc = rtf(ot, ow, oc)
299 if ow*oc > 512 || ow*oc < 128 || ow < 8 || ow > 64 {
300 return
301 }
302
303 if ot == "float" && ow < 32 {
304 return
305 }
306 if ot == baseType && ow == width && oc == count && strings.Contains(t.Name(), "convert_helpers") {
307 return
308 }
309 }
310
311 ob := ow * oc
312 isSmall := (b <= 128) && (ob <= 128)
313 switch filter {
314 case filterSmallOnly:
315 if !isSmall {
316 return
317 }
318 case filterLarge:
319 if isSmall {
320 return
321 }
322 }
323
324 ovType := fmt.Sprintf("%s%dx%d", strings.ToUpper(ot[:1])+ot[1:], ow, oc)
325 oeType := fmt.Sprintf("%s%d", ot, ow)
326 oEType := fmt.Sprintf("%s%d", strings.ToUpper(ot[:1])+ot[1:], ow)
327
328 wxc := fmt.Sprintf("%dx%d", width, count)
329 BaseType := strings.ToUpper(baseType[:1]) + baseType[1:]
330 vType := fmt.Sprintf("%s%s", BaseType, wxc)
331 eType := fmt.Sprintf("%s%d", baseType, width)
332
333 bxc := fmt.Sprintf("%dx%d", 8, count*(width/8))
334 aOrAn := "a"
335 if strings.Contains("aeiou", baseType[:1]) {
336 aOrAn = "an"
337 }
338 oxFF := fmt.Sprintf("0x%x", uint64((1<<count)-1))
339 err := t.Execute(out, templateData{
340 VType: vType,
341 AOrAn: aOrAn,
342 EWidth: width,
343 Vwidth: b,
344 Count: count,
345 WxC: wxc,
346 BxC: bxc,
347 Base: BaseType,
348 Etype: eType,
349 OxFF: oxFF,
350 OVType: ovType,
351 OEtype: oeType,
352 OCount: oc,
353 OEType: oEType,
354 })
355 if err != nil {
356 panic(fmt.Errorf("template execute failed, %v", err))
357 }
358 }
359
360
361
362 func (sat shapeAndTemplate) forTemplates(out io.Writer, filter shapeFilter) {
363 t, s := sat.t, sat.s
364 vecs := s.vecs
365 ints := s.ints
366 uints := s.uints
367 floats := s.floats
368 for _, v := range vecs {
369 for _, w := range ints {
370 c := v / w
371 oneTemplate(t, "int", w, c, out, sat.s.output, filter)
372 }
373 for _, w := range uints {
374 c := v / w
375 oneTemplate(t, "uint", w, c, out, sat.s.output, filter)
376 }
377 for _, w := range floats {
378 c := v / w
379 oneTemplate(t, "float", w, c, out, sat.s.output, filter)
380 }
381 }
382 }
383
384 func prologue(s, ba string, out io.Writer) {
385 fmt.Fprintf(out,
386 `// Code generated by '%s'; DO NOT EDIT.
387
388 //go:build goexperiment.simd
389
390 package archsimd
391
392 `, s)
393 }
394
395 func ssaPrologue(s string, out io.Writer) {
396 fmt.Fprintf(out,
397 `// Code generated by '%s'; DO NOT EDIT.
398
399 package ssa
400
401 `, s)
402 }
403
404 func unsafePrologue(s, ba string, out io.Writer) {
405 fmt.Fprintf(out,
406 `// Code generated by '%s'; DO NOT EDIT.
407
408 //go:build goexperiment.simd
409
410 package archsimd
411
412 import "unsafe"
413
414 `, s)
415 }
416
417 func testPrologue(t, s, ba string, out io.Writer) {
418 fmt.Fprintf(out,
419 `// Code generated by '%s'; DO NOT EDIT.
420
421 //go:build goexperiment.simd && %s
422
423 // This file contains functions testing %s.
424 // Each function in this file is specialized for a
425 // particular simd type <BaseType><Width>x<Count>.
426
427 package simd_test
428
429 import (
430 "simd/archsimd"
431 "testing"
432 )
433
434 `, s, ba, t)
435 }
436
437 func curryTestPrologue(t string) func(s, ba string, out io.Writer) {
438 return func(s, ba string, out io.Writer) {
439 testPrologue(t, s, ba, out)
440 }
441 }
442
443 func templateOf(name, temp string) shapeAndTemplate {
444 return shapeAndTemplate{s: allShapes,
445 t: template.Must(template.New(name).Parse(temp))}
446 }
447
448 func shapedTemplateOf(s *shapes, name, temp string) shapeAndTemplate {
449 return shapeAndTemplate{s: s,
450 t: template.Must(template.New(name).Parse(temp))}
451 }
452
453 const sliceTemplateText = `
454 // Load{{.VType}} loads {{.AOrAn}} {{.VType}} from a slice of elements.
455 // If s does not have at least {{.Count}} elements, it panics.
456 func Load{{.VType}}(s []{{.Etype}}) {{.VType}} {
457 return Load{{.VType}}Array((*[{{.Count}}]{{.Etype}})(s))
458 }
459
460 // Store stores the elements of x into a slice.
461 // If s does not have at least {{.Count}} elements, it panics.
462 func (x {{.VType}}) Store(s []{{.Etype}}) {
463 x.StoreArray((*[{{.Count}}]{{.Etype}})(s))
464 }
465 `
466
467 var sliceTemplate = templateOf("slice", sliceTemplateText)
468 var sliceTemplateArm64 = shapedTemplateOf(arm64Shapes, "arm64_slice", sliceTemplateText)
469
470 const unaryTestTemplate = `
471 // test{{.VType}}Unary tests the simd unary method f against the expected behavior generated by want
472 func test{{.VType}}Unary(t *testing.T, f func(_ archsimd.{{.VType}}) archsimd.{{.VType}}, want func(_ []{{.Etype}}) []{{.Etype}}) {
473 n := {{.Count}}
474 t.Helper()
475 forSlice(t, {{.Etype}}s, n, func(x []{{.Etype}}) bool {
476 t.Helper()
477 a := archsimd.Load{{.VType}}(x)
478 g := make([]{{.Etype}}, n)
479 f(a).Store(g)
480 w := want(x)
481 return checkSlicesLogInput(t, g, w, 0.0, func() {t.Helper(); t.Logf("x=%v", x)})
482 })
483 }
484 `
485
486 var unaryTemplate = templateOf("unary_helpers", unaryTestTemplate)
487 var unaryTemplateArm64 = shapedTemplateOf(arm64Shapes, "arm64_unary_helpers", unaryTestTemplate)
488
489 var unaryFlakyTemplate = shapedTemplateOf(unaryFlaky, "unary_flaky_helpers", `
490 // test{{.VType}}UnaryFlaky tests the simd unary method f against the expected behavior generated by want,
491 // but using a flakiness parameter because we haven't exactly figured out how simd floating point works
492 func test{{.VType}}UnaryFlaky(t *testing.T, f func(x archsimd.{{.VType}}) archsimd.{{.VType}}, want func(x []{{.Etype}}) []{{.Etype}}, flakiness float64) {
493 n := {{.Count}}
494 t.Helper()
495 forSlice(t, {{.Etype}}s, n, func(x []{{.Etype}}) bool {
496 t.Helper()
497 a := archsimd.Load{{.VType}}(x)
498 g := make([]{{.Etype}}, n)
499 f(a).Store(g)
500 w := want(x)
501 return checkSlicesLogInput(t, g, w, flakiness, func() {t.Helper(); t.Logf("x=%v", x)})
502 })
503 }
504 `)
505
506 var convertTemplate = templateOf("convert_helpers", `
507 // test{{.VType}}ConvertTo{{.OEType}} tests the simd conversion method f against the expected behavior generated by want.
508 // This is for count-preserving conversions, so if there is a change in size, then there is a change in vector width,
509 // (extended to at least 128 bits, or truncated to at most 512 bits).
510 func test{{.VType}}ConvertTo{{.OEType}}(t *testing.T, f func(x archsimd.{{.VType}}) archsimd.{{.OVType}}, want func(x []{{.Etype}}) []{{.OEtype}}) {
511 n := {{.Count}}
512 t.Helper()
513 forSlice(t, {{.Etype}}s, n, func(x []{{.Etype}}) bool {
514 t.Helper()
515 a := archsimd.Load{{.VType}}(x)
516 g := make([]{{.OEtype}}, {{.OCount}})
517 f(a).Store(g)
518 w := want(x)
519 return checkSlicesLogInput(t, g, w, 0.0, func() {t.Helper(); t.Logf("x=%v", x)})
520 })
521 }
522 `)
523
524 var (
525
526
527
528 unaryToInt8 = convertTemplate.target("int", 8)
529 unaryToUint8 = convertTemplate.target("uint", 8)
530 unaryToInt16 = convertTemplate.target("int", 16)
531 unaryToUint16 = convertTemplate.target("uint", 16)
532 unaryToInt32 = convertTemplate.target("int", 32)
533 unaryToUint32 = convertTemplate.target("uint", 32)
534 unaryToInt64 = convertTemplate.target("int", 64)
535 unaryToUint64 = convertTemplate.target("uint", 64)
536 unaryToFloat32 = convertTemplate.target("float", 32)
537 unaryToFloat64 = convertTemplate.target("float", 64)
538 )
539
540 var convertLoTemplate = shapedTemplateOf(allShapes, "convert_lo_helpers", `
541 // test{{.VType}}ConvertLoTo{{.OVType}} tests the simd conversion method f against the expected behavior generated by want.
542 // This converts only the low {{.OCount}} elements.
543 func test{{.VType}}ConvertLoTo{{.OVType}}(t *testing.T, f func(x archsimd.{{.VType}}) archsimd.{{.OVType}}, want func(x []{{.Etype}}) []{{.OEtype}}) {
544 n := {{.Count}}
545 t.Helper()
546 forSlice(t, {{.Etype}}s, n, func(x []{{.Etype}}) bool {
547 t.Helper()
548 a := archsimd.Load{{.VType}}(x)
549 g := make([]{{.OEtype}}, {{.OCount}})
550 f(a).Store(g)
551 w := want(x)
552 return checkSlicesLogInput(t, g, w, 0.0, func() {t.Helper(); t.Logf("x=%v", x)})
553 })
554 }
555 `)
556
557 var (
558
559
560
561
562
563 unaryToInt64x2 = convertLoTemplate.targetFixed("int", 64, 2)
564 unaryToInt64x4 = convertLoTemplate.targetFixed("int", 64, 4)
565 unaryToUint64x2 = convertLoTemplate.targetFixed("uint", 64, 2)
566 unaryToUint64x4 = convertLoTemplate.targetFixed("uint", 64, 4)
567 unaryToInt32x4 = convertLoTemplate.targetFixed("int", 32, 4)
568 unaryToInt32x8 = convertLoTemplate.targetFixed("int", 32, 8)
569 unaryToUint32x4 = convertLoTemplate.targetFixed("uint", 32, 4)
570 unaryToUint32x8 = convertLoTemplate.targetFixed("uint", 32, 8)
571 unaryToInt16x8 = convertLoTemplate.targetFixed("int", 16, 8)
572 unaryToUint16x8 = convertLoTemplate.targetFixed("uint", 16, 8)
573 unaryToFloat64x2 = convertLoTemplate.targetFixed("float", 64, 2)
574 unaryToFloat64x4 = convertLoTemplate.targetFixed("float", 64, 4)
575 )
576
577 const binaryTestTemplate = `
578 // test{{.VType}}Binary tests the simd binary method f against the expected behavior generated by want
579 func test{{.VType}}Binary(t *testing.T, f func(_, _ archsimd.{{.VType}}) archsimd.{{.VType}}, want func(_, _ []{{.Etype}}) []{{.Etype}}) {
580 n := {{.Count}}
581 t.Helper()
582 forSlicePair(t, {{.Etype}}s, n, func(x, y []{{.Etype}}) bool {
583 t.Helper()
584 a := archsimd.Load{{.VType}}(x)
585 b := archsimd.Load{{.VType}}(y)
586 g := make([]{{.Etype}}, n)
587 f(a, b).Store(g)
588 w := want(x, y)
589 return checkSlicesLogInput(t, g, w, 0.0, func() {t.Helper(); t.Logf("x=%v", x); t.Logf("y=%v", y); })
590 })
591 }
592 `
593
594 var binaryTemplate = templateOf("binary_helpers", binaryTestTemplate)
595 var binaryTemplateArm64 = shapedTemplateOf(arm64Shapes, "arm64_binary_helpers", binaryTestTemplate)
596
597
598
599 var shiftAllTestTemplate = shapedTemplateOf(integerShapes, "shift_all_helpers", `
600 // test{{.VType}}ShiftAll tests a shift-all method (unary + scalar uint64).
601 func test{{.VType}}ShiftAll(t *testing.T, f func(_ archsimd.{{.VType}}, _ uint64) archsimd.{{.VType}}, want func(_ []{{.Etype}}, _ uint64) []{{.Etype}}) {
602 n := {{.Count}}
603 t.Helper()
604 forSlice(t, {{.Etype}}s, n, func(x []{{.Etype}}) bool {
605 t.Helper()
606 for _, amt := range testShiftAllAmts {
607 a := archsimd.Load{{.VType}}(x)
608 g := make([]{{.Etype}}, n)
609 f(a, amt).Store(g)
610 w := want(x, amt)
611 if !checkSlicesLogInput(t, g, w, 0.0, func() { t.Helper(); t.Logf("x=%v, amt=%d", x, amt) }) {
612 return false
613 }
614 }
615 return true
616 })
617 }
618 `)
619
620 var shiftMixedTestTemplateArm64 = shapedTemplateOf(arm64UintToIntShapes, "arm64_shift_mixed_helpers", `
621 // test{{.VType}}Shift tests a shift-like method where the first operand is {{.VType}}
622 // and the second operand is {{.OVType}} (mixed-type shift).
623 func test{{.VType}}Shift(t *testing.T, f func(_ archsimd.{{.VType}}, _ archsimd.{{.OVType}}) archsimd.{{.VType}}, want func(_ []{{.Etype}}, _ []{{.OEtype}}) []{{.Etype}}) {
624 n := {{.Count}}
625 t.Helper()
626 forSliceMixed(t, {{.Etype}}s, {{.OEtype}}s, n, func(x []{{.Etype}}, y []{{.OEtype}}) bool {
627 t.Helper()
628 a := archsimd.Load{{.VType}}(x)
629 b := archsimd.Load{{.OVType}}(y)
630 g := make([]{{.Etype}}, n)
631 f(a, b).Store(g)
632 w := want(x, y)
633 return checkSlicesLogInput(t, g, w, 0.0, func() { t.Helper(); t.Logf("x=%v", x); t.Logf("y=%v", y) })
634 })
635 }
636 `)
637
638 const ternaryTestTemplateText = `
639 // test{{.VType}}Ternary tests the simd ternary method f against the expected behavior generated by want
640 func test{{.VType}}Ternary(t *testing.T, f func(_, _, _ archsimd.{{.VType}}) archsimd.{{.VType}}, want func(_, _, _ []{{.Etype}}) []{{.Etype}}) {
641 n := {{.Count}}
642 t.Helper()
643 forSliceTriple(t, {{.Etype}}s, n, func(x, y, z []{{.Etype}}) bool {
644 t.Helper()
645 a := archsimd.Load{{.VType}}(x)
646 b := archsimd.Load{{.VType}}(y)
647 c := archsimd.Load{{.VType}}(z)
648 g := make([]{{.Etype}}, n)
649 f(a, b, c).Store(g)
650 w := want(x, y, z)
651 return checkSlicesLogInput(t, g, w, 0.0, func() {t.Helper(); t.Logf("x=%v", x); t.Logf("y=%v", y); t.Logf("z=%v", z); })
652 })
653 }
654 `
655
656 const ternaryFlakyTestTemplateText = `
657 // test{{.VType}}TernaryFlaky tests the simd ternary method f against the expected behavior generated by want,
658 // but using a flakiness parameter because we haven't exactly figured out how simd floating point works
659 func test{{.VType}}TernaryFlaky(t *testing.T, f func(x, y, z archsimd.{{.VType}}) archsimd.{{.VType}}, want func(x, y, z []{{.Etype}}) []{{.Etype}}, flakiness float64) {
660 n := {{.Count}}
661 t.Helper()
662 forSliceTriple(t, {{.Etype}}s, n, func(x, y, z []{{.Etype}}) bool {
663 t.Helper()
664 a := archsimd.Load{{.VType}}(x)
665 b := archsimd.Load{{.VType}}(y)
666 c := archsimd.Load{{.VType}}(z)
667 g := make([]{{.Etype}}, n)
668 f(a, b, c).Store(g)
669 w := want(x, y, z)
670 return checkSlicesLogInput(t, g, w, flakiness, func() {t.Helper(); t.Logf("x=%v", x); t.Logf("y=%v", y); t.Logf("z=%v", z); })
671 })
672 }
673 `
674
675 var ternaryTemplate = templateOf("ternary_helpers", ternaryTestTemplateText)
676 var ternaryFlakyTemplate = shapedTemplateOf(ternaryFlaky, "ternary_helpers", ternaryFlakyTestTemplateText)
677
678 const reduceTestTemplateText = `
679 func test{{.VType}}Reduce(t *testing.T, f func(_ archsimd.{{.VType}}) {{.Etype}}, want func(_ []{{.Etype}}) {{.Etype}}) {
680 n := {{.Count}}
681 t.Helper()
682 forSlice(t, {{.Etype}}s, n, func(x []{{.Etype}}) bool {
683 t.Helper()
684 a := archsimd.Load{{.VType}}(x)
685 g := f(a)
686 w := want(x)
687 {{- if eq .Base "Float" }}
688 if g != w && !(math.IsNaN(float64(g)) && math.IsNaN(float64(w))) {
689 {{- else}}
690 if g != w {
691 {{- end}}
692 t.Errorf("got %v, want %v, input %v", g, w, x)
693 return false
694 }
695 return true
696 })
697 }
698 `
699
700 var reduceTestTemplateArm64 = shapedTemplateOf(arm64ReduceAllShapes, "reduce_arm64_helpers", reduceTestTemplateText)
701
702 func reduceTestPrologue(s, ba string, out io.Writer) {
703 fmt.Fprintf(out,
704 `// Code generated by '%s'; DO NOT EDIT.
705
706 //go:build goexperiment.simd && %s
707
708 // This file contains functions testing %s.
709 // Each function in this file is specialized for a
710 // particular simd type <BaseType><Width>x<Count>.
711
712 package simd_test
713
714 import (
715 "math"
716 "simd/archsimd"
717 "testing"
718 )
719 `, "tmplgen", ba, "simd reduce methods")
720 }
721
722 var compareTemplate = templateOf("compare_helpers", `
723 // test{{.VType}}Compare tests the simd comparison method f against the expected behavior generated by want
724 func test{{.VType}}Compare(t *testing.T, f func(_, _ archsimd.{{.VType}}) archsimd.Mask{{.WxC}}, want func(_, _ []{{.Etype}}) []int64) {
725 n := {{.Count}}
726 t.Helper()
727 forSlicePair(t, {{.Etype}}s, n, func(x, y []{{.Etype}}) bool {
728 t.Helper()
729 a := archsimd.Load{{.VType}}(x)
730 b := archsimd.Load{{.VType}}(y)
731 g := make([]int{{.EWidth}}, n)
732 f(a, b).ToInt{{.WxC}}().Store(g)
733 w := want(x, y)
734 return checkSlicesLogInput(t, s64(g), w, 0.0, func() {t.Helper(); t.Logf("x=%v", x); t.Logf("y=%v", y); })
735 })
736 }
737 `)
738
739 var compareUnaryTemplate = shapedTemplateOf(floatShapes, "compare_unary_helpers", `
740 // test{{.VType}}UnaryCompare tests the simd unary comparison method f against the expected behavior generated by want
741 func test{{.VType}}UnaryCompare(t *testing.T, f func(x archsimd.{{.VType}}) archsimd.Mask{{.WxC}}, want func(x []{{.Etype}}) []int64) {
742 n := {{.Count}}
743 t.Helper()
744 forSlice(t, {{.Etype}}s, n, func(x []{{.Etype}}) bool {
745 t.Helper()
746 a := archsimd.Load{{.VType}}(x)
747 g := make([]int{{.EWidth}}, n)
748 f(a).ToInt{{.WxC}}().Store(g)
749 w := want(x)
750 return checkSlicesLogInput(t, s64(g), w, 0.0, func() {t.Helper(); t.Logf("x=%v", x)})
751 })
752 }
753 `)
754
755
756 var compareMaskedTemplate = templateOf("comparemasked_helpers", `
757 // test{{.VType}}CompareMasked tests the simd masked comparison method f against the expected behavior generated by want
758 // The mask is applied to the output of want; anything not in the mask, is zeroed.
759 func test{{.VType}}CompareMasked(t *testing.T,
760 f func(_, _ archsimd.{{.VType}}, m archsimd.Mask{{.WxC}}) archsimd.Mask{{.WxC}},
761 want func(_, _ []{{.Etype}}) []int64) {
762 n := {{.Count}}
763 t.Helper()
764 forSlicePairMasked(t, {{.Etype}}s, n, func(x, y []{{.Etype}}, m []bool) bool {
765 t.Helper()
766 a := archsimd.Load{{.VType}}(x)
767 b := archsimd.Load{{.VType}}(y)
768 k := archsimd.LoadInt{{.WxC}}(toVect[int{{.EWidth}}](m)).ToMask()
769 g := make([]int{{.EWidth}}, n)
770 f(a, b, k).ToInt{{.WxC}}().Store(g)
771 w := want(x, y)
772 for i := range m {
773 if !m[i] {
774 w[i] = 0
775 }
776 }
777 return checkSlicesLogInput(t, s64(g), w, 0.0, func() {t.Helper(); t.Logf("x=%v", x); t.Logf("y=%v", y); t.Logf("m=%v", m); })
778 })
779 }
780 `)
781
782 var avx512MaskedLoadSliceTemplate = shapedTemplateOf(avx512Shapes, "avx 512 load slice part", `
783 // Load{{.VType}}Part loads a {{.VType}} from the slice s, it returns the loaded vector and the
784 // number of elements loaded.
785 // If s has fewer than {{.Count}} elements, the remaining elements of the vector are filled with zeroes.
786 // If s has {{.Count}} or more elements, the function is equivalent to Load{{.VType}}.
787 func Load{{.VType}}Part(s []{{.Etype}}) ({{.VType}}, int) {
788 l := len(s)
789 if l >= {{.Count}} {
790 return Load{{.VType}}(s), {{.Count}}
791 }
792 if l == 0 {
793 var x {{.VType}}
794 return x, 0
795 }
796 mask := Mask{{.WxC}}FromBits({{.OxFF}} >> ({{.Count}} - l))
797 return Load{{.VType}}Array(pa{{.VType}}(s)).Masked(mask), l
798 }
799
800 // StorePart stores the {{.Count}} elements of x into the slice s.
801 // It stores as many elements as will fit in s.
802 // If s has {{.Count}} or more elements, the method is equivalent to x.Store.
803 func (x {{.VType}}) StorePart(s []{{.Etype}}) int {
804 l := len(s)
805 if l >= {{.Count}} {
806 x.Store(s)
807 return {{.Count}}
808 }
809 if l == 0 {
810 return 0
811 }
812 mask := Mask{{.WxC}}FromBits({{.OxFF}} >> ({{.Count}} - l))
813 x.StoreArrayMasked(pa{{.VType}}(s), mask)
814 return l
815 }
816 `)
817
818 var avx2MaskedLoadSliceTemplate = shapedTemplateOf(avx2MaskedLoadShapes, "avx 2 load slice part", `
819 // Load{{.VType}}Part loads a {{.VType}} from the slice s, it returns the loaded vector and the
820 // number of elements loaded.
821 // If s has fewer than {{.Count}} elements, the remaining elements of the vector are filled with zeroes.
822 // If s has {{.Count}} or more elements, the function is equivalent to Load{{.VType}}.
823 func Load{{.VType}}Part(s []{{.Etype}}) ({{.VType}}, int) {
824 l := len(s)
825 if l >= {{.Count}} {
826 return Load{{.VType}}(s), {{.Count}}
827 }
828 if l == 0 {
829 var x {{.VType}}
830 return x, 0
831 }
832 mask := vecMask{{.EWidth}}[len(vecMask{{.EWidth}})/2-l:]
833 return Load{{.VType}}Array(pa{{.VType}}(s)).Masked(LoadInt{{.WxC}}(mask).asMask()), l
834 }
835
836 // StorePart stores the {{.Count}} elements of x into the slice s.
837 // It stores as many elements as will fit in s.
838 // If s has {{.Count}} or more elements, the method is equivalent to x.Store.
839 func (x {{.VType}}) StorePart(s []{{.Etype}}) int {
840 l := len(s)
841 if l >= {{.Count}} {
842 x.Store(s)
843 return {{.Count}}
844 }
845 if l == 0 {
846 return 0
847 }
848 mask := vecMask{{.EWidth}}[len(vecMask{{.EWidth}})/2-l:]
849 x.StoreArrayMasked(pa{{.VType}}(s), LoadInt{{.WxC}}(mask).asMask())
850 return l
851 }
852 `)
853
854 var avx2SmallLoadSliceTemplate = shapedTemplateOf(avx2SmallLoadPunShapes, "avx 2 small load slice part", `
855 // Load{{.VType}}Part loads a {{.VType}} from the slice s, it returns the loaded vector and the
856 // number of elements loaded.
857 // If s has fewer than {{.Count}} elements, the remaining elements of the vector are filled with zeroes.
858 // If s has {{.Count}} or more elements, the function is equivalent to Load{{.VType}}.
859 func Load{{.VType}}Part(s []{{.Etype}}) ({{.VType}}, int) {
860 if len(s) == 0 {
861 var zero {{.VType}}
862 return zero, 0
863 }
864 t := unsafe.Slice((*int{{.EWidth}})(unsafe.Pointer(&s[0])), len(s))
865 v, l := LoadInt{{.WxC}}Part(t)
866 return v.As{{.VType}}(), l
867 }
868
869 // StorePart stores the {{.Count}} elements of x into the slice s.
870 // It stores as many elements as will fit in s.
871 // If s has {{.Count}} or more elements, the method is equivalent to x.Store.
872 func (x {{.VType}}) StorePart(s []{{.Etype}}) int {
873 if len(s) == 0 {
874 return 0
875 }
876 t := unsafe.Slice((*int{{.EWidth}})(unsafe.Pointer(&s[0])), len(s))
877 return x.AsInt{{.WxC}}().StorePart(t)
878 }
879 `)
880
881 func (t templateData) CPUfeature() string {
882 switch t.Vwidth {
883 case 128:
884 return "AVX"
885 case 256:
886 return "AVX2"
887 case 512:
888 return "AVX512"
889 }
890 panic(fmt.Errorf("unexpected vector width %d", t.Vwidth))
891 }
892
893 var avx2SignedComparisonsTemplate = shapedTemplateOf(avx2SignedComparisons, "avx2 signed comparisons", `
894 // Less returns a mask whose elements indicate whether x < y.
895 //
896 // Emulated, CPU Feature: {{.CPUfeature}}
897 func (x {{.VType}}) Less(y {{.VType}}) Mask{{.WxC}} {
898 return y.Greater(x)
899 }
900
901 // GreaterEqual returns a mask whose elements indicate whether x >= y.
902 //
903 // Emulated, CPU Feature: {{.CPUfeature}}
904 func (x {{.VType}}) GreaterEqual(y {{.VType}}) Mask{{.WxC}} {
905 ones := x.Equal(x).ToInt{{.WxC}}()
906 return y.Greater(x).ToInt{{.WxC}}().Xor(ones).asMask()
907 }
908
909 // LessEqual returns a mask whose elements indicate whether x <= y.
910 //
911 // Emulated, CPU Feature: {{.CPUfeature}}
912 func (x {{.VType}}) LessEqual(y {{.VType}}) Mask{{.WxC}} {
913 ones := x.Equal(x).ToInt{{.WxC}}()
914 return x.Greater(y).ToInt{{.WxC}}().Xor(ones).asMask()
915 }
916
917 // NotEqual returns a mask whose elements indicate whether x != y.
918 //
919 // Emulated, CPU Feature: {{.CPUfeature}}
920 func (x {{.VType}}) NotEqual(y {{.VType}}) Mask{{.WxC}} {
921 ones := x.Equal(x).ToInt{{.WxC}}()
922 return x.Equal(y).ToInt{{.WxC}}().Xor(ones).asMask()
923 }
924 `)
925
926 var intRotateAllTemplate = sgutil.TemplateNamed("intRotateAll", `
927 // RotateAllLeft rotates all elements left by the specified amount
928 //
929 // Emulated
930 func (x {{.VType}}) RotateAllLeft(dist uint64) {{.VType}} {
931 dist = dist & ({{.EWidth}}-1)
932 ndist := {{.EWidth}} - dist
933 return x.ToBits().ShiftAllLeft(dist).Or(x.ToBits().ShiftAllRight(ndist)).BitsToInt{{.EWidth}}()
934 }
935
936 // RotateAllRight rotates all elements right by the specified amount
937 //
938 // Emulated
939 func (x {{.VType}}) RotateAllRight(dist uint64) {{.VType}} {
940 dist = dist & ({{.EWidth}}-1)
941 ndist := {{.EWidth}} - dist
942 return x.ToBits().ShiftAllLeft(ndist).Or(x.ToBits().ShiftAllRight(dist)).BitsToInt{{.EWidth}}()
943 }
944 `)
945
946 var uintRotateAllTemplate = sgutil.TemplateNamed("intRotateAll", `
947 // RotateAllLeft rotates all elements left by the specified amount
948 //
949 // Emulated
950 func (x {{.VType}}) RotateAllLeft(dist uint64) {{.VType}} {
951 dist = dist & ({{.EWidth}}-1)
952 ndist := {{.EWidth}} - dist
953 return x.ShiftAllLeft(dist).Or(x.ShiftAllRight(ndist))
954 }
955
956 // RotateAllRight rotates all elements right by the specified amount
957 //
958 // Emulated
959 func (x {{.VType}}) RotateAllRight(dist uint64) {{.VType}} {
960 dist = dist & ({{.EWidth}}-1)
961 ndist := {{.EWidth}} - dist
962 return x.ShiftAllLeft(ndist).Or(x.ShiftAllRight(dist))
963 }
964 `)
965
966 var bitWiseIntTemplate = shapedTemplateOf(intShapes, "bitwise int complement", `
967 // Not returns the bitwise complement of x.
968 //
969 // Emulated, CPU Feature: {{.CPUfeature}}
970 func (x {{.VType}}) Not() {{.VType}} {
971 return x.Xor(x.Equal(x).ToInt{{.WxC}}())
972 }
973
974 // Neg returns the element-wise negation of x.
975 //
976 // Emulated, CPU Feature: {{.CPUfeature}}
977 func (x {{.VType}}) Neg() {{.VType}} {
978 var zero {{.VType}}
979 return zero.Sub(x)
980 }
981
982 `)
983
984 var bitWiseUintTemplate = shapedTemplateOf(uintShapes, "bitwise uint complement", `
985 // Not returns the bitwise complement of x.
986 //
987 // Emulated, CPU Feature: {{.CPUfeature}}
988 func (x {{.VType}}) Not() {{.VType}} {
989 return x.Xor(x.Equal(x).ToInt{{.WxC}}().As{{.VType}}())
990 }
991 `)
992
993
994
995
996
997
998 func (t templateData) CPUfeatureAVX2if8() string {
999 if t.EWidth == 8 {
1000 return "AVX2"
1001 }
1002 return t.CPUfeature()
1003 }
1004
1005 var avx2UnsignedComparisonsTemplate = shapedTemplateOf(avx2UnsignedComparisons, "avx2 unsigned comparisons", `
1006 // Greater returns a mask whose elements indicate whether x > y.
1007 //
1008 // Emulated, CPU Feature: {{.CPUfeatureAVX2if8}}
1009 func (x {{.VType}}) Greater(y {{.VType}}) Mask{{.WxC}} {
1010 a, b := x.AsInt{{.WxC}}(), y.AsInt{{.WxC}}()
1011 {{- if eq .EWidth 8}}
1012 signs := BroadcastInt{{.WxC}}(-1 << ({{.EWidth}}-1))
1013 {{- else}}
1014 ones := x.Equal(x).ToInt{{.WxC}}()
1015 signs := ones.ShiftAllLeft({{.EWidth}}-1)
1016 {{- end }}
1017 return a.Xor(signs).Greater(b.Xor(signs))
1018 }
1019
1020 // Less returns a mask whose elements indicate whether x < y.
1021 //
1022 // Emulated, CPU Feature: {{.CPUfeatureAVX2if8}}
1023 func (x {{.VType}}) Less(y {{.VType}}) Mask{{.WxC}} {
1024 a, b := x.AsInt{{.WxC}}(), y.AsInt{{.WxC}}()
1025 {{- if eq .EWidth 8}}
1026 signs := BroadcastInt{{.WxC}}(-1 << ({{.EWidth}}-1))
1027 {{- else}}
1028 ones := x.Equal(x).ToInt{{.WxC}}()
1029 signs := ones.ShiftAllLeft({{.EWidth}}-1)
1030 {{- end }}
1031 return b.Xor(signs).Greater(a.Xor(signs))
1032 }
1033
1034 // GreaterEqual returns a mask whose elements indicate whether x >= y.
1035 //
1036 // Emulated, CPU Feature: {{.CPUfeatureAVX2if8}}
1037 func (x {{.VType}}) GreaterEqual(y {{.VType}}) Mask{{.WxC}} {
1038 a, b := x.AsInt{{.WxC}}(), y.AsInt{{.WxC}}()
1039 ones := x.Equal(x).ToInt{{.WxC}}()
1040 {{- if eq .EWidth 8}}
1041 signs := BroadcastInt{{.WxC}}(-1 << ({{.EWidth}}-1))
1042 {{- else}}
1043 signs := ones.ShiftAllLeft({{.EWidth}}-1)
1044 {{- end }}
1045 return b.Xor(signs).Greater(a.Xor(signs)).ToInt{{.WxC}}().Xor(ones).asMask()
1046 }
1047
1048 // LessEqual returns a mask whose elements indicate whether x <= y.
1049 //
1050 // Emulated, CPU Feature: {{.CPUfeatureAVX2if8}}
1051 func (x {{.VType}}) LessEqual(y {{.VType}}) Mask{{.WxC}} {
1052 a, b := x.AsInt{{.WxC}}(), y.AsInt{{.WxC}}()
1053 ones := x.Equal(x).ToInt{{.WxC}}()
1054 {{- if eq .EWidth 8}}
1055 signs := BroadcastInt{{.WxC}}(-1 << ({{.EWidth}}-1))
1056 {{- else}}
1057 signs := ones.ShiftAllLeft({{.EWidth}}-1)
1058 {{- end }}
1059 return a.Xor(signs).Greater(b.Xor(signs)).ToInt{{.WxC}}().Xor(ones).asMask()
1060 }
1061
1062 // NotEqual returns a mask whose elements indicate whether x != y.
1063 //
1064 // Emulated, CPU Feature: {{.CPUfeature}}
1065 func (x {{.VType}}) NotEqual(y {{.VType}}) Mask{{.WxC}} {
1066 a, b := x.AsInt{{.WxC}}(), y.AsInt{{.WxC}}()
1067 ones := x.Equal(x).ToInt{{.WxC}}()
1068 return a.Equal(b).ToInt{{.WxC}}().Xor(ones).asMask()
1069 }
1070 `)
1071
1072 var unsafePATemplate = templateOf("unsafe PA helper", `
1073 // pa{{.VType}} returns a type-unsafe pointer to array that can
1074 // only be used with partial load/store operations that only
1075 // access the known-safe portions of the array.
1076 //
1077 //go:nocheckptr
1078 func pa{{.VType}}(s []{{.Etype}}) *[{{.Count}}]{{.Etype}} {
1079 return (*[{{.Count}}]{{.Etype}})(unsafe.Pointer(&s[0]))
1080 }
1081 `)
1082
1083 var avx2MaskedTemplate = shapedTemplateOf(avx2Shapes, "avx2 .Masked methods", `
1084 // Masked returns x but with elements zeroed where mask is false.
1085 //
1086 // Emulated, CPU Feature: {{.CPUfeature}}
1087 func (x {{.VType}}) Masked(mask Mask{{.WxC}}) {{.VType}} {
1088 im := mask.ToInt{{.WxC}}()
1089 {{- if eq .Base "Int" }}
1090 return im.And(x)
1091 {{- else}}
1092 return x.AsInt{{.WxC}}().And(im).As{{.VType}}()
1093 {{- end -}}
1094 }
1095
1096 // Merge returns x but with elements set to y where mask is false.
1097 //
1098 // Emulated, CPU Feature: {{.CPUfeature}}
1099 //
1100 // Deprecated: use x.IfElse(mask, y)
1101 //
1102 //go:fix inline
1103 func (x {{.VType}}) Merge(y {{.VType}}, mask Mask{{.WxC}}) {{.VType}} {
1104 return x.IfElse(mask, y)
1105 }
1106
1107 // IfElse returns x but with elements set to y where mask is false.
1108 //
1109 // Emulated, CPU Feature: {{.CPUfeature}}
1110 func (x {{.VType}}) IfElse(mask Mask{{.WxC}}, y {{.VType}}) {{.VType}} {
1111 {{- if eq .BxC .WxC -}}
1112 im := mask.ToInt{{.BxC}}()
1113 {{- else}}
1114 im := mask.ToInt{{.WxC}}().AsInt{{.BxC}}()
1115 {{- end -}}
1116 {{- if and (eq .Base "Int") (eq .BxC .WxC) }}
1117 return y.blend(x, im)
1118 {{- else}}
1119 ix := x.AsInt{{.BxC}}()
1120 iy := y.AsInt{{.BxC}}()
1121 return iy.blend(ix, im).As{{.VType}}()
1122 {{- end -}}
1123 }
1124 `)
1125
1126
1127 var avx512MaskedTemplate = shapedTemplateOf(avx512Shapes, "avx512 .Masked methods", `
1128 // Masked returns x but with elements zeroed where mask is false.
1129 //
1130 // Emulated, CPU Feature: AVX512
1131 func (x {{.VType}}) Masked(mask Mask{{.WxC}}) {{.VType}} {
1132 im := mask.ToInt{{.WxC}}()
1133 {{- if eq .Base "Int" }}
1134 return im.And(x)
1135 {{- else}}
1136 return x.AsInt{{.WxC}}().And(im).As{{.VType}}()
1137 {{- end -}}
1138 }
1139
1140 // Merge returns x but with elements set to y where mask is false.
1141 //
1142 // Emulated, CPU Feature: AVX512
1143 //
1144 // Deprecated: use x.IfElse(mask, y)
1145 //
1146 //go:fix inline
1147 func (x {{.VType}}) Merge(y {{.VType}}, mask Mask{{.WxC}}) {{.VType}} {
1148 return x.IfElse(mask, y)
1149 }
1150
1151 // IfElse returns x but with elements set to y where mask is false.
1152 //
1153 // Emulated, CPU Feature: AVX512
1154 func (x {{.VType}}) IfElse(mask Mask{{.WxC}}, y {{.VType}}) {{.VType}} {
1155 {{- if eq .Base "Int" }}
1156 return y.blendMasked(x, mask)
1157 {{- else}}
1158 ix := x.AsInt{{.WxC}}()
1159 iy := y.AsInt{{.WxC}}()
1160 return iy.blendMasked(ix, mask).As{{.VType}}()
1161 {{- end -}}
1162 }
1163 `)
1164
1165 func (t templateData) CPUfeatureBC() string {
1166 switch t.Vwidth {
1167 case 128:
1168 return "AVX2"
1169 case 256:
1170 return "AVX2"
1171 case 512:
1172 if t.EWidth <= 16 {
1173 return "AVX512BW"
1174 }
1175 return "AVX512F"
1176 }
1177 panic(fmt.Errorf("unexpected vector width %d", t.Vwidth))
1178 }
1179
1180 var broadcastTemplate = templateOf("Broadcast functions", `
1181 // Broadcast{{.VType}} returns a vector with the input
1182 // x assigned to all elements of the output.
1183 //
1184 // Emulated, CPU Feature: {{.CPUfeatureBC}}
1185 func Broadcast{{.VType}}(x {{.Etype}}) {{.VType}} {
1186 var z {{.As128BitVec }}
1187 return z.SetElem(0, x).broadcast1To{{.Count}}()
1188 }
1189 `)
1190
1191 var broadcastTemplateArm64 = shapedTemplateOf(arm64Shapes, "arm64_broadcast", `
1192 // Broadcast{{.VType}} returns a vector with the input
1193 // x assigned to all elements of the output.
1194 func Broadcast{{.VType}}(x {{.Etype}}) {{.VType}} {
1195 var z {{.VType}}
1196 return z.SetElem(0, x).broadcast1To{{.Count}}()
1197 }
1198 `)
1199
1200 var stringTemplateArm64 = shapedTemplateOf(arm64Shapes, "arm64_String methods", `
1201 // String returns a string representation of SIMD vector x.
1202 func (x {{.VType}}) String() string {
1203 var s [{{.Count}}]{{.Etype}}
1204 x.StoreArray(&s)
1205 return sliceToString(s[:])
1206 }
1207 `)
1208
1209 var getHiTemplateArm64 = shapedTemplateOf(arm64Shapes, "arm64_HiToLo methods", `
1210 // HiToLo returns a vector with the upper 64 bits zeroed and the lower
1211 // 64 bits replaced with the upper 64 bits of x.
1212 func (x {{.VType}}) HiToLo() {{.VType}} {
1213 var z {{.VType}}
1214 {{- if and (eq .Base "Float") (eq .EWidth 64)}}
1215 return z.SetElem(0, x.GetElem(1))
1216 {{- else if (eq .EWidth 64)}}
1217 {{- if (eq .Base "Uint")}}
1218 return z.BitsToFloat64().SetElem(0, x.BitsToFloat64().GetElem(1)).ToBits()
1219 {{- else}}
1220 return z.ToBits().BitsToFloat64().SetElem(0, x.ToBits().BitsToFloat64().GetElem(1)).ToBits().BitsTo{{.Base}}{{.EWidth}}()
1221 {{- end}}
1222 {{- else}}
1223 {{- if (eq .Base "Uint")}}
1224 return z.ReshapeToUint64s().BitsToFloat64().SetElem(0, x.ReshapeToUint64s().BitsToFloat64().GetElem(1)).ToBits().ReshapeToUint{{.EWidth}}s()
1225 {{- else}}
1226 return z.ToBits().ReshapeToUint64s().BitsToFloat64().SetElem(0, x.ToBits().ReshapeToUint64s().BitsToFloat64().GetElem(1)).ToBits().ReshapeToUint{{.EWidth}}s().BitsTo{{.Base}}{{.EWidth}}()
1227 {{- end}}
1228 {{- end}}
1229 }
1230 `)
1231
1232 var reduceSumTemplateArm64 = shapedTemplateOf(arm64ReduceIntegerShapes, "arm64_ReduceSum methods", `
1233 // ReduceSum reduces x by summing all elements.
1234 //
1235 // Emulated, CPU Feature: NEON
1236 func (x {{.VType}}) ReduceSum() {{.Etype}} {
1237 return x.reduceSum().GetElem(0)
1238 }
1239 `)
1240
1241 var reduceMinMaxTemplateArm64 = shapedTemplateOf(arm64ReduceAllShapes, "arm64_ReduceMax/Min methods", `
1242 // ReduceMax reduces x by taking the maximum of all elements.
1243 //
1244 // Emulated, CPU Feature: NEON
1245 func (x {{.VType}}) ReduceMax() {{.Etype}} {
1246 return x.reduceMax().GetElem(0)
1247 }
1248
1249 // ReduceMin reduces x by taking the minimum of all elements.
1250 //
1251 // Emulated, CPU Feature: NEON
1252 func (x {{.VType}}) ReduceMin() {{.Etype}} {
1253 return x.reduceMin().GetElem(0)
1254 }
1255 `)
1256
1257 var maskCvtTemplate = shapedTemplateOf(intShapes, "Mask conversions", `
1258 // ToMask returns a mask whose i'th element is set if x[i] is non-zero.
1259 func (from {{.Base}}{{.WxC}}) ToMask() (to Mask{{.WxC}}) {
1260 return from.NotEqual({{.Base}}{{.WxC}}{})
1261 }
1262 `)
1263
1264 var arm64MaskCvtTemplate = shapedTemplateOf(arm64IntShapes, "Mask conversions", `
1265 // ToMask returns a mask whose i'th element is set if x[i] is non-zero.
1266 func (from {{.Base}}{{.WxC}}) ToMask() (to Mask{{.WxC}}) {
1267 return from.NotEqual({{.Base}}{{.WxC}}{})
1268 }
1269 `)
1270
1271
1272
1273
1274
1275 var arm64LessTemplate = shapedTemplateOf(arm64Shapes, "arm64_less", `
1276 // Less returns a mask whose elements indicate whether x < y.
1277 func (x {{.VType}}) Less(y {{.VType}}) Mask{{.WxC}} {
1278 return y.Greater(x)
1279 }
1280 `)
1281
1282 var arm64LessEqualTemplate = shapedTemplateOf(arm64Shapes, "arm64_less_equal", `
1283 // LessEqual returns a mask whose elements indicate whether x <= y.
1284 func (x {{.VType}}) LessEqual(y {{.VType}}) Mask{{.WxC}} {
1285 return y.GreaterEqual(x)
1286 }
1287 `)
1288
1289 var arm64NotEqualTemplate = shapedTemplateOf(arm64Shapes, "arm64_not_equal", `
1290 // NotEqual returns a mask whose elements indicate whether x != y.
1291 func (x {{.VType}}) NotEqual(y {{.VType}}) Mask{{.WxC}} {
1292 return x.Equal(y).Not()
1293 }
1294 `)
1295
1296
1297
1298 var arm64MaskedMergeTemplate = shapedTemplateOf(arm64Shapes, "arm64_masked_merge", `
1299 // Masked returns x but with elements zeroed where mask is false.
1300 func (x {{.VType}}) Masked(mask Mask{{.WxC}}) {{.VType}} {
1301 im := mask.ToInt{{.WxC}}()
1302 {{- if eq .Base "Int" }}
1303 return im.And(x)
1304 {{- else if eq .Base "Uint" }}
1305 return im.And(x.BitsToInt{{.EWidth}}()).ToBits()
1306 {{- else }}
1307 return im.And(x.ToBits().BitsToInt{{.EWidth}}()).ToBits().BitsTo{{.Base}}{{.EWidth}}()
1308 {{- end }}
1309 }
1310
1311 // IfElse returns x but with elements set to y where mask is false.
1312 func (x {{.VType}}) IfElse(mask Mask{{.WxC}}, y {{.VType}}) {{.VType}} {
1313 {{- if eq .WxC "8x16" }}
1314 {{- if eq .Base "Int" }}
1315 return x.bitSelect(y, mask.ToInt8x16())
1316 {{- else if eq .Base "Uint" }}
1317 return x.BitsToInt8().bitSelect(y.BitsToInt8(), mask.ToInt8x16()).ToBits()
1318 {{- else }}
1319 return x.ToBits().BitsToInt8().bitSelect(y.ToBits().BitsToInt8(), mask.ToInt8x16()).ToBits().BitsTo{{.Base}}{{.EWidth}}()
1320 {{- end }}
1321 {{- else if eq .Base "Uint" }}
1322 im := mask.ToInt{{.WxC}}().ToBits().ReshapeToUint8s().BitsToInt8()
1323 ix := x.ReshapeToUint8s().BitsToInt8()
1324 iy := y.ReshapeToUint8s().BitsToInt8()
1325 return ix.bitSelect(iy, im).ToBits().ReshapeToUint{{.EWidth}}s()
1326 {{- else }}
1327 im := mask.ToInt{{.WxC}}().ToBits().ReshapeToUint8s().BitsToInt8()
1328 ix := x.ToBits().ReshapeToUint8s().BitsToInt8()
1329 iy := y.ToBits().ReshapeToUint8s().BitsToInt8()
1330 return ix.bitSelect(iy, im).ToBits().ReshapeToUint{{.EWidth}}s().BitsTo{{.Base}}{{.EWidth}}()
1331 {{- end }}
1332 }
1333 `)
1334
1335 var compareTemplateArm64 = shapedTemplateOf(arm64Shapes, "arm64_compare_helpers", `
1336 // test{{.VType}}Compare tests the simd comparison method f against the expected behavior generated by want
1337 func test{{.VType}}Compare(t *testing.T, f func(_, _ archsimd.{{.VType}}) archsimd.Mask{{.WxC}}, want func(_, _ []{{.Etype}}) []int64) {
1338 n := {{.Count}}
1339 t.Helper()
1340 forSlicePair(t, {{.Etype}}s, n, func(x, y []{{.Etype}}) bool {
1341 t.Helper()
1342 a := archsimd.Load{{.VType}}(x)
1343 b := archsimd.Load{{.VType}}(y)
1344 g := make([]int{{.EWidth}}, n)
1345 f(a, b).ToInt{{.WxC}}().Store(g)
1346 w := want(x, y)
1347 return checkSlicesLogInput(t, s64(g), w, 0.0, func() {t.Helper(); t.Logf("x=%v", x); t.Logf("y=%v", y); })
1348 })
1349 }
1350 `)
1351
1352 var arm64MaskToString = shapedTemplateOf(arm64IntShapes, "arm64_maskToString", `
1353 // String returns a string representation of SIMD mask x.
1354 func (x Mask{{.WxC}}) String() string {
1355 var s [{{.Count}}]{{.Etype}}
1356 x.ToInt{{.WxC}}().Neg().StoreArray(&s)
1357 return sliceToString(s[:])
1358 }
1359 `)
1360
1361 var stringTemplate = shapedTemplateOf(allShapes, "String methods", `
1362 // String returns a string representation of SIMD vector x.
1363 func (x {{.VType}}) String() string {
1364 var s [{{.Count}}]{{.Etype}}
1365 x.StoreArray(&s)
1366 return sliceToString(s[:])
1367 }
1368 `)
1369
1370 var maskToString = shapedTemplateOf(intShapes, "maskToString", `
1371 // String returns a string representation of SIMD mask x.
1372 func (x Mask{{.WxC}}) String() string {
1373 var s [{{.Count}}]{{.Etype}}
1374 x.ToInt{{.WxC}}().Neg().StoreArray(&s)
1375 return sliceToString(s[:])
1376 }
1377 `)
1378
1379 const SIMD = "../../"
1380 const TD = "../../internal/simd_test/"
1381 const SSA = "../../../../cmd/compile/internal/ssa/"
1382
1383 func main() {
1384 sl := flag.String("sl", SIMD+"slice_gen_amd64.go", "file name for slice operations")
1385 cm := flag.String("cm", SIMD+"compare_gen_amd64.go", "file name for comparison operations")
1386 mm := flag.String("mm", SIMD+"maskmerge_gen_amd64.go", "file name for mask/merge operations")
1387 op := flag.String("op", SIMD+"other_gen_amd64.go", "file name for other operations")
1388 ush := flag.String("ush", SIMD+"unsafe_helpers.go", "file name for unsafe helpers")
1389 bh := flag.String("bh", TD+"binary_helpers_%W_test.go", "file name for binary test helpers")
1390 uh := flag.String("uh", TD+"unary_helpers_%W_test.go", "file name for unary test helpers")
1391 cvh := flag.String("cvh", TD+"convert_helpers_%W_test.go", "file name for conversion test helpers")
1392 th := flag.String("th", TD+"ternary_helpers_%W_test.go", "file name for ternary test helpers")
1393 ch := flag.String("ch", TD+"compare_helpers_%W_test.go", "file name for compare test helpers")
1394 cmh := flag.String("cmh", TD+"comparemasked_helpers_test.go", "file name for compare-masked test helpers")
1395 sh := flag.String("sh", TD+"shift_helpers_%W_test.go", "file name for shift test helpers")
1396
1397 slArm64 := flag.String("slArm64", SIMD+"slice_gen_arm64.go", "file name for ARM64 slice operations")
1398 opArm64 := flag.String("opArm64", SIMD+"other_gen_arm64.go", "file name for ARM64 other operations")
1399 shArm64 := flag.String("shArm64", TD+"shift_helpers_arm64_test.go", "file name for ARM64 shift test helpers")
1400 cmArm64 := flag.String("cmArm64", SIMD+"compare_gen_arm64.go", "file name for ARM64 comparison operations")
1401 mmArm64 := flag.String("mmArm64", SIMD+"maskmerge_gen_arm64.go", "file name for ARM64 mask/merge operations")
1402 rhArm64 := flag.String("rhArm64", TD+"reduce_helpers_arm64_test.go", "file name for ARM64 reduce test helpers")
1403 flag.Parse()
1404
1405 if *sl != "" {
1406 one(*sl, unsafePrologue,
1407 sliceTemplate,
1408 avx512MaskedLoadSliceTemplate,
1409 avx2MaskedLoadSliceTemplate,
1410 avx2SmallLoadSliceTemplate,
1411 )
1412 }
1413 if *cm != "" {
1414 one(*cm, prologue,
1415 avx2SignedComparisonsTemplate,
1416 avx2UnsignedComparisonsTemplate,
1417 )
1418 }
1419 if *mm != "" {
1420 one(*mm, prologue,
1421 avx2MaskedTemplate,
1422 avx512MaskedTemplate,
1423 )
1424 }
1425 if *op != "" {
1426 one(*op, prologue,
1427 broadcastTemplate,
1428 maskCvtTemplate,
1429 bitWiseIntTemplate,
1430 bitWiseUintTemplate,
1431 stringTemplate,
1432 maskToString,
1433 shapeAndTemplate{amdIntShiftAllShapes, intRotateAllTemplate},
1434 shapeAndTemplate{amdUintShiftAllShapes, uintRotateAllTemplate},
1435 )
1436 }
1437 if *ush != "" {
1438 one(*ush, unsafePrologue, unsafePATemplate)
1439 }
1440 if *uh != "" {
1441 one(*uh, curryTestPrologue("unary simd methods"), unaryTemplate)
1442 }
1443 if *cvh != "" {
1444 one(*cvh, curryTestPrologue("conversion simd methods"),
1445 unaryToInt8, unaryToUint8, unaryToInt16, unaryToUint16,
1446 unaryToInt32, unaryToUint32, unaryToInt64, unaryToUint64,
1447 unaryToFloat32, unaryToFloat64,
1448 unaryToInt64x2, unaryToInt64x4,
1449 unaryToUint64x2, unaryToUint64x4,
1450 unaryToInt32x4, unaryToInt32x8,
1451 unaryToUint32x4, unaryToUint32x8,
1452 unaryToInt16x8, unaryToUint16x8,
1453 unaryToFloat64x2, unaryToFloat64x4,
1454 unaryFlakyTemplate,
1455 )
1456 }
1457 if *bh != "" {
1458 one(*bh, curryTestPrologue("binary simd methods"), binaryTemplate)
1459 }
1460 if *th != "" {
1461 one(*th, curryTestPrologue("ternary simd methods"), ternaryTemplate, ternaryFlakyTemplate)
1462 }
1463 if *ch != "" {
1464 one(*ch, curryTestPrologue("simd methods that compare two operands"), compareTemplate, compareUnaryTemplate)
1465 }
1466 if *cmh != "" {
1467 one(*cmh, curryTestPrologue("simd methods that compare two operands under a mask"), compareMaskedTemplate)
1468 }
1469 if *sh != "" {
1470 one(*sh, curryTestPrologue("shift simd methods"),
1471 shiftAllTestTemplate,
1472 )
1473 }
1474
1475
1476 if *slArm64 != "" {
1477 one(*slArm64, prologue, sliceTemplateArm64)
1478 }
1479 if *opArm64 != "" {
1480 one(*opArm64, prologue,
1481 broadcastTemplateArm64,
1482 stringTemplateArm64,
1483 getHiTemplateArm64,
1484 arm64MaskCvtTemplate,
1485 shapeAndTemplate{neonIntShiftAllShapes, intRotateAllTemplate},
1486 shapeAndTemplate{neonUintShiftAllShapes, uintRotateAllTemplate},
1487 reduceSumTemplateArm64,
1488 reduceMinMaxTemplateArm64)
1489 }
1490 if *shArm64 != "" {
1491 oneArch(*shArm64, "arm64", curryTestPrologue("shift simd methods"), filterAll,
1492 shiftMixedTestTemplateArm64,
1493 )
1494 }
1495 if *cmArm64 != "" {
1496 one(*cmArm64, prologue,
1497 arm64LessTemplate,
1498 arm64LessEqualTemplate,
1499 arm64NotEqualTemplate,
1500 )
1501 }
1502 if *mmArm64 != "" {
1503 one(*mmArm64, prologue,
1504 arm64MaskedMergeTemplate,
1505 arm64MaskToString,
1506 )
1507 }
1508 if *rhArm64 != "" {
1509 oneArch(*rhArm64, "arm64", reduceTestPrologue, filterAll, reduceTestTemplateArm64)
1510 }
1511
1512 nonTemplateRewrites(SSA+"tern_helpers.go", ssaPrologue, classifyBooleanSIMD, ternOpForLogical)
1513
1514 }
1515
1516 func ternOpForLogical(out io.Writer) {
1517 fmt.Fprintf(out, `
1518 func ternOpForLogical(op Op) Op {
1519 switch op {
1520 `)
1521
1522 intShapes.forAllShapes(func(seq int, t, upperT string, w, c int, out io.Writer) {
1523 wt, ct := w, c
1524 if wt < 32 {
1525 wt = 32
1526 ct = (w * c) / wt
1527 }
1528 fmt.Fprintf(out, "case OpAndInt%[1]dx%[2]d, OpOrInt%[1]dx%[2]d, OpXorInt%[1]dx%[2]d,OpAndNotInt%[1]dx%[2]d: return OpternInt%dx%d\n", w, c, wt, ct)
1529 fmt.Fprintf(out, "case OpAndUint%[1]dx%[2]d, OpOrUint%[1]dx%[2]d, OpXorUint%[1]dx%[2]d,OpAndNotUint%[1]dx%[2]d: return OpternUint%dx%d\n", w, c, wt, ct)
1530 }, out)
1531
1532 fmt.Fprintf(out, `
1533 }
1534 return op
1535 }
1536 `)
1537
1538 }
1539
1540 func classifyBooleanSIMD(out io.Writer) {
1541 fmt.Fprintf(out, `
1542 type SIMDLogicalOP uint8
1543 const (
1544 // boolean simd operations, for reducing expression to VPTERNLOG* instructions
1545 // sloInterior is set for non-root nodes in logical-op expression trees.
1546 // the operations are even-numbered.
1547 sloInterior SIMDLogicalOP = 1
1548 sloNone SIMDLogicalOP = 2 * iota
1549 sloAnd
1550 sloOr
1551 sloAndNot
1552 sloXor
1553 sloNot
1554 )
1555 func classifyBooleanSIMD(v *Value) SIMDLogicalOP {
1556 switch v.Op {
1557 case `)
1558 intShapes.forAllShapes(func(seq int, t, upperT string, w, c int, out io.Writer) {
1559 op := "And"
1560 if seq > 0 {
1561 fmt.Fprintf(out, ",Op%s%s%dx%d", op, upperT, w, c)
1562 } else {
1563 fmt.Fprintf(out, "Op%s%s%dx%d", op, upperT, w, c)
1564 }
1565 seq++
1566 }, out)
1567
1568 fmt.Fprintf(out, `:
1569 return sloAnd
1570
1571 case `)
1572 intShapes.forAllShapes(func(seq int, t, upperT string, w, c int, out io.Writer) {
1573 op := "Or"
1574 if seq > 0 {
1575 fmt.Fprintf(out, ",Op%s%s%dx%d", op, upperT, w, c)
1576 } else {
1577 fmt.Fprintf(out, "Op%s%s%dx%d", op, upperT, w, c)
1578 }
1579 seq++
1580 }, out)
1581
1582 fmt.Fprintf(out, `:
1583 return sloOr
1584
1585 case `)
1586 intShapes.forAllShapes(func(seq int, t, upperT string, w, c int, out io.Writer) {
1587 op := "AndNot"
1588 if seq > 0 {
1589 fmt.Fprintf(out, ",Op%s%s%dx%d", op, upperT, w, c)
1590 } else {
1591 fmt.Fprintf(out, "Op%s%s%dx%d", op, upperT, w, c)
1592 }
1593 seq++
1594 }, out)
1595
1596 fmt.Fprintf(out, `:
1597 return sloAndNot
1598 `)
1599
1600
1601
1602
1603
1604 intShapes.forAllShapes(
1605 func(seq int, t, upperT string, w, c int, out io.Writer) {
1606 fmt.Fprintf(out, "case OpXor%s%dx%d: ", upperT, w, c)
1607 fmt.Fprintf(out, `
1608 if y := v.Args[1]; y.Op == OpEqual%s%dx%d &&
1609 y.Args[0] == y.Args[1] {
1610 return sloNot
1611 }
1612 `, upperT, w, c)
1613 fmt.Fprintf(out, "return sloXor\n")
1614 }, out)
1615
1616 fmt.Fprintf(out, `
1617 }
1618 return sloNone
1619 }
1620 `)
1621 }
1622
1623
1624
1625 func numberLines(data []byte) string {
1626 var buf bytes.Buffer
1627 r := bytes.NewReader(data)
1628 s := bufio.NewScanner(r)
1629 for i := 1; s.Scan(); i++ {
1630 fmt.Fprintf(&buf, "%d: %s\n", i, s.Text())
1631 }
1632 return buf.String()
1633 }
1634
1635 func nonTemplateRewrites(filename string, prologue func(s string, out io.Writer), rewrites ...func(out io.Writer)) {
1636 if filename == "" {
1637 return
1638 }
1639
1640 ofile := os.Stdout
1641
1642 if filename != "-" {
1643 var err error
1644 ofile, err = os.Create(filename)
1645 if err != nil {
1646 fmt.Fprintf(os.Stderr, "Could not create the output file %s for the generated code, %v", filename, err)
1647 os.Exit(1)
1648 }
1649 }
1650
1651 out := new(bytes.Buffer)
1652
1653 prologue("tmplgen", out)
1654 for _, rewrite := range rewrites {
1655 rewrite(out)
1656 }
1657
1658 b, err := format.Source(out.Bytes())
1659 if err != nil {
1660 fmt.Fprintf(os.Stderr, "There was a problem formatting the generated code for %s, %v\n", filename, err)
1661 fmt.Fprintf(os.Stderr, "%s\n", numberLines(out.Bytes()))
1662 fmt.Fprintf(os.Stderr, "There was a problem formatting the generated code for %s, %v\n", filename, err)
1663 os.Exit(1)
1664 } else {
1665 ofile.Write(b)
1666 ofile.Close()
1667 }
1668
1669 }
1670
1671 func one(filename string, prologue func(s, buildArch string, out io.Writer), sats ...shapeAndTemplate) {
1672 if filename == "" {
1673 return
1674 }
1675
1676 if strings.Contains(filename, "%W") {
1677 smallFile := strings.ReplaceAll(filename, "%W", "128")
1678 largeFile := strings.ReplaceAll(filename, "%W", "wider")
1679 oneArch(smallFile, "(amd64 || wasm || arm64)", prologue, filterSmallOnly, sats...)
1680 oneArch(largeFile, "amd64", prologue, filterLarge, sats...)
1681 return
1682 }
1683 oneArch(filename, "amd64", prologue, filterAll, sats...)
1684 }
1685
1686 func oneArch(filename, buildArch string, prologue func(s, buildArch string, out io.Writer), filter shapeFilter, sats ...shapeAndTemplate) {
1687
1688 ofile := os.Stdout
1689
1690 if filename != "-" {
1691 var err error
1692 ofile, err = os.Create(filename)
1693 if err != nil {
1694 fmt.Fprintf(os.Stderr, "Could not create the output file %s for the generated code, %v", filename, err)
1695 os.Exit(1)
1696 }
1697 }
1698
1699 out := new(bytes.Buffer)
1700
1701 prologue("tmplgen", buildArch, out)
1702 for _, sat := range sats {
1703 sat.forTemplates(out, filter)
1704 }
1705
1706 b, err := format.Source(out.Bytes())
1707 if err != nil {
1708 fmt.Fprintf(os.Stderr, "There was a problem formatting the generated code for %s, %v\n", filename, err)
1709 fmt.Fprintf(os.Stderr, "%s\n", numberLines(out.Bytes()))
1710 fmt.Fprintf(os.Stderr, "There was a problem formatting the generated code for %s, %v\n", filename, err)
1711 os.Exit(1)
1712 } else {
1713 ofile.Write(b)
1714 ofile.Close()
1715 }
1716
1717 }
1718
View as plain text