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