1
2
3
4
5 package sve
6
7 import (
8 "cmp"
9 "fmt"
10 "log"
11 "slices"
12 "strings"
13
14 "simd/archsimd/_gen/unify"
15 )
16
17
18 func asComment(text string, width int) string {
19 text = strings.TrimSpace(text)
20 text = strings.ReplaceAll(text, "&", "&")
21 text = strings.ReplaceAll(text, "\n", " ")
22 words := strings.Fields(text)
23 var lines []string
24 line := ""
25 for _, w := range words {
26 if line != "" {
27 line += " "
28 }
29 line += w
30 if len(line) >= width {
31 lines = append(lines, "// "+line)
32 line = ""
33 }
34 }
35 if line != "" {
36 lines = append(lines, "// "+line)
37 }
38 return strings.Join(lines, "\n")
39 }
40
41
42
43 var mixedWidthLogged = map[string]bool{}
44
45
46
47
48 func (op *Operand) emit() *unify.Value {
49 var db unify.DefBuilder
50 db.Add("class", unify.NewValue(unify.NewStringExact(op.Class)))
51 if op.BaseType != "" {
52 db.Add("base", unify.NewValue(unify.NewStringExact(op.BaseType)))
53 }
54 switch {
55 case op.Bits > 0:
56
57 db.Add("bits", unify.NewValue(unify.NewStringExact(fmt.Sprint(op.Bits))))
58 if op.Lanes > 0 {
59 db.Add("lanes", unify.NewValue(unify.NewStringExact(fmt.Sprint(op.Lanes))))
60 }
61 case op.Class == "vreg" || op.Class == "mask":
62
63
64
65
66 db.Add("bits", unify.NewValue(unify.NewStringExact("scalable")))
67 }
68 if op.ElemBits > 0 {
69 db.Add("elemBits", unify.NewValue(unify.NewStringExact(fmt.Sprint(op.ElemBits))))
70 }
71 if op.Predication != "" {
72
73
74 db.Add("predication", unify.NewValue(unify.NewStringExact(op.Predication)))
75 }
76 if op.governing {
77
78 db.Add("governing", unify.NewValue(unify.NewStringExact("true")))
79 }
80 if op.isList {
81
82
83 db.Add("listNumber", unify.NewValue(unify.NewStringExact("0")))
84 }
85 if op.regName != "" {
86
87 db.Add("regName", unify.NewValue(unify.NewStringExact(op.regName)))
88 }
89
90
91
92
93
94
95
96
97 names := make([]*unify.Value, len(op.predRegName))
98 for i, n := range op.predRegName {
99 names[i] = unify.NewValue(unify.NewStringExact(n))
100 }
101 db.Add("predRegName", unify.NewValue(unify.NewTuple(names...)))
102 db.Add("asmPos", unify.NewValue(unify.NewStringExact(fmt.Sprint(op.AsmPos))))
103 return unify.NewValue(db.Build())
104 }
105
106
107
108
109 func pickRegNames(variants []predVariant, idx int, sel func(predVariant) []string) []string {
110 if len(variants) == 0 {
111 return nil
112 }
113 out := make([]string, len(variants))
114 for i, pv := range variants {
115 names := sel(pv)
116 if idx >= len(names) {
117 panic(fmt.Sprintf("operand %d has no counterpart in predicated encoding %d", idx, i))
118 }
119 out[i] = names[idx]
120 }
121 return out
122 }
123
124
125
126
127
128
129
130 func (inst *Instruction) emitOne(asm string, ops []Operand, widthAgnostic bool) *unify.Value {
131 var db unify.DefBuilder
132 db.Add("asm", unify.NewValue(unify.NewStringExact(asm)))
133 db.Add("goarch", unify.NewValue(unify.NewStringExact("arm64")))
134
135
136
137
138 feature := inst.cpuFeature()
139 unpred := ""
140 for _, pv := range inst.predVariants {
141 if pv.cpuFeature == "SVE" && feature == "SVE2" {
142 unpred = feature
143 feature = pv.cpuFeature
144 }
145 }
146 db.Add("cpuFeature", unify.NewValue(unify.NewStringExact(feature)))
147 if unpred != "" {
148 db.Add("unpredCPUFeature", unify.NewValue(unify.NewStringExact(unpred)))
149 }
150 if doc := inst.documentation(); doc != "" {
151 db.Add("details", unify.NewValue(unify.NewStringExact(asComment(doc, 80))))
152 }
153 if widthAgnostic {
154 db.Add("widthAgnostic", unify.NewValue(unify.NewStringExact("true")))
155 }
156
157
158
159
160
161
162 var inOps, outOps []Operand
163 var outIdx, inIdx int
164 for _, op := range ops {
165 switch {
166 case op.governing:
167
168
169 inOps = append(inOps, op)
170 case op.role == "destination":
171 op.predRegName = pickRegNames(inst.predVariants, outIdx, func(pv predVariant) []string { return pv.outRegNames })
172 outIdx++
173 outOps = append(outOps, op)
174 default:
175 op.predRegName = pickRegNames(inst.predVariants, inIdx, func(pv predVariant) []string { return pv.inRegNames })
176 inIdx++
177 inOps = append(inOps, op)
178 }
179 }
180 priority := map[string]int{"immediate": 0, "vreg": 1, "greg": 1, "memory": 1, "mask": 2}
181 slices.SortStableFunc(inOps, func(a, b Operand) int {
182 pa := priority[a.Class]
183 pb := priority[b.Class]
184 if pa != pb {
185 return cmp.Compare(pa, pb)
186 }
187 return cmp.Compare(a.AsmPos, b.AsmPos)
188 })
189
190 var ins, outs []*unify.Value
191 for i := range inOps {
192 ins = append(ins, inOps[i].emit())
193 }
194 for i := range outOps {
195 outs = append(outs, outOps[i].emit())
196 }
197 db.Add("in", unify.NewValue(unify.NewTuple(ins...)))
198 var inVar []*unify.Value
199 for _, pv := range inst.predVariants {
200
201 var pdb unify.DefBuilder
202 pdb.Add("class", unify.NewValue(unify.NewStringExact("mask")))
203 pdb.Add("bits", unify.NewValue(unify.NewStringExact("scalable")))
204 pdb.Add("predication", unify.NewValue(unify.NewStringExact(pv.quals)))
205 pdb.Add("asmPos", unify.NewValue(unify.NewStringExact(fmt.Sprint(pv.predAsmPos))))
206 inVar = append(inVar, unify.NewValue(pdb.Build()))
207 }
208 db.Add("inVariant", unify.NewValue(unify.NewTuple(inVar...)))
209 db.Add("out", unify.NewValue(unify.NewTuple(outs...)))
210 return unify.NewValue(db.Build())
211 }
212
213
214
215
216 func (inst *Instruction) emitAll() []*unify.Value {
217
218
219 defs, _, _ := inst.classify()
220 return defs
221 }
222
223
224 func lookup(rows []arngRow, size string) (int, bool) {
225 for _, r := range rows {
226 if r.size == size {
227 return r.bits, true
228 }
229 }
230 return 0, false
231 }
232
233
234
235
236
237
238
239
240
241 func (inst *Instruction) emitVariants(template []Operand) []*unify.Value {
242 asm := inst.goOpPrefix() + inst.mnemonic()
243
244 links := arngLinks(template)
245 tables := map[string][]arngRow{}
246 for _, l := range links {
247 tables[l] = inst.resolveArrangementTable(l)
248 }
249
250
251
252 var sizes []string
253 if len(links) > 0 {
254 for _, r := range tables[links[0]] {
255 sizes = append(sizes, r.size)
256 }
257 } else {
258 sizes = []string{""}
259 }
260
261 signs := inst.integerSignedness(template)
262
263
264
265 preds := predicationVariants(template)
266
267
268
269
270
271
272 widths := []int{0}
273 widthAgnostic := len(links) == 0 && inst.bitwise()
274 if widthAgnostic {
275 widths = []int{8, 16, 32, 64}
276 }
277
278 var defs []*unify.Value
279 for _, sign := range signs {
280 for _, size := range sizes {
281 ops := make([]Operand, len(template))
282 copy(ops, template)
283 skip := false
284 for i := range ops {
285 eb := ops[i].fixedElem
286 if ops[i].fixedBits > 0 {
287
288
289 eb = ops[i].fixedBits
290 } else if l := ops[i].arngLink; l != "" {
291 b, ok := lookup(tables[l], size)
292 if !ok {
293
294
295 skip = true
296 break
297 }
298 eb = b
299 }
300 base := sign
301 if inst.laneIsFloat(&ops[i]) {
302 base = "float"
303 if eb > 0 && eb < 16 {
304
305 skip = true
306 break
307 }
308 }
309 ops[i].instantiate(base, eb)
310 }
311 if skip {
312 continue
313 }
314 for _, pred := range preds {
315 variant := make([]Operand, len(ops))
316 copy(variant, ops)
317 elem := 0
318 mixedWidths := false
319 for i := range variant {
320 if variant[i].Class == "vreg" && variant[i].ElemBits > 0 {
321 if elem == 0 {
322 elem = variant[i].ElemBits
323 } else if variant[i].ElemBits != elem {
324 mixedWidths = true
325 }
326 }
327 }
328 for i := range variant {
329 if variant[i].Class != "mask" {
330 continue
331 }
332 if variant[i].governing {
333 variant[i].Predication = pred
334 }
335 if variant[i].ElemBits == 0 {
336
337
338 if mixedWidths && !mixedWidthLogged[inst.mnemonic()] {
339 mixedWidthLogged[inst.mnemonic()] = true
340 log.Printf("sve: %s: operands have mixed element widths; predicate width provisionally %d — derive esize from the pseudocode before generating an API from this def",
341 inst.mnemonic(), elem)
342 }
343 variant[i].ElemBits = elem
344 }
345 }
346 for _, w := range widths {
347 v := variant
348 if w > 0 {
349 v = make([]Operand, len(variant))
350 copy(v, variant)
351 for i := range v {
352 if v[i].Class == "vreg" || v[i].Class == "mask" {
353 v[i].ElemBits = w
354 }
355 }
356 }
357 defs = append(defs, inst.emitOne(asm, v, widthAgnostic))
358 }
359 }
360 }
361 }
362 return defs
363 }
364
View as plain text