1
2
3
4
5
6
7 package main
8
9 import (
10 "bytes"
11 "cmp"
12 "flag"
13 "fmt"
14 "go/types"
15 "log"
16 "maps"
17 "os"
18 "slices"
19 "strings"
20
21 "simd/archsimd/_gen/gentools"
22 "simd/archsimd/_gen/specgen"
23 "simd/archsimd/_gen/specgen/specexpr"
24 )
25
26 func main() {
27 genOpts := gentools.RegisterFlags(nil)
28
29 flag.Usage = func() {
30 w := flag.CommandLine.Output()
31 fmt.Fprintf(w, "usage: refgen [flags]\n")
32 flag.CommandLine.PrintDefaults()
33 }
34
35 flag.Parse()
36 if flag.NArg() != 0 {
37 flag.Usage()
38 os.Exit(1)
39 }
40 specDir := specgen.MustFindSpecDir(genOpts.GOROOT)
41
42 funcs, err := specgen.Load(specDir, nil)
43 if err != nil {
44 fmt.Fprintf(os.Stderr, "%s\n", err.Error())
45 os.Exit(1)
46 }
47
48 var files gentools.Files
49 defer files.FlushOrExit()
50
51 src := &srcWriter{Buffer: files.NewGoFile("simd/internal/simdref/simdref.go")}
52
53 fmt.Fprintf(src, `// Code generated by 'refgen'. DO NOT EDIT.
54
55 package simdref
56
57 import "simd/internal/spec"
58
59 `)
60
61
62 vecTypeSet := make(map[specexpr.Vector]bool)
63 for _, fn := range funcs {
64 if fn.Recv.Type != nil {
65 vecTypeSet[fn.Recv.Type.(specexpr.Vector)] = true
66 }
67 }
68 vecTypes := slices.SortedFunc(maps.Keys(vecTypeSet), func(a, b specexpr.Vector) int {
69 if a.Elem != b.Elem {
70 return cmp.Compare(a.Elem.String(), b.Elem.String())
71 }
72 cmp, ok := a.Width.Compare(b.Width)
73 if ok {
74 return cmp
75 }
76 _, aScale := a.Width.(specexpr.ScalableWidth)
77 _, bScale := b.Width.(specexpr.ScalableWidth)
78 if !aScale && bScale {
79 return 1
80 }
81 return -1
82 })
83 fmt.Fprintf(src, "type (\n")
84 for _, vec := range vecTypes {
85 elem := vec.Elem.String()
86 if vec.Elem.Base == "Mask" {
87 elem = fmt.Sprintf("spec.Mask%d", vec.Elem.Bits)
88 }
89
90 fmt.Fprintf(src, "\t%s struct { v []%s }\n", vec.String(), elem)
91 }
92 fmt.Fprintf(src, ")\n\n")
93
94
95 var args []string
96 for _, fn := range funcs {
97 fmt.Fprintf(src, "%s {\n", fn.Decl())
98 src.id = 0
99
100 specName, specSig, typeArgs := fn.SpecFunc()
101 specInst, err := types.Instantiate(nil, specSig, typeArgs, false)
102 if err != nil {
103 panic(fmt.Sprintf("instantiating spec function %s: %s", specName, err))
104 }
105
106 specParams := specInst.(*types.Signature).Params()
107 args = args[:0]
108 if fn.Recv.Type != nil {
109 args = append(args, toSpec(fn.Recv.Type, specParams.At(len(args)).Type(), fn.Recv.Name, src))
110 }
111 for _, in := range fn.In {
112 args = append(args, toSpec(in.Type, specParams.At(len(args)).Type(), in.Name, src))
113 }
114
115 call := formatCall(specName, typeArgs, args)
116
117 specResults := specInst.(*types.Signature).Results()
118 switch len(fn.Out) {
119 case 0:
120 fmt.Fprintf(src, "\t%s\n", call)
121 case 1:
122 fmt.Fprintf(src, "\treturn %s\n", fromSpec(fn.Out[0].Type, specResults.At(0).Type(), call, src))
123 default:
124 var tmps []string
125 var res []string
126 for i := range fn.Out {
127 tmp := fmt.Sprintf("r%d", i+1)
128 tmps = append(tmps, tmp)
129 res = append(res, fromSpec(fn.Out[i].Type, specResults.At(i).Type(), tmp, src))
130 }
131 fmt.Fprintf(src, "\t%s := %s\n", strings.Join(tmps, ", "), call)
132 fmt.Fprintf(src, "\treturn %s\n", strings.Join(res, ", "))
133 }
134 fmt.Fprintf(src, "}\n\n")
135 }
136 }
137
138 type srcWriter struct {
139 *bytes.Buffer
140 id int
141 }
142
143 func (w *srcWriter) genIdent() string {
144 ident := fmt.Sprintf("tmp%d", w.id)
145 w.id++
146 return ident
147 }
148
149 func formatCall(specName string, typeArgs []types.Type, args []string) string {
150 var callBuf bytes.Buffer
151 fmt.Fprintf(&callBuf, "spec.%s[", specName)
152 for i, typeArg := range typeArgs {
153 if i > 0 {
154 callBuf.WriteString(", ")
155 }
156 types.WriteType(&callBuf, typeArg, specQualifier)
157 }
158 fmt.Fprintf(&callBuf, "](%s)", strings.Join(args, ", "))
159 return callBuf.String()
160 }
161
162 func specQualifier(pkg *types.Package) string {
163 if pkg.Path() == "simd/internal/spec" {
164 return "spec"
165 }
166 return ""
167 }
168
169
170
171 func toSpec(t specexpr.Type, tt types.Type, val string, src *srcWriter) string {
172 arrayOrSlice := func(tElem specexpr.Type, ttElem types.Type) string {
173 if eVal := toSpec(tElem, ttElem, val, src); eVal == val {
174
175 return val
176 }
177
178 tmp := src.genIdent()
179 fmt.Fprintf(src, "var %s %s\n", tmp, types.TypeString(tt, specQualifier))
180 fmt.Fprintf(src, "for i := range %s {\n", val)
181 eVal := fromSpec(tElem, ttElem, "("+val+")[i]", src)
182 fmt.Fprintf(src, "\t%s[i] = %s\n", tmp, eVal)
183 fmt.Fprintf(src, "}\n")
184 return tmp
185 }
186
187 switch t := t.(type) {
188 case specexpr.Vector:
189 return val + ".v"
190 case specexpr.Basic:
191 switch tt := tt.(type) {
192 case *types.Named:
193 if tt.Obj().Name() == "UintN" {
194 return "spec.UintN(" + val + ")"
195 }
196 }
197 return val
198 case specexpr.Slice:
199 tt := tt.Underlying().(*types.Slice)
200 return arrayOrSlice(t.Elem, tt.Elem())
201 case specexpr.Array:
202 switch tt := tt.(type) {
203 case *types.Named:
204 if tt.Obj().Name() == "Array" {
205 return toSpec(t.Elem, tt.TypeArgs().At(0), val, src) + "[:]"
206 }
207 }
208 tt := tt.Underlying().(*types.Array)
209 return arrayOrSlice(t.Elem, tt.Elem())
210 case specexpr.Pointer:
211 tt := tt.(*types.Pointer)
212 eVal := toSpec(t.Elem, tt.Elem(), val, src)
213 tmp := src.genIdent()
214 fmt.Fprintf(src, "var %s %s = %s\n", tmp, types.TypeString(tt.Elem(), specQualifier), eVal)
215 return "&" + tmp
216 }
217 log.Fatalf("unexpected specexpr type %s (%T)", t, t)
218 panic("not reachable")
219 }
220
221
222
223 func fromSpec(t specexpr.Type, tt types.Type, val string, src *srcWriter) string {
224 arrayOrSlice := func(tElem specexpr.Type, ttElem types.Type) string {
225 if eVal := fromSpec(tElem, ttElem, val, src); eVal == val {
226
227 return val
228 }
229
230 tmp := src.genIdent()
231 fmt.Fprintf(src, "var %s %s\n", tmp, t)
232 fmt.Fprintf(src, "for i := range %s {\n", val)
233 eVal := fromSpec(tElem, ttElem, "("+val+")[i]", src)
234 fmt.Fprintf(src, "\t%s[i] = %s\n", tmp, eVal)
235 fmt.Fprintf(src, "}\n")
236 return tmp
237 }
238
239 switch t := t.(type) {
240 case specexpr.Vector:
241 return fmt.Sprintf("%s{%s}", t, val)
242 case specexpr.Basic:
243 switch tt := tt.(type) {
244 case *types.Named:
245 if tt.Obj().Name() == "UintN" {
246 return fmt.Sprintf("%s(%s)", t, val)
247 }
248 }
249 return val
250 case specexpr.Slice:
251 tt := tt.Underlying().(*types.Slice)
252 return arrayOrSlice(t.Elem, tt.Elem())
253 case specexpr.Array:
254 switch tt := tt.(type) {
255 case *types.Named:
256 if tt.Obj().Name() == "Array" {
257 return fmt.Sprintf("(%s)(%s)", t, fromSpec(t.Elem, tt.TypeArgs().At(0), val, src))
258 }
259 }
260 tt := tt.Underlying().(*types.Array)
261 return arrayOrSlice(t.Elem, tt.Elem())
262 case specexpr.Pointer:
263 tt := tt.(*types.Pointer)
264 eVal := fromSpec(t.Elem, tt.Elem(), val, src)
265 tmp := src.genIdent()
266 fmt.Fprintf(src, "var %s %s = %s\n", tmp, t, eVal)
267 return "&" + tmp
268 }
269 log.Fatalf("unexpected specexpr type %s (%T)", t, t)
270 panic("not reachable")
271 }
272
View as plain text