1
2
3
4
5 package main
6
7 import (
8 "bytes"
9 "fmt"
10 "log"
11 "sort"
12 "strings"
13 "text/template"
14 )
15
16 var (
17 ssaTemplates = template.Must(template.New("simdSSA").Parse(`{{define "header"}}{{.GeneratedHeader}}
18 package {{.Arch}}
19
20 import (
21 "cmd/compile/internal/ssa"
22 "cmd/compile/internal/ssagen"
23 "cmd/internal/obj"
24 "cmd/internal/obj/{{.ObjArch}}"
25 )
26
27 func ssaGenSIMDValue(s *ssagen.State, v *ssa.Value) bool {
28 var p *obj.Prog
29 switch v.Op {{"{"}}{{end}}
30 {{define "case"}}
31 case {{.Cases}}:
32 p = {{.Helper}}(s, v{{if .Arrangement}}, {{.Arrangement}}{{end}})
33 {{end}}
34 {{define "footer"}}
35 default:
36 // Unknown reg shape
37 return false
38 }
39 {{end}}
40 {{define "zeroing"}}
41 // Masked operation are always compiled with zeroing.
42 switch v.Op {
43 case {{.}}:
44 x86.ParseSuffix(p, "Z")
45 }
46 {{end}}
47 {{define "ending"}}
48 // Ensure p is marked as used (may not be used in all generated code paths)
49 _ = p
50 return true
51 }
52 {{end}}`))
53 )
54
55 type tplSSAData struct {
56 Cases string
57 Helper string
58 Arrangement string
59 }
60
61 type tplSSAHeader struct {
62 Arch string
63 ObjArch string
64 GeneratedHeader string
65 }
66
67
68
69 func getArrangementFromOp(archInfo ArchInfo, caseStr string) string {
70 for _, a := range archInfo.Arrangements {
71 if strings.Contains(caseStr, a) {
72 return archInfo.Arch + ".ARNG_" + a
73 }
74 }
75 return ""
76 }
77
78
79
80 func writeSIMDSSA(ops []Operation) *bytes.Buffer {
81 archInfo := CurrentArch()
82 var ZeroingMask []string
83 regInfoKeys := archInfo.RegInfoKeys
84 regInfoSet := map[string][]string{}
85 for _, key := range regInfoKeys {
86 regInfoSet[key] = []string{}
87 }
88
89 seen := map[string]struct{}{}
90 allUnseen := make(map[string][]Operation)
91 allUnseenCaseStr := make(map[string][]string)
92
93 computeBaseRegShape := func(op Operation, mem memShape, shapeIn inShape, shapeOut outShape, immOpArg string, immType immShape) (string, error) {
94 regShape, err := op.regShape(mem)
95 if err != nil {
96 return "", err
97 }
98 if regShape == "v01load" {
99 regShape = "vload"
100 }
101 if shapeOut == OneVregOutAtIn {
102 regShape += "ResultInArg0"
103 } else if shapeOut == OneVregOutScalar {
104 regShape += "Scalar"
105 }
106 if shapeIn == OneImmIn || shapeIn == OneKmaskImmIn {
107 if immOpArg != "" {
108 regShape += "Imm"
109 regShape += immOpArg
110 } else if immType == VarImmLim {
111 regShape += "Imm"
112 } else {
113 regShape += "Imm8"
114 }
115 }
116 if shapeIn == VlistIn {
117 regShape += "List"
118 }
119 regShape, err = rewriteVecAsScalarRegInfo(op, regShape)
120 if err != nil {
121 return "", err
122 }
123 return regShape, nil
124 }
125 registerRegShape := func(regShape string, caseStr string, op Operation) {
126 if _, ok := regInfoSet[regShape]; !ok {
127 allUnseen[regShape] = append(allUnseen[regShape], op)
128 allUnseenCaseStr[regShape] = append(allUnseenCaseStr[regShape], caseStr)
129 }
130 regInfoSet[regShape] = append(regInfoSet[regShape], caseStr)
131 }
132 classifyOp := func(op Operation, maskType maskShape, shapeIn inShape, shapeOut outShape, caseStr string, mem memShape, immOpArg string, immType immShape) error {
133 regShape, err := computeBaseRegShape(op, mem, shapeIn, shapeOut, immOpArg, immType)
134 if err != nil {
135 return err
136 }
137
138 if op.HiHalfAsm != nil {
139 kind := op.hiHalfKind()
140 if kind != "" {
141 regShape += capitalizeFirst(kind)
142 }
143 }
144 registerRegShape(regShape, caseStr, op)
145 if mem == NoMem && op.hasMaskedMerging(maskType, shapeOut) {
146 regShapeMerging := regShape
147 if shapeOut != OneVregOutAtIn {
148
149
150 newIn := make([]Operand, len(op.In), len(op.In)+1)
151 copy(newIn, op.In)
152 op.In = newIn
153 op.In = append(op.In, op.Out[0])
154 op.sortOperand()
155 regShapeMerging, err = op.regShape(mem)
156 regShapeMerging += "ResultInArg0"
157 }
158 if err != nil {
159 return err
160 }
161 registerRegShape(regShapeMerging, caseStr+"Merging", op)
162 }
163 return nil
164 }
165
166
167 classifyHiHalfOp := func(op Operation, kind string, caseStr string, immOpArg string, immType immShape) error {
168 shapeIn, shapeOut, _, _, _, _ := op.shape()
169 regShape, err := computeBaseRegShape(op, NoMem, shapeIn, shapeOut, immOpArg, immType)
170 if err != nil {
171 return err
172 }
173 regShape = hiHalfLoweringRegShape(regShape, kind, true)
174 registerRegShape(regShape, caseStr, op)
175 return nil
176 }
177 for _, op := range ops {
178 shapeIn, shapeOut, maskType, immType, gOp, immOpArg := op.shape()
179 asm := machineOpName(maskType, gOp)
180 if _, ok := seen[asm]; ok {
181 continue
182 }
183 seen[asm] = struct{}{}
184 caseStr := fmt.Sprintf("ssa.Op%s%s", archInfo.ArchUpper, asm)
185 isZeroMasking := false
186 if shapeIn == OneKmaskIn || shapeIn == OneKmaskImmIn {
187 if gOp.Zeroing == nil || *gOp.Zeroing {
188 ZeroingMask = append(ZeroingMask, caseStr)
189 isZeroMasking = true
190 }
191 }
192 if err := classifyOp(op, maskType, shapeIn, shapeOut, caseStr, NoMem, immOpArg, immType); err != nil {
193 panic(err)
194 }
195
196
197
198
199 if gOp.HiHalfAsm != nil {
200 kind := op.hiHalfKind()
201 if kind != "" {
202 asm2 := hiHalfOpName(*gOp.HiHalfAsm, gOp)
203 caseStr2 := fmt.Sprintf("ssa.Op%s%s", archInfo.ArchUpper, asm2)
204 if _, ok2 := seen[asm2]; !ok2 {
205 seen[asm2] = struct{}{}
206 if err := classifyHiHalfOp(op, kind, caseStr2, immOpArg, immType); err != nil {
207 panic(err)
208 }
209 }
210 }
211 }
212
213 if op.MemFeatures != nil && *op.MemFeatures == "vbcst" {
214
215 op = rewriteLastVregToMem(op)
216
217
218
219 if err := classifyOp(op, maskType, shapeIn, shapeOut, caseStr+"load", VregMemIn, immOpArg, immType); err != nil {
220 if *Verbose {
221 log.Printf("Seen error: %e", err)
222 }
223 } else if isZeroMasking {
224 ZeroingMask = append(ZeroingMask, caseStr+"load")
225 }
226 }
227 }
228 if len(allUnseen) != 0 {
229 allKeys := make([]string, 0)
230 for k := range allUnseen {
231 allKeys = append(allKeys, k)
232 }
233 panic(fmt.Errorf("unsupported register constraint for prog, please update gen_simdssa.go and amd64/ssa.go: %+v\nAll keys: %v\n, cases: %v\n", allUnseen, allKeys, allUnseenCaseStr))
234 }
235
236 buffer := new(bytes.Buffer)
237
238 headerData := tplSSAHeader{
239 Arch: archInfo.Arch,
240 ObjArch: archInfo.ObjArch,
241 GeneratedHeader: archInfo.GeneratedHeader,
242 }
243 if err := ssaTemplates.ExecuteTemplate(buffer, "header", headerData); err != nil {
244 panic(fmt.Errorf("failed to execute header template: %w", err))
245 }
246
247 for _, regShape := range regInfoKeys {
248
249 cases := regInfoSet[regShape]
250 if len(cases) == 0 {
251 continue
252 }
253
254
255 arrangementGroups := make(map[string][]string)
256 for _, caseStr := range cases {
257 arrangement := getArrangementFromOp(archInfo, caseStr)
258 arrangementGroups[arrangement] = append(arrangementGroups[arrangement], caseStr)
259 }
260
261
262 var arrangements []string
263 for arrangement := range arrangementGroups {
264 arrangements = append(arrangements, arrangement)
265 }
266 sort.Strings(arrangements)
267
268
269 for _, arrangement := range arrangements {
270 groupCases := arrangementGroups[arrangement]
271 data := tplSSAData{
272 Cases: strings.Join(groupCases, ",\n\t\t"),
273 Helper: "simd" + capitalizeFirst(regShape),
274 }
275 if arrangement != "" {
276 data.Arrangement = arrangement
277 }
278 if err := ssaTemplates.ExecuteTemplate(buffer, "case", data); err != nil {
279 panic(fmt.Errorf("failed to execute case template for %s: %w", regShape, err))
280 }
281 }
282 }
283
284 if err := ssaTemplates.ExecuteTemplate(buffer, "footer", nil); err != nil {
285 panic(fmt.Errorf("failed to execute footer template: %w", err))
286 }
287
288 if len(ZeroingMask) != 0 {
289 if err := ssaTemplates.ExecuteTemplate(buffer, "zeroing", strings.Join(ZeroingMask, ",\n\t\t")); err != nil {
290 panic(fmt.Errorf("failed to execute footer template: %w", err))
291 }
292 }
293
294 if err := ssaTemplates.ExecuteTemplate(buffer, "ending", headerData); err != nil {
295 panic(fmt.Errorf("failed to execute ending template: %w", err))
296 }
297
298 return buffer
299 }
300
View as plain text