Source file src/simd/archsimd/_gen/unify/env.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 unify
     6  
     7  import (
     8  	"fmt"
     9  	"iter"
    10  	"reflect"
    11  	"strings"
    12  	"sync/atomic"
    13  )
    14  
    15  // An envSet is an immutable set of environments, where each environment is a
    16  // mapping from [ident]s to [Value]s.
    17  //
    18  // To keep this compact, we use an algebraic representation similar to
    19  // relational algebra. The atoms are zero, unit, or a singular binding:
    20  //
    21  // - A singular binding {x: v} is an environment set consisting of a single
    22  // environment that binds a single ident x to a single value v.
    23  //
    24  // - Zero (0) is the empty set.
    25  //
    26  // - Unit (1) is an environment set consisting of a single, empty environment
    27  // (no bindings).
    28  //
    29  // From these, we build up more complex sets of environments using sums and
    30  // cross products:
    31  //
    32  // - A sum, E + F, is simply the union of the two environment sets: E ∪ F
    33  //
    34  // - A cross product, E ⨯ F, is the Cartesian product of the two environment
    35  // sets, followed by joining each pair of environments: {e ⊕ f | (e, f) ∊ E ⨯ F}
    36  //
    37  // The join of two environments, e ⊕ f, is an environment that contains all of
    38  // the bindings in either e or f. To detect bugs, it is an error if an
    39  // identifier is bound in both e and f (however, see below for what we could do
    40  // differently).
    41  //
    42  // Environment sets form a commutative semiring and thus obey the usual
    43  // commutative semiring rules:
    44  //
    45  //	e + 0 = e
    46  //	e ⨯ 0 = 0
    47  //	e ⨯ 1 = e
    48  //	e + f = f + e
    49  //	e ⨯ f = f ⨯ e
    50  //
    51  // Furthermore, environments sets are additively and multiplicatively idempotent
    52  // because + and ⨯ are themselves defined in terms of sets:
    53  //
    54  //	e + e = e
    55  //	e ⨯ e = e
    56  //
    57  // # Examples
    58  //
    59  // To represent {{x: 1, y: 1}, {x: 2, y: 2}}, we build the two environments and
    60  // sum them:
    61  //
    62  //	({x: 1} ⨯ {y: 1}) + ({x: 2} ⨯ {y: 2})
    63  //
    64  // If we add a third variable z that can be 1 or 2, independent of x and y, we
    65  // get four logical environments:
    66  //
    67  //	{x: 1, y: 1, z: 1}
    68  //	{x: 2, y: 2, z: 1}
    69  //	{x: 1, y: 1, z: 2}
    70  //	{x: 2, y: 2, z: 2}
    71  //
    72  // This could be represented as a sum of all four environments, but because z is
    73  // independent, we can use a more compact representation:
    74  //
    75  //	(({x: 1} ⨯ {y: 1}) + ({x: 2} ⨯ {y: 2})) ⨯ ({z: 1} + {z: 2})
    76  //
    77  // # Generalized cross product
    78  //
    79  // While cross-product is currently restricted to disjoint environments, we
    80  // could generalize the definition of joining two environments to:
    81  //
    82  //	{xₖ: vₖ} ⊕ {xₖ: wₖ} = {xₖ: vₖ ∩ wₖ} (where unbound idents are bound to the [Top] value, ⟙)
    83  //
    84  // where v ∩ w is the unification of v and w. This itself could be coarsened to
    85  //
    86  //	v ∩ w = v if w = ⟙
    87  //	      = w if v = ⟙
    88  //	      = v if v = w
    89  //	      = 0 otherwise
    90  //
    91  // We could use this rule to implement substitution. For example, E ⨯ {x: 1}
    92  // narrows environment set E to only environments in which x is bound to 1. But
    93  // we currently don't do this.
    94  type envSet struct {
    95  	root *envExpr
    96  }
    97  
    98  type envExpr struct {
    99  	// TODO: A tree-based data structure for this may not be ideal, since it
   100  	// involves a lot of walking to find things and we often have to do deep
   101  	// rewrites anyway for partitioning. Would some flattened array-style
   102  	// representation be better, possibly combined with an index of ident uses?
   103  	// We could even combine that with an immutable array abstraction (ala
   104  	// Clojure) that could enable more efficient construction operations.
   105  
   106  	kind envExprKind
   107  	mask [2]uint64
   108  
   109  	// For envBinding
   110  	id  *ident
   111  	val *Value
   112  
   113  	// For sum or product. Len must be >= 2 and none of the elements can have
   114  	// the same kind as this node.
   115  	operands []*envExpr
   116  }
   117  
   118  type envExprKind byte
   119  
   120  const (
   121  	envZero envExprKind = iota
   122  	envUnit
   123  	envProduct
   124  	envSum
   125  	envBinding
   126  )
   127  
   128  var (
   129  	// topEnv is the unit value (multiplicative identity) of a [envSet].
   130  	topEnv = envSet{envExprUnit}
   131  	// bottomEnv is the zero value (additive identity) of a [envSet].
   132  	bottomEnv = envSet{envExprZero}
   133  
   134  	envExprZero = &envExpr{kind: envZero}
   135  	envExprUnit = &envExpr{kind: envUnit}
   136  )
   137  
   138  // bind binds id to each of vals in e.
   139  //
   140  // Its panics if id is already bound in e.
   141  //
   142  // Environments are typically initially constructed by starting with [topEnv]
   143  // and calling bind one or more times.
   144  func (e envSet) bind(id *ident, vals ...*Value) envSet {
   145  	if e.isEmpty() {
   146  		return bottomEnv
   147  	}
   148  
   149  	// TODO: If any of vals are _, should we just drop that val? We're kind of
   150  	// inconsistent about whether an id missing from e means id is invalid or
   151  	// means id is _.
   152  
   153  	// Check that id isn't present in e.
   154  	for range e.root.bindings(id) {
   155  		panic("id " + id.name + " already present in environment")
   156  	}
   157  
   158  	// Create a sum of all the values.
   159  	bindings := make([]*envExpr, 0, 1)
   160  	for _, val := range vals {
   161  		bindings = append(bindings, &envExpr{kind: envBinding, mask: id.mask, id: id, val: val})
   162  	}
   163  
   164  	// Multiply it in.
   165  	return envSet{newEnvExprProduct(e.root, newEnvExprSum(bindings...))}
   166  }
   167  
   168  func (e envSet) isEmpty() bool {
   169  	return e.root.kind == envZero
   170  }
   171  
   172  // bindings yields all [envBinding] nodes in e with the given id. If id is nil,
   173  // it yields all binding nodes.
   174  func (e *envExpr) bindings(id *ident) iter.Seq[*envExpr] {
   175  	// This is just a pre-order walk and it happens this is the only thing we
   176  	// need a pre-order walk for.
   177  	return func(yield func(*envExpr) bool) {
   178  		var rec func(e *envExpr) bool
   179  		rec = func(e *envExpr) bool {
   180  			if id != nil && (e.mask[0]&id.mask[0] == 0 && e.mask[1]&id.mask[1] == 0) {
   181  				return true
   182  			}
   183  			if e.kind == envBinding && (id == nil || e.id == id) {
   184  				if !yield(e) {
   185  					return false
   186  				}
   187  			}
   188  			for _, o := range e.operands {
   189  				if !rec(o) {
   190  					return false
   191  				}
   192  			}
   193  			return true
   194  		}
   195  		rec(e)
   196  	}
   197  }
   198  
   199  // newEnvExprProduct constructs a product node from exprs, performing
   200  // simplifications. It does NOT check that bindings are disjoint.
   201  func newEnvExprProduct(exprs ...*envExpr) *envExpr {
   202  	factors := make([]*envExpr, 0, 2)
   203  	for _, expr := range exprs {
   204  		switch expr.kind {
   205  		case envZero:
   206  			return envExprZero
   207  		case envUnit:
   208  			// No effect on product
   209  		case envProduct:
   210  			factors = append(factors, expr.operands...)
   211  		default:
   212  			factors = append(factors, expr)
   213  		}
   214  	}
   215  
   216  	if len(factors) == 0 {
   217  		return envExprUnit
   218  	} else if len(factors) == 1 {
   219  		return factors[0]
   220  	}
   221  	var mask [2]uint64
   222  	for _, f := range factors {
   223  		mask[0] |= f.mask[0]
   224  		mask[1] |= f.mask[1]
   225  	}
   226  	return &envExpr{kind: envProduct, mask: mask, operands: factors}
   227  }
   228  
   229  // newEnvExprSum constructs a sum node from exprs, performing simplifications.
   230  func newEnvExprSum(exprs ...*envExpr) *envExpr {
   231  	// TODO: If all of envs are products (or bindings), factor any common terms.
   232  	// E.g., x * y + x * z ==> x * (y + z). This is easy to do for binding
   233  	// terms, but harder to do for more general terms.
   234  
   235  	var have smallSet[*envExpr]
   236  	terms := make([]*envExpr, 0, 2)
   237  	for _, expr := range exprs {
   238  		switch expr.kind {
   239  		case envZero:
   240  			// No effect on sum
   241  		case envSum:
   242  			for _, expr1 := range expr.operands {
   243  				if have.Add(expr1) {
   244  					terms = append(terms, expr1)
   245  				}
   246  			}
   247  		default:
   248  			if have.Add(expr) {
   249  				terms = append(terms, expr)
   250  			}
   251  		}
   252  	}
   253  
   254  	if len(terms) == 0 {
   255  		return envExprZero
   256  	} else if len(terms) == 1 {
   257  		return terms[0]
   258  	}
   259  	var mask [2]uint64
   260  	for _, t := range terms {
   261  		mask[0] |= t.mask[0]
   262  		mask[1] |= t.mask[1]
   263  	}
   264  	return &envExpr{kind: envSum, mask: mask, operands: terms}
   265  }
   266  
   267  func crossEnvs(env1, env2 envSet) envSet {
   268  	// Confirm that envs have disjoint idents.
   269  	var ids1 smallSet[*ident]
   270  	for e := range env1.root.bindings(nil) {
   271  		ids1.Add(e.id)
   272  	}
   273  	for e := range env2.root.bindings(nil) {
   274  		if ids1.Has(e.id) {
   275  			panic(fmt.Sprintf("%s bound on both sides of cross-product", e.id.name))
   276  		}
   277  	}
   278  
   279  	return envSet{newEnvExprProduct(env1.root, env2.root)}
   280  }
   281  
   282  func unionEnvs(envs ...envSet) envSet {
   283  	exprs := make([]*envExpr, len(envs))
   284  	for i := range envs {
   285  		exprs[i] = envs[i].root
   286  	}
   287  	return envSet{newEnvExprSum(exprs...)}
   288  }
   289  
   290  // envPartition is a subset of an env where id is bound to value in all
   291  // deterministic environments.
   292  type envPartition struct {
   293  	id    *ident
   294  	value *Value
   295  	env   envSet
   296  }
   297  
   298  // partitionBy splits e by distinct bindings of id and removes id from each
   299  // partition.
   300  //
   301  // If there are environments in e where id is not bound, they will not be
   302  // reflected in any partition.
   303  //
   304  // It panics if e is bottom, since attempting to partition an empty environment
   305  // set almost certainly indicates a bug.
   306  func (e envSet) partitionBy(id *ident) []envPartition {
   307  	if e.isEmpty() {
   308  		// We could return zero partitions, but getting here at all almost
   309  		// certainly indicates a bug.
   310  		panic("cannot partition empty environment set")
   311  	}
   312  
   313  	// Emit a partition for each value of id.
   314  	var seen smallSet[*Value]
   315  	var parts []envPartition
   316  	for n := range e.root.bindings(id) {
   317  		if !seen.Add(n.val) {
   318  			// Already emitted a partition for this value.
   319  			continue
   320  		}
   321  
   322  		parts = append(parts, envPartition{
   323  			id:    id,
   324  			value: n.val,
   325  			env:   envSet{e.root.substitute(id, n.val)},
   326  		})
   327  	}
   328  
   329  	return parts
   330  }
   331  
   332  // substitute replaces bindings of id to val with 1 and bindings of id to any
   333  // other value with 0 and simplifies the result.
   334  func (e *envExpr) substitute(id *ident, val *Value) *envExpr {
   335  	if e.mask[0]&id.mask[0] == 0 && e.mask[1]&id.mask[1] == 0 {
   336  		return e
   337  	}
   338  	switch e.kind {
   339  	default:
   340  		panic("bad kind")
   341  
   342  	case envZero, envUnit:
   343  		return e
   344  
   345  	case envBinding:
   346  		if e.id != id {
   347  			return e
   348  		} else if e.val != val {
   349  			return envExprZero
   350  		} else {
   351  			return envExprUnit
   352  		}
   353  
   354  	case envProduct, envSum:
   355  		// Substitute each operand. Sometimes, this won't change anything, so we
   356  		// build the new operands list lazily.
   357  		var nOperands []*envExpr
   358  		for i, op := range e.operands {
   359  			nOp := op.substitute(id, val)
   360  			if nOperands == nil && op != nOp {
   361  				// Operand diverged; initialize nOperands.
   362  				nOperands = make([]*envExpr, 0, len(e.operands))
   363  				nOperands = append(nOperands, e.operands[:i]...)
   364  			}
   365  			if nOperands != nil {
   366  				nOperands = append(nOperands, nOp)
   367  			}
   368  		}
   369  		if nOperands == nil {
   370  			// Nothing changed.
   371  			return e
   372  		}
   373  		if e.kind == envProduct {
   374  			return newEnvExprProduct(nOperands...)
   375  		} else {
   376  			return newEnvExprSum(nOperands...)
   377  		}
   378  	}
   379  }
   380  
   381  // A smallSet is a set optimized for stack allocation when small.
   382  type smallSet[T comparable] struct {
   383  	array [32]T
   384  	n     int
   385  
   386  	m map[T]struct{}
   387  }
   388  
   389  // Has returns whether val is in set.
   390  func (s *smallSet[T]) Has(val T) bool {
   391  	arr := s.array[:s.n]
   392  	for i := range arr {
   393  		if arr[i] == val {
   394  			return true
   395  		}
   396  	}
   397  	_, ok := s.m[val]
   398  	return ok
   399  }
   400  
   401  // Add adds val to the set and returns true if it was added (not already
   402  // present).
   403  func (s *smallSet[T]) Add(val T) bool {
   404  	// Test for presence.
   405  	if s.Has(val) {
   406  		return false
   407  	}
   408  
   409  	// Add it
   410  	if s.n < len(s.array) {
   411  		s.array[s.n] = val
   412  		s.n++
   413  	} else {
   414  		if s.m == nil {
   415  			s.m = make(map[T]struct{})
   416  		}
   417  		s.m[val] = struct{}{}
   418  	}
   419  	return true
   420  }
   421  
   422  type ident struct {
   423  	_    [0]func() // Not comparable (only compare *ident)
   424  	mask [2]uint64
   425  	name string
   426  }
   427  
   428  var identCounter atomic.Uint64
   429  
   430  func newIdent(name string) *ident {
   431  	c := identCounter.Add(1)
   432  	var mask [2]uint64
   433  	bit := c % 128
   434  	mask[bit/64] = 1 << (bit % 64)
   435  	return &ident{
   436  		mask: mask,
   437  		name: name,
   438  	}
   439  }
   440  
   441  type Var struct {
   442  	id *ident
   443  }
   444  
   445  func (d Var) Exact() bool {
   446  	// These can't appear in concrete Values.
   447  	panic("Exact called on non-concrete Value")
   448  }
   449  
   450  func (d Var) WhyNotExact() string {
   451  	// These can't appear in concrete Values.
   452  	return "WhyNotExact called on non-concrete Value"
   453  }
   454  
   455  func (d Var) decode(rv reflect.Value) error {
   456  	return &inexactError{"var", rv.Type().String()}
   457  }
   458  
   459  func (d Var) unify(w *Value, e envSet, swap bool, uf *unifier) (Domain, envSet, error) {
   460  	// TODO: Vars from !sums in the input can have a huge number of values.
   461  	// Unifying these could be way more efficient with some indexes over any
   462  	// exact values we can pull out, like Def fields that are exact Strings.
   463  	// Maybe we try to produce an array of yes/no/maybe matches and then we only
   464  	// have to do deeper evaluation of the maybes. We could probably cache this
   465  	// on an envTerm. It may also help to special-case Var/Var unification to
   466  	// pick which one to index versus enumerate.
   467  
   468  	if vd, ok := w.Domain.(Var); ok && d.id == vd.id {
   469  		// Unifying $x with $x results in $x. If we descend into this we'll have
   470  		// problems because we strip $x out of the environment to keep ourselves
   471  		// honest and then can't find it on the other side.
   472  		//
   473  		// TODO: I'm not positive this is the right fix.
   474  		return vd, e, nil
   475  	}
   476  
   477  	// We need to unify w with the value of d in each possible environment. We
   478  	// can save some work by grouping environments by the value of d, since
   479  	// there will be a lot of redundancy here.
   480  	var nEnvs []envSet
   481  	envParts := e.partitionBy(d.id)
   482  	for i, envPart := range envParts {
   483  		exit := uf.enterVar(d.id, i)
   484  		// Each branch logically gets its own copy of the initial environment
   485  		// (narrowed down to just this binding of the variable), and each branch
   486  		// may result in different changes to that starting environment.
   487  		res, e2, err := w.unify(envPart.value, envPart.env, swap, uf)
   488  		exit.exit()
   489  		if err != nil {
   490  			return nil, envSet{}, err
   491  		}
   492  		if res.Domain == nil {
   493  			// This branch entirely failed to unify, so it's gone.
   494  			continue
   495  		}
   496  		nEnv := e2.bind(d.id, res)
   497  		nEnvs = append(nEnvs, nEnv)
   498  	}
   499  
   500  	if len(nEnvs) == 0 {
   501  		// All branches failed
   502  		return nil, bottomEnv, nil
   503  	}
   504  
   505  	// The effect of this is entirely captured in the environment. We can return
   506  	// back the same Bind node.
   507  	return d, unionEnvs(nEnvs...), nil
   508  }
   509  
   510  // An identPrinter maps [ident]s to unique string names.
   511  type identPrinter struct {
   512  	ids   map[*ident]string
   513  	idGen map[string]int
   514  }
   515  
   516  func (p *identPrinter) unique(id *ident) string {
   517  	if p.ids == nil {
   518  		p.ids = make(map[*ident]string)
   519  		p.idGen = make(map[string]int)
   520  	}
   521  
   522  	name, ok := p.ids[id]
   523  	if !ok {
   524  		gen := p.idGen[id.name]
   525  		p.idGen[id.name]++
   526  		if gen == 0 {
   527  			name = id.name
   528  		} else {
   529  			name = fmt.Sprintf("%s#%d", id.name, gen)
   530  		}
   531  		p.ids[id] = name
   532  	}
   533  
   534  	return name
   535  }
   536  
   537  func (p *identPrinter) slice(ids []*ident) string {
   538  	var strs []string
   539  	for _, id := range ids {
   540  		strs = append(strs, p.unique(id))
   541  	}
   542  	return fmt.Sprintf("[%s]", strings.Join(strs, ", "))
   543  }
   544  

View as plain text