Source file src/simd/archsimd/_gen/simdgen/gen_simdMachineOps.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  	"bytes"
     9  	"fmt"
    10  	"log"
    11  	"sort"
    12  	"strings"
    13  )
    14  
    15  const simdMachineOpsTmpl = `
    16  package main
    17  
    18  func simd{{.ArchUpper}}Ops({{.RegInfoParams}}) []opData {
    19  	return []opData{
    20  {{- range .OpsData }}
    21  		{name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
    22  {{- end }}
    23  {{- range .OpsDataImm }}
    24  		{name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}", aux: "UInt8",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
    25  {{- end }}
    26  {{- range .OpsDataLoad}}
    27  		{name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}", aux: "SymOff", symEffect: "Read"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
    28  {{- end}}
    29  {{- range .OpsDataImmLoad}}
    30  		{name: "{{.OpName}}", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}",{{if .Comm}} commutative: true,{{end}} typ: "{{.Type}}", aux: "SymValAndOff", symEffect: "Read"{{if .ResultInArg0}}, resultInArg0: true{{end}}},
    31  {{- end}}
    32  {{- range .OpsDataMerging }}
    33  		{name: "{{.OpName}}Merging", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}", typ: "{{.Type}}", resultInArg0: true},
    34  {{- end }}
    35  {{- range .OpsDataImmMerging }}
    36  		{name: "{{.OpName}}Merging", argLength: {{.OpInLen}}, reg: {{.RegInfo}}, asm: "{{.Asm}}", aux: "UInt8", typ: "{{.Type}}", resultInArg0: true},
    37  {{- end }}
    38  	}
    39  }
    40  `
    41  
    42  // writeSIMDMachineOps generates the machine ops and writes it to simdAMD64ops.go
    43  // within the specified directory.
    44  func writeSIMDMachineOps(ops []Operation) *bytes.Buffer {
    45  	t := templateOf(simdMachineOpsTmpl, "simdAMD64Ops")
    46  	buffer := new(bytes.Buffer)
    47  	buffer.WriteString(generatedHeader())
    48  
    49  	type opData struct {
    50  		OpName       string
    51  		Asm          string
    52  		OpInLen      int
    53  		RegInfo      string
    54  		Comm         bool
    55  		Type         string
    56  		ResultInArg0 bool
    57  	}
    58  	type machineOpsData struct {
    59  		ArchUpper         string
    60  		RegInfoParams     string
    61  		OpsData           []opData
    62  		OpsDataImm        []opData
    63  		OpsDataLoad       []opData
    64  		OpsDataImmLoad    []opData
    65  		OpsDataMerging    []opData
    66  		OpsDataImmMerging []opData
    67  	}
    68  
    69  	archInfo := CurrentArch()
    70  
    71  	regInfoSet := archInfo.RegInfoSet
    72  	opsData := make([]opData, 0)
    73  	opsDataImm := make([]opData, 0)
    74  	opsDataLoad := make([]opData, 0)
    75  	opsDataImmLoad := make([]opData, 0)
    76  	opsDataMerging := make([]opData, 0)
    77  	opsDataImmMerging := make([]opData, 0)
    78  
    79  	// Determine the "best" version of an instruction to use
    80  	best := make(map[string]Operation)
    81  	var mOpOrder []string
    82  	countOverrides := func(s []Operand) int {
    83  		a := 0
    84  		for _, o := range s {
    85  			if o.OverwriteBase != nil {
    86  				a++
    87  			}
    88  		}
    89  		return a
    90  	}
    91  	for _, op := range ops {
    92  		_, _, maskType, _, gOp, _ := op.shape()
    93  		asm := machineOpName(maskType, gOp)
    94  		other, ok := best[asm]
    95  		if !ok {
    96  			best[asm] = op
    97  			mOpOrder = append(mOpOrder, asm)
    98  			continue
    99  		}
   100  		if !op.Commutative && other.Commutative { // if there's a non-commutative version of the op, it wins.
   101  			best[asm] = op
   102  			continue
   103  		}
   104  		// see if "op" is better than "other"
   105  		if countOverrides(op.In)+countOverrides(op.Out) < countOverrides(other.In)+countOverrides(other.Out) {
   106  			best[asm] = op
   107  		}
   108  	}
   109  
   110  	regInfoErrs := make([]error, 0)
   111  	regInfoMissing := make(map[string]bool, 0)
   112  	for _, asm := range mOpOrder {
   113  		op := best[asm]
   114  		shapeIn, shapeOut, maskType, _, gOp, _ := op.shape()
   115  
   116  		// TODO: all our masked operations are now zeroing, we need to generate machine ops with merging masks, maybe copy
   117  		// one here with a name suffix "Merging". The rewrite rules will need them.
   118  		makeRegInfo := func(op Operation, mem memShape) (string, error) {
   119  			regInfo, err := op.regShape(mem)
   120  			if err != nil {
   121  				panic(err)
   122  			}
   123  			regInfo, err = rewriteVecAsScalarRegInfo(op, regInfo)
   124  			if err != nil {
   125  				if mem == NoMem || mem == InvalidMem {
   126  					panic(err)
   127  				}
   128  				return "", err
   129  			}
   130  			if regInfo == "v01load" {
   131  				regInfo = "vload"
   132  			}
   133  			// Makes AVX512 operations use upper registers
   134  			if strings.Contains(op.CPUFeature, "AVX512") {
   135  				regInfo = strings.ReplaceAll(regInfo, "v", "w")
   136  			}
   137  			if _, ok := regInfoSet[regInfo]; !ok {
   138  				regInfoErrs = append(regInfoErrs, fmt.Errorf("unsupported register constraint, please update the template and AMD64Ops.go: %s.  Op is %s", regInfo, op))
   139  				regInfoMissing[regInfo] = true
   140  			}
   141  			return regInfo, nil
   142  		}
   143  		regInfo, err := makeRegInfo(op, NoMem)
   144  		if err != nil {
   145  			panic(err)
   146  		}
   147  		var outType string
   148  		if shapeOut == OneVregOut || shapeOut == OneVregOutAtIn || shapeOut == OneVregOutScalar || gOp.Out[0].OverwriteClass != nil {
   149  			// If class overwrite is happening, that's not really a mask but a vreg.
   150  			outType = fmt.Sprintf("Vec%d", *gOp.Out[0].Bits)
   151  		} else if shapeOut == OneGregOut {
   152  			outType = gOp.GoType() // this is a straight Go type, not a VecNNN type
   153  		} else if shapeOut == OneKmaskOut {
   154  			outType = "Mask"
   155  		} else {
   156  			panic(fmt.Errorf("simdgen does not recognize this output shape: %d", shapeOut))
   157  		}
   158  		resultInArg0 := false
   159  		if shapeOut == OneVregOutAtIn {
   160  			resultInArg0 = true
   161  		}
   162  		var memOpData *opData
   163  		regInfoMerging := regInfo
   164  		hasMerging := false
   165  		if op.MemFeatures != nil && *op.MemFeatures == "vbcst" {
   166  			// Right now we only have vbcst case
   167  			// Make a full vec memory variant.
   168  			opMem := rewriteLastVregToMem(op)
   169  			regInfo, err := makeRegInfo(opMem, VregMemIn)
   170  			if err != nil {
   171  				// Just skip it if it's non nill.
   172  				// an error could be triggered by [checkVecAsScalar].
   173  				// TODO: make [checkVecAsScalar] aware of mem ops.
   174  				if *Verbose {
   175  					log.Printf("Seen error: %e", err)
   176  				}
   177  			} else {
   178  				memOpData = &opData{asm + "load", gOp.Asm, len(gOp.In) + 1, regInfo, false, outType, resultInArg0}
   179  			}
   180  		}
   181  		hasMerging = gOp.hasMaskedMerging(maskType, shapeOut)
   182  		if hasMerging && !resultInArg0 {
   183  			// We have to copy the slice here because the sort will be visible from other
   184  			// aliases when no reslicing is happening.
   185  			newIn := make([]Operand, len(op.In), len(op.In)+1)
   186  			copy(newIn, op.In)
   187  			op.In = newIn
   188  			op.In = append(op.In, op.Out[0])
   189  			op.sortOperand()
   190  			regInfoMerging, err = makeRegInfo(op, NoMem)
   191  			if err != nil {
   192  				panic(err)
   193  			}
   194  		}
   195  
   196  		if shapeIn == OneImmIn || shapeIn == OneKmaskImmIn {
   197  			opsDataImm = append(opsDataImm, opData{asm, gOp.Asm, len(gOp.In), regInfo, gOp.Commutative, outType, resultInArg0})
   198  			if memOpData != nil {
   199  				if *op.MemFeatures != "vbcst" {
   200  					panic("simdgen only knows vbcst for mem ops for now")
   201  				}
   202  				opsDataImmLoad = append(opsDataImmLoad, *memOpData)
   203  			}
   204  			if hasMerging {
   205  				mergingLen := len(gOp.In)
   206  				if !resultInArg0 {
   207  					mergingLen++
   208  				}
   209  				opsDataImmMerging = append(opsDataImmMerging, opData{asm, gOp.Asm, mergingLen, regInfoMerging, gOp.Commutative, outType, resultInArg0})
   210  			}
   211  		} else {
   212  			opsData = append(opsData, opData{asm, gOp.Asm, len(gOp.In), regInfo, gOp.Commutative, outType, resultInArg0})
   213  			if memOpData != nil {
   214  				if *op.MemFeatures != "vbcst" {
   215  					panic("simdgen only knows vbcst for mem ops for now")
   216  				}
   217  				opsDataLoad = append(opsDataLoad, *memOpData)
   218  			}
   219  			if hasMerging {
   220  				mergingLen := len(gOp.In)
   221  				if !resultInArg0 {
   222  					mergingLen++
   223  				}
   224  				opsDataMerging = append(opsDataMerging, opData{asm, gOp.Asm, mergingLen, regInfoMerging, gOp.Commutative, outType, resultInArg0})
   225  			}
   226  		}
   227  		// Generate hi-half "2" variant machine op
   228  		if gOp.HiHalfAsm != nil {
   229  			opsDataTarget := &opsData
   230  			if shapeIn == OneImmIn || shapeIn == OneKmaskImmIn {
   231  				opsDataTarget = &opsDataImm
   232  			}
   233  			kind := op.hiHalfKind()
   234  			asm2Name := hiHalfOpName(*gOp.HiHalfAsm, gOp)
   235  			argLen2 := len(gOp.In)
   236  			regInfo2 := regInfo
   237  			resultInArg02 := false
   238  			if kind == "narrow" {
   239  				argLen2++ // extra vreg input for destination
   240  				regInfo2 = hiHalfRegShape2(regInfo, kind)
   241  				resultInArg02 = true
   242  			}
   243  			if _, ok := regInfoSet[regInfo2]; !ok {
   244  				regInfoErrs = append(regInfoErrs, fmt.Errorf("unsupported hi-half register constraint: %s for op %s", regInfo2, asm2Name))
   245  				regInfoMissing[regInfo2] = true
   246  			} else {
   247  				*opsDataTarget = append(*opsDataTarget, opData{asm2Name, *gOp.HiHalfAsm, argLen2, regInfo2, gOp.Commutative, outType, resultInArg02})
   248  			}
   249  		}
   250  	}
   251  	if len(regInfoErrs) != 0 {
   252  		for _, e := range regInfoErrs {
   253  			log.Printf("Errors: %e\n", e)
   254  		}
   255  		panic(fmt.Errorf("these regInfo unseen: %v", regInfoMissing))
   256  	}
   257  	sort.Slice(opsData, func(i, j int) bool {
   258  		return compareNatural(opsData[i].OpName, opsData[j].OpName) < 0
   259  	})
   260  	sort.Slice(opsDataImm, func(i, j int) bool {
   261  		return compareNatural(opsDataImm[i].OpName, opsDataImm[j].OpName) < 0
   262  	})
   263  	sort.Slice(opsDataLoad, func(i, j int) bool {
   264  		return compareNatural(opsDataLoad[i].OpName, opsDataLoad[j].OpName) < 0
   265  	})
   266  	sort.Slice(opsDataImmLoad, func(i, j int) bool {
   267  		return compareNatural(opsDataImmLoad[i].OpName, opsDataImmLoad[j].OpName) < 0
   268  	})
   269  	sort.Slice(opsDataMerging, func(i, j int) bool {
   270  		return compareNatural(opsDataMerging[i].OpName, opsDataMerging[j].OpName) < 0
   271  	})
   272  	sort.Slice(opsDataImmMerging, func(i, j int) bool {
   273  		return compareNatural(opsDataImmMerging[i].OpName, opsDataImmMerging[j].OpName) < 0
   274  	})
   275  
   276  	err := t.Execute(buffer, machineOpsData{
   277  		ArchUpper:         archInfo.ArchUpper,
   278  		RegInfoParams:     archInfo.RegInfoParams,
   279  		OpsData:           opsData,
   280  		OpsDataImm:        opsDataImm,
   281  		OpsDataLoad:       opsDataLoad,
   282  		OpsDataImmLoad:    opsDataImmLoad,
   283  		OpsDataMerging:    opsDataMerging,
   284  		OpsDataImmMerging: opsDataImmMerging,
   285  	})
   286  	if err != nil {
   287  		panic(fmt.Errorf("failed to execute template: %w", err))
   288  	}
   289  
   290  	return buffer
   291  }
   292  

View as plain text