1
2
3
4
5 package main
6
7 import (
8 "bytes"
9 "fmt"
10 "slices"
11 "strings"
12 "text/template"
13 "unicode"
14
15 "simd/archsimd/_gen/sgutil"
16 )
17
18 type tplRuleData struct {
19 TplName string
20 GoOp string
21 GoType string
22 Args string
23 Asm string
24 ArgsOut string
25 MaskInConvert string
26 MaskOutConvert string
27 ElementSize int
28 Size int
29 ArgsLoadAddr string
30 ArgsAddr string
31 FeatCheck string
32 RuleCond string
33 RuleOut string
34 RuleArgs string
35 Rule string
36 }
37
38
39
40 type ruleTemplateMap struct {
41 sgutil.InsertMap[string, *template.Template]
42 }
43
44
45
46
47 func (rtm *ruleTemplateMap) Add(name string, templ string) *ruleTemplateMap {
48
49 templ += " // {{.TplName}}\n"
50 ct := sgutil.TemplateNamed(name, templ)
51 rtm.InsertMap.Put(name, ct)
52 return rtm
53 }
54
55 var (
56
57
58
59
60
61 ruleTemplates = new(ruleTemplateMap).
62 Add("masksftimm", `({{.Asm}} x (MOVQconst [c]) mask) => ({{.Asm}}const [amd64CapAVXShift(c)] x mask)`).
63 Add("sftimm", `({{.Asm}} x (MOVQconst [c])) => ({{.Asm}}const [amd64CapAVXShift(c)] x)`).
64 Add("maskInMaskOut", `({{.GoOp}}{{.GoType}} {{.Args}} mask) => ({{.MaskOutConvert}} ({{.Asm}} {{.ArgsOut}} ({{.MaskInConvert}} <types.TypeMask> mask)))`).
65 Add("maskOut", `({{.GoOp}}{{.GoType}} {{.Args}}) => ({{.MaskOutConvert}} ({{.Asm}} {{.ArgsOut}}))`).
66 Add("maskIn", `({{.GoOp}}{{.GoType}} {{.Args}} mask) => ({{.Asm}} {{.ArgsOut}} ({{.MaskInConvert}} <types.TypeMask> mask))`).
67 Add("pureVreg", `({{.GoOp}}{{.GoType}} {{.Args}}) => ({{.Asm}} {{.ArgsOut}})`).
68 Add("vregMem", `({{.Asm}} {{.ArgsLoadAddr}}) && canMergeLoad(v, l) && clobber(l) => ({{.Asm}}load {{.ArgsAddr}})`).
69 Add("vregMemFeatCheck", `({{.Asm}} {{.ArgsLoadAddr}}) && {{.FeatCheck}} && canMergeLoad(v, l) && clobber(l) => ({{.Asm}}load {{.ArgsAddr}})`).
70 Add("asmRule", `({{.Asm}} {{.Args}}) {{.RuleCond}} => {{.RuleOut}}`).
71 Add("specialLower", `{{.Rule}}`)
72 )
73
74 func (d tplRuleData) MaskOptimization(asmCheck map[string]bool) string {
75 asmNoMask := d.Asm
76 if i := strings.Index(asmNoMask, "Masked"); i == -1 {
77 return ""
78 }
79 asmNoMask = strings.ReplaceAll(asmNoMask, "Masked", "")
80 if asmCheck[asmNoMask] == false {
81 return ""
82 }
83
84 for _, nope := range []string{"VMOVDQU", "VPCOMPRESS", "VCOMPRESS", "VPEXPAND", "VEXPAND", "VPBLENDM", "VMOVUP"} {
85 if strings.HasPrefix(asmNoMask, nope) {
86 return ""
87 }
88 }
89
90 size := asmNoMask[len(asmNoMask)-3:]
91 if strings.HasSuffix(asmNoMask, "const") {
92 sufLen := len("128const")
93 size = asmNoMask[len(asmNoMask)-sufLen:][:3]
94 }
95 switch size {
96 case "128", "256", "512":
97 default:
98 panic("Unexpected operation size on " + d.Asm)
99 }
100
101 switch d.ElementSize {
102 case 8, 16, 32, 64:
103 default:
104 panic(fmt.Errorf("Unexpected operation width %d on %v", d.ElementSize, d.Asm))
105 }
106
107 return fmt.Sprintf("(VMOVDQU%dMasked%s (%s %s) mask) => (%s %s mask)\n", d.ElementSize, size, asmNoMask, d.Args, d.Asm, d.Args)
108 }
109
110 func compareTplRuleData(x, y tplRuleData) int {
111 if c := compareNatural(x.GoOp, y.GoOp); c != 0 {
112 return c
113 }
114 if c := compareNatural(x.GoType, y.GoType); c != 0 {
115 return c
116 }
117 if c := compareNatural(x.Args, y.Args); c != 0 {
118 return c
119 }
120 if x.TplName == y.TplName {
121 return 0
122 }
123 return ruleTemplates.Compare(x.TplName, y.TplName)
124 }
125
126
127
128
129
130
131
132
133 func parseAsmRule(rule string) (bool, string, string) {
134 arrowIndex := strings.Index(rule, "=>")
135 if arrowIndex == -1 {
136 return false, "", ""
137 }
138
139 condPart := rule[:arrowIndex]
140 outPart := rule[arrowIndex+len("=>"):]
141
142
143 cond := strings.TrimPrefix(condPart, "if")
144 if cond == condPart || len(cond) == 0 || !unicode.IsSpace(rune(cond[0])) {
145 return false, "", ""
146 }
147
148
149 cond = strings.TrimSpace(cond)
150 out := strings.TrimSpace(outPart)
151 if cond == "" || out == "" {
152 return false, "", ""
153 }
154
155 return true, cond, out
156 }
157
158
159
160
161
162
163
164
165
166
167
168 func parseArgsMatchRule(rule string) (bool, bool, string) {
169 if strings.HasPrefix(rule, "earlymatch ") {
170 return true, true, rule[len("earlymatch "):]
171 } else if strings.HasPrefix(rule, "match ") {
172 return true, false, rule[len("match "):]
173 } else {
174 return false, false, rule
175 }
176 }
177
178
179
180
181
182
183
184
185
186 func expandFormatSpecifiers(s string, elemBits int) string {
187 elemLetters := map[int]string{8: "B", 16: "H", 32: "S", 64: "D"}
188 arrangements := map[int]string{8: "16B", 16: "8H", 32: "4S", 64: "2D"}
189 s = strings.ReplaceAll(s, "%s", elemLetters[elemBits])
190 s = strings.ReplaceAll(s, "%a", arrangements[elemBits])
191 s = strings.ReplaceAll(s, "%b", fmt.Sprintf("%d", elemBits))
192 return s
193 }
194
195
196
197 func writeSIMDRules(ops []Operation) *bytes.Buffer {
198 buffer := new(bytes.Buffer)
199 buffer.WriteString(generatedHeader() + "\n")
200
201
202 maskedMergeOpts := make(map[string]string)
203 s2n := map[int]string{8: "B", 16: "W", 32: "D", 64: "Q"}
204 asmCheck := map[string]bool{}
205 sftimmCheck := map[string]bool{}
206 var allData []tplRuleData
207 var optData []tplRuleData
208 var memOptData []tplRuleData
209 memOpSeen := make(map[string]bool)
210 ruleDone := make(map[string]struct{})
211
212 for _, opr := range ops {
213 opInShape, opOutShape, maskType, immType, gOp, _ := opr.shape()
214 asm := machineOpName(maskType, gOp)
215 vregInCnt := len(gOp.In)
216 if maskType == OneMask {
217 vregInCnt--
218 }
219
220 data := tplRuleData{
221 GoOp: gOp.Go,
222 Asm: asm,
223 }
224
225 if vregInCnt == 1 {
226 data.Args = "x"
227 data.ArgsOut = data.Args
228 } else if vregInCnt == 2 {
229 data.Args = "x y"
230 data.ArgsOut = data.Args
231 } else if vregInCnt == 3 {
232 data.Args = "x y z"
233 data.ArgsOut = data.Args
234 } else {
235 panic(fmt.Errorf("simdgen does not support more than 3 vreg in inputs"))
236 }
237 if immType == ConstImm {
238 data.ArgsOut = fmt.Sprintf("[%s] %s", *opr.In[0].Const, data.ArgsOut)
239 } else if immType == VarImm || immType == VarImmLim {
240 data.Args = fmt.Sprintf("[a] %s", data.Args)
241 data.ArgsOut = fmt.Sprintf("[a] %s", data.ArgsOut)
242 } else if immType == ConstVarImm {
243 data.Args = fmt.Sprintf("[a] %s", data.Args)
244 data.ArgsOut = fmt.Sprintf("[a+%s] %s", *opr.In[0].Const, data.ArgsOut)
245 }
246
247 goType := func(op Operation) string {
248 if op.OperandOrder != nil {
249 switch *op.OperandOrder {
250 case "21Type1", "231Type1":
251
252 return *op.In[1].Go
253 }
254 }
255 return *op.In[0].Go
256 }
257 var tplName string
258
259 if opOutShape == OneVregOut || opOutShape == OneVregOutAtIn || opOutShape == OneVregOutScalar || gOp.Out[0].OverwriteClass != nil {
260 switch opInShape {
261 case OneImmIn:
262 tplName = "pureVreg"
263 data.GoType = goType(gOp)
264 case PureVregIn, VlistIn:
265 tplName = "pureVreg"
266 data.GoType = goType(gOp)
267 case OneKmaskImmIn:
268 fallthrough
269 case OneKmaskIn:
270 tplName = "maskIn"
271 data.GoType = goType(gOp)
272 rearIdx := len(gOp.In) - 1
273
274 width := *gOp.In[rearIdx].ElemBits
275 data.MaskInConvert = fmt.Sprintf("VPMOVVec%dx%dToM", width, *gOp.In[rearIdx].Lanes)
276 data.ElementSize = width
277 case PureKmaskIn:
278 panic(fmt.Errorf("simdgen does not support pure k mask instructions, they should be generated by compiler optimizations"))
279 }
280 } else if opOutShape == OneGregOut {
281 tplName = "pureVreg"
282 data.GoType = goType(gOp)
283 } else {
284
285 data.MaskOutConvert = fmt.Sprintf("VPMOVMToVec%dx%d", *gOp.Out[0].ElemBits, *gOp.In[0].Lanes)
286 switch opInShape {
287 case OneImmIn:
288 fallthrough
289 case PureVregIn:
290 tplName = "maskOut"
291 data.GoType = goType(gOp)
292 case OneKmaskImmIn:
293 fallthrough
294 case OneKmaskIn:
295 tplName = "maskInMaskOut"
296 data.GoType = goType(gOp)
297 rearIdx := len(gOp.In) - 1
298 data.MaskInConvert = fmt.Sprintf("VPMOVVec%dx%dToM", *gOp.In[rearIdx].ElemBits, *gOp.In[rearIdx].Lanes)
299 case PureKmaskIn:
300 panic(fmt.Errorf("simdgen does not support pure k mask instructions, they should be generated by compiler optimizations"))
301 }
302 }
303
304 if gOp.SpecialLower != nil {
305 if *gOp.SpecialLower == "sftimm" {
306 if !sftimmCheck[data.Asm] {
307 sftimmCheck[data.Asm] = true
308 sftImmData := data
309 if tplName == "maskIn" {
310 sftImmData.TplName = "masksftimm"
311 } else {
312 sftImmData.TplName = "sftimm"
313 }
314 allData = append(allData, sftImmData)
315 asmCheck[sftImmData.Asm+"const"] = true
316 }
317 } else if ok, cond, out := parseAsmRule(*gOp.SpecialLower); ok {
318 if _, done := ruleDone[data.Asm]; !done {
319 ruleDone[data.Asm] = struct{}{}
320 optData := data
321 optData.TplName = "asmRule"
322 optData.RuleCond = cond
323 if cond != "" {
324 optData.RuleCond = "&& " + cond
325 }
326 optData.RuleOut = out
327 if maskType == OneMask {
328 optData.Args += " mask"
329 }
330 allData = append(allData, optData)
331 }
332 } else if ok, isEarly, rest := parseArgsMatchRule(*gOp.SpecialLower); ok {
333 key := data.Asm
334 if isEarly {
335 key = data.GoOp + data.GoType
336 }
337 if _, done := ruleDone[key]; !done {
338 ruleDone[key] = struct{}{}
339
340 elemBits := 0
341 for _, in := range gOp.In {
342 if in.ElemBits != nil {
343 elemBits = *in.ElemBits
344 break
345 }
346 }
347 optData := data
348 optData.Rule = rest
349 optData.Rule = expandFormatSpecifiers(optData.Rule, elemBits)
350 optData.TplName = "specialLower"
351
352 optData.Rule = strings.ReplaceAll(optData.Rule, "%g", optData.GoOp+optData.GoType)
353
354 optData.Rule = strings.ReplaceAll(optData.Rule, "%h", optData.Asm)
355
356 allData = append(allData, optData)
357 }
358 if isEarly {
359 continue
360 }
361 } else {
362 panic("simdgen sees unknown special lower " + *gOp.SpecialLower + ", maybe implement it?")
363 }
364 }
365 if gOp.MemFeatures != nil && *gOp.MemFeatures == "vbcst" {
366
367 selected := true
368 for _, a := range gOp.In {
369 if a.TreatLikeAScalarOfSize != nil {
370 selected = false
371 break
372 }
373 }
374 if _, ok := memOpSeen[data.Asm]; ok {
375 selected = false
376 }
377 if selected {
378 memOpSeen[data.Asm] = true
379 lastVreg := gOp.In[vregInCnt-1]
380
381 if lastVreg.Class != "vreg" {
382 panic(fmt.Errorf("simdgen expects vbcst replaced operand to be a vreg, but %v found", lastVreg))
383 }
384 memOpData := data
385
386 origArgs := data.Args[:len(data.Args)-1]
387
388 immArg := ""
389 immArgCombineOff := " [off] "
390 if immType != NoImm && immType != InvalidImm {
391 _, after, found := strings.Cut(origArgs, "]")
392 if found {
393 origArgs = after
394 }
395 immArg = "[c] "
396 immArgCombineOff = " [makeValAndOff(int32(uint8(c)),off)] "
397 }
398 memOpData.ArgsLoadAddr = immArg + origArgs + fmt.Sprintf("l:(VMOVDQUload%d {sym} [off] ptr mem)", *lastVreg.Bits)
399
400 memOpData.ArgsAddr = "{sym}" + immArgCombineOff + origArgs + "ptr"
401 if maskType == OneMask {
402 memOpData.ArgsAddr += " mask"
403 memOpData.ArgsLoadAddr += " mask"
404 }
405 memOpData.ArgsAddr += " mem"
406 if gOp.MemFeaturesData != nil {
407 _, feat2 := getVbcstData(*gOp.MemFeaturesData)
408 knownFeatChecks := map[string]string{
409 "AVX": "v.Block.CPUfeatures.hasFeature(CPUavx)",
410 "AVX2": "v.Block.CPUfeatures.hasFeature(CPUavx2)",
411 "AVX512": "v.Block.CPUfeatures.hasFeature(CPUavx512)",
412 }
413 memOpData.FeatCheck = knownFeatChecks[feat2]
414 memOpData.TplName = "vregMemFeatCheck"
415 } else {
416 memOpData.TplName = "vregMem"
417 }
418 memOptData = append(memOptData, memOpData)
419 asmCheck[memOpData.Asm+"load"] = true
420 }
421 }
422
423 if gOp.hasMaskedMerging(maskType, opOutShape) {
424
425 maskElem := gOp.In[len(gOp.In)-1]
426 if maskElem.Bits == nil {
427 panic("mask has no bits")
428 }
429 if maskElem.ElemBits == nil {
430 panic("mask has no elemBits")
431 }
432 if maskElem.Lanes == nil {
433 panic("mask has no lanes")
434 }
435 switch *maskElem.Bits {
436 case 128, 256:
437
438 noMaskName := machineOpName(NoMask, gOp)
439 ruleExisting, ok := maskedMergeOpts[noMaskName]
440 rule := fmt.Sprintf("(VPBLENDVB%d dst (%s %s) mask) && v.Block.CPUfeatures.hasFeature(CPUavx512) => (%sMerging dst %s (VPMOVVec%dx%dToM <types.TypeMask> mask))\n",
441 *maskElem.Bits, noMaskName, data.Args, data.Asm, data.Args, *maskElem.ElemBits, *maskElem.Lanes)
442 if ok && ruleExisting != rule {
443 panic(fmt.Sprintf("multiple masked merge rules for one op:\n%s\n%s\n", ruleExisting, rule))
444 } else {
445 maskedMergeOpts[noMaskName] = rule
446 }
447 case 512:
448
449 noMaskName := machineOpName(NoMask, gOp)
450 ruleExisting, ok := maskedMergeOpts[noMaskName]
451 rule := fmt.Sprintf("(VPBLENDM%sMasked%d dst (%s %s) mask) => (%sMerging dst %s mask)\n",
452 s2n[*maskElem.ElemBits], *maskElem.Bits, noMaskName, data.Args, data.Asm, data.Args)
453 if ok && ruleExisting != rule {
454 panic(fmt.Sprintf("multiple masked merge rules for one op:\n%s\n%s\n", ruleExisting, rule))
455 } else {
456 maskedMergeOpts[noMaskName] = rule
457 }
458 }
459 }
460
461 if tplName == "pureVreg" && data.Args == data.ArgsOut {
462 data.Args = "..."
463 data.ArgsOut = "..."
464 }
465 data.TplName = tplName
466 if opr.NoGenericOps != nil && *opr.NoGenericOps == "true" ||
467 opr.SkipMaskedMethod() {
468 optData = append(optData, data)
469 continue
470 }
471 allData = append(allData, data)
472 asmCheck[data.Asm] = true
473 }
474
475 slices.SortFunc(allData, compareTplRuleData)
476
477 hiHalfRules := generateHiHalfFoldingRules(ops)
478 for _, rule := range hiHalfRules {
479 buffer.WriteString(rule)
480 }
481
482 for _, data := range allData {
483 tpl := ruleTemplates.Get(data.TplName)
484 if tpl == nil {
485 panic(fmt.Errorf("template %s not found", data.TplName))
486 }
487 if err := tpl.Execute(buffer, data); err != nil {
488 panic(fmt.Errorf("failed to execute template %s for %s: %w", data.TplName, data.GoOp+data.GoType, err))
489 }
490 }
491
492 seen := make(map[string]bool)
493
494 for _, data := range optData {
495 if data.TplName == "maskIn" {
496 rule := data.MaskOptimization(asmCheck)
497 if seen[rule] {
498 continue
499 }
500 seen[rule] = true
501 buffer.WriteString(rule)
502 }
503 }
504
505 maskedMergeOptsRules := []string{}
506 for asm, rule := range maskedMergeOpts {
507 if !asmCheck[asm] {
508 continue
509 }
510 maskedMergeOptsRules = append(maskedMergeOptsRules, rule)
511 }
512 slices.Sort(maskedMergeOptsRules)
513 for _, rule := range maskedMergeOptsRules {
514 buffer.WriteString(rule)
515 }
516
517 for _, data := range memOptData {
518 tpl := ruleTemplates.Get(data.TplName)
519 if tpl == nil {
520 panic(fmt.Errorf("template %s not found", data.TplName))
521 }
522 if err := tpl.Execute(buffer, data); err != nil {
523 panic(fmt.Errorf("failed to execute template %s for %s: %w", data.TplName, data.Asm, err))
524 }
525 }
526
527 return buffer
528 }
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551 func generateHiHalfFoldingRules(ops []Operation) []string {
552 seen := make(map[string]bool)
553 var rules []string
554
555 for _, opr := range ops {
556 if opr.HiHalfAsm == nil {
557 continue
558 }
559 kind := opr.hiHalfKind()
560 if kind == "" {
561 continue
562 }
563 _, _, maskType, immType, gOp, _ := opr.shape()
564 asm := machineOpName(maskType, gOp)
565 asm2 := hiHalfOpName(*gOp.HiHalfAsm, gOp)
566
567 if seen[asm] {
568 continue
569 }
570 seen[asm] = true
571
572 vregInCnt := 0
573 for _, in := range gOp.In {
574 if in.Class == "vreg" {
575 vregInCnt++
576 }
577 }
578 hasImm := immType == VarImm || immType == VarImmLim || immType == ConstVarImm
579
580 switch kind {
581 case "narrow":
582 switch vregInCnt {
583 case 1:
584
585
586
587
588
589
590 if hasImm {
591 rules = append(rules, fmt.Sprintf("(VMOVDins0 [1] dst (VDUPDextr [0] (%s [c] y))) => (%s dst [c] y)\n", asm, asm2))
592 } else {
593 rules = append(rules, fmt.Sprintf("(VMOVDins0 [1] dst (VDUPDextr [0] (%s y))) => (%s dst y)\n", asm, asm2))
594 }
595 default:
596 panic("unsupported yet folding narrow ops cases")
597 }
598 case "long":
599 switch vregInCnt {
600 case 1:
601
602 if hasImm {
603 rules = append(rules, fmt.Sprintf("(%s [a] (VDUPDextr [1] x)) => (%s [a] x)\n", asm, asm2))
604 } else {
605 rules = append(rules, fmt.Sprintf("(%s (VDUPDextr [1] x)) => (%s x)\n", asm, asm2))
606 }
607 case 2:
608
609 rules = append(rules, fmt.Sprintf("(%s (VDUPDextr [1] x) (VDUPDextr [1] y)) => (%s x y)\n", asm, asm2))
610 default:
611 panic("unsupported yet folding long ops cases")
612 }
613 }
614 }
615
616 slices.Sort(rules)
617 return rules
618 }
619
View as plain text