1
2
3
4
5 package main
6
7 import (
8 "bytes"
9 "fmt"
10 "log"
11 "sort"
12 "strings"
13 )
14
15 const simdMachineOpsTmpl = `
16 package main
17
18 func simd{{.ArchUpper}}Ops({{.RegInfoParams}}) []opData {
19 return []opData{
20 {{- range .OpsData }}
21 {name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
22 {{- end }}
23 {{- range .OpsDataImm }}
24 {name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}", aux: "UInt8",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
25 {{- end }}
26 {{- range .OpsDataLoad}}
27 {name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}", aux: "SymOff", symEffect: "Read"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
28 {{- end}}
29 {{- range .OpsDataImmLoad}}
30 {name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}", aux: "SymValAndOff", symEffect: "Read"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
31 {{- end}}
32 {{- range .OpsDataMerging }}
33 {name: "{{.OpName}}Merging", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}", typ: "{{.Type}}", resultInArg0: true},
34 {{- end }}
35 {{- range .OpsDataImmMerging }}
36 {name: "{{.OpName}}Merging", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}", aux: "UInt8", typ: "{{.Type}}", resultInArg0: true},
37 {{- end }}
38 }
39 }
40 `
41
42
43
44 func writeSIMDMachineOps(ops []Operation) *bytes.Buffer {
45 t := templateOf(simdMachineOpsTmpl, "simdAMD64Ops")
46 buffer := new(bytes.Buffer)
47 buffer.WriteString(generatedHeader())
48
49 type opData struct {
50 OpName string
51 Asm string
52 OpInLen int
53 RegInfo string
54 Comm bool
55 Type string
56 ResultInArg0 bool
57 }
58 type machineOpsData struct {
59 ArchUpper string
60 RegInfoParams string
61 OpsData []opData
62 OpsDataImm []opData
63 OpsDataLoad []opData
64 OpsDataImmLoad []opData
65 OpsDataMerging []opData
66 OpsDataImmMerging []opData
67 }
68
69 archInfo := CurrentArch()
70
71 regInfoSet := archInfo.RegInfoSet
72 opsData := make([]opData, 0)
73 opsDataImm := make([]opData, 0)
74 opsDataLoad := make([]opData, 0)
75 opsDataImmLoad := make([]opData, 0)
76 opsDataMerging := make([]opData, 0)
77 opsDataImmMerging := make([]opData, 0)
78
79
80 best := make(map[string]Operation)
81 var mOpOrder []string
82 countOverrides := func(s []Operand) int {
83 a := 0
84 for _, o := range s {
85 if o.OverwriteBase != nil {
86 a++
87 }
88 }
89 return a
90 }
91 for _, op := range ops {
92 _, _, maskType, _, gOp, _ := op.shape()
93 asm := machineOpName(maskType, gOp)
94 other, ok := best[asm]
95 if !ok {
96 best[asm] = op
97 mOpOrder = append(mOpOrder, asm)
98 continue
99 }
100 if !op.Commutative && other.Commutative {
101 best[asm] = op
102 continue
103 }
104
105 if countOverrides(op.In)+countOverrides(op.Out) < countOverrides(other.In)+countOverrides(other.Out) {
106 best[asm] = op
107 }
108 }
109
110 regInfoErrs := make([]error, 0)
111 regInfoMissing := make(map[string]bool, 0)
112 for _, asm := range mOpOrder {
113 op := best[asm]
114 shapeIn, shapeOut, maskType, _, gOp, _ := op.shape()
115
116
117
118 makeRegInfo := func(op Operation, mem memShape) (string, error) {
119 regInfo, err := op.regShape(mem)
120 if err != nil {
121 panic(err)
122 }
123 regInfo, err = rewriteVecAsScalarRegInfo(op, regInfo)
124 if err != nil {
125 if mem == NoMem || mem == InvalidMem {
126 panic(err)
127 }
128 return "", err
129 }
130 if regInfo == "v01load" {
131 regInfo = "vload"
132 }
133
134 if strings.Contains(op.CPUFeature, "AVX512") {
135 regInfo = strings.ReplaceAll(regInfo, "v", "w")
136 }
137 if _, ok := regInfoSet[regInfo]; !ok {
138 regInfoErrs = append(regInfoErrs, fmt.Errorf("unsupported register constraint, please update the template and AMD64Ops.go: %s. Op is %s", regInfo, op))
139 regInfoMissing[regInfo] = true
140 }
141 return regInfo, nil
142 }
143 regInfo, err := makeRegInfo(op, NoMem)
144 if err != nil {
145 panic(err)
146 }
147 var outType string
148 if shapeOut == OneVregOut || shapeOut == OneVregOutAtIn || shapeOut == OneVregOutScalar || gOp.Out[0].OverwriteClass != nil {
149
150 outType = fmt.Sprintf("Vec%d", *gOp.Out[0].Bits)
151 } else if shapeOut == OneGregOut {
152 outType = gOp.GoType()
153 } else if shapeOut == OneKmaskOut {
154 outType = "Mask"
155 } else {
156 panic(fmt.Errorf("simdgen does not recognize this output shape: %d", shapeOut))
157 }
158 resultInArg0 := false
159 if shapeOut == OneVregOutAtIn {
160 resultInArg0 = true
161 }
162 var memOpData *opData
163 regInfoMerging := regInfo
164 hasMerging := false
165 if op.MemFeatures != nil && *op.MemFeatures == "vbcst" {
166
167
168 opMem := rewriteLastVregToMem(op)
169 regInfo, err := makeRegInfo(opMem, VregMemIn)
170 if err != nil {
171
172
173
174 if *Verbose {
175 log.Printf("Seen error: %e", err)
176 }
177 } else {
178 memOpData = &opData{asm + "load", gOp.Asm, len(gOp.In) + 1, regInfo, false, outType, resultInArg0}
179 }
180 }
181 hasMerging = gOp.hasMaskedMerging(maskType, shapeOut)
182 if hasMerging && !resultInArg0 {
183
184
185 newIn := make([]Operand, len(op.In), len(op.In)+1)
186 copy(newIn, op.In)
187 op.In = newIn
188 op.In = append(op.In, op.Out[0])
189 op.sortOperand()
190 regInfoMerging, err = makeRegInfo(op, NoMem)
191 if err != nil {
192 panic(err)
193 }
194 }
195
196 if shapeIn == OneImmIn || shapeIn == OneKmaskImmIn {
197 opsDataImm = append(opsDataImm, opData{asm, gOp.Asm, len(gOp.In), regInfo, gOp.Commutative, outType, resultInArg0})
198 if memOpData != nil {
199 if *op.MemFeatures != "vbcst" {
200 panic("simdgen only knows vbcst for mem ops for now")
201 }
202 opsDataImmLoad = append(opsDataImmLoad, *memOpData)
203 }
204 if hasMerging {
205 mergingLen := len(gOp.In)
206 if !resultInArg0 {
207 mergingLen++
208 }
209 opsDataImmMerging = append(opsDataImmMerging, opData{asm, gOp.Asm, mergingLen, regInfoMerging, gOp.Commutative, outType, resultInArg0})
210 }
211 } else {
212 opsData = append(opsData, opData{asm, gOp.Asm, len(gOp.In), regInfo, gOp.Commutative, outType, resultInArg0})
213 if memOpData != nil {
214 if *op.MemFeatures != "vbcst" {
215 panic("simdgen only knows vbcst for mem ops for now")
216 }
217 opsDataLoad = append(opsDataLoad, *memOpData)
218 }
219 if hasMerging {
220 mergingLen := len(gOp.In)
221 if !resultInArg0 {
222 mergingLen++
223 }
224 opsDataMerging = append(opsDataMerging, opData{asm, gOp.Asm, mergingLen, regInfoMerging, gOp.Commutative, outType, resultInArg0})
225 }
226 }
227
228 if gOp.HiHalfAsm != nil {
229 opsDataTarget := &opsData
230 if shapeIn == OneImmIn || shapeIn == OneKmaskImmIn {
231 opsDataTarget = &opsDataImm
232 }
233 kind := op.hiHalfKind()
234 asm2Name := hiHalfOpName(*gOp.HiHalfAsm, gOp)
235 argLen2 := len(gOp.In)
236 regInfo2 := regInfo
237 resultInArg02 := false
238 if kind == "narrow" {
239 argLen2++
240 regInfo2 = hiHalfRegShape2(regInfo, kind)
241 resultInArg02 = true
242 }
243 if _, ok := regInfoSet[regInfo2]; !ok {
244 regInfoErrs = append(regInfoErrs, fmt.Errorf("unsupported hi-half register constraint: %s for op %s", regInfo2, asm2Name))
245 regInfoMissing[regInfo2] = true
246 } else {
247 *opsDataTarget = append(*opsDataTarget, opData{asm2Name, *gOp.HiHalfAsm, argLen2, regInfo2, gOp.Commutative, outType, resultInArg02})
248 }
249 }
250 }
251 if len(regInfoErrs) != 0 {
252 for _, e := range regInfoErrs {
253 log.Printf("Errors: %e\n", e)
254 }
255 panic(fmt.Errorf("these regInfo unseen: %v", regInfoMissing))
256 }
257 sort.Slice(opsData, func(i, j int) bool {
258 return compareNatural(opsData[i].OpName, opsData[j].OpName) < 0
259 })
260 sort.Slice(opsDataImm, func(i, j int) bool {
261 return compareNatural(opsDataImm[i].OpName, opsDataImm[j].OpName) < 0
262 })
263 sort.Slice(opsDataLoad, func(i, j int) bool {
264 return compareNatural(opsDataLoad[i].OpName, opsDataLoad[j].OpName) < 0
265 })
266 sort.Slice(opsDataImmLoad, func(i, j int) bool {
267 return compareNatural(opsDataImmLoad[i].OpName, opsDataImmLoad[j].OpName) < 0
268 })
269 sort.Slice(opsDataMerging, func(i, j int) bool {
270 return compareNatural(opsDataMerging[i].OpName, opsDataMerging[j].OpName) < 0
271 })
272 sort.Slice(opsDataImmMerging, func(i, j int) bool {
273 return compareNatural(opsDataImmMerging[i].OpName, opsDataImmMerging[j].OpName) < 0
274 })
275
276 err := t.Execute(buffer, machineOpsData{
277 ArchUpper: archInfo.ArchUpper,
278 RegInfoParams: archInfo.RegInfoParams,
279 OpsData: opsData,
280 OpsDataImm: opsDataImm,
281 OpsDataLoad: opsDataLoad,
282 OpsDataImmLoad: opsDataImmLoad,
283 OpsDataMerging: opsDataMerging,
284 OpsDataImmMerging: opsDataImmMerging,
285 })
286 if err != nil {
287 panic(fmt.Errorf("failed to execute template: %w", err))
288 }
289
290 return buffer
291 }
292
View as plain text