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  	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  	// Define all vector types
    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  	// Define functions
    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  // toSpec returns an expression that converts val from the specref Go type for t
   170  // to spec type tt. It may write statements to src.
   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  			// Easy case: the values don't need to change.
   175  			return val
   176  		}
   177  		// Hard case: we need to map each element
   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  // fromSpec returns an expression that converts val from the spec package type
   222  // tt to the specref Go type for t. It may write statements to src.
   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  			// Easy case: the values don't need to change.
   227  			return val
   228  		}
   229  		// Hard case: we need to map each element
   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