Source file src/simd/archsimd/internal/simd_test/simulation_helpers_test.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  //go:build goexperiment.simd
     6  
     7  package simd_test
     8  
     9  import (
    10  	"math"
    11  	"math/bits"
    12  	"unsafe"
    13  )
    14  
    15  func rotl[T unsigned](x T, dist uint64) T {
    16  	size := uint64(unsafe.Sizeof(x)) * 8
    17  	dist = dist & (size - 1)
    18  	if dist == 0 {
    19  		return x
    20  	}
    21  	return (x << dist) | (x >> (size - dist))
    22  }
    23  
    24  func rotr[T unsigned](x T, dist uint64) T {
    25  	size := uint64(unsafe.Sizeof(x)) * 8
    26  	dist = dist & (size - 1)
    27  	if dist == 0 {
    28  		return x
    29  	}
    30  	return (x >> dist) | (x << (size - dist))
    31  }
    32  
    33  // rotlOfSlice returns a slice simulation of a left rotate
    34  // of a specified distance.
    35  func rotlOfSlice[T unsigned](dist uint64) func(x []T) []T {
    36  	return map1[T](func(x T) T { return rotl(x, dist) })
    37  }
    38  
    39  // rotrOfSlice returns a slice simulation of a right rotate
    40  // of a specified distance.
    41  func rotrOfSlice[T unsigned](dist uint64) func(x []T) []T {
    42  	return map1[T](func(x T) T { return rotr(x, dist) })
    43  }
    44  
    45  func curry2[T, U, V any](f func(T, U) V, y U) func(x T) V {
    46  	return func(x T) V { return f(x, y) }
    47  }
    48  
    49  func curry1[T, U, V any](f func(T, U) V, x T) func(y U) V {
    50  	return func(y U) V { return f(x, y) }
    51  }
    52  
    53  func less[T number](x, y T) bool {
    54  	return x < y
    55  }
    56  func lessEqual[T number](x, y T) bool {
    57  	return x <= y
    58  }
    59  func greater[T number](x, y T) bool {
    60  	return x > y
    61  }
    62  func greaterEqual[T number](x, y T) bool {
    63  	return x >= y
    64  }
    65  func equal[T number](x, y T) bool {
    66  	return x == y
    67  }
    68  func notEqual[T number](x, y T) bool {
    69  	return x != y
    70  }
    71  
    72  func isNaN[T float](x T) bool {
    73  	return x != x
    74  }
    75  
    76  func abs[T number](x T) T {
    77  	// TODO this will need a non-standard FP-equality test.
    78  	if x == 0 { // true if x is -0.
    79  		return 0 // this is not a negative zero
    80  	}
    81  	if x < 0 {
    82  		return -x
    83  	}
    84  	return x
    85  }
    86  
    87  func neg[T number](x T) T {
    88  	return -x
    89  }
    90  
    91  func onesCount[T integer](x T) T {
    92  	size := uint64(unsafe.Sizeof(x)) * 8
    93  	return T(bits.OnesCount64(uint64(x) & ((1 << size) - 1)))
    94  }
    95  
    96  func ceil[T float](x T) T {
    97  	return T(math.Ceil(float64(x)))
    98  }
    99  func floor[T float](x T) T {
   100  	return T(math.Floor(float64(x)))
   101  }
   102  func not[T integer](x T) T {
   103  	return ^x
   104  }
   105  func round[T float](x T) T {
   106  	return T(math.RoundToEven(float64(x)))
   107  }
   108  func sqrt[T float](x T) T {
   109  	return T(math.Sqrt(float64(x)))
   110  }
   111  func trunc[T float](x T) T {
   112  	return T(math.Trunc(float64(x)))
   113  }
   114  
   115  func add[T number](x, y T) T {
   116  	return x + y
   117  }
   118  
   119  func sub[T number](x, y T) T {
   120  	return x - y
   121  }
   122  
   123  func max_[T number](x, y T) T { // "max" lands in infinite recursion
   124  	return max(x, y)
   125  }
   126  
   127  func min_[T number](x, y T) T { // "min" lands in infinite recursion
   128  	return min(x, y)
   129  }
   130  
   131  // Also mulLow for integers
   132  func mul[T number](x, y T) T {
   133  	return x * y
   134  }
   135  
   136  func div[T number](x, y T) T {
   137  	return x / y
   138  }
   139  
   140  func and[T integer](x, y T) T {
   141  	return x & y
   142  }
   143  
   144  func andNotI[T integer](x, y T) T {
   145  	return x & ^y // order corrected to match expectations
   146  }
   147  
   148  func orI[T integer](x, y T) T {
   149  	return x | y
   150  }
   151  
   152  func orNotI[T integer](x, y T) T {
   153  	return x | ^y
   154  }
   155  
   156  func xorI[T integer](x, y T) T {
   157  	return x ^ y
   158  }
   159  
   160  func ima[T integer](x, y, z T) T {
   161  	return x*y + z
   162  }
   163  
   164  func fma[T float](x, y, z T) T {
   165  	return T(math.FMA(float64(x), float64(y), float64(z)))
   166  }
   167  
   168  func toUint8[T number](x T) uint8 {
   169  	return uint8(x)
   170  }
   171  
   172  func toUint16[T number](x T) uint16 {
   173  	return uint16(x)
   174  }
   175  
   176  func toUint64[T number](x T) uint64 {
   177  	return uint64(x)
   178  }
   179  
   180  func toUint32[T number](x T) uint32 {
   181  	return uint32(x)
   182  }
   183  
   184  func toInt8[T number](x T) int8 {
   185  	return int8(x)
   186  }
   187  
   188  func toInt16[T number](x T) int16 {
   189  	return int16(x)
   190  }
   191  
   192  func toInt32[T number](x T) int32 {
   193  	return int32(x)
   194  }
   195  
   196  func toInt64[T number](x T) int64 {
   197  	return int64(x)
   198  }
   199  
   200  func toFloat32[T number](x T) float32 {
   201  	return float32(x)
   202  }
   203  
   204  func toFloat64[T number](x T) float64 {
   205  	return float64(x)
   206  }
   207  
   208  // X86 specific behavior for conversion from float to int32.
   209  // If the value cannot be represented as int32, it returns -0x80000000.
   210  func floatToInt32_x86[T float](x T) int32 {
   211  	switch y := (any(x)).(type) {
   212  	case float32:
   213  		if y != y || y < math.MinInt32 ||
   214  			y >= math.MaxInt32 { // float32(MaxInt32) == 0x80000000, actually overflows
   215  			return -0x80000000
   216  		}
   217  	case float64:
   218  		if y != y || y < math.MinInt32 ||
   219  			y > math.MaxInt32 { // float64(MaxInt32) is exact, no overflow
   220  			return -0x80000000
   221  		}
   222  	}
   223  	return int32(x)
   224  }
   225  
   226  // X86 specific behavior for conversion from float to int64.
   227  // If the value cannot be represented as int64, it returns -0x80000000_00000000.
   228  func floatToInt64_x86[T float](x T) int64 {
   229  	switch y := (any(x)).(type) {
   230  	case float32:
   231  		if y != y || y < math.MinInt64 ||
   232  			y >= math.MaxInt64 { // float32(MaxInt64) == 0x80000000_00000000, actually overflows
   233  			return -0x80000000_00000000
   234  		}
   235  	case float64:
   236  		if y != y || y < math.MinInt64 ||
   237  			y >= math.MaxInt64 { // float64(MaxInt64) == 0x80000000_00000000, also overflows
   238  			return -0x80000000_00000000
   239  		}
   240  	}
   241  	return int64(x)
   242  }
   243  
   244  // X86 specific behavior for conversion from float to uint32.
   245  // If the value cannot be represented as uint32, it returns 1<<32 - 1.
   246  func floatToUint32_x86[T float](x T) uint32 {
   247  	switch y := (any(x)).(type) {
   248  	case float32:
   249  		if y < 0 || y > math.MaxUint32 || y != y {
   250  			return 1<<32 - 1
   251  		}
   252  	case float64:
   253  		if y < 0 || y > math.MaxUint32 || y != y {
   254  			return 1<<32 - 1
   255  		}
   256  	}
   257  	return uint32(x)
   258  }
   259  
   260  // X86 specific behavior for conversion from float to uint64.
   261  // If the value cannot be represented as uint64, it returns 1<<64 - 1.
   262  func floatToUint64_x86[T float](x T) uint64 {
   263  	switch y := (any(x)).(type) {
   264  	case float32:
   265  		if y < 0 || y > math.MaxUint64 || y != y {
   266  			return 1<<64 - 1
   267  		}
   268  	case float64:
   269  		if y < 0 || y > math.MaxUint64 || y != y {
   270  			return 1<<64 - 1
   271  		}
   272  	}
   273  	return uint64(x)
   274  }
   275  
   276  func ceilResidueForPrecision[T float](i int) func(T) T {
   277  	f := 1.0
   278  	for i > 0 {
   279  		f *= 2
   280  		i--
   281  	}
   282  	return func(x T) T {
   283  		y := float64(x)
   284  		if math.IsInf(float64(x*T(f)), 0) {
   285  			return 0
   286  		}
   287  		// TODO sort out the rounding issues when T === float32
   288  		return T(y - math.Ceil(y*f)/f)
   289  	}
   290  }
   291  
   292  // Slice versions of all these elementwise operations
   293  
   294  func addSlice[T number](x, y []T) []T {
   295  	return map2[T](add)(x, y)
   296  }
   297  
   298  func subSlice[T number](x, y []T) []T {
   299  	return map2[T](sub)(x, y)
   300  }
   301  
   302  func maxSlice[T number](x, y []T) []T {
   303  	return map2[T](max_)(x, y)
   304  }
   305  
   306  func minSlice[T number](x, y []T) []T {
   307  	return map2[T](min_)(x, y)
   308  }
   309  
   310  // mulLow for integers
   311  func mulSlice[T number](x, y []T) []T {
   312  	return map2[T](mul)(x, y)
   313  }
   314  
   315  func divSlice[T number](x, y []T) []T {
   316  	return map2[T](div)(x, y)
   317  }
   318  
   319  func andSlice[T integer](x, y []T) []T {
   320  	return map2[T](and)(x, y)
   321  }
   322  
   323  func andNotSlice[T integer](x, y []T) []T {
   324  	return map2[T](andNotI)(x, y)
   325  }
   326  
   327  func orSlice[T integer](x, y []T) []T {
   328  	return map2[T](orI)(x, y)
   329  }
   330  
   331  func orNotSlice[T integer](x, y []T) []T {
   332  	return map2[T](orNotI)(x, y)
   333  }
   334  
   335  func xorSlice[T integer](x, y []T) []T {
   336  	return map2[T](xorI)(x, y)
   337  }
   338  
   339  func lessSlice[T number](x, y []T) []int64 {
   340  	return mapCompare[T](less)(x, y)
   341  }
   342  
   343  func lessEqualSlice[T number](x, y []T) []int64 {
   344  	return mapCompare[T](lessEqual)(x, y)
   345  }
   346  
   347  func greaterSlice[T number](x, y []T) []int64 {
   348  	return mapCompare[T](greater)(x, y)
   349  }
   350  
   351  func greaterEqualSlice[T number](x, y []T) []int64 {
   352  	return mapCompare[T](greaterEqual)(x, y)
   353  }
   354  
   355  func equalSlice[T number](x, y []T) []int64 {
   356  	return mapCompare[T](equal)(x, y)
   357  }
   358  
   359  func notEqualSlice[T number](x, y []T) []int64 {
   360  	return mapCompare[T](notEqual)(x, y)
   361  }
   362  
   363  func isNaNSlice[T float](x []T) []int64 {
   364  	return map1[T](func(x T) int64 {
   365  		if isNaN(x) {
   366  			return -1
   367  		}
   368  		return 0
   369  	})(x)
   370  }
   371  
   372  func ceilSlice[T float](x []T) []T {
   373  	return map1[T](ceil)(x)
   374  }
   375  
   376  func floorSlice[T float](x []T) []T {
   377  	return map1[T](floor)(x)
   378  }
   379  
   380  func notSlice[T integer](x []T) []T {
   381  	return map1[T](not)(x)
   382  }
   383  
   384  func roundSlice[T float](x []T) []T {
   385  	return map1[T](round)(x)
   386  }
   387  
   388  // lanewiseSlice is the common helper for interleave, deinterleave, and transpose
   389  // simulations. It handles lane computation, allocation, and iteration.
   390  // laneBits is the lane size in bits (128 for NEON/x86 128-bit, 0 for whole-input/SVE).
   391  // hi selects the half-lane offset (offHalf = 0 or half, for interleave hi/lo).
   392  // odd selects the single-element offset (offOne = 0 or 1, for deinterleave/transpose odd/even).
   393  // body receives (out, x, y, base, i, half, offHalf, offOne) for each pair within each lane
   394  // and performs the operation-specific element assignment.
   395  func lanewiseSlice[T number](laneBits int, hi bool, odd bool, body func(out, x, y []T, base, i, half, offHalf, offOne int)) func(x, y []T) []T {
   396  	return func(x, y []T) []T {
   397  		lane := laneBits / (8 * int(unsafe.Sizeof(x[0])))
   398  		if lane == 0 || lane > len(x) {
   399  			lane = len(x)
   400  		}
   401  		half := lane / 2
   402  		offHalf := 0
   403  		if hi {
   404  			offHalf = half
   405  		}
   406  		offOne := 0
   407  		if odd {
   408  			offOne = 1
   409  		}
   410  		out := make([]T, len(x))
   411  		for base := 0; base < len(x); base += lane {
   412  			for i := 0; i < half; i++ {
   413  				body(out, x, y, base, i, half, offHalf, offOne)
   414  			}
   415  		}
   416  		return out
   417  	}
   418  }
   419  
   420  func interleaveSlice[T number](laneBits int, hi bool) func(x, y []T) []T {
   421  	return lanewiseSlice(laneBits, hi, false, func(out, x, y []T, base, i, half, offHalf, _ int) {
   422  		out[base+2*i] = x[base+offHalf+i]
   423  		out[base+2*i+1] = y[base+offHalf+i]
   424  	})
   425  }
   426  
   427  func deinterleaveSlice[T number](laneBits int, odd bool) func(x, y []T) []T {
   428  	return lanewiseSlice(laneBits, false, odd, func(out, x, y []T, base, i, half, _, offOne int) {
   429  		out[base+i] = x[base+2*i+offOne]
   430  		out[base+half+i] = y[base+2*i+offOne]
   431  	})
   432  }
   433  
   434  func transposeSlice[T number](laneBits int, odd bool) func(x, y []T) []T {
   435  	return lanewiseSlice(laneBits, false, odd, func(out, x, y []T, base, i, half, _, offOne int) {
   436  		out[base+2*i] = x[base+2*i+offOne]
   437  		out[base+2*i+1] = y[base+2*i+offOne]
   438  	})
   439  }
   440  
   441  func sqrtSlice[T float](x []T) []T {
   442  	return map1[T](sqrt)(x)
   443  }
   444  
   445  func truncSlice[T float](x []T) []T {
   446  	return map1[T](trunc)(x)
   447  }
   448  
   449  func imaSlice[T integer](x, y, z []T) []T {
   450  	return map3[T](ima)(x, y, z)
   451  }
   452  
   453  func fmaSlice[T float](x, y, z []T) []T {
   454  	return map3[T](fma)(x, y, z)
   455  }
   456  
   457  // reduceSlice reduces x using fn as the combining operation.
   458  func reduceSlice[T number](x []T, fn func(a, b T) T) T {
   459  	acc := x[0]
   460  	for _, v := range x[1:] {
   461  		acc = fn(acc, v)
   462  	}
   463  	return acc
   464  }
   465  
   466  func satToInt8[T integer](x T) int8 {
   467  	var m int8 = -128
   468  	var M int8 = 127
   469  	if T(M) < T(m) { // expecting T being a larger type
   470  		panic("bad input type")
   471  	}
   472  	if x < T(m) {
   473  		return m
   474  	}
   475  	if x > T(M) {
   476  		return M
   477  	}
   478  	return int8(x)
   479  }
   480  
   481  func satToUint8[T integer](x T) uint8 {
   482  	var M uint8 = 255
   483  	if T(M) < 0 { // expecting T being a larger type
   484  		panic("bad input type")
   485  	}
   486  	if x < 0 {
   487  		return 0
   488  	}
   489  	if x > T(M) {
   490  		return M
   491  	}
   492  	return uint8(x)
   493  }
   494  
   495  func satToInt16[T integer](x T) int16 {
   496  	var m int16 = -32768
   497  	var M int16 = 32767
   498  	if T(M) < T(m) { // expecting T being a larger type
   499  		panic("bad input type")
   500  	}
   501  	if x < T(m) {
   502  		return m
   503  	}
   504  	if x > T(M) {
   505  		return M
   506  	}
   507  	return int16(x)
   508  }
   509  
   510  func satToUint16[T integer](x T) uint16 {
   511  	var M uint16 = 65535
   512  	if T(M) < 0 { // expecting T being a larger type
   513  		panic("bad input type")
   514  	}
   515  	if x < 0 {
   516  		return 0
   517  	}
   518  	if x > T(M) {
   519  		return M
   520  	}
   521  	return uint16(x)
   522  }
   523  
   524  func satToInt32[T integer](x T) int32 {
   525  	var m int32 = -1 << 31
   526  	var M int32 = 1<<31 - 1
   527  	if T(M) < T(m) { // expecting T being a larger type
   528  		panic("bad input type")
   529  	}
   530  	if x < T(m) {
   531  		return m
   532  	}
   533  	if x > T(M) {
   534  		return M
   535  	}
   536  	return int32(x)
   537  }
   538  
   539  func satToUint32[T integer](x T) uint32 {
   540  	var M uint32 = 1<<32 - 1
   541  	if T(M) < 0 { // expecting T being a larger type
   542  		panic("bad input type")
   543  	}
   544  	if x < 0 {
   545  		return 0
   546  	}
   547  	if x > T(M) {
   548  		return M
   549  	}
   550  	return uint32(x)
   551  }
   552  
   553  // shiftAmount extracts the signed shift amount from the least significant byte of s.
   554  // ARM64 SSHL/USHL use only bits [7:0] of the shift amount element, sign-extended.
   555  func shiftAmount[T integer](s T) int8 {
   556  	return int8(uint8(s))
   557  }
   558  
   559  // shiftBy shifts x by signed amount: positive = left, negative = right.
   560  func shiftBy[T integer](x T, amt int8) T {
   561  	a := int(amt)
   562  	if a > 0 {
   563  		return x << uint(a)
   564  	}
   565  	if a < 0 {
   566  		return x >> uint(-a)
   567  	}
   568  	return x
   569  }
   570  
   571  // shiftSaturatingSigned shifts x by signed amount with signed saturation on overflow.
   572  func shiftSaturatingSigned[T signed](x T, amt int8) T {
   573  	a := int(amt)
   574  	if a > 0 {
   575  		r := x << uint(a)
   576  		if r>>uint(a) != x { // overflow
   577  			bits := uint(unsafe.Sizeof(x)) * 8
   578  			if x >= 0 {
   579  				return ^T(0) ^ (T(1) << (bits - 1)) // MaxSigned
   580  			}
   581  			return T(1) << (bits - 1) // MinSigned
   582  		}
   583  		return r
   584  	}
   585  	if a < 0 {
   586  		return x >> uint(-a)
   587  	}
   588  	return x
   589  }
   590  
   591  // shiftSaturatingUnsigned shifts x by signed amount with unsigned saturation on overflow.
   592  func shiftSaturatingUnsigned[T unsigned](x T, amt int8) T {
   593  	a := int(amt)
   594  	if a > 0 {
   595  		r := x << uint(a)
   596  		if r>>uint(a) != x { // overflow
   597  			return ^T(0) // MaxUnsigned
   598  		}
   599  		return r
   600  	}
   601  	if a < 0 {
   602  		return x >> uint(-a)
   603  	}
   604  	return x
   605  }
   606  
   607  // Slice versions for shift operations
   608  
   609  // shiftSlice applies shiftBy element-wise using same-type slices.
   610  func shiftSlice[T integer](x, y []T) []T {
   611  	return map2(func(a, b T) T { return shiftBy(a, shiftAmount(b)) })(x, y)
   612  }
   613  
   614  // shiftMixedSlice applies shiftBy element-wise using mixed-type slices (unsigned data, signed amounts).
   615  func shiftMixedSlice[D integer, S integer](x []D, y []S) []D {
   616  	r := make([]D, len(x))
   617  	for i := range r {
   618  		r[i] = shiftBy(x[i], shiftAmount(y[i]))
   619  	}
   620  	return r
   621  }
   622  
   623  // shiftSaturatingSignedSlice applies saturating shift element-wise (same-type).
   624  func shiftSaturatingSignedSlice[T signed](x, y []T) []T {
   625  	return map2(func(a, b T) T { return shiftSaturatingSigned(a, shiftAmount(b)) })(x, y)
   626  }
   627  
   628  // shiftSaturatingUnsignedSlice applies saturating shift element-wise (mixed-type).
   629  func shiftSaturatingUnsignedSlice[D unsigned, S integer](x []D, y []S) []D {
   630  	r := make([]D, len(x))
   631  	for i := range r {
   632  		r[i] = shiftSaturatingUnsigned(x[i], shiftAmount(y[i]))
   633  	}
   634  	return r
   635  }
   636  
   637  // Slice versions for const shift operations (same constant amount for all elements)
   638  
   639  // shiftLeftByConstSlice shifts all elements left by constant amount.
   640  func shiftLeftByConstSlice[T integer](x []T, amt uint64) []T {
   641  	return map1(func(a T) T { return a << amt })(x)
   642  }
   643  
   644  // shiftRightByConstSlice shifts all elements right by constant amount.
   645  // Signed types use arithmetic shift, unsigned types use logical shift.
   646  func shiftRightByConstSlice[T integer](x []T, amt uint64) []T {
   647  	return map1(func(a T) T { return a >> amt })(x)
   648  }
   649  
   650  // shiftLeftSaturatingByConstSlice shifts all elements left by constant amount with signed saturation.
   651  func shiftLeftSaturatingByConstSlice[T signed](x []T, amt uint64) []T {
   652  	return map1(func(a T) T { return shiftSaturatingSigned(a, int8(amt)) })(x)
   653  }
   654  
   655  // shiftLeftSaturatingUByConstSlice shifts all elements left by constant amount with unsigned saturation.
   656  func shiftLeftSaturatingUByConstSlice[T unsigned](x []T, amt uint64) []T {
   657  	return map1(func(a T) T { return shiftSaturatingUnsigned(a, int8(amt)) })(x)
   658  }
   659  
   660  // shiftAllLeftSlice shifts all elements left by the same amount.
   661  func shiftAllLeftSlice[T integer](x []T, amt uint64) []T {
   662  	return map1(func(a T) T { return a << amt })(x)
   663  }
   664  
   665  // shiftAllRightSlice shifts all elements right by the same amount.
   666  // Signed types use arithmetic shift, unsigned types use logical shift.
   667  func shiftAllRightSlice[T integer](x []T, amt uint64) []T {
   668  	return map1(func(a T) T { return a >> amt })(x)
   669  }
   670  
   671  // ARM64-specific float-to-int conversion saturation helpers.
   672  // ARM64 uses IEEE 754 saturation: out-of-range values clamp to min/max of the target type.
   673  // NaN converts to 0. Negative values convert to 0 for unsigned types.
   674  
   675  func floatToInt32_arm64[T float](x T) int32 {
   676  	if x != x { // NaN
   677  		return 0
   678  	}
   679  	if x >= math.MaxInt32 {
   680  		return math.MaxInt32
   681  	}
   682  	if x < math.MinInt32 {
   683  		return math.MinInt32
   684  	}
   685  	return int32(x)
   686  }
   687  
   688  func floatToInt64_arm64[T float](x T) int64 {
   689  	if x != x { // NaN
   690  		return 0
   691  	}
   692  	if x >= math.MaxInt64 {
   693  		return math.MaxInt64
   694  	}
   695  	if x < math.MinInt64 {
   696  		return math.MinInt64
   697  	}
   698  	return int64(x)
   699  }
   700  
   701  func floatToUint32_arm64[T float](x T) uint32 {
   702  	if x != x || x < 0 { // NaN or negative
   703  		return 0
   704  	}
   705  	if x >= math.MaxUint32 {
   706  		return math.MaxUint32
   707  	}
   708  	return uint32(x)
   709  }
   710  
   711  func floatToUint64_arm64[T float](x T) uint64 {
   712  	if x != x || x < 0 { // NaN or negative
   713  		return 0
   714  	}
   715  	if x >= math.MaxUint64 {
   716  		return math.MaxUint64
   717  	}
   718  	return uint64(x)
   719  }
   720  

View as plain text