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/token"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
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
|
||||
}
|
||||
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 {
|
||||
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, "\trng := rand.New(rand.NewSource(1))\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 {
|
||||
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))
|
||||
for _, p := range params {
|
||||
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\tgot := %s(%s)\n", name, strings.Join(aArgs, ", "))
|
||||
fmt.Fprintf(&out, "\t\twant := %sPortable(%s)\n", name, strings.Join(bArgs, ", "))
|
||||
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}\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 {
|
||||
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
|
||||
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
|
||||
if fd.Type.Results != nil {
|
||||
@@ -207,78 +227,63 @@ func resultDecl(results []sigResult) string {
|
||||
return strings.Join(parts, ", ")
|
||||
}
|
||||
|
||||
// seedArg renders the argument expression referencing the seeded variables.
|
||||
func seedArg(p sigParam) string {
|
||||
if len(p.Names) == 0 {
|
||||
return typeSeedExpr(p.Type, "")
|
||||
// genParamSeed emits the seeding statements for one parameter and returns
|
||||
// the kernel-side (A) and reference-side (B) argument expressions, plus the
|
||||
// names of any slice variables written in place (compared after the calls).
|
||||
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 {
|
||||
case strings.HasPrefix(typ, "[]"):
|
||||
if base == "" {
|
||||
base = "arg"
|
||||
}
|
||||
return base + "Slice"
|
||||
case strings.HasPrefix(typ, "*"):
|
||||
if base == "" {
|
||||
base = "arg"
|
||||
}
|
||||
return "&" + base + "Elem"
|
||||
case isSlice:
|
||||
v := uniqueName(seen, name)
|
||||
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)
|
||||
fmt.Fprintf(out, "\t\tfor i := range %sA {\n", v)
|
||||
fmt.Fprintf(out, "\t\t\tw%s := %s(rng.Intn(256))\n", v, goCast(elem))
|
||||
fmt.Fprintf(out, "\t\t\t%sA[i] = w%s\n", v, v)
|
||||
fmt.Fprintf(out, "\t\t\t%sB[i] = w%s\n", v, v)
|
||||
fmt.Fprintf(out, "\t\t}\n")
|
||||
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:
|
||||
if base == "" {
|
||||
base = "arg"
|
||||
}
|
||||
return base + "Scalar"
|
||||
v := uniqueName(seen, name)
|
||||
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 v + "A", v + "B", nil
|
||||
}
|
||||
}
|
||||
|
||||
// genParamSeed emits the seeding statements for one parameter.
|
||||
func genParamSeed(out *strings.Builder, p sigParam, written map[string]bool) {
|
||||
names := sortedUnique(p.Names)
|
||||
for _, n := range names {
|
||||
switch {
|
||||
case strings.HasPrefix(p.Type, "[]"):
|
||||
if written[n+"Slice"] {
|
||||
continue
|
||||
}
|
||||
written[n+"Slice"] = true
|
||||
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, "[]")))
|
||||
case strings.HasPrefix(p.Type, "*"):
|
||||
if written[n+"Elem"] {
|
||||
continue
|
||||
}
|
||||
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)
|
||||
// uniqueName de-duplicates seeded variable names when one kernel takes two
|
||||
// parameters of the same name (impossible in Go) or a name repeats across
|
||||
// kernels in one file.
|
||||
func uniqueName(seen map[string]bool, base string) string {
|
||||
if base == "" {
|
||||
base = "arg"
|
||||
}
|
||||
if !seen[base] {
|
||||
seen[base] = true
|
||||
return base
|
||||
}
|
||||
for i := 2; ; i++ {
|
||||
cand := fmt.Sprintf("%s%d", base, i)
|
||||
if !seen[cand] {
|
||||
seen[cand] = true
|
||||
return cand
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func sortedUnique(names []string) []string {
|
||||
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.
|
||||
// goCast returns the conversion turning rng.Intn into the element type.
|
||||
func goCast(elem string) string {
|
||||
switch elem {
|
||||
case "byte", "uint8":
|
||||
|
||||
Reference in New Issue
Block a user