Source file src/simd/archsimd/_gen/cmd/refgen/main.go

     1  // Copyright 2026 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  // refgen produces a reference implementation of the SIMD API backed by the spec
     6  // implementation.
     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  	// Define all vector types
    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  	// Define functions
   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  // toSpec returns an expression that converts val from the specref Go type for t
   175  // to spec type tt. It may write statements to src.
   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  			// Easy case: the values don't need to change.
   180  			return val
   181  		}
   182  		// Hard case: we need to map each element
   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  // fromSpec returns an expression that converts val from the spec package type
   227  // tt to the specref Go type for t. It may write statements to src.
   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  			// Easy case: the values don't need to change.
   232  			return val
   233  		}
   234  		// Hard case: we need to map each element
   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