Source file src/simd/archsimd/_gen/simdgen/gen_simdssa.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  	"text/template"
    14  )
    15  
    16  var (
    17  	ssaTemplates = template.Must(template.New("simdSSA").Parse(`{{define "header"}}{{.GeneratedHeader}}
    18  package {{.Arch}}
    19  
    20  import (
    21  	"cmd/compile/internal/ssa"
    22  	"cmd/compile/internal/ssagen"
    23  	"cmd/internal/obj"
    24  	"cmd/internal/obj/{{.ObjArch}}"
    25  )
    26  
    27  func ssaGenSIMDValue(s *ssagen.State, v *ssa.Value) bool {
    28  	var p *obj.Prog
    29  	switch v.Op {{"{"}}{{end}}
    30  {{define "case"}}
    31  	case {{.Cases}}:
    32  		p = {{.Helper}}(s, v{{if .Arrangement}}, {{.Arrangement}}{{end}})
    33  {{end}}
    34  {{define "footer"}}
    35  	default:
    36  		// Unknown reg shape
    37  		return false
    38  	}
    39  {{end}}
    40  {{define "zeroing"}}
    41  	// Masked operation are always compiled with zeroing.
    42  	switch v.Op {
    43  	case {{.}}:
    44  		x86.ParseSuffix(p, "Z")
    45  	}
    46  {{end}}
    47  {{define "ending"}}
    48  	// Ensure p is marked as used (may not be used in all generated code paths)
    49  	_ = p
    50  	return true
    51  }
    52  {{end}}`))
    53  )
    54  
    55  type tplSSAData struct {
    56  	Cases       string
    57  	Helper      string
    58  	Arrangement string // Optional arrangement constant for ARM64, e.g. arm64.ARNG_4S
    59  }
    60  
    61  type tplSSAHeader struct {
    62  	Arch            string
    63  	ObjArch         string
    64  	GeneratedHeader string
    65  }
    66  
    67  // getArrangementFromOp extracts the arrangement constant from an SSA op name for ARM64.
    68  // For example, "ssa.OpARM64VFADD4S" returns "arm64.ARNG_4S".
    69  func getArrangementFromOp(archInfo ArchInfo, caseStr string) string {
    70  	for _, a := range archInfo.Arrangements {
    71  		if strings.Contains(caseStr, a) {
    72  			return archInfo.Arch + ".ARNG_" + a
    73  		}
    74  	}
    75  	return ""
    76  }
    77  
    78  // writeSIMDSSA generates the ssa to prog lowering codes and writes it to simdssa.go
    79  // within the specified directory.
    80  func writeSIMDSSA(ops []Operation) *bytes.Buffer {
    81  	archInfo := CurrentArch()
    82  	var ZeroingMask []string
    83  	regInfoKeys := archInfo.RegInfoKeys
    84  	regInfoSet := map[string][]string{}
    85  	for _, key := range regInfoKeys {
    86  		regInfoSet[key] = []string{}
    87  	}
    88  
    89  	seen := map[string]struct{}{}
    90  	allUnseen := make(map[string][]Operation)
    91  	allUnseenCaseStr := make(map[string][]string)
    92  	// computeBaseRegShape computes the base regShape for an op.
    93  	computeBaseRegShape := func(op Operation, mem memShape, shapeIn inShape, shapeOut outShape, immOpArg string, immType immShape) (string, error) {
    94  		regShape, err := op.regShape(mem)
    95  		if err != nil {
    96  			return "", err
    97  		}
    98  		if regShape == "v01load" {
    99  			regShape = "vload"
   100  		}
   101  		if shapeOut == OneVregOutAtIn {
   102  			regShape += "ResultInArg0"
   103  		} else if shapeOut == OneVregOutScalar {
   104  			regShape += "Scalar"
   105  		}
   106  		if shapeIn == OneImmIn || shapeIn == OneKmaskImmIn {
   107  			if immOpArg != "" {
   108  				regShape += "Imm"
   109  				regShape += immOpArg
   110  			} else if immType == VarImmLim {
   111  				regShape += "Imm" // limited range immediate (ImmMax set)
   112  			} else {
   113  				regShape += "Imm8" // full 8-bit range (0-255)
   114  			}
   115  		}
   116  		if shapeIn == VlistIn {
   117  			regShape += "List"
   118  		}
   119  		regShape, err = rewriteVecAsScalarRegInfo(op, regShape)
   120  		if err != nil {
   121  			return "", err
   122  		}
   123  		return regShape, nil
   124  	}
   125  	registerRegShape := func(regShape string, caseStr string, op Operation) {
   126  		if _, ok := regInfoSet[regShape]; !ok {
   127  			allUnseen[regShape] = append(allUnseen[regShape], op)
   128  			allUnseenCaseStr[regShape] = append(allUnseenCaseStr[regShape], caseStr)
   129  		}
   130  		regInfoSet[regShape] = append(regInfoSet[regShape], caseStr)
   131  	}
   132  	classifyOp := func(op Operation, maskType maskShape, shapeIn inShape, shapeOut outShape, caseStr string, mem memShape, immOpArg string, immType immShape) error {
   133  		regShape, err := computeBaseRegShape(op, mem, shapeIn, shapeOut, immOpArg, immType)
   134  		if err != nil {
   135  			return err
   136  		}
   137  		// For hi-half base ops, append the kind suffix for lowering dispatch.
   138  		if op.HiHalfAsm != nil {
   139  			kind := op.hiHalfKind()
   140  			if kind != "" {
   141  				regShape += capitalizeFirst(kind) // e.g., "v11Imm" + "Narrow" = "v11ImmNarrow"
   142  			}
   143  		}
   144  		registerRegShape(regShape, caseStr, op)
   145  		if mem == NoMem && op.hasMaskedMerging(maskType, shapeOut) {
   146  			regShapeMerging := regShape
   147  			if shapeOut != OneVregOutAtIn {
   148  				// We have to copy the slice here because the sort will be visible from other
   149  				// aliases when no reslicing is happening.
   150  				newIn := make([]Operand, len(op.In), len(op.In)+1)
   151  				copy(newIn, op.In)
   152  				op.In = newIn
   153  				op.In = append(op.In, op.Out[0])
   154  				op.sortOperand()
   155  				regShapeMerging, err = op.regShape(mem)
   156  				regShapeMerging += "ResultInArg0"
   157  			}
   158  			if err != nil {
   159  				return err
   160  			}
   161  			registerRegShape(regShapeMerging, caseStr+"Merging", op)
   162  		}
   163  		return nil
   164  	}
   165  	// classifyHiHalfOp computes the lowering dispatch regShape for a hi-half "2" variant.
   166  	// It derives the regShape from the base op's shape and applies hi-half transformation.
   167  	classifyHiHalfOp := func(op Operation, kind string, caseStr string, immOpArg string, immType immShape) error {
   168  		shapeIn, shapeOut, _, _, _, _ := op.shape()
   169  		regShape, err := computeBaseRegShape(op, NoMem, shapeIn, shapeOut, immOpArg, immType)
   170  		if err != nil {
   171  			return err
   172  		}
   173  		regShape = hiHalfLoweringRegShape(regShape, kind, true)
   174  		registerRegShape(regShape, caseStr, op)
   175  		return nil
   176  	}
   177  	for _, op := range ops {
   178  		shapeIn, shapeOut, maskType, immType, gOp, immOpArg := op.shape()
   179  		asm := machineOpName(maskType, gOp)
   180  		if _, ok := seen[asm]; ok {
   181  			continue
   182  		}
   183  		seen[asm] = struct{}{}
   184  		caseStr := fmt.Sprintf("ssa.Op%s%s", archInfo.ArchUpper, asm)
   185  		isZeroMasking := false
   186  		if shapeIn == OneKmaskIn || shapeIn == OneKmaskImmIn {
   187  			if gOp.Zeroing == nil || *gOp.Zeroing {
   188  				ZeroingMask = append(ZeroingMask, caseStr)
   189  				isZeroMasking = true
   190  			}
   191  		}
   192  		if err := classifyOp(op, maskType, shapeIn, shapeOut, caseStr, NoMem, immOpArg, immType); err != nil {
   193  			panic(err)
   194  		}
   195  
   196  		// Generate hi-half "2" variant SSA lowering case.
   197  		// The base op gets a suffixed regShape (e.g., "v11ImmNarrow"), and
   198  		// the "2" variant gets a derived regShape (e.g., "v21ImmNarrow2").
   199  		if gOp.HiHalfAsm != nil {
   200  			kind := op.hiHalfKind()
   201  			if kind != "" {
   202  				asm2 := hiHalfOpName(*gOp.HiHalfAsm, gOp)
   203  				caseStr2 := fmt.Sprintf("ssa.Op%s%s", archInfo.ArchUpper, asm2)
   204  				if _, ok2 := seen[asm2]; !ok2 {
   205  					seen[asm2] = struct{}{}
   206  					if err := classifyHiHalfOp(op, kind, caseStr2, immOpArg, immType); err != nil {
   207  						panic(err)
   208  					}
   209  				}
   210  			}
   211  		}
   212  
   213  		if op.MemFeatures != nil && *op.MemFeatures == "vbcst" {
   214  			// Make a full vec memory variant
   215  			op = rewriteLastVregToMem(op)
   216  			// Ignore the error
   217  			// an error could be triggered by [checkVecAsScalar].
   218  			// TODO: make [checkVecAsScalar] aware of mem ops.
   219  			if err := classifyOp(op, maskType, shapeIn, shapeOut, caseStr+"load", VregMemIn, immOpArg, immType); err != nil {
   220  				if *Verbose {
   221  					log.Printf("Seen error: %e", err)
   222  				}
   223  			} else if isZeroMasking {
   224  				ZeroingMask = append(ZeroingMask, caseStr+"load")
   225  			}
   226  		}
   227  	}
   228  	if len(allUnseen) != 0 {
   229  		allKeys := make([]string, 0)
   230  		for k := range allUnseen {
   231  			allKeys = append(allKeys, k)
   232  		}
   233  		panic(fmt.Errorf("unsupported register constraint for prog, please update gen_simdssa.go and amd64/ssa.go: %+v\nAll keys: %v\n, cases: %v\n", allUnseen, allKeys, allUnseenCaseStr))
   234  	}
   235  
   236  	buffer := new(bytes.Buffer)
   237  
   238  	headerData := tplSSAHeader{
   239  		Arch:            archInfo.Arch,
   240  		ObjArch:         archInfo.ObjArch,
   241  		GeneratedHeader: archInfo.GeneratedHeader,
   242  	}
   243  	if err := ssaTemplates.ExecuteTemplate(buffer, "header", headerData); err != nil {
   244  		panic(fmt.Errorf("failed to execute header template: %w", err))
   245  	}
   246  
   247  	for _, regShape := range regInfoKeys {
   248  		// Stable traversal of regInfoSet
   249  		cases := regInfoSet[regShape]
   250  		if len(cases) == 0 {
   251  			continue
   252  		}
   253  
   254  		// Group cases by arrangement (for ARM64)
   255  		arrangementGroups := make(map[string][]string)
   256  		for _, caseStr := range cases {
   257  			arrangement := getArrangementFromOp(archInfo, caseStr)
   258  			arrangementGroups[arrangement] = append(arrangementGroups[arrangement], caseStr)
   259  		}
   260  
   261  		// Sort arrangement keys for deterministic output
   262  		var arrangements []string
   263  		for arrangement := range arrangementGroups {
   264  			arrangements = append(arrangements, arrangement)
   265  		}
   266  		sort.Strings(arrangements)
   267  
   268  		// Generate cases for each arrangement group in sorted order
   269  		for _, arrangement := range arrangements {
   270  			groupCases := arrangementGroups[arrangement]
   271  			data := tplSSAData{
   272  				Cases:  strings.Join(groupCases, ",\n\t\t"),
   273  				Helper: "simd" + capitalizeFirst(regShape),
   274  			}
   275  			if arrangement != "" {
   276  				data.Arrangement = arrangement
   277  			}
   278  			if err := ssaTemplates.ExecuteTemplate(buffer, "case", data); err != nil {
   279  				panic(fmt.Errorf("failed to execute case template for %s: %w", regShape, err))
   280  			}
   281  		}
   282  	}
   283  
   284  	if err := ssaTemplates.ExecuteTemplate(buffer, "footer", nil); err != nil {
   285  		panic(fmt.Errorf("failed to execute footer template: %w", err))
   286  	}
   287  
   288  	if len(ZeroingMask) != 0 {
   289  		if err := ssaTemplates.ExecuteTemplate(buffer, "zeroing", strings.Join(ZeroingMask, ",\n\t\t")); err != nil {
   290  			panic(fmt.Errorf("failed to execute footer template: %w", err))
   291  		}
   292  	}
   293  
   294  	if err := ssaTemplates.ExecuteTemplate(buffer, "ending", headerData); err != nil {
   295  		panic(fmt.Errorf("failed to execute ending template: %w", err))
   296  	}
   297  
   298  	return buffer
   299  }
   300  

View as plain text