fix(gasm): scaffold shared param names and two-sided seed sets
This commit is contained in:
+82
-77
@@ -9,7 +9,6 @@ import (
|
|||||||
"go/parser"
|
"go/parser"
|
||||||
"go/token"
|
"go/token"
|
||||||
"os"
|
"os"
|
||||||
"sort"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
gasmast "sourcedock.dev/petrbalvin/gasm-devkit/ast"
|
gasmast "sourcedock.dev/petrbalvin/gasm-devkit/ast"
|
||||||
@@ -39,6 +38,11 @@ bodies, place the file in the kernel's package, and run it in CI.
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
rest := fs.Args()
|
rest := fs.Args()
|
||||||
|
// The first positional word is the scaffold style; "differential" is the
|
||||||
|
// only one today.
|
||||||
|
if len(rest) > 0 && rest[0] == "differential" {
|
||||||
|
rest = rest[1:]
|
||||||
|
}
|
||||||
if len(rest) != 1 {
|
if len(rest) != 1 {
|
||||||
return fmt.Errorf("usage: gasm scaffold differential <file.s>")
|
return fmt.Errorf("usage: gasm scaffold differential <file.s>")
|
||||||
}
|
}
|
||||||
@@ -76,23 +80,30 @@ bodies, place the file in the kernel's package, and run it in CI.
|
|||||||
fmt.Fprintf(&out, "func Test%sDifferential(t *testing.T) {\n", strings.ToUpper(name[:1])+name[1:])
|
fmt.Fprintf(&out, "func Test%sDifferential(t *testing.T) {\n", strings.ToUpper(name[:1])+name[1:])
|
||||||
fmt.Fprintf(&out, "\trng := rand.New(rand.NewSource(1))\n")
|
fmt.Fprintf(&out, "\trng := rand.New(rand.NewSource(1))\n")
|
||||||
fmt.Fprintf(&out, "\tfor iter := 0; iter < 1000; iter++ {\n")
|
fmt.Fprintf(&out, "\tfor iter := 0; iter < 1000; iter++ {\n")
|
||||||
written := map[string]bool{}
|
// Seed two independent argument sets per iteration: the kernel runs
|
||||||
|
// on set A, the portable reference on set B, so in-place writes
|
||||||
|
// through pointer/slice arguments cannot contaminate the other side.
|
||||||
|
var sliceNames []string
|
||||||
|
seen := map[string]bool{}
|
||||||
|
aArgs := make([]string, 0, len(params))
|
||||||
|
bArgs := make([]string, 0, len(params))
|
||||||
for _, p := range params {
|
for _, p := range params {
|
||||||
genParamSeed(&out, p, written)
|
a, b, slices := genParamSeed(&out, p, seen)
|
||||||
|
aArgs = append(aArgs, a)
|
||||||
|
bArgs = append(bArgs, b)
|
||||||
|
sliceNames = append(sliceNames, slices...)
|
||||||
}
|
}
|
||||||
argList := make([]string, 0, len(params))
|
fmt.Fprintf(&out, "\t\tgot := %s(%s)\n", name, strings.Join(aArgs, ", "))
|
||||||
for _, p := range params {
|
fmt.Fprintf(&out, "\t\twant := %sPortable(%s)\n", name, strings.Join(bArgs, ", "))
|
||||||
argList = append(argList, seedArg(p))
|
|
||||||
}
|
|
||||||
fmt.Fprintf(&out, "\t\tgot := %s(%s)\n", name, strings.Join(argList, ", "))
|
|
||||||
wantArgs := make([]string, 0, len(params))
|
|
||||||
for _, p := range params {
|
|
||||||
wantArgs = append(wantArgs, seedArg(p))
|
|
||||||
}
|
|
||||||
fmt.Fprintf(&out, "\t\twant := %sPortable(%s)\n", name, strings.Join(wantArgs, ", "))
|
|
||||||
fmt.Fprintf(&out, "\t\tif !bytes.Equal(outputBytes(got), outputBytes(want)) {\n")
|
fmt.Fprintf(&out, "\t\tif !bytes.Equal(outputBytes(got), outputBytes(want)) {\n")
|
||||||
fmt.Fprintf(&out, "\t\t\tt.Fatalf(\"iter %%d: kernel diverges from the portable spec\", iter)\n")
|
fmt.Fprintf(&out, "\t\t\tt.Fatalf(\"iter %%d: kernel diverges from the portable spec\", iter)\n")
|
||||||
fmt.Fprintf(&out, "\t\t}\n\t}\n}\n\n")
|
fmt.Fprintf(&out, "\t\t}\n")
|
||||||
|
for _, s := range sliceNames {
|
||||||
|
fmt.Fprintf(&out, "\t\tif !bytes.Equal(outputBytes(%sA), outputBytes(%sB)) {\n", s, s)
|
||||||
|
fmt.Fprintf(&out, "\t\t\tt.Fatalf(\"iter %%d: kernel mutated %%s differently\", iter, %q)\n", s)
|
||||||
|
fmt.Fprintf(&out, "\t\t}\n")
|
||||||
|
}
|
||||||
|
fmt.Fprintf(&out, "\t}\n}\n\n")
|
||||||
}
|
}
|
||||||
if kernels == 0 {
|
if kernels == 0 {
|
||||||
return fmt.Errorf("%s: no // func signatures found; add one doc comment per kernel", path)
|
return fmt.Errorf("%s: no // func signatures found; add one doc comment per kernel", path)
|
||||||
@@ -144,7 +155,16 @@ func parseSig(doc string) ([]sigParam, []sigResult, bool) {
|
|||||||
}
|
}
|
||||||
var params []sigParam
|
var params []sigParam
|
||||||
for _, field := range fd.Type.Params.List {
|
for _, field := range fd.Type.Params.List {
|
||||||
params = append(params, sigParam{Names: identNames(field.Names), Type: exprString(field.Type)})
|
typ := exprString(field.Type)
|
||||||
|
if len(field.Names) == 0 {
|
||||||
|
params = append(params, sigParam{Names: []string{""}, Type: typ})
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Shared names (`L, result *byte`) expand to one entry per name:
|
||||||
|
// every name is a separate argument at the call site.
|
||||||
|
for _, n := range field.Names {
|
||||||
|
params = append(params, sigParam{Names: []string{n.Name}, Type: typ})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
var results []sigResult
|
var results []sigResult
|
||||||
if fd.Type.Results != nil {
|
if fd.Type.Results != nil {
|
||||||
@@ -207,78 +227,63 @@ func resultDecl(results []sigResult) string {
|
|||||||
return strings.Join(parts, ", ")
|
return strings.Join(parts, ", ")
|
||||||
}
|
}
|
||||||
|
|
||||||
// seedArg renders the argument expression referencing the seeded variables.
|
// genParamSeed emits the seeding statements for one parameter and returns
|
||||||
func seedArg(p sigParam) string {
|
// the kernel-side (A) and reference-side (B) argument expressions, plus the
|
||||||
if len(p.Names) == 0 {
|
// names of any slice variables written in place (compared after the calls).
|
||||||
return typeSeedExpr(p.Type, "")
|
func genParamSeed(out *strings.Builder, p sigParam, seen map[string]bool) (aArg, bArg string, slices []string) {
|
||||||
|
name := p.Names[0]
|
||||||
|
elem := strings.TrimPrefix(p.Type, "*")
|
||||||
|
isSlice := strings.HasPrefix(p.Type, "[]")
|
||||||
|
if isSlice {
|
||||||
|
elem = strings.TrimPrefix(p.Type, "[]")
|
||||||
}
|
}
|
||||||
return typeSeedExpr(p.Type, p.Names[0])
|
|
||||||
}
|
|
||||||
|
|
||||||
// typeSeedExpr returns the expression passed for one parameter, given the
|
|
||||||
// seeded variable base name ("" for nameless parameters).
|
|
||||||
func typeSeedExpr(typ, base string) string {
|
|
||||||
switch {
|
switch {
|
||||||
case strings.HasPrefix(typ, "[]"):
|
case isSlice:
|
||||||
if base == "" {
|
v := uniqueName(seen, name)
|
||||||
base = "arg"
|
fmt.Fprintf(out, "\t\t%sA := make([]%s, 1+rng.Intn(512))\n", v, elem)
|
||||||
}
|
fmt.Fprintf(out, "\t\t%sB := make([]%s, len(%sA))\n", v, elem, v)
|
||||||
return base + "Slice"
|
fmt.Fprintf(out, "\t\tfor i := range %sA {\n", v)
|
||||||
case strings.HasPrefix(typ, "*"):
|
fmt.Fprintf(out, "\t\t\tw%s := %s(rng.Intn(256))\n", v, goCast(elem))
|
||||||
if base == "" {
|
fmt.Fprintf(out, "\t\t\t%sA[i] = w%s\n", v, v)
|
||||||
base = "arg"
|
fmt.Fprintf(out, "\t\t\t%sB[i] = w%s\n", v, v)
|
||||||
}
|
fmt.Fprintf(out, "\t\t}\n")
|
||||||
return "&" + base + "Elem"
|
return v, v, []string{v}
|
||||||
|
case strings.HasPrefix(p.Type, "*"):
|
||||||
|
v := uniqueName(seen, name)
|
||||||
|
fmt.Fprintf(out, "\t\tvar %sA, %sB %s\n", v, v, elem)
|
||||||
|
fmt.Fprintf(out, "\t\tw%s := %s(rng.Intn(256))\n", v, goCast(elem))
|
||||||
|
fmt.Fprintf(out, "\t\t%sA = w%s\n", v, v)
|
||||||
|
fmt.Fprintf(out, "\t\t%sB = w%s\n", v, v)
|
||||||
|
return "&" + v + "A", "&" + v + "B", nil
|
||||||
default:
|
default:
|
||||||
if base == "" {
|
v := uniqueName(seen, name)
|
||||||
base = "arg"
|
fmt.Fprintf(out, "\t\tw%s := %s(rng.Intn(512))\n", v, goCast(""))
|
||||||
}
|
fmt.Fprintf(out, "\t\tvar %sA, %sB %s = w%s, w%s\n", v, v, p.Type, v, v)
|
||||||
return base + "Scalar"
|
return v + "A", v + "B", nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// genParamSeed emits the seeding statements for one parameter.
|
// uniqueName de-duplicates seeded variable names when one kernel takes two
|
||||||
func genParamSeed(out *strings.Builder, p sigParam, written map[string]bool) {
|
// parameters of the same name (impossible in Go) or a name repeats across
|
||||||
names := sortedUnique(p.Names)
|
// kernels in one file.
|
||||||
for _, n := range names {
|
func uniqueName(seen map[string]bool, base string) string {
|
||||||
switch {
|
if base == "" {
|
||||||
case strings.HasPrefix(p.Type, "[]"):
|
base = "arg"
|
||||||
if written[n+"Slice"] {
|
}
|
||||||
continue
|
if !seen[base] {
|
||||||
}
|
seen[base] = true
|
||||||
written[n+"Slice"] = true
|
return base
|
||||||
fmt.Fprintf(out, "\t\t%sSlice := make([]%s, 1+rng.Intn(512))\n", n, strings.TrimPrefix(p.Type, "[]"))
|
}
|
||||||
fmt.Fprintf(out, "\t\tfor i := range %sSlice {\n\t\t\t%sSlice[i] = %s(rng.Intn(256))\n\t\t}\n", n, strings.TrimPrefix(p.Type, "[]"), goCast(strings.TrimPrefix(p.Type, "[]")))
|
for i := 2; ; i++ {
|
||||||
case strings.HasPrefix(p.Type, "*"):
|
cand := fmt.Sprintf("%s%d", base, i)
|
||||||
if written[n+"Elem"] {
|
if !seen[cand] {
|
||||||
continue
|
seen[cand] = true
|
||||||
}
|
return cand
|
||||||
written[n+"Elem"] = true
|
|
||||||
fmt.Fprintf(out, "\t\tvar %sElem %s\n\t\t%sElem = %s(rng.Intn(256))\n", n, strings.TrimPrefix(p.Type, "*"), n, goCast(strings.TrimPrefix(p.Type, "*")))
|
|
||||||
default:
|
|
||||||
if written[n+"Scalar"] {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
written[n+"Scalar"] = true
|
|
||||||
fmt.Fprintf(out, "\t\t%sScalar := rng.Intn(512)\n", n)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func sortedUnique(names []string) []string {
|
// goCast returns the conversion turning rng.Intn into the element type.
|
||||||
seen := map[string]bool{}
|
|
||||||
var out []string
|
|
||||||
for _, n := range names {
|
|
||||||
if !seen[n] {
|
|
||||||
seen[n] = true
|
|
||||||
out = append(out, n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
sort.Strings(out)
|
|
||||||
return out
|
|
||||||
}
|
|
||||||
|
|
||||||
// goCast returns the conversion turning rng.Intn(256) into the element type.
|
|
||||||
func goCast(elem string) string {
|
func goCast(elem string) string {
|
||||||
switch elem {
|
switch elem {
|
||||||
case "byte", "uint8":
|
case "byte", "uint8":
|
||||||
|
|||||||
Reference in New Issue
Block a user