blob: b664b0f460aae2dd998667ac0a313c0cdffb99f1 [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{
"v11",
"v21",
"v2k",
"v2kv",
"v2kk",
"vkv",
"v31",
"v3kv",
"v11Imm8",
"vkvImm8",
"v21Imm8",
"v2kImm8",
"v2kkImm8",
"v31ResultInArg0",
"v3kvResultInArg0",
"vfpv",
"vfpkv",
"vgpvImm8",
"vgpImm8",
"v2kvImm8",
}
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 == "v21" {
regShape = "vfpv"
} else if regShape == "v2kv" {
regShape = "vfpkv"
} 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
}