blob: ffb172a6441bb03184a07dc36096ae5ca4153c92 [file]
// Copyright 2025 The Go Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
package main
import (
"bytes"
"fmt"
"strings"
"text/template"
)
var (
ssaTemplates = template.Must(template.New("simdSSA").Parse(`
{{define "header"}}// Code generated by x/arch/internal/simdgen using 'go run . -xedPath $XED_PATH -o godefs -goroot $GOROOT go.yaml types.yaml categories.yaml'; DO NOT EDIT.
package amd64
import (
"cmd/compile/internal/ssa"
"cmd/compile/internal/ssagen"
"cmd/internal/obj"
"cmd/internal/obj/x86"
)
func ssaGenSIMDValue(s *ssagen.State, v *ssa.Value) bool {
var p *obj.Prog
switch v.Op {{"{"}}{{end}}
{{define "case"}}
case {{.Cases}}:
p = {{.Helper}}(s, v)
{{end}}
{{define "footer"}}
default:
// Unknown reg shape
return false
}
{{end}}
{{define "zeroing"}}
// Masked operation are always compiled with zeroing.
switch v.Op {
case {{.}}:
x86.ParseSuffix(p, "Z")
}
{{end}}
{{define "ending"}}
return true
}
{{end}}`))
)
type tplSSAData struct {
Cases string
Helper string
}
// writeSIMDSSA generates the ssa to prog lowering codes and writes it to simdssa.go
// within the specified directory.
func writeSIMDSSA(ops []Operation) *bytes.Buffer {
var ZeroingMask []string
regInfoKeys := []string{
"fp11",
"fp21",
"fp2k",
"fp2kfp",
"fp2kk",
"fpkfp",
"fp31",
"fp3kfp",
"fp11Imm8",
"fpkfpImm8",
"fp21Imm8",
"fp2kImm8",
"fp2kkImm8",
"fp31ResultInArg0",
"fp3kfpResultInArg0",
"fpXfp",
"fpXkfp",
"fpgpfpImm8",
"fpgpImm8",
"fp2kfpImm8",
}
regInfoSet := map[string][]string{}
for _, key := range regInfoKeys {
regInfoSet[key] = []string{}
}
seen := map[string]struct{}{}
allUnseen := make(map[string][]Operation)
for _, op := range ops {
asm := op.Asm
shapeIn, shapeOut, maskType, _, _, _, gOp := op.shape()
if maskType == 2 {
asm += "Masked"
}
asm = fmt.Sprintf("%s%d", asm, gOp.VectorWidth())
if _, ok := seen[asm]; ok {
continue
}
seen[asm] = struct{}{}
caseStr := fmt.Sprintf("ssa.OpAMD64%s", asm)
if shapeIn == OneKmaskIn || shapeIn == OneKmaskImmIn {
if gOp.Zeroing == nil {
ZeroingMask = append(ZeroingMask, caseStr)
}
}
regShape, err := op.regShape()
if err != nil {
panic(err)
}
if shapeOut == OneVregOutAtIn {
regShape += "ResultInArg0"
}
if shapeIn == OneImmIn || shapeIn == OneKmaskImmIn {
regShape += "Imm8"
}
idx, err := checkVecAsScalar(op)
if err != nil {
panic(err)
}
if idx != -1 {
if regShape == "fp21" {
regShape = "fpXfp"
} else if regShape == "fp2kfp" {
regShape = "fpXkfp"
} else {
panic(fmt.Errorf("simdgen does not recognize uses of treatLikeAScalarOfSize with op regShape %s in op: %s", regShape, op))
}
}
if _, ok := regInfoSet[regShape]; !ok {
allUnseen[regShape] = append(allUnseen[regShape], op)
}
regInfoSet[regShape] = append(regInfoSet[regShape], caseStr)
}
if len(allUnseen) != 0 {
panic(fmt.Errorf("unsupported register constraint for prog, please update gen_simdssa.go and amd64/ssa.go: %+v", allUnseen))
}
buffer := new(bytes.Buffer)
if err := ssaTemplates.ExecuteTemplate(buffer, "header", nil); err != nil {
panic(fmt.Errorf("failed to execute header template: %w", err))
}
for _, regShape := range regInfoKeys {
// Stable traversal of regInfoSet
cases := regInfoSet[regShape]
if len(cases) == 0 {
continue
}
data := tplSSAData{
Cases: strings.Join(cases, ",\n\t\t"),
Helper: "simd" + capitalizeFirst(regShape),
}
if err := ssaTemplates.ExecuteTemplate(buffer, "case", data); err != nil {
panic(fmt.Errorf("failed to execute case template for %s: %w", regShape, err))
}
}
if err := ssaTemplates.ExecuteTemplate(buffer, "footer", nil); err != nil {
panic(fmt.Errorf("failed to execute footer template: %w", err))
}
if len(ZeroingMask) != 0 {
if err := ssaTemplates.ExecuteTemplate(buffer, "zeroing", strings.Join(ZeroingMask, ",\n\t\t")); err != nil {
panic(fmt.Errorf("failed to execute footer template: %w", err))
}
}
if err := ssaTemplates.ExecuteTemplate(buffer, "ending", nil); err != nil {
panic(fmt.Errorf("failed to execute footer template: %w", err))
}
return buffer
}