Source file src/cmd/compile/internal/ssacompile/mem2reg.go

     1  // Copyright 2026 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 ssacompile
     6  
     7  // This file contains the mem2reg pass.
     8  // mem2reg promotes memory operations to register operations.
     9  // This is a classic compiler optimization that can significantly
    10  // improve the performance of generated code by reducing the number
    11  // of memory accesses.
    12  //
    13  // The algorithm identifies memory stores (OpStore) to stack slots
    14  // that are followed by loads (OpLoad) from the same stack slot,
    15  // with no intervening stores to that slot. The loads are then
    16  // replaced with the value that was stored.
    17  
    18  import (
    19  	"cmd/compile/internal/base"
    20  	"cmd/compile/internal/ir"
    21  	"cmd/compile/internal/ssa"
    22  	"cmd/compile/internal/ssa/ssaop"
    23  	"cmd/compile/internal/types"
    24  	"internal/buildcfg"
    25  	"slices"
    26  )
    27  
    28  type useType int
    29  
    30  const (
    31  	useLoad   useType = 1 << iota // used as the address to load from
    32  	useStore                      // used as the address to store to
    33  	useOffset                     // pointer arithmetic (OpOffPtr)
    34  	useCopy                       // uses propagated through a copy
    35  	useZero                       // used in OpZero
    36  	useMove                       // used in OpMove
    37  	useOther                      // includes escaping uses and OpPtrIndex
    38  )
    39  
    40  // useTable wraps the shared reverse-use table (see uses.go) so that
    41  // queries about values created after the table was built - their IDs
    42  // are out of range - report no uses instead of panicking. The mem2reg
    43  // and decomposeAddr passes ask about such values after their rewrites.
    44  type useTable struct {
    45  	useInfo
    46  }
    47  
    48  // of returns the uses of v recorded in u, or nil if v postdates the table.
    49  func (u *useTable) of(v *ssa.Value) []*ssa.Value {
    50  	if int(v.ID) >= len(u.starts) {
    51  		return nil
    52  	}
    53  	return u.get(v)
    54  }
    55  
    56  // classifyUses analyzes how a pointer is used and returns a bitmask of use types.
    57  func classifyUses(v *ssa.Value, uses *useTable) useType {
    58  	q := []*ssa.Value{v}
    59  	var u useType
    60  	for len(q) > 0 {
    61  		curr := q[len(q)-1]
    62  		q = q[:len(q)-1]
    63  		for _, use := range uses.of(curr) {
    64  			switch use.Op {
    65  			case ssaop.OpCopy:
    66  				u |= useCopy
    67  				q = append(q, use)
    68  			case ssaop.OpOffPtr:
    69  				u |= useOffset
    70  				q = append(q, use)
    71  			case ssaop.OpLoad:
    72  				u |= useLoad
    73  			case ssaop.OpStore:
    74  				if curr == use.Args[1] {
    75  					// Storing the address as value.
    76  					u |= useOther
    77  				} else {
    78  					u |= useStore
    79  				}
    80  			case ssaop.OpZero:
    81  				u |= useZero
    82  			case ssaop.OpMove:
    83  				u |= useMove
    84  			default:
    85  				// TODO: handle more cases?
    86  				// e.g. AddPtr, SubPtr, PtrIndex, NilCheck, EqPtr, NeqPtr
    87  				u |= useOther
    88  			}
    89  		}
    90  	}
    91  	return u
    92  }
    93  
    94  // variableDemographic summarizes how the address of an auto variable is
    95  // used, over all LocalAddr uses of that variable.
    96  type variableDemographic struct {
    97  	// ls is true when the address of this variable is only used by Load
    98  	// and Store, and no VarLive is applied on the name.
    99  	ls bool
   100  	// lszmco is true when the address of this variable is only used by
   101  	// Load, Store, Zero, Move, Copy and OffPtr, and no VarLive is
   102  	// applied on the name.
   103  	lszmco  bool
   104  	varDefs []*ssa.Value // All the VarDefs on this name.
   105  }
   106  
   107  // variableDemographics analyzes f and returns the variable demographics data of f.
   108  //
   109  //	demographics is the map from names to their demographics.
   110  //	localAddrs is the LocalAddr ssa values of the names in demographics.
   111  //	u tracks the uses of the values in f (see uses.go); it is nil, and
   112  //	nothing is allocated, when the function has no candidate variables.
   113  //
   114  // The caller must free u (u.free(f)) when non-nil.
   115  func variableDemographics(f *ssa.Func) (
   116  	demographics map[ssa.Aux]*variableDemographic,
   117  	localAddrs []*ssa.Value,
   118  	u *useTable) {
   119  	// Fast path: most functions have no candidate LocalAddr at all, and
   120  	// for them the pass chain has nothing to do. Detect that with a scan
   121  	// that allocates nothing, so those functions pay (almost) nothing.
   122  	found := false
   123  	for _, b := range f.Blocks {
   124  		for _, c := range b.ControlValues() {
   125  			if c.Op == ssaop.OpOffPtr || c.Op == ssaop.OpLocalAddr {
   126  				f.Fatalf("unexpected pointer value in block control")
   127  			}
   128  		}
   129  		for _, v := range b.Values {
   130  			if v.Op == ssaop.OpLocalAddr {
   131  				if n := v.Aux.(*ir.Name); n.Class == ir.PAUTO || isABIInternalParam(f, n) {
   132  					found = true
   133  					break
   134  				}
   135  			}
   136  		}
   137  		if found {
   138  			break
   139  		}
   140  	}
   141  	if !found {
   142  		return nil, nil, nil
   143  	}
   144  
   145  	// Build the reverse-use table using the shared cache-backed
   146  	// infrastructure from uses.go.
   147  	u = &useTable{uses(f)}
   148  	varLives := map[ssa.Aux]bool{}
   149  	demographics = make(map[ssa.Aux]*variableDemographic)
   150  	for _, b := range f.Blocks {
   151  		for _, v := range b.Values {
   152  			switch v.Op {
   153  			case ssaop.OpVarLive:
   154  				varLives[v.Aux] = true
   155  			case ssaop.OpVarDef:
   156  				d := demographics[v.Aux]
   157  				if d == nil {
   158  					d = &variableDemographic{ls: true, lszmco: true}
   159  					demographics[v.Aux] = d
   160  				}
   161  				d.varDefs = append(d.varDefs, v)
   162  			}
   163  		}
   164  	}
   165  	// Compute the use pattern of variables and the memory access demographic patterns.
   166  	for _, b := range f.Blocks {
   167  		for _, v := range b.Values {
   168  			if v.Op == ssaop.OpLocalAddr {
   169  				if n := v.Aux.(*ir.Name); n.Class == ir.PAUTO || isABIInternalParam(f, n) {
   170  					ut := classifyUses(v, u)
   171  					d := demographics[n]
   172  					if d == nil {
   173  						d = &variableDemographic{ls: true, lszmco: true}
   174  						demographics[n] = d
   175  					}
   176  					if ut&^(useLoad|useStore) != 0 {
   177  						d.ls = false
   178  					}
   179  					if ut&useOther != 0 {
   180  						d.lszmco = false
   181  					}
   182  					localAddrs = append(localAddrs, v)
   183  				}
   184  			}
   185  		}
   186  	}
   187  	// Names that have a VarLive applied are expected by the runtime to stay
   188  	// in memory (e.g. to be kept alive across a call for the GC), so they do
   189  	// not qualify for promotion.
   190  	for n := range varLives {
   191  		if d := demographics[n]; d != nil {
   192  			d.ls = false
   193  			d.lszmco = false
   194  		}
   195  	}
   196  	return
   197  }
   198  
   199  var reinterpretOpMap = map[[2]types.Kind]ssaop.Op{
   200  	{types.TUINT32, types.TFLOAT32}: ssaop.OpI32AsF32,
   201  	{types.TINT32, types.TFLOAT32}:  ssaop.OpI32AsF32,
   202  	{types.TFLOAT32, types.TUINT32}: ssaop.OpF32AsI32,
   203  	{types.TFLOAT32, types.TINT32}:  ssaop.OpF32AsI32,
   204  	{types.TUINT64, types.TFLOAT64}: ssaop.OpI64AsF64,
   205  	{types.TINT64, types.TFLOAT64}:  ssaop.OpI64AsF64,
   206  	{types.TFLOAT64, types.TUINT64}: ssaop.OpF64AsI64,
   207  	{types.TFLOAT64, types.TINT64}:  ssaop.OpF64AsI64,
   208  }
   209  
   210  // reinterpretOp returns a t1 => t2 reinterpret op if available in [reinterpretOpMap].
   211  // Otherwise it returns OpInvalid.
   212  func reinterpretOp(t1, t2 *types.Type) ssaop.Op {
   213  	if buildcfg.GOARCH != "amd64" {
   214  		// Currently only amd64 has the proper lowering rules for these ops.
   215  		return ssaop.OpInvalid
   216  	}
   217  	if op, ok := reinterpretOpMap[[2]types.Kind{t1.Kind(), t2.Kind()}]; ok {
   218  		return op
   219  	}
   220  
   221  	return ssaop.OpInvalid
   222  }
   223  
   224  // copyCompatibleType reports whether a value of type t1 can directly stand in
   225  // for a value of type t2 (same size and compatible representation).
   226  func copyCompatibleType(t1, t2 *types.Type) bool {
   227  	if t1.Size() != t2.Size() {
   228  		return false
   229  	}
   230  	if t1.IsInteger() {
   231  		return t2.IsInteger()
   232  	}
   233  	if ssa.IsPtr(t1) {
   234  		return ssa.IsPtr(t2)
   235  	}
   236  	return t1.Compare(t2) == types.CMPeq
   237  }
   238  
   239  // sortLoadStores sorts loadStores (Loads and Stores in block bb) into
   240  // program order, i.e. the order given by bb's store chain. This is needed
   241  // because there is no guarantee of their order in b.Values before the
   242  // schedule pass.
   243  //
   244  // memoryOrders is a per-block cache, populated on demand:
   245  // memoryOrders[b.ID][v.ID] is the position of the memory value v in b's
   246  // store chain, in even increments. Position 0 is the memory coming into
   247  // the block (InitMem, a memory Phi, or any memory defined outside bb);
   248  // each memory-producing value in bb (a Store, but also any other memory
   249  // generator such as a call) gets the position of its memory arg plus
   250  // two. A Load sorts at the position of its memory arg plus one - odd,
   251  // so after the value that produced that memory and before whatever
   252  // consumes it, which orders Stores before Loads of their output memory
   253  // with a plain integer comparison.
   254  func sortLoadStores(bb *ssa.Block, loadStores []*ssa.Value, memoryOrders map[ssa.ID]map[ssa.ID]int) []*ssa.Value {
   255  	memOrder, ok := memoryOrders[bb.ID]
   256  	if !ok {
   257  		memOrder = make(map[ssa.ID]int)
   258  		var computeDepth func(v *ssa.Value) int
   259  		computeDepth = func(v *ssa.Value) int {
   260  			if d, ok := memOrder[v.ID]; ok {
   261  				return d
   262  			}
   263  			// Starting point
   264  			if v.Block != bb || v.Op == ssaop.OpInitMem || v.Op == ssaop.OpPhi {
   265  				memOrder[v.ID] = 0
   266  				return 0
   267  			}
   268  			// Search backwards
   269  			d := computeDepth(v.MemoryArg()) + 2
   270  			memOrder[v.ID] = d
   271  			return d
   272  		}
   273  		for _, v := range bb.Values {
   274  			if v.Type.IsMemory() {
   275  				computeDepth(v)
   276  			}
   277  		}
   278  		memoryOrders[bb.ID] = memOrder
   279  	}
   280  	key := func(v *ssa.Value) int {
   281  		if v.Op == ssaop.OpLoad {
   282  			return memOrder[v.MemoryArg().ID] + 1
   283  		}
   284  		return memOrder[v.ID]
   285  	}
   286  	slices.SortFunc(loadStores, func(a, b *ssa.Value) int {
   287  		return key(a) - key(b)
   288  	})
   289  	return loadStores
   290  }
   291  
   292  func mem2reg(f *ssa.Func) {
   293  	changed := false
   294  	if base.Flag.N != 0 {
   295  		return
   296  	}
   297  	st := f.NewStats("mem2reg")
   298  
   299  	// Get demographics
   300  	demographics, localAddrs, uses := variableDemographics(f)
   301  	if uses == nil {
   302  		// No candidate variables; nothing to do.
   303  		return
   304  	}
   305  	defer uses.free(f)
   306  	memoryOrders := make(map[ssa.ID]map[ssa.ID]int)
   307  
   308  	// First, we need to group all LocalAddr/Addrs that point to the same variable name
   309  	// varGrouped[n] = all localAddrs that have n as their Aux field and have only load and stores.
   310  	varGrouped := make(map[*ir.Name][]*ssa.Value)
   311  	for _, v := range localAddrs {
   312  		// TODO: mem2reg currently only supports load and stores on the LocalAddr directly.
   313  		// However as a small step further it could also handle copy-derived addresses:
   314  		// OffPtr [0] LocalAddr
   315  		// Copy    LocalAddr
   316  		if d, ok := demographics[v.Aux]; !ok || !d.ls {
   317  			continue
   318  		}
   319  		n := v.Aux.(*ir.Name)
   320  		if n.Class != ir.PAUTO {
   321  			// Not handling params right now.
   322  			continue
   323  		}
   324  		varGrouped[n] = append(varGrouped[n], v)
   325  	}
   326  	namesOrdered := []*ir.Name{}
   327  	for n, vag := range varGrouped {
   328  		namesOrdered = append(namesOrdered, n)
   329  		slices.SortFunc(vag, func(a, b *ssa.Value) int { return int(a.ID - b.ID) })
   330  	}
   331  	slices.SortFunc(namesOrdered, func(a, b *ir.Name) int {
   332  		return int(varGrouped[a][0].ID - varGrouped[b][0].ID)
   333  	})
   334  	// Utility functions
   335  	removeStore := func(v *ssa.Value) {
   336  		changed = true
   337  		v.SetArgs1(v.MemoryArg())
   338  		v.Aux = nil
   339  		v.AuxInt = 0
   340  		v.Op = ssaop.OpCopy
   341  	}
   342  	type loadCandidate struct {
   343  		l *ssa.Value // The load
   344  		v *ssa.Value // The value to replace the load
   345  		// if v needs to be reinterpreted to match l's type, the op required. OpInvalid otherwise.
   346  		// if l.Type != v.Type, this op must not be OpInvalid.
   347  		reinterpret ssaop.Op
   348  	}
   349  	replaceLoad := func(lc loadCandidate) {
   350  		changed = true
   351  		if lc.reinterpret == ssaop.OpInvalid {
   352  			if !copyCompatibleType(lc.l.Type, lc.v.Type) {
   353  				f.Fatalf("mem2reg: load is being replaced by a value of an incompatible type")
   354  			}
   355  			lc.l.SetArgs1(lc.v)
   356  		} else {
   357  			// TODO: if we want to be more cautious, check that the reinterpret op will
   358  			// actually make the type right.
   359  			lc.l.SetArgs1(lc.l.Block.NewValue1(lc.l.Pos, lc.reinterpret, lc.l.Type, lc.v))
   360  		}
   361  		lc.l.Aux = nil
   362  		lc.l.AuxInt = 0
   363  		lc.l.Op = ssaop.OpCopy
   364  	}
   365  	// Now we should start the promotion
   366  	// Simple case one - all uses of v is within the same block.
   367  	storeCands := []*ssa.Value{} // These stores are to be overwritten
   368  	loadCands := []loadCandidate{}
   369  	vaUses := []*ssa.Value{}
   370  NextVar:
   371  	for _, n := range namesOrdered {
   372  		vag := varGrouped[n]
   373  		var block *ssa.Block
   374  		for _, va := range vag {
   375  			// Walk through all uses of v and check if they are within the same block
   376  			for _, use := range uses.of(va) {
   377  				if block == nil {
   378  					block = use.Block
   379  				}
   380  				if use.Block != block {
   381  					// Across control flow, needs DF and Phi.
   382  					continue NextVar
   383  				}
   384  			}
   385  		}
   386  		// The uses are all within the same block and are loads/stores.
   387  		storeCands = storeCands[:0]
   388  		loadCands = loadCands[:0]
   389  		vaUses = vaUses[:0]
   390  		for _, va := range vag {
   391  			vaUses = append(vaUses, uses.of(va)...)
   392  		}
   393  		vaUses = sortLoadStores(block, vaUses, memoryOrders)
   394  		var curV *ssa.Value // The current value of n.
   395  		for _, v := range vaUses {
   396  			switch v.Op {
   397  			case ssaop.OpLoad:
   398  				if curV == nil {
   399  					f.Fatalf("mem2reg sees a load from an auto variable before any store")
   400  				}
   401  				reinter := ssaop.OpInvalid
   402  				if !copyCompatibleType(v.Type, curV.Type) {
   403  					reinter = reinterpretOp(curV.Type, v.Type)
   404  					if reinter == ssaop.OpInvalid {
   405  						// Not something we can optimize, bailout the variable and
   406  						// continue to the next variable.
   407  						st.Record("incompatible types in single block case", 1)
   408  						delete(varGrouped, n)
   409  						continue NextVar
   410  					}
   411  				}
   412  				loadCands = append(loadCands, loadCandidate{
   413  					l:           v,
   414  					v:           curV,
   415  					reinterpret: reinter,
   416  				})
   417  			case ssaop.OpStore:
   418  				curV = v.Args[1]
   419  				storeCands = append(storeCands, v)
   420  			default:
   421  				f.Fatalf("should only has load or store uses")
   422  			}
   423  		}
   424  		// Do the replacement.
   425  		for _, v := range storeCands {
   426  			// Remove the stores.
   427  			removeStore(v)
   428  		}
   429  		for _, v := range demographics[n].varDefs {
   430  			// Also remove the VarDefs.
   431  			removeStore(v)
   432  		}
   433  		for _, lc := range loadCands {
   434  			// Replace the loads with the corresponding value.
   435  			replaceLoad(lc)
   436  		}
   437  		// This var is done, remove it from varGrouped.
   438  		if f.Pass.Debug > 1 {
   439  			f.Warnl(n.Pos(), "promoted %v in single block case", n)
   440  		}
   441  		delete(varGrouped, n)
   442  		st.Record("promoted variable in single block case", 1)
   443  	}
   444  
   445  	// TODO: the rest of variables needs analysis across basic blocks, implement this.
   446  	if changed {
   447  		deadcode(f)
   448  	}
   449  }
   450  

View as plain text