Source file src/simd/archsimd/_gen/simdgen/gen_simdrules.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  	"slices"
    11  	"strings"
    12  	"text/template"
    13  	"unicode"
    14  
    15  	"simd/archsimd/_gen/sgutil"
    16  )
    17  
    18  type tplRuleData struct {
    19  	TplName        string // e.g. "sftimm"
    20  	GoOp           string // e.g. "ShiftAllLeft"
    21  	GoType         string // e.g. "Uint32x8"
    22  	Args           string // e.g. "x y"
    23  	Asm            string // e.g. "VPSLLD256"
    24  	ArgsOut        string // e.g. "x y"
    25  	MaskInConvert  string // e.g. "VPMOVVec32x8ToM"
    26  	MaskOutConvert string // e.g. "VPMOVMToVec32x8"
    27  	ElementSize    int    // e.g. 32
    28  	Size           int    // e.g. 128
    29  	ArgsLoadAddr   string // [Args] with its last vreg arg being a concrete "(VMOVDQUload* ptr mem)", and might contain mask.
    30  	ArgsAddr       string // [Args] with its last vreg arg being replaced by "ptr", and might contain mask, and with a "mem" at the end.
    31  	FeatCheck      string // e.g. "v.Block.CPUfeatures.hasFeature(CPUavx512)" -- for a ssa/_gen rules file.
    32  	RuleCond       string // e.g. "a==0" -- condition for asmRule or argsMatchRule.
    33  	RuleOut        string // e.g. "y" -- output of an asmRule or argsMatchRule.
    34  	RuleArgs       string // custom args pattern for argsMatchRule.
    35  	Rule           string // the whole rule
    36  }
    37  
    38  // Helper type to make template map initialization less repetitive
    39  // (and also remove a chance for errors.)
    40  type ruleTemplateMap struct {
    41  	sgutil.InsertMap[string, *template.Template]
    42  }
    43  
    44  // Add creates a template named "name" after appending "// {{.TmplName}}\n" to the
    45  // template, and returns the input so that additions may be chained.
    46  // This helps make template initialization easy to order and easy to read.
    47  func (rtm *ruleTemplateMap) Add(name string, templ string) *ruleTemplateMap {
    48  	// Append debugging comment AND end of line
    49  	templ += " // {{.TplName}}\n"
    50  	ct := sgutil.TemplateNamed(name, templ)
    51  	rtm.InsertMap.Put(name, ct)
    52  	return rtm
    53  }
    54  
    55  var (
    56  	// ORDER MATTERS.  These should appear in most-to-least-specific order
    57  	// TODO: vregMemFeatCheck is not necessarily in the right place; there was an order,
    58  	// this was copied from it, but vregMemFeatCheck was not in it.  It's not clear
    59  	// it was ever used.
    60  	// It's also not clear that this is strictly most-to-least-specific?
    61  	ruleTemplates = new(ruleTemplateMap).
    62  		Add("masksftimm", `({{.Asm}} x (MOVQconst [c]) mask) => ({{.Asm}}const [amd64CapAVXShift(c)] x mask)`).
    63  		Add("sftimm", `({{.Asm}} x (MOVQconst [c])) => ({{.Asm}}const [amd64CapAVXShift(c)] x)`).
    64  		Add("maskInMaskOut", `({{.GoOp}}{{.GoType}} {{.Args}} mask) => ({{.MaskOutConvert}} ({{.Asm}} {{.ArgsOut}} ({{.MaskInConvert}} <types.TypeMask> mask)))`).
    65  		Add("maskOut", `({{.GoOp}}{{.GoType}} {{.Args}}) => ({{.MaskOutConvert}} ({{.Asm}} {{.ArgsOut}}))`).
    66  		Add("maskIn", `({{.GoOp}}{{.GoType}} {{.Args}} mask) => ({{.Asm}} {{.ArgsOut}} ({{.MaskInConvert}} <types.TypeMask> mask))`).
    67  		Add("pureVreg", `({{.GoOp}}{{.GoType}} {{.Args}}) => ({{.Asm}} {{.ArgsOut}})`).
    68  		Add("vregMem", `({{.Asm}} {{.ArgsLoadAddr}}) && canMergeLoad(v, l) && clobber(l) => ({{.Asm}}load {{.ArgsAddr}})`).
    69  		Add("vregMemFeatCheck", `({{.Asm}} {{.ArgsLoadAddr}}) && {{.FeatCheck}} && canMergeLoad(v, l) && clobber(l) => ({{.Asm}}load {{.ArgsAddr}})`).
    70  		Add("asmRule", `({{.Asm}} {{.Args}}) {{.RuleCond}} => {{.RuleOut}}`).
    71  		Add("specialLower", `{{.Rule}}`)
    72  )
    73  
    74  func (d tplRuleData) MaskOptimization(asmCheck map[string]bool) string {
    75  	asmNoMask := d.Asm
    76  	if i := strings.Index(asmNoMask, "Masked"); i == -1 {
    77  		return ""
    78  	}
    79  	asmNoMask = strings.ReplaceAll(asmNoMask, "Masked", "")
    80  	if asmCheck[asmNoMask] == false {
    81  		return ""
    82  	}
    83  
    84  	for _, nope := range []string{"VMOVDQU", "VPCOMPRESS", "VCOMPRESS", "VPEXPAND", "VEXPAND", "VPBLENDM", "VMOVUP"} {
    85  		if strings.HasPrefix(asmNoMask, nope) {
    86  			return ""
    87  		}
    88  	}
    89  
    90  	size := asmNoMask[len(asmNoMask)-3:]
    91  	if strings.HasSuffix(asmNoMask, "const") {
    92  		sufLen := len("128const")
    93  		size = asmNoMask[len(asmNoMask)-sufLen:][:3]
    94  	}
    95  	switch size {
    96  	case "128", "256", "512":
    97  	default:
    98  		panic("Unexpected operation size on " + d.Asm)
    99  	}
   100  
   101  	switch d.ElementSize {
   102  	case 8, 16, 32, 64:
   103  	default:
   104  		panic(fmt.Errorf("Unexpected operation width %d on %v", d.ElementSize, d.Asm))
   105  	}
   106  
   107  	return fmt.Sprintf("(VMOVDQU%dMasked%s (%s %s) mask) => (%s %s mask)\n", d.ElementSize, size, asmNoMask, d.Args, d.Asm, d.Args)
   108  }
   109  
   110  func compareTplRuleData(x, y tplRuleData) int {
   111  	if c := compareNatural(x.GoOp, y.GoOp); c != 0 {
   112  		return c
   113  	}
   114  	if c := compareNatural(x.GoType, y.GoType); c != 0 {
   115  		return c
   116  	}
   117  	if c := compareNatural(x.Args, y.Args); c != 0 {
   118  		return c
   119  	}
   120  	if x.TplName == y.TplName {
   121  		return 0
   122  	}
   123  	return ruleTemplates.Compare(x.TplName, y.TplName)
   124  }
   125  
   126  // parseAsmRule tries to parse given string as it would be asmRule:
   127  // if <cond> => <out>
   128  // Return false, "", "" if not matched, otherwise true, condition, rule output.
   129  // For example:
   130  // rule:"if a==0 => (VADD4S x y)" can be used to provide addional
   131  // lowering for an instruction which doesn't support zero immediate encoding:
   132  // (VUSRA4S [a] x y) && a==0 => (VADD4S x y)
   133  func parseAsmRule(rule string) (bool, string, string) {
   134  	arrowIndex := strings.Index(rule, "=>")
   135  	if arrowIndex == -1 {
   136  		return false, "", ""
   137  	}
   138  
   139  	condPart := rule[:arrowIndex]
   140  	outPart := rule[arrowIndex+len("=>"):]
   141  
   142  	// Check if condPart starts with "if" followed by at least one space.
   143  	cond := strings.TrimPrefix(condPart, "if")
   144  	if cond == condPart || len(cond) == 0 || !unicode.IsSpace(rune(cond[0])) {
   145  		return false, "", ""
   146  	}
   147  
   148  	// Trim any spaces around <cond> and <out>.
   149  	cond = strings.TrimSpace(cond)
   150  	out := strings.TrimSpace(outPart)
   151  	if cond == "" || out == "" {
   152  		return false, "", ""
   153  	}
   154  
   155  	return true, cond, out
   156  }
   157  
   158  // parseArgsMatchRule tries to parse given string as it would be asmRule with custom arguments to match:
   159  // match <args> [&& <cond>] => <out>
   160  // earlymatch <args> [&& <cond>] => <out>
   161  // Return false, "", "", "", false if not matched, otherwise true, args, cond, rule output, isEarly.
   162  // For example:
   163  // rule:"match [0] (VMOV%sins [0] _ (MOVDconst [c])) && uint64(c)<= 255 => (VMOVI%a [uint8(c)])" can be used to provide addional
   164  // lowering for a broadcast to use immediate source:
   165  // (VDUPSbcast [0] (VMOVSins [0] _ (MOVDconst [c]))) && uint64(c)<= 255 => (VMOVI4S [uint8(c)])
   166  // The specifiers currently supported are only arm64-specific, it may be generalized in the expandFormatSpecifiers function in future.
   167  // The "earlymatch" variant uses GoOp instead of Asm on the left-hand side, replacing the default lowering rule.
   168  func parseArgsMatchRule(rule string) (bool, bool, string) {
   169  	if strings.HasPrefix(rule, "earlymatch ") {
   170  		return true, true, rule[len("earlymatch "):]
   171  	} else if strings.HasPrefix(rule, "match ") {
   172  		return true, false, rule[len("match "):]
   173  	} else {
   174  		return false, false, rule
   175  	}
   176  }
   177  
   178  // expandFormatSpecifiers replaces format specifiers in s with concrete values
   179  // derived from elemBits (the element size in bits for the current vector type).
   180  //
   181  // Supported specifiers:
   182  //
   183  //	%s - lane size letter (B, H, S, D)
   184  //	%a - arrangement suffix (16B, 8H, 4S, 2D) for neon instructions
   185  //	%b - bits per lane (8, 16, 32, 64)
   186  func expandFormatSpecifiers(s string, elemBits int) string {
   187  	elemLetters := map[int]string{8: "B", 16: "H", 32: "S", 64: "D"}
   188  	arrangements := map[int]string{8: "16B", 16: "8H", 32: "4S", 64: "2D"}
   189  	s = strings.ReplaceAll(s, "%s", elemLetters[elemBits])
   190  	s = strings.ReplaceAll(s, "%a", arrangements[elemBits])
   191  	s = strings.ReplaceAll(s, "%b", fmt.Sprintf("%d", elemBits))
   192  	return s
   193  }
   194  
   195  // writeSIMDRules generates the lowering and rewrite rules for ssa and writes it to simdAMD64.rules
   196  // within the specified directory.
   197  func writeSIMDRules(ops []Operation) *bytes.Buffer {
   198  	buffer := new(bytes.Buffer)
   199  	buffer.WriteString(generatedHeader() + "\n")
   200  
   201  	// asm -> masked merging rules
   202  	maskedMergeOpts := make(map[string]string)
   203  	s2n := map[int]string{8: "B", 16: "W", 32: "D", 64: "Q"}
   204  	asmCheck := map[string]bool{}    // for masked merge optimizations.
   205  	sftimmCheck := map[string]bool{} // deduplicate sftimm rules
   206  	var allData []tplRuleData
   207  	var optData []tplRuleData    // for mask peephole optimizations, and other misc
   208  	var memOptData []tplRuleData // for memory peephole optimizations
   209  	memOpSeen := make(map[string]bool)
   210  	ruleDone := make(map[string]struct{})
   211  
   212  	for _, opr := range ops {
   213  		opInShape, opOutShape, maskType, immType, gOp, _ := opr.shape()
   214  		asm := machineOpName(maskType, gOp)
   215  		vregInCnt := len(gOp.In)
   216  		if maskType == OneMask {
   217  			vregInCnt--
   218  		}
   219  
   220  		data := tplRuleData{
   221  			GoOp: gOp.Go,
   222  			Asm:  asm,
   223  		}
   224  
   225  		if vregInCnt == 1 {
   226  			data.Args = "x"
   227  			data.ArgsOut = data.Args
   228  		} else if vregInCnt == 2 {
   229  			data.Args = "x y"
   230  			data.ArgsOut = data.Args
   231  		} else if vregInCnt == 3 {
   232  			data.Args = "x y z"
   233  			data.ArgsOut = data.Args
   234  		} else {
   235  			panic(fmt.Errorf("simdgen does not support more than 3 vreg in inputs"))
   236  		}
   237  		if immType == ConstImm {
   238  			data.ArgsOut = fmt.Sprintf("[%s] %s", *opr.In[0].Const, data.ArgsOut)
   239  		} else if immType == VarImm || immType == VarImmLim {
   240  			data.Args = fmt.Sprintf("[a] %s", data.Args)
   241  			data.ArgsOut = fmt.Sprintf("[a] %s", data.ArgsOut)
   242  		} else if immType == ConstVarImm {
   243  			data.Args = fmt.Sprintf("[a] %s", data.Args)
   244  			data.ArgsOut = fmt.Sprintf("[a+%s] %s", *opr.In[0].Const, data.ArgsOut)
   245  		}
   246  
   247  		goType := func(op Operation) string {
   248  			if op.OperandOrder != nil {
   249  				switch *op.OperandOrder {
   250  				case "21Type1", "231Type1":
   251  					// Permute uses operand[1] for method receiver.
   252  					return *op.In[1].Go
   253  				}
   254  			}
   255  			return *op.In[0].Go
   256  		}
   257  		var tplName string
   258  		// If class overwrite is happening, that's not really a mask but a vreg.
   259  		if opOutShape == OneVregOut || opOutShape == OneVregOutAtIn || opOutShape == OneVregOutScalar || gOp.Out[0].OverwriteClass != nil {
   260  			switch opInShape {
   261  			case OneImmIn:
   262  				tplName = "pureVreg"
   263  				data.GoType = goType(gOp)
   264  			case PureVregIn, VlistIn:
   265  				tplName = "pureVreg"
   266  				data.GoType = goType(gOp)
   267  			case OneKmaskImmIn:
   268  				fallthrough
   269  			case OneKmaskIn:
   270  				tplName = "maskIn"
   271  				data.GoType = goType(gOp)
   272  				rearIdx := len(gOp.In) - 1
   273  				// Mask is at the end.
   274  				width := *gOp.In[rearIdx].ElemBits
   275  				data.MaskInConvert = fmt.Sprintf("VPMOVVec%dx%dToM", width, *gOp.In[rearIdx].Lanes)
   276  				data.ElementSize = width
   277  			case PureKmaskIn:
   278  				panic(fmt.Errorf("simdgen does not support pure k mask instructions, they should be generated by compiler optimizations"))
   279  			}
   280  		} else if opOutShape == OneGregOut {
   281  			tplName = "pureVreg" // TODO this will be wrong
   282  			data.GoType = goType(gOp)
   283  		} else {
   284  			// OneKmaskOut case
   285  			data.MaskOutConvert = fmt.Sprintf("VPMOVMToVec%dx%d", *gOp.Out[0].ElemBits, *gOp.In[0].Lanes)
   286  			switch opInShape {
   287  			case OneImmIn:
   288  				fallthrough
   289  			case PureVregIn:
   290  				tplName = "maskOut"
   291  				data.GoType = goType(gOp)
   292  			case OneKmaskImmIn:
   293  				fallthrough
   294  			case OneKmaskIn:
   295  				tplName = "maskInMaskOut"
   296  				data.GoType = goType(gOp)
   297  				rearIdx := len(gOp.In) - 1
   298  				data.MaskInConvert = fmt.Sprintf("VPMOVVec%dx%dToM", *gOp.In[rearIdx].ElemBits, *gOp.In[rearIdx].Lanes)
   299  			case PureKmaskIn:
   300  				panic(fmt.Errorf("simdgen does not support pure k mask instructions, they should be generated by compiler optimizations"))
   301  			}
   302  		}
   303  
   304  		if gOp.SpecialLower != nil {
   305  			if *gOp.SpecialLower == "sftimm" {
   306  				if !sftimmCheck[data.Asm] {
   307  					sftimmCheck[data.Asm] = true
   308  					sftImmData := data
   309  					if tplName == "maskIn" {
   310  						sftImmData.TplName = "masksftimm"
   311  					} else {
   312  						sftImmData.TplName = "sftimm"
   313  					}
   314  					allData = append(allData, sftImmData)
   315  					asmCheck[sftImmData.Asm+"const"] = true
   316  				}
   317  			} else if ok, cond, out := parseAsmRule(*gOp.SpecialLower); ok {
   318  				if _, done := ruleDone[data.Asm]; !done {
   319  					ruleDone[data.Asm] = struct{}{}
   320  					optData := data
   321  					optData.TplName = "asmRule"
   322  					optData.RuleCond = cond
   323  					if cond != "" {
   324  						optData.RuleCond = "&& " + cond
   325  					}
   326  					optData.RuleOut = out
   327  					if maskType == OneMask {
   328  						optData.Args += " mask"
   329  					}
   330  					allData = append(allData, optData)
   331  				}
   332  			} else if ok, isEarly, rest := parseArgsMatchRule(*gOp.SpecialLower); ok {
   333  				key := data.Asm
   334  				if isEarly {
   335  					key = data.GoOp + data.GoType
   336  				}
   337  				if _, done := ruleDone[key]; !done {
   338  					ruleDone[key] = struct{}{}
   339  					// Get elemBits from the operation's inputs.
   340  					elemBits := 0
   341  					for _, in := range gOp.In {
   342  						if in.ElemBits != nil {
   343  							elemBits = *in.ElemBits
   344  							break
   345  						}
   346  					}
   347  					optData := data
   348  					optData.Rule = rest
   349  					optData.Rule = expandFormatSpecifiers(optData.Rule, elemBits)
   350  					optData.TplName = "specialLower"
   351  					// %g is the generic operation name
   352  					optData.Rule = strings.ReplaceAll(optData.Rule, "%g", optData.GoOp+optData.GoType)
   353  					// %h is the hardware operation name
   354  					optData.Rule = strings.ReplaceAll(optData.Rule, "%h", optData.Asm)
   355  
   356  					allData = append(allData, optData)
   357  				}
   358  				if isEarly {
   359  					continue
   360  				}
   361  			} else {
   362  				panic("simdgen sees unknown special lower " + *gOp.SpecialLower + ", maybe implement it?")
   363  			}
   364  		}
   365  		if gOp.MemFeatures != nil && *gOp.MemFeatures == "vbcst" {
   366  			// sanity check
   367  			selected := true
   368  			for _, a := range gOp.In {
   369  				if a.TreatLikeAScalarOfSize != nil {
   370  					selected = false
   371  					break
   372  				}
   373  			}
   374  			if _, ok := memOpSeen[data.Asm]; ok {
   375  				selected = false
   376  			}
   377  			if selected {
   378  				memOpSeen[data.Asm] = true
   379  				lastVreg := gOp.In[vregInCnt-1]
   380  				// sanity check
   381  				if lastVreg.Class != "vreg" {
   382  					panic(fmt.Errorf("simdgen expects vbcst replaced operand to be a vreg, but %v found", lastVreg))
   383  				}
   384  				memOpData := data
   385  				// Remove the last vreg from the arg and change it to a load.
   386  				origArgs := data.Args[:len(data.Args)-1]
   387  				// Prepare imm args.
   388  				immArg := ""
   389  				immArgCombineOff := " [off] "
   390  				if immType != NoImm && immType != InvalidImm {
   391  					_, after, found := strings.Cut(origArgs, "]")
   392  					if found {
   393  						origArgs = after
   394  					}
   395  					immArg = "[c] "
   396  					immArgCombineOff = " [makeValAndOff(int32(uint8(c)),off)] "
   397  				}
   398  				memOpData.ArgsLoadAddr = immArg + origArgs + fmt.Sprintf("l:(VMOVDQUload%d {sym} [off] ptr mem)", *lastVreg.Bits)
   399  				// Remove the last vreg from the arg and change it to "ptr".
   400  				memOpData.ArgsAddr = "{sym}" + immArgCombineOff + origArgs + "ptr"
   401  				if maskType == OneMask {
   402  					memOpData.ArgsAddr += " mask"
   403  					memOpData.ArgsLoadAddr += " mask"
   404  				}
   405  				memOpData.ArgsAddr += " mem"
   406  				if gOp.MemFeaturesData != nil {
   407  					_, feat2 := getVbcstData(*gOp.MemFeaturesData)
   408  					knownFeatChecks := map[string]string{
   409  						"AVX":    "v.Block.CPUfeatures.hasFeature(CPUavx)",
   410  						"AVX2":   "v.Block.CPUfeatures.hasFeature(CPUavx2)",
   411  						"AVX512": "v.Block.CPUfeatures.hasFeature(CPUavx512)",
   412  					}
   413  					memOpData.FeatCheck = knownFeatChecks[feat2]
   414  					memOpData.TplName = "vregMemFeatCheck"
   415  				} else {
   416  					memOpData.TplName = "vregMem"
   417  				}
   418  				memOptData = append(memOptData, memOpData)
   419  				asmCheck[memOpData.Asm+"load"] = true
   420  			}
   421  		}
   422  		// Generate the masked merging optimization rules
   423  		if gOp.hasMaskedMerging(maskType, opOutShape) {
   424  			// TODO: handle customized operand order and special lower.
   425  			maskElem := gOp.In[len(gOp.In)-1]
   426  			if maskElem.Bits == nil {
   427  				panic("mask has no bits")
   428  			}
   429  			if maskElem.ElemBits == nil {
   430  				panic("mask has no elemBits")
   431  			}
   432  			if maskElem.Lanes == nil {
   433  				panic("mask has no lanes")
   434  			}
   435  			switch *maskElem.Bits {
   436  			case 128, 256:
   437  				// VPBLENDVB cases.
   438  				noMaskName := machineOpName(NoMask, gOp)
   439  				ruleExisting, ok := maskedMergeOpts[noMaskName]
   440  				rule := fmt.Sprintf("(VPBLENDVB%d dst (%s %s) mask) && v.Block.CPUfeatures.hasFeature(CPUavx512) => (%sMerging dst %s (VPMOVVec%dx%dToM <types.TypeMask> mask))\n",
   441  					*maskElem.Bits, noMaskName, data.Args, data.Asm, data.Args, *maskElem.ElemBits, *maskElem.Lanes)
   442  				if ok && ruleExisting != rule {
   443  					panic(fmt.Sprintf("multiple masked merge rules for one op:\n%s\n%s\n", ruleExisting, rule))
   444  				} else {
   445  					maskedMergeOpts[noMaskName] = rule
   446  				}
   447  			case 512:
   448  				// VPBLENDM[BWDQ] cases.
   449  				noMaskName := machineOpName(NoMask, gOp)
   450  				ruleExisting, ok := maskedMergeOpts[noMaskName]
   451  				rule := fmt.Sprintf("(VPBLENDM%sMasked%d dst (%s %s) mask) => (%sMerging dst %s mask)\n",
   452  					s2n[*maskElem.ElemBits], *maskElem.Bits, noMaskName, data.Args, data.Asm, data.Args)
   453  				if ok && ruleExisting != rule {
   454  					panic(fmt.Sprintf("multiple masked merge rules for one op:\n%s\n%s\n", ruleExisting, rule))
   455  				} else {
   456  					maskedMergeOpts[noMaskName] = rule
   457  				}
   458  			}
   459  		}
   460  
   461  		if tplName == "pureVreg" && data.Args == data.ArgsOut {
   462  			data.Args = "..."
   463  			data.ArgsOut = "..."
   464  		}
   465  		data.TplName = tplName
   466  		if opr.NoGenericOps != nil && *opr.NoGenericOps == "true" ||
   467  			opr.SkipMaskedMethod() {
   468  			optData = append(optData, data)
   469  			continue
   470  		}
   471  		allData = append(allData, data)
   472  		asmCheck[data.Asm] = true
   473  	}
   474  
   475  	slices.SortFunc(allData, compareTplRuleData)
   476  
   477  	hiHalfRules := generateHiHalfFoldingRules(ops)
   478  	for _, rule := range hiHalfRules {
   479  		buffer.WriteString(rule)
   480  	}
   481  
   482  	for _, data := range allData {
   483  		tpl := ruleTemplates.Get(data.TplName)
   484  		if tpl == nil {
   485  			panic(fmt.Errorf("template %s not found", data.TplName))
   486  		}
   487  		if err := tpl.Execute(buffer, data); err != nil {
   488  			panic(fmt.Errorf("failed to execute template %s for %s: %w", data.TplName, data.GoOp+data.GoType, err))
   489  		}
   490  	}
   491  
   492  	seen := make(map[string]bool)
   493  
   494  	for _, data := range optData {
   495  		if data.TplName == "maskIn" {
   496  			rule := data.MaskOptimization(asmCheck)
   497  			if seen[rule] {
   498  				continue
   499  			}
   500  			seen[rule] = true
   501  			buffer.WriteString(rule)
   502  		}
   503  	}
   504  
   505  	maskedMergeOptsRules := []string{}
   506  	for asm, rule := range maskedMergeOpts {
   507  		if !asmCheck[asm] {
   508  			continue
   509  		}
   510  		maskedMergeOptsRules = append(maskedMergeOptsRules, rule)
   511  	}
   512  	slices.Sort(maskedMergeOptsRules)
   513  	for _, rule := range maskedMergeOptsRules {
   514  		buffer.WriteString(rule)
   515  	}
   516  
   517  	for _, data := range memOptData {
   518  		tpl := ruleTemplates.Get(data.TplName)
   519  		if tpl == nil {
   520  			panic(fmt.Errorf("template %s not found", data.TplName))
   521  		}
   522  		if err := tpl.Execute(buffer, data); err != nil {
   523  			panic(fmt.Errorf("failed to execute template %s for %s: %w", data.TplName, data.Asm, err))
   524  		}
   525  	}
   526  
   527  	return buffer
   528  }
   529  
   530  // Note: SetHi was removed (temporarily?) from 1.27, but may return in some form
   531  // and these optimizations may apply to the equivalent idioms.
   532  //
   533  // generateHiHalfFoldingRules generates folding rules that combine SetHi/HiToLo patterns
   534  // with narrow/long base operations into the hi-half "2" variant instruction.
   535  //
   536  // x.HiToLo() is lowered as (VDUPDextr[1] x)
   537  //
   538  // x.SetHi(lo) is lowered as (VMOVDins0 [1] x (VDUPDextr [0] lo))
   539  //
   540  // Narrow (e.g., SHRN): SetHi wraps the narrow result.
   541  //
   542  //	(VMOVDins0 [1] dst y:(VSHRN4S [a] x)) => (VSHRN2_4S dst [a] x)
   543  //
   544  // Unary Long (e.g., USHLL) getting high half (HiToLo) as input:
   545  //
   546  //	(VUSHLL4H [a] (VDUPDextr [1] x)) => (VUSHLL2_4H [a] x)
   547  //
   548  // Binary Long (e.g., UMULL, SMULL), both inputs from HiToLo:
   549  //
   550  //	(VUMULL4H (VDUPDextr [1] x) (VDUPDextr [1] y)) => (VUMULL2_4H x y)
   551  func generateHiHalfFoldingRules(ops []Operation) []string {
   552  	seen := make(map[string]bool)
   553  	var rules []string
   554  
   555  	for _, opr := range ops {
   556  		if opr.HiHalfAsm == nil {
   557  			continue
   558  		}
   559  		kind := opr.hiHalfKind()
   560  		if kind == "" {
   561  			continue
   562  		}
   563  		_, _, maskType, immType, gOp, _ := opr.shape()
   564  		asm := machineOpName(maskType, gOp)
   565  		asm2 := hiHalfOpName(*gOp.HiHalfAsm, gOp)
   566  
   567  		if seen[asm] {
   568  			continue
   569  		}
   570  		seen[asm] = true
   571  
   572  		vregInCnt := 0
   573  		for _, in := range gOp.In {
   574  			if in.Class == "vreg" {
   575  				vregInCnt++
   576  			}
   577  		}
   578  		hasImm := immType == VarImm || immType == VarImmLim || immType == ConstVarImm
   579  
   580  		switch kind {
   581  		case "narrow":
   582  			switch vregInCnt {
   583  			case 1:
   584  				// Note, at least for 1.27, SetHi is not supported, because it is an emulated instruction,
   585  				// not part of the scalable set, and though it is intended to mimic the similar instruction
   586  				// on amd64, it is actually slightly different.
   587  				//
   588  				// Narrow: the value used by SetHi may be produced with asm2 (hiHalf) variant.
   589  				// TODO: Maybe have a separate rule: (VMOVDins0 [c] dst (VDUPDextr [0] src)) => (VMOVDins0 [c] dst src)
   590  				if hasImm {
   591  					rules = append(rules, fmt.Sprintf("(VMOVDins0 [1] dst (VDUPDextr [0] (%s [c] y))) => (%s dst [c] y)\n", asm, asm2))
   592  				} else {
   593  					rules = append(rules, fmt.Sprintf("(VMOVDins0 [1] dst (VDUPDextr [0] (%s y))) => (%s dst y)\n", asm, asm2))
   594  				}
   595  			default:
   596  				panic("unsupported yet folding narrow ops cases")
   597  			}
   598  		case "long":
   599  			switch vregInCnt {
   600  			case 1:
   601  				// Long: input from HiToLo (VDUPDextr [1]), prefer hiHalf variant.
   602  				if hasImm {
   603  					rules = append(rules, fmt.Sprintf("(%s [a] (VDUPDextr [1] x)) => (%s [a] x)\n", asm, asm2))
   604  				} else {
   605  					rules = append(rules, fmt.Sprintf("(%s (VDUPDextr [1] x)) => (%s x)\n", asm, asm2))
   606  				}
   607  			case 2:
   608  				// Binary Long: both inputs are from HiToLo, fold into hiHalf variant.
   609  				rules = append(rules, fmt.Sprintf("(%s (VDUPDextr [1] x) (VDUPDextr [1] y)) => (%s x y)\n", asm, asm2))
   610  			default:
   611  				panic("unsupported yet folding long ops cases")
   612  			}
   613  		}
   614  	}
   615  
   616  	slices.Sort(rules)
   617  	return rules
   618  }
   619  

View as plain text