Source file src/simd/archsimd/_gen/simdgen/gen_utility.go

     1  // Copyright 2025 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  package main
     6  
     7  import (
     8  	"bufio"
     9  	"bytes"
    10  	"fmt"
    11  	"go/format"
    12  	"log"
    13  	"os"
    14  	"path/filepath"
    15  	"reflect"
    16  	"slices"
    17  	"sort"
    18  	"strings"
    19  	"text/template" // NOLINT
    20  	"unicode"
    21  )
    22  
    23  func templateOf(temp, name string) *template.Template {
    24  	t, err := template.New(name).Parse(temp)
    25  	if err != nil {
    26  		panic(fmt.Errorf("failed to parse template %s: %w", name, err))
    27  	}
    28  	return t
    29  }
    30  
    31  func createPath(goroot string, file string) (*os.File, error) {
    32  	fp := filepath.Join(goroot, file)
    33  	dir := filepath.Dir(fp)
    34  	err := os.MkdirAll(dir, 0755)
    35  	if err != nil {
    36  		return nil, fmt.Errorf("failed to create directory %s: %w", dir, err)
    37  	}
    38  	f, err := os.Create(fp)
    39  	if err != nil {
    40  		return nil, fmt.Errorf("failed to create file %s: %w", fp, err)
    41  	}
    42  	return f, nil
    43  }
    44  
    45  func formatWriteAndClose(out *bytes.Buffer, goroot string, file string) {
    46  	b, err := format.Source(out.Bytes())
    47  	if err != nil {
    48  		fmt.Fprintf(os.Stderr, "%v\n", err)
    49  		fmt.Fprintf(os.Stderr, "%s\n", numberLines(out.Bytes()))
    50  		fmt.Fprintf(os.Stderr, "%v\n", err)
    51  		panic(err)
    52  	} else {
    53  		writeAndClose(b, goroot, file)
    54  	}
    55  }
    56  
    57  func writeAndClose(b []byte, goroot string, file string) {
    58  	ofile, err := createPath(goroot, file)
    59  	if err != nil {
    60  		panic(err)
    61  	}
    62  	ofile.Write(b)
    63  	ofile.Close()
    64  }
    65  
    66  // numberLines takes a slice of bytes, and returns a string where each line
    67  // is numbered, starting from 1.
    68  func numberLines(data []byte) string {
    69  	var buf bytes.Buffer
    70  	r := bytes.NewReader(data)
    71  	s := bufio.NewScanner(r)
    72  	for i := 1; s.Scan(); i++ {
    73  		fmt.Fprintf(&buf, "%d: %s\n", i, s.Text())
    74  	}
    75  	return buf.String()
    76  }
    77  
    78  type inShape uint8
    79  type outShape uint8
    80  type maskShape uint8
    81  type immShape uint8
    82  type memShape uint8
    83  
    84  const (
    85  	InvalidIn     inShape = iota
    86  	PureVregIn            // vector register input only
    87  	OneKmaskIn            // vector and kmask input
    88  	OneImmIn              // vector and immediate input
    89  	OneKmaskImmIn         // vector, kmask, and immediate inputs
    90  	PureKmaskIn           // only mask inputs.
    91  	VlistIn               // vector list input (e.g. TBL [v1 v2 v3] vidx)
    92  )
    93  
    94  const (
    95  	InvalidOut       outShape = iota
    96  	NoOut                     // no output
    97  	OneVregOut                // (one) vector register output
    98  	OneGregOut                // (one) general register output
    99  	OneKmaskOut               // mask output
   100  	OneVregOutAtIn            // the first input is also the output
   101  	OneVregOutScalar          // (one) vector register output scalar in lane 0, other lanes zeroed
   102  )
   103  
   104  const (
   105  	InvalidMask maskShape = iota
   106  	NoMask                // no mask
   107  	OneMask               // with mask (K1 to K7)
   108  	AllMasks              // a K mask instruction (K0-K7)
   109  )
   110  
   111  const (
   112  	InvalidImm  immShape = iota
   113  	NoImm                // no immediate
   114  	ConstImm             // const only immediate
   115  	VarImm               // pure imm argument provided by the users
   116  	ConstVarImm          // a combination of user arg and const
   117  	VarImmLim            // pure imm argument provided by the users, up to maximum in op.ImmMax
   118  )
   119  
   120  const (
   121  	InvalidMem memShape = iota
   122  	NoMem
   123  	VregMemIn // The instruction contains a mem input which is loading a vreg.
   124  )
   125  
   126  // opShape returns the several integers describing the shape of the operation,
   127  // and modified versions of the op:
   128  //
   129  // opNoImm is op with its inputs excluding the const imm.
   130  //
   131  // This function does not modify op.
   132  func (op *Operation) shape() (shapeIn inShape, shapeOut outShape, maskType maskShape, immType immShape,
   133  	opNoImm Operation, immOpIdx string) {
   134  	if len(op.Out) > 1 {
   135  		panic(fmt.Errorf("simdgen only supports 1 output: %s", op))
   136  	}
   137  	var outputReg int
   138  	if len(op.Out) == 1 {
   139  		outputReg = op.Out[0].AsmPos
   140  		if op.Out[0].Class == "vreg" {
   141  			shapeOut = OneVregOut
   142  			if op.Out[0].TreatLikeAScalarOfSize != nil {
   143  				shapeOut = OneVregOutScalar
   144  			}
   145  		} else if op.Out[0].Class == "greg" {
   146  			shapeOut = OneGregOut
   147  			if op.Out[0].TreatLikeAScalarOfSize != nil {
   148  				shapeOut = OneVregOutScalar
   149  			}
   150  		} else if op.Out[0].Class == "mask" {
   151  			shapeOut = OneKmaskOut
   152  		} else {
   153  			panic(fmt.Errorf("simdgen only supports output of class vreg or mask: %s", op))
   154  		}
   155  	} else {
   156  		shapeOut = NoOut
   157  		// TODO: are these only Load/Stores?
   158  		// We manually supported two Load and Store, are those enough?
   159  		panic(fmt.Errorf("simdgen only supports 1 output: %s", op))
   160  	}
   161  	hasImm := false
   162  	immAsmPos := -1
   163  	maskCount := 0
   164  	hasVreg := false
   165  	hasListIn := false
   166  	for _, in := range op.In {
   167  		if in.ListNumber != nil {
   168  			hasListIn = true
   169  		}
   170  		if in.AsmPos == outputReg {
   171  			if shapeOut != OneVregOutAtIn && in.AsmPos == 0 && in.Class == "vreg" {
   172  				shapeOut = OneVregOutAtIn
   173  			} else if in.Class != "immediate" { // on arm64 immediate may be indexing into output register, e.g. INS greg, vreg[1]
   174  				panic(fmt.Errorf("simdgen only supports output and input sharing the same position case of \"the first input is vreg and the only output\": %s", op))
   175  			}
   176  		}
   177  		if in.Class == "immediate" {
   178  			// A manual check on XED data found that AMD64 SIMD instructions at most
   179  			// have 1 immediates. So we don't need to check this here.
   180  			if *in.Bits != 8 {
   181  				panic(fmt.Errorf("simdgen only supports immediates of 8 bits: %s", op))
   182  			}
   183  			hasImm = true
   184  			immAsmPos = in.AsmPos
   185  			if immAsmPos == outputReg {
   186  				immOpIdx += "Out"
   187  			}
   188  		} else if in.Class == "mask" {
   189  			maskCount++
   190  		} else {
   191  			if immAsmPos == in.AsmPos {
   192  				immOpIdx += fmt.Sprintf("In%d", in.AsmPos)
   193  			}
   194  			hasVreg = true
   195  		}
   196  	}
   197  	opNoImm = *op
   198  
   199  	removeImm := func(o *Operation) {
   200  		o.In = o.In[1:]
   201  		if o.In[0].Class == "immediate" {
   202  			o.In = o.In[1:] // e.g. arm64 INS Vn[0], Vd[imm]
   203  		}
   204  	}
   205  	if hasImm {
   206  		removeImm(&opNoImm)
   207  		if op.In[0].Const != nil {
   208  			if op.In[0].ImmOffset != nil {
   209  				immType = ConstVarImm
   210  			} else {
   211  				immType = ConstImm
   212  			}
   213  		} else if op.In[0].ImmOffset != nil {
   214  			if op.In[0].ImmMax != nil {
   215  				immType = VarImmLim
   216  			} else {
   217  				immType = VarImm
   218  			}
   219  		} else {
   220  			panic(fmt.Errorf("simdgen requires imm to have at least one of ImmOffset or Const set: %s", op))
   221  		}
   222  	} else {
   223  		immType = NoImm
   224  	}
   225  	if maskCount == 0 {
   226  		maskType = NoMask
   227  	} else {
   228  		maskType = OneMask
   229  	}
   230  	checkPureMask := func() bool {
   231  		if hasImm {
   232  			panic(fmt.Errorf("simdgen does not support immediates in pure mask operations: %s", op))
   233  		}
   234  		if hasVreg {
   235  			panic(fmt.Errorf("simdgen does not support more than 1 mask in non-pure mask operations: %s", op))
   236  		}
   237  		return false
   238  	}
   239  	if !hasImm && maskCount == 0 {
   240  		shapeIn = PureVregIn
   241  		if hasListIn {
   242  			shapeIn = VlistIn
   243  		}
   244  	} else if !hasImm && maskCount > 0 {
   245  		if maskCount == 1 {
   246  			shapeIn = OneKmaskIn
   247  		} else {
   248  			if checkPureMask() {
   249  				return
   250  			}
   251  			shapeIn = PureKmaskIn
   252  			maskType = AllMasks
   253  		}
   254  	} else if hasImm && maskCount == 0 {
   255  		shapeIn = OneImmIn
   256  	} else {
   257  		if maskCount == 1 {
   258  			shapeIn = OneKmaskImmIn
   259  		} else {
   260  			checkPureMask()
   261  			return
   262  		}
   263  	}
   264  	return
   265  }
   266  
   267  // regShape returns a string representation of the register shape.
   268  func (op *Operation) regShape(mem memShape) (string, error) {
   269  	_, _, _, _, gOp, _ := op.shape()
   270  	var regInfo, fixedName string
   271  	var vRegInCnt, gRegInCnt, kMaskInCnt, vRegOutCnt, gRegOutCnt, kMaskOutCnt, memInCnt, memOutCnt int
   272  	for i, in := range gOp.In {
   273  		switch in.Class {
   274  		case "vreg":
   275  			vRegInCnt++
   276  		case "greg":
   277  			gRegInCnt++
   278  		case "mask":
   279  			kMaskInCnt++
   280  		case "memory":
   281  			if mem != VregMemIn {
   282  				panic("simdgen only knows VregMemIn in regShape")
   283  			}
   284  			memInCnt++
   285  			vRegInCnt++
   286  		}
   287  		if in.FixedReg != nil {
   288  			fixedName = fmt.Sprintf("%sAtIn%d", *in.FixedReg, i)
   289  		}
   290  	}
   291  	for i, out := range gOp.Out {
   292  		// If class overwrite is happening, that's not really a mask but a vreg.
   293  		if out.Class == "vreg" || out.OverwriteClass != nil {
   294  			vRegOutCnt++
   295  		} else if out.Class == "greg" {
   296  			gRegOutCnt++
   297  		} else if out.Class == "mask" {
   298  			kMaskOutCnt++
   299  		} else if out.Class == "memory" {
   300  			if mem != VregMemIn {
   301  				panic("simdgen only knows VregMemIn in regShape")
   302  			}
   303  			vRegOutCnt++
   304  			memOutCnt++
   305  		}
   306  		if out.FixedReg != nil {
   307  			fixedName = fmt.Sprintf("%sAtIn%d", *out.FixedReg, i)
   308  		}
   309  	}
   310  	var inRegs, inMasks, outRegs, outMasks string
   311  
   312  	rmAbbrev := func(s string, i int) string {
   313  		if i == 0 {
   314  			return ""
   315  		}
   316  		if i == 1 {
   317  			return s
   318  		}
   319  		return fmt.Sprintf("%s%d", s, i)
   320  
   321  	}
   322  
   323  	inRegs = rmAbbrev("v", vRegInCnt)
   324  	inRegs += rmAbbrev("gp", gRegInCnt)
   325  	inMasks = rmAbbrev("k", kMaskInCnt)
   326  
   327  	outRegs = rmAbbrev("v", vRegOutCnt)
   328  	outRegs += rmAbbrev("gp", gRegOutCnt)
   329  	outMasks = rmAbbrev("k", kMaskOutCnt)
   330  
   331  	if kMaskInCnt == 0 && kMaskOutCnt == 0 && gRegInCnt == 0 && gRegOutCnt == 0 {
   332  		// For pure v we can abbreviate it as v%d%d.
   333  		regInfo = fmt.Sprintf("v%d%d", vRegInCnt, vRegOutCnt)
   334  	} else if kMaskInCnt == 0 && kMaskOutCnt == 0 {
   335  		regInfo = fmt.Sprintf("%s%s", inRegs, outRegs)
   336  	} else {
   337  		regInfo = fmt.Sprintf("%s%s%s%s", inRegs, inMasks, outRegs, outMasks)
   338  	}
   339  	if memInCnt > 0 {
   340  		if memInCnt == 1 {
   341  			regInfo += "load"
   342  		} else {
   343  			panic("simdgen does not understand more than 1 mem op as of now")
   344  		}
   345  	}
   346  	if memOutCnt > 0 {
   347  		panic("simdgen does not understand memory as output as of now")
   348  	}
   349  	regInfo += fixedName
   350  	return regInfo, nil
   351  }
   352  
   353  // sortOperand sorts op.In by putting immediates first, then vreg, and mask the last.
   354  // TODO: verify that this is a safe assumption of the prog structure.
   355  // from my observation looks like in asm, imms are always the first,
   356  // masks are always the last, with vreg in between.
   357  func (op *Operation) sortOperand() {
   358  	priority := map[string]int{"immediate": 0, "vreg": 1, "greg": 1, "mask": 2}
   359  	sort.SliceStable(op.In, func(i, j int) bool {
   360  		pi := priority[op.In[i].Class]
   361  		pj := priority[op.In[j].Class]
   362  		if pi != pj {
   363  			return pi < pj
   364  		}
   365  		return op.In[i].AsmPos < op.In[j].AsmPos
   366  	})
   367  }
   368  
   369  // adjustAsm adjusts the asm to make it align with Go's assembler.
   370  func (op *Operation) adjustAsm() {
   371  	if op.Asm == "VCVTTPD2DQ" || op.Asm == "VCVTTPD2UDQ" ||
   372  		op.Asm == "VCVTQQ2PS" || op.Asm == "VCVTUQQ2PS" ||
   373  		op.Asm == "VCVTPD2PS" {
   374  		switch *op.In[0].Bits {
   375  		case 128:
   376  			op.Asm += "X"
   377  		case 256:
   378  			op.Asm += "Y"
   379  		}
   380  	}
   381  }
   382  
   383  // goNormalType returns the Go type name for the result of an Op that
   384  // does not return a vector, i.e., that returns a result in a general
   385  // register.  Currently there's only one family of Ops in Go's simd library
   386  // that does this (GetElem), and so this is specialized to work for that,
   387  // but the problem (mismatch between hardware register width and Go type
   388  // width) seems likely to recur if there are any other cases.
   389  func (op Operation) goNormalType() string {
   390  	if op.Go == "GetElem" {
   391  		// GetElem returns an element of the vector into a general register
   392  		// but as far as the hardware is concerned, that result is either 32
   393  		// or 64 bits wide, no matter what the vector element width is.
   394  		// This is not "wrong" but it is not the right answer for Go source code.
   395  		// To get the Go type right, combine the base type ("int", "uint", "float"),
   396  		// with the input vector element width in bits (8,16,32,64).
   397  
   398  		at := 0 // proper value of at depends on whether immediate was stripped or not
   399  		if op.In[at].Class == "immediate" {
   400  			at++
   401  		}
   402  		return fmt.Sprintf("%s%d", *op.Out[0].Base, *op.In[at].ElemBits)
   403  	}
   404  	panic(fmt.Errorf("Implement goNormalType for %v", op))
   405  }
   406  
   407  // SSAType returns the string for the type reference in SSA generation,
   408  // for example in the intrinsics generating template.
   409  func (op Operation) SSAType() string {
   410  	if op.Out[0].Class == "greg" {
   411  		return fmt.Sprintf("types.Types[types.T%s]", strings.ToUpper(op.goNormalType()))
   412  	}
   413  	return fmt.Sprintf("types.TypeVec%d", *op.Out[0].Bits)
   414  }
   415  
   416  // GoType returns the Go type returned by this operation (relative to the simd package),
   417  // for example "int32" or "Int8x16".  This is used in a template.
   418  func (op Operation) GoType() string {
   419  	if op.Out[0].Class == "greg" {
   420  		return op.goNormalType()
   421  	}
   422  	return *op.Out[0].Go
   423  }
   424  
   425  // ImmName returns the name to use for an operation's immediate operand.
   426  // This can be overridden in the yaml with "name" on an operand,
   427  // otherwise, for now, "constant"
   428  func (op Operation) ImmName() string {
   429  	return op.Op0Name("constant")
   430  }
   431  
   432  func (op Operation) ImmType() string {
   433  	if strings.Contains(op.Go, "Shift") || strings.Contains(op.Go, "Rotate") {
   434  		return "uint64"
   435  	}
   436  	return "uint8"
   437  }
   438  
   439  func (o Operand) OpName(s string) string {
   440  	if n := o.Name; n != nil {
   441  		return *n
   442  	}
   443  	if o.Class == "mask" {
   444  		return "mask"
   445  	}
   446  	return s
   447  }
   448  
   449  func (o Operand) OpNameAndType(s string) string {
   450  	return o.OpName(s) + " " + *o.Go
   451  }
   452  
   453  // GoExported returns [Go] with first character capitalized.
   454  func (op Operation) GoExported() string {
   455  	return capitalizeFirst(op.Go)
   456  }
   457  
   458  // DocumentationExported returns [Documentation] with method name capitalized.
   459  func (op Operation) DocumentationExported() string {
   460  	return strings.ReplaceAll(op.Documentation, op.Go, op.GoExported())
   461  }
   462  
   463  // Op0Name returns the name to use for the 0 operand,
   464  // if any is present, otherwise the parameter is used.
   465  func (op Operation) Op0Name(s string) string {
   466  	return op.In[0].OpName(s)
   467  }
   468  
   469  // Op1Name returns the name to use for the 1 operand,
   470  // if any is present, otherwise the parameter is used.
   471  func (op Operation) Op1Name(s string) string {
   472  	return op.In[1].OpName(s)
   473  }
   474  
   475  // Op2Name returns the name to use for the 2 operand,
   476  // if any is present, otherwise the parameter is used.
   477  func (op Operation) Op2Name(s string) string {
   478  	return op.In[2].OpName(s)
   479  }
   480  
   481  // Op3Name returns the name to use for the 3 operand,
   482  // if any is present, otherwise the parameter is used.
   483  func (op Operation) Op3Name(s string) string {
   484  	return op.In[3].OpName(s)
   485  }
   486  
   487  // Op0NameAndType returns the name and type to use for
   488  // the 0 operand, if a name is provided, otherwise
   489  // the parameter value is used as the default.
   490  func (op Operation) Op0NameAndType(s string) string {
   491  	return op.In[0].OpNameAndType(s)
   492  }
   493  
   494  // Op1NameAndType returns the name and type to use for
   495  // the 1 operand, if a name is provided, otherwise
   496  // the parameter value is used as the default.
   497  func (op Operation) Op1NameAndType(s string) string {
   498  	return op.In[1].OpNameAndType(s)
   499  }
   500  
   501  // Op2NameAndType returns the name and type to use for
   502  // the 2 operand, if a name is provided, otherwise
   503  // the parameter value is used as the default.
   504  func (op Operation) Op2NameAndType(s string) string {
   505  	return op.In[2].OpNameAndType(s)
   506  }
   507  
   508  // Op3NameAndType returns the name and type to use for
   509  // the 3 operand, if a name is provided, otherwise
   510  // the parameter value is used as the default.
   511  func (op Operation) Op3NameAndType(s string) string {
   512  	return op.In[3].OpNameAndType(s)
   513  }
   514  
   515  // Op4NameAndType returns the name and type to use for
   516  // the 4 operand, if a name is provided, otherwise
   517  // the parameter value is used as the default.
   518  func (op Operation) Op4NameAndType(s string) string {
   519  	return op.In[4].OpNameAndType(s)
   520  }
   521  
   522  var immClasses []string = []string{"BAD0Imm", "BAD1Imm", "op1Imm", "op2Imm", "op3Imm", "op4Imm"}
   523  var classes []string = []string{"BAD0", "op1", "op2", "op3", "op4"}
   524  
   525  // classifyOp returns a classification string, modified operation, and perhaps error based
   526  // on the stub and intrinsic shape for the operation.
   527  // The classification string is in the regular expression set "op[1234](Imm(8)?)?(_<order>)?"
   528  // where the "<order>" suffix is optionally attached to the Operation in its input yaml.
   529  // The classification string is used to select a template or a clause of a template
   530  // for intrinsics declaration and the ssagen intrinisics glue code in the compiler.
   531  func classifyOp(op Operation) (string, Operation, error) {
   532  	_, _, _, immType, gOp, _ := op.shape()
   533  
   534  	var class string
   535  
   536  	if immType == VarImm || immType == VarImmLim || immType == ConstVarImm {
   537  		switch l := len(op.In); l {
   538  		case 1:
   539  			return "", op, fmt.Errorf("simdgen does not recognize this operation of only immediate input: %s", op)
   540  		case 2, 3, 4, 5:
   541  			if immType == VarImmLim {
   542  				if len(op.In)-len(gOp.In) == 2 {
   543  					class = immClasses[l-1] // arm64: do not account const 0 imm in INS Vn[0], Vd[imm]
   544  				} else {
   545  					class = immClasses[l] // known immediate maximum value
   546  				}
   547  			} else {
   548  				// No known maximum: default to full uint8 range via the "8"-suffixed
   549  				// template variants (e.g. "op2Imm"+"8" → "op2Imm8", mapping to
   550  				// opLen2Imm8 which hardcodes immMax=255).
   551  				class = immClasses[l] + "8"
   552  			}
   553  		default:
   554  			return "", op, fmt.Errorf("simdgen does not recognize this operation of input length %d: %s", len(op.In), op)
   555  		}
   556  		if order := op.OperandOrder; order != nil {
   557  			class += "_" + *order
   558  		}
   559  		return class, op, nil
   560  	} else {
   561  		switch l := len(gOp.In); l {
   562  		case 1, 2, 3, 4:
   563  			class = classes[l]
   564  		default:
   565  			return "", op, fmt.Errorf("simdgen does not recognize this operation of input length %d: %s", len(op.In), op)
   566  		}
   567  		if order := op.OperandOrder; order != nil {
   568  			class += "_" + *order
   569  		}
   570  		return class, gOp, nil
   571  	}
   572  }
   573  
   574  func checkVecAsScalar(op Operation) (idx int, err error) {
   575  	idx = -1
   576  	sSize := 0
   577  	for i, o := range op.In {
   578  		if o.TreatLikeAScalarOfSize != nil {
   579  			if idx == -1 {
   580  				idx = i
   581  				sSize = *o.TreatLikeAScalarOfSize
   582  				if sSize == 0 && CurrentArch().Arch == "arm64" {
   583  					sSize = *o.ElemBits // treating lane 0 as element-sized fp scalar, e.g. INS Vn[0], Vd[imm]
   584  				}
   585  			} else {
   586  				err = fmt.Errorf("simdgen only supports one TreatLikeAScalarOfSize in the arg list: %s", op)
   587  				return
   588  			}
   589  		}
   590  	}
   591  	if idx >= 0 {
   592  		if sSize != 8 && sSize != 16 && sSize != 32 && sSize != 64 {
   593  			err = fmt.Errorf("simdgen does not recognize this uint size: %d, %s", sSize, op)
   594  			return
   595  		}
   596  	}
   597  	return
   598  }
   599  
   600  func rewriteVecAsScalarRegInfo(op Operation, regInfo string) (string, error) {
   601  	idx, err := checkVecAsScalar(op)
   602  	if err != nil {
   603  		return "", err
   604  	}
   605  	if idx != -1 {
   606  		if regInfo == "v21" {
   607  			regInfo = "vfpv"
   608  		} else if regInfo == "v2kv" {
   609  			regInfo = "vfpkv"
   610  		} else if regInfo == "v31" {
   611  			regInfo = "v2fpv"
   612  		} else if regInfo == "v3kv" {
   613  			regInfo = "v2fpkv"
   614  		} else if regInfo == "v21ResultInArg0ImmOutIn1" {
   615  			regInfo = "vfpvResultInArg0ImmOutIn1"
   616  		} else {
   617  			return "", fmt.Errorf("simdgen does not recognize uses of treatLikeAScalarOfSize with op regShape %s in op: %s", regInfo, op)
   618  		}
   619  	}
   620  	return regInfo, nil
   621  }
   622  
   623  func rewriteLastVregToMem(op Operation) Operation {
   624  	newIn := make([]Operand, len(op.In))
   625  	lastVregIdx := -1
   626  	for i := range len(op.In) {
   627  		newIn[i] = op.In[i]
   628  		if op.In[i].Class == "vreg" {
   629  			lastVregIdx = i
   630  		}
   631  	}
   632  	// vbcst operations put their mem op always as the last vreg.
   633  	if lastVregIdx == -1 {
   634  		panic("simdgen cannot find one vreg in the mem op vreg original")
   635  	}
   636  	newIn[lastVregIdx].Class = "memory"
   637  	op.In = newIn
   638  
   639  	return op
   640  }
   641  
   642  // dedup is deduping operations in the full structure level.
   643  func dedup(ops []Operation) (deduped []Operation) {
   644  	for _, op := range ops {
   645  		seen := false
   646  		for _, dop := range deduped {
   647  			if reflect.DeepEqual(op, dop) {
   648  				seen = true
   649  				break
   650  			}
   651  		}
   652  		if !seen {
   653  			deduped = append(deduped, op)
   654  		}
   655  	}
   656  	return
   657  }
   658  
   659  func (op Operation) GenericName() string {
   660  	if op.OperandOrder != nil {
   661  		switch *op.OperandOrder {
   662  		case "21Type1", "231Type1":
   663  			// Permute uses operand[1] for method receiver.
   664  			return op.Go + *op.In[1].Go
   665  		}
   666  	}
   667  	if op.In[0].Class == "immediate" {
   668  		if op.In[1].Class == "immediate" {
   669  			return op.Go + *op.In[2].Go // e.g. arm64 INS Vn[0], Vd[imm]
   670  		}
   671  		return op.Go + *op.In[1].Go
   672  	}
   673  	return op.Go + *op.In[0].Go
   674  }
   675  
   676  // dedupGodef is deduping operations in [Op.Go]+[*Op.In[0].Go] level.
   677  // By deduping, it means picking the least advanced architecture that satisfy the requirement:
   678  // AVX512 will be least preferred.
   679  // If FlagNoDedup is set, it will report the duplicates to the console.
   680  func dedupGodef(ops []Operation) ([]Operation, error) {
   681  	seen := map[string][]Operation{}
   682  	for _, op := range ops {
   683  		_, _, _, _, gOp, _ := op.shape()
   684  
   685  		gN := gOp.GenericName()
   686  		seen[gN] = append(seen[gN], op)
   687  	}
   688  	if *FlagReportDup {
   689  		for gName, dup := range seen {
   690  			if len(dup) > 1 {
   691  				log.Printf("Duplicate for %s:\n", gName)
   692  				for _, op := range dup {
   693  					log.Printf("%s\n", op)
   694  				}
   695  			}
   696  		}
   697  		return ops, nil
   698  	}
   699  	isAVX512 := func(op Operation) bool {
   700  		return strings.Contains(op.CPUFeature, "AVX512")
   701  	}
   702  	deduped := []Operation{}
   703  	for _, dup := range seen {
   704  		if len(dup) > 1 {
   705  			slices.SortFunc(dup, func(i, j Operation) int {
   706  				// Put non-AVX512 candidates at the beginning
   707  				if !isAVX512(i) && isAVX512(j) {
   708  					return -1
   709  				}
   710  				if isAVX512(i) && !isAVX512(j) {
   711  					return 1
   712  				}
   713  				if i.CPUFeature != j.CPUFeature {
   714  					return strings.Compare(i.CPUFeature, j.CPUFeature)
   715  				}
   716  				// Weirdly Intel sometimes has duplicated definitions for the same instruction,
   717  				// this confuses the XED mem-op merge logic: [MemFeature] will only be attached to an instruction
   718  				// for only once, which means that for essentially duplicated instructions only one will have the
   719  				// proper [MemFeature] set. We have to make this sort deterministic for [MemFeature].
   720  				if i.MemFeatures != nil && j.MemFeatures == nil {
   721  					return -1
   722  				}
   723  				if i.MemFeatures == nil && j.MemFeatures != nil {
   724  					return 1
   725  				}
   726  				if i.Commutative != j.Commutative {
   727  					if j.Commutative {
   728  						return -1
   729  					}
   730  					return 1
   731  				}
   732  				// Their order does not matter anymore, at least for now.
   733  				return 0
   734  			})
   735  		}
   736  		deduped = append(deduped, dup[0])
   737  	}
   738  	slices.SortFunc(deduped, compareOperations)
   739  	return deduped, nil
   740  }
   741  
   742  // Copy op.ConstImm to op.In[0].Const
   743  // This is a hack to reduce the size of defs we need for const imm operations.
   744  func copyConstImm(ops []Operation) error {
   745  	for _, op := range ops {
   746  		if op.ConstImm == nil {
   747  			continue
   748  		}
   749  		_, _, _, immType, _, _ := op.shape()
   750  
   751  		if immType == ConstImm || immType == ConstVarImm {
   752  			op.In[0].Const = op.ConstImm
   753  			// If the immediate operand is tagged with name:"@", it is fully constant;
   754  			// clear ImmOffset to ensure this is treated as ConstImm (no aux field).
   755  			if op.In[0].Name != nil && *op.In[0].Name == "@" {
   756  				op.In[0].ImmOffset = nil
   757  			}
   758  		}
   759  		// Otherwise, just not port it - e.g. {VPCMP[BWDQ] imm=0} and {VPCMPEQ[BWDQ]} are
   760  		// the same operations "Equal", [dedupgodef] should be able to distinguish them.
   761  	}
   762  	return nil
   763  }
   764  
   765  func capitalizeFirst(s string) string {
   766  	if s == "" {
   767  		return ""
   768  	}
   769  	// Convert the string to a slice of runes to handle multi-byte characters correctly.
   770  	r := []rune(s)
   771  	r[0] = unicode.ToUpper(r[0])
   772  	return string(r)
   773  }
   774  
   775  // overwrite corrects some errors due to:
   776  //   - The XED data is wrong
   777  //   - Go's SIMD API requirement, for example AVX2 compares should also produce masks.
   778  //     This rewrite has strict constraints, please see the error message.
   779  //     These constraints are also explointed in [writeSIMDRules], [writeSIMDMachineOps]
   780  //     and [writeSIMDSSA], please be careful when updating these constraints.
   781  func overwrite(ops []Operation) error {
   782  	hasClassOverwrite := false
   783  	overwrite := func(op []Operand, idx int, o Operation) error {
   784  		if op[idx].OverwriteElementBits != nil {
   785  			if op[idx].ElemBits == nil {
   786  				panic(fmt.Errorf("ElemBits is nil at operand %d of %v", idx, o))
   787  			}
   788  			*op[idx].ElemBits = *op[idx].OverwriteElementBits
   789  			*op[idx].Lanes = *op[idx].Bits / *op[idx].ElemBits
   790  			*op[idx].Go = fmt.Sprintf("%s%dx%d", capitalizeFirst(*op[idx].Base), *op[idx].ElemBits, *op[idx].Lanes)
   791  		}
   792  		if CurrentArch().Arch == "arm64" && op[idx].OverwriteClass != nil && *op[idx].OverwriteClass == "greg" {
   793  			if op[idx].OverwriteBase == nil {
   794  				panic(fmt.Errorf("simdgen: [OverwriteClass] must be set together with [OverwriteBase]: %s", op[idx]))
   795  			}
   796  			oBase := *op[idx].OverwriteBase
   797  			oClass := *op[idx].OverwriteClass
   798  			if oBase != "float" {
   799  				panic(fmt.Errorf("simdgen: [Class] overwrite must set [OverwriteBase] to float: %s", op[idx]))
   800  			}
   801  			if op[idx].Class != "vreg" {
   802  				panic(fmt.Errorf("simdgen: [Class] overwrite must be overwriting [Class] from vreg: %s", op[idx]))
   803  			}
   804  			// The low lane of vreg (with other lanes zeroed) also represents a regular floating point greg.
   805  			// This is supposed to be used only by special instructions like float GetElem
   806  			// and floating point vector reduction across-lanes to lane 0 like FMINNMV.
   807  			hasClassOverwrite = true
   808  			*op[idx].Base = oBase
   809  			op[idx].Class = oClass
   810  			*op[idx].Go = fmt.Sprintf("float%d", *op[idx].ElemBits)
   811  		} else if op[idx].OverwriteClass != nil {
   812  			if op[idx].OverwriteBase == nil {
   813  				panic(fmt.Errorf("simdgen: [OverwriteClass] must be set together with [OverwriteBase]: %s", op[idx]))
   814  			}
   815  			oBase := *op[idx].OverwriteBase
   816  			oClass := *op[idx].OverwriteClass
   817  			if oClass != "mask" {
   818  				panic(fmt.Errorf("simdgen: [Class] overwrite only supports overwriting to mask: %s", op[idx]))
   819  			}
   820  			if oBase != "int" {
   821  				panic(fmt.Errorf("simdgen: [Class] overwrite must set [OverwriteBase] to int: %s", op[idx]))
   822  			}
   823  			if op[idx].Class != "vreg" {
   824  				panic(fmt.Errorf("simdgen: [Class] overwrite must be overwriting [Class] from vreg: %s", op[idx]))
   825  			}
   826  			hasClassOverwrite = true
   827  			*op[idx].Base = oBase
   828  			op[idx].Class = oClass
   829  			*op[idx].Go = fmt.Sprintf("Mask%dx%d", *op[idx].ElemBits, *op[idx].Lanes)
   830  		} else if op[idx].OverwriteBase != nil {
   831  			oBase := *op[idx].OverwriteBase
   832  			*op[idx].Go = strings.ReplaceAll(*op[idx].Go, capitalizeFirst(*op[idx].Base), capitalizeFirst(oBase))
   833  			if op[idx].Class == "greg" {
   834  				*op[idx].Go = strings.ReplaceAll(*op[idx].Go, *op[idx].Base, oBase)
   835  			}
   836  			*op[idx].Base = oBase
   837  		} else if op[idx].OverwriteBits != nil {
   838  			if op[idx].Class != "greg" {
   839  				panic(fmt.Errorf("simdgen: [OverwriteBits] is only supported for greg int: %s", op[idx]))
   840  			}
   841  			*op[idx].Bits = *op[idx].OverwriteBits
   842  			*op[idx].Go = fmt.Sprintf("%s%d", *op[idx].Base, *op[idx].Bits)
   843  		}
   844  		return nil
   845  	}
   846  	for i, o := range ops {
   847  		hasClassOverwrite = false
   848  		for j := range ops[i].In {
   849  			if err := overwrite(ops[i].In, j, o); err != nil {
   850  				return err
   851  			}
   852  			if hasClassOverwrite {
   853  				return fmt.Errorf("simdgen does not support [OverwriteClass] in inputs: %s", ops[i])
   854  			}
   855  		}
   856  		for j := range ops[i].Out {
   857  			if err := overwrite(ops[i].Out, j, o); err != nil {
   858  				return err
   859  			}
   860  		}
   861  		if hasClassOverwrite {
   862  			for _, in := range ops[i].In {
   863  				if in.Class == "mask" {
   864  					return fmt.Errorf("simdgen only supports [OverwriteClass] for operations without mask inputs")
   865  				}
   866  			}
   867  		}
   868  	}
   869  	return nil
   870  }
   871  
   872  // reportXEDInconsistency reports potential XED inconsistencies.
   873  // We can add more fields to [Operation] to enable more checks and implement it here.
   874  // Supported checks:
   875  // [NameAndSizeCheck]: NAME[BWDQ] should set the elemBits accordingly.
   876  // This check is useful to find inconsistencies, then we can add overwrite fields to
   877  // those defs to correct them manually.
   878  func reportXEDInconsistency(ops []Operation) error {
   879  	for _, o := range ops {
   880  		if o.NameAndSizeCheck != nil {
   881  			suffixSizeMap := map[byte]int{'B': 8, 'W': 16, 'D': 32, 'Q': 64}
   882  			checkOperand := func(opr Operand) error {
   883  				if opr.ElemBits == nil {
   884  					return fmt.Errorf("simdgen expects elemBits to be set when performing NameAndSizeCheck")
   885  				}
   886  				if v, ok := suffixSizeMap[o.Asm[len(o.Asm)-1]]; !ok {
   887  					return fmt.Errorf("simdgen expects asm to end with [BWDQ] when performing NameAndSizeCheck")
   888  				} else {
   889  					if v != *opr.ElemBits {
   890  						return fmt.Errorf("simdgen finds NameAndSizeCheck inconsistency in def: %s", o)
   891  					}
   892  				}
   893  				return nil
   894  			}
   895  			for _, in := range o.In {
   896  				if in.Class != "vreg" && in.Class != "mask" {
   897  					continue
   898  				}
   899  				if in.TreatLikeAScalarOfSize != nil {
   900  					// This is an irregular operand, don't check it.
   901  					continue
   902  				}
   903  				if err := checkOperand(in); err != nil {
   904  					return err
   905  				}
   906  			}
   907  			for _, out := range o.Out {
   908  				if err := checkOperand(out); err != nil {
   909  					return err
   910  				}
   911  			}
   912  		}
   913  	}
   914  	return nil
   915  }
   916  
   917  func (o *Operation) hasMaskedMerging(maskType maskShape, outType outShape) bool {
   918  	if o.SpecialLower != nil {
   919  		// asmRule and argsMatchRule/earlyMatchRule should not affect masked merging
   920  		ok, _, _ := parseAsmRule(*o.SpecialLower)
   921  		if !ok {
   922  			ok, _, _ = parseArgsMatchRule(*o.SpecialLower)
   923  			if !ok {
   924  				return false
   925  			}
   926  		}
   927  	}
   928  	// BLEND and VMOVDQU are not user-facing ops so we should filter them out.
   929  	return o.OperandOrder == nil && maskType == OneMask && outType == OneVregOut &&
   930  		len(o.InVariant) == 1 && !strings.Contains(o.Asm, "BLEND") && !strings.Contains(o.Asm, "VMOVDQU")
   931  }
   932  
   933  func getVbcstData(s string) (string, string) {
   934  	feat1, feat2, found := strings.Cut(s, ";")
   935  	if !found || !strings.HasPrefix(feat1, "feat1=") || !strings.HasPrefix(feat2, "feat2=") {
   936  		panic(fmt.Sprintf("unexpected format for vbcst data: %s", s))
   937  	}
   938  	return strings.TrimPrefix(feat1, "feat1="), strings.TrimPrefix(feat2, "feat2=")
   939  }
   940  
   941  func (o Operation) String() string {
   942  	return pprints(o)
   943  }
   944  
   945  func (op Operand) String() string {
   946  	return pprints(op)
   947  }
   948  
   949  // hiHalfOpName constructs the SSA machine op name for a hi-half "2" variant.
   950  // For example: hiHalfAsm="VSHRN2", arrangement="4S" → "VSHRN2_4S".
   951  func hiHalfOpName(hiHalfAsm string, gOp Operation) string {
   952  	return hiHalfAsm + "_" + *gOp.Arrangement
   953  }
   954  
   955  // hiHalfRegShape derives the regShape for hi-half "2" variant ops from the base regShape.
   956  // For narrow "2": the "2" variant has an extra destination input (resultInArg0),
   957  // so vreg input count increases by 1 (e.g., "v11Imm" → "v21Imm").
   958  // For long "2": same shape as base.
   959  func hiHalfRegShape2(baseRegShape string, kind string) string {
   960  	if kind == "narrow" {
   961  		if len(baseRegShape) > 2 && baseRegShape[0] == 'v' {
   962  			inCnt := int(baseRegShape[1] - '0')
   963  			outCnt := int(baseRegShape[2] - '0')
   964  			rest := baseRegShape[3:]
   965  			return fmt.Sprintf("v%d%d%s", inCnt+1, outCnt, rest)
   966  		}
   967  		panic(fmt.Sprintf("hiHalfRegShape2: unexpected regShape %q for narrow kind", baseRegShape))
   968  	}
   969  	return baseRegShape
   970  }
   971  
   972  // hiHalfLoweringRegShape computes the lowering dispatch regShape for a hi-half operation.
   973  // Base ops get a suffix like "Narrow" or "Long" appended to their standard regShape.
   974  // "2" variant ops get a derived regShape + "2" suffix.
   975  func hiHalfLoweringRegShape(baseRegShape string, kind string, isVariant2 bool) string {
   976  	suffix := capitalizeFirst(kind)
   977  	if isVariant2 {
   978  		suffix += "2"
   979  		// For narrow "2", the regShape changes (extra vreg input for resultInArg0)
   980  		baseRegShape = hiHalfRegShape2(baseRegShape, kind)
   981  	}
   982  	return baseRegShape + suffix
   983  }
   984  

View as plain text