From f8a975d0473efd92e3e89e7919fdedf912706b9a Mon Sep 17 00:00:00 2001 From: zhenghaoz Date: Tue, 11 Aug 2026 23:41:53 +0800 Subject: [PATCH] [verified] perf: accelerate distance kernels with SIMD --- go.mod | 10 +- go.sum | 20 +- internal/ailego/vector_math.go | 69 ++- internal/ailego/vector_math_prevalidated.go | 43 ++ internal/ailego/vector_math_scalar.go | 110 ---- internal/ailego/vector_math_test.go | 47 ++ internal/floats/Makefile | 18 + internal/floats/floats.go | 70 +++ internal/floats/floats_amd64.go | 97 ++++ internal/floats/floats_amd64_test.go | 37 ++ internal/floats/floats_arm64.go | 62 +++ internal/floats/floats_arm64_test.go | 30 ++ internal/floats/floats_avx.go | 20 + internal/floats/floats_avx.s | 517 +++++++++++++++++++ internal/floats/floats_avx512.go | 20 + internal/floats/floats_avx512.s | 528 ++++++++++++++++++++ internal/floats/floats_neon.go | 20 + internal/floats/floats_neon.s | 234 +++++++++ internal/floats/floats_test.go | 156 ++++++ internal/floats/src/.gitignore | 2 + internal/floats/src/floats_avx.c | 90 ++++ internal/floats/src/floats_avx512.c | 90 ++++ internal/floats/src/floats_neon.c | 86 ++++ 23 files changed, 2226 insertions(+), 150 deletions(-) create mode 100644 internal/ailego/vector_math_prevalidated.go delete mode 100644 internal/ailego/vector_math_scalar.go create mode 100644 internal/floats/Makefile create mode 100644 internal/floats/floats.go create mode 100644 internal/floats/floats_amd64.go create mode 100644 internal/floats/floats_amd64_test.go create mode 100644 internal/floats/floats_arm64.go create mode 100644 internal/floats/floats_arm64_test.go create mode 100644 internal/floats/floats_avx.go create mode 100644 internal/floats/floats_avx.s create mode 100644 internal/floats/floats_avx512.go create mode 100644 internal/floats/floats_avx512.s create mode 100644 internal/floats/floats_neon.go create mode 100644 internal/floats/floats_neon.s create mode 100644 internal/floats/floats_test.go create mode 100644 internal/floats/src/.gitignore create mode 100644 internal/floats/src/floats_avx.c create mode 100644 internal/floats/src/floats_avx512.c create mode 100644 internal/floats/src/floats_neon.c diff --git a/go.mod b/go.mod index 3b6d2eb..63233e2 100644 --- a/go.mod +++ b/go.mod @@ -33,6 +33,9 @@ require ( github.com/golang/protobuf v1.5.3 // indirect github.com/golang/snappy v0.0.5-0.20231225225746-43d5d4cd4e0e // indirect github.com/google/uuid v1.6.0 // indirect + github.com/gorse-io/goat v0.2.1-0.20260618151728-201cbcf325ad // indirect + github.com/inconshreveable/mousetrap v1.1.0 // indirect + github.com/klauspost/asmfmt v1.3.2 // indirect github.com/klauspost/compress v1.17.11 // indirect github.com/kr/pretty v0.3.1 // indirect github.com/kr/text v0.2.0 // indirect @@ -49,9 +52,14 @@ require ( github.com/prometheus/common v0.42.0 // indirect github.com/prometheus/procfs v0.10.1 // indirect github.com/rogpeppe/go-internal v1.16.0 // indirect + github.com/samber/lo v1.53.0 // indirect + github.com/spf13/cobra v1.10.2 // indirect + github.com/spf13/pflag v1.0.10 // indirect github.com/twpayne/go-geom v1.6.1 // indirect golang.org/x/exp v0.0.0-20230626212559-97b1e661b5df // indirect - golang.org/x/text v0.22.0 // indirect + golang.org/x/text v0.36.0 // indirect google.golang.org/protobuf v1.34.2 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect ) + +tool github.com/gorse-io/goat diff --git a/go.sum b/go.sum index 55404e5..a5e063e 100644 --- a/go.sum +++ b/go.sum @@ -42,6 +42,7 @@ github.com/cockroachdb/swiss v0.0.0-20251224182025-b0f6560f979b h1:VXvSNzmr8hMj8 github.com/cockroachdb/swiss v0.0.0-20251224182025-b0f6560f979b/go.mod h1:yBRu/cnL4ks9bgy4vAASdjIW+/xMlFwuHKqtmh3GZQg= github.com/cockroachdb/tokenbucket v0.0.0-20230807174530-cc333fc44b06 h1:zuQyyAKVxetITBuuhv3BI9cMrmStnpT18zmgmTxunpo= github.com/cockroachdb/tokenbucket v0.0.0-20230807174530-cc333fc44b06/go.mod h1:7nc4anLGjupUW/PeY5qiNYsdNXj7zopG+eqsS7To5IQ= +github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= @@ -67,10 +68,16 @@ github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/gorse-io/goat v0.2.1-0.20260618151728-201cbcf325ad h1:HESEie1ivHg8yI9Qnej4Cu+AP8w8gQ9gqpYcnV5zC7Q= +github.com/gorse-io/goat v0.2.1-0.20260618151728-201cbcf325ad/go.mod h1:NjwKvGIyhFveMjnUdYrngJYvYLqK0aK7Qdp7QJWkBLo= github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM= github.com/hexops/gotextdiff v1.0.3/go.mod h1:pSWU5MAI3yDq+fZBTazCSJysOMbxWL1BSow5/V2vxeg= +github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= +github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= +github.com/klauspost/asmfmt v1.3.2 h1:4Ri7ox3EwapiOjCki+hw14RyKk201CN4rzyCJRFLpK4= +github.com/klauspost/asmfmt v1.3.2/go.mod h1:AG8TuvYojzulgDAMCnYn50l/5QV3Bs/tp6j0HLHbNSE= github.com/klauspost/compress v1.17.11 h1:In6xLpyWOi1+C7tXUUWv2ot1QvBjxevKAaI6IXrJmUc= github.com/klauspost/compress v1.17.11/go.mod h1:pMDklpSncoRMuLFrf1W9Ss9KT+0rH90U12bZKk7uwG0= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= @@ -109,6 +116,14 @@ github.com/prometheus/procfs v0.10.1/go.mod h1:nwNm2aOCAYw8uTR/9bWRREkZFxAUcWzPH github.com/rogpeppe/go-internal v1.9.0/go.mod h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs= github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g= github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs= +github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= +github.com/samber/lo v1.53.0 h1:t975lj2py4kJPQ6haz1QMgtId2gtmfktACxIXArw3HM= +github.com/samber/lo v1.53.0/go.mod h1:4+MXEGsJzbKGaUEQFKBq2xtfuznW9oz/WrgyzMzRoM0= +github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU= +github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4= +github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= +github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk= +github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/twpayne/go-geom v1.6.1 h1:iLE+Opv0Ihm/ABIcvQFGIiFBXd76oBIar9drAwHFhR4= @@ -117,6 +132,7 @@ github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZ github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= +go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= @@ -141,8 +157,8 @@ golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= -golang.org/x/text v0.22.0 h1:bofq7m3/HAFvbF51jz3Q9wLg3jkvSPuiZu/pD1XwgtM= -golang.org/x/text v0.22.0/go.mod h1:YRoo4H8PVmsu+E3Ou7cqLVH8oXWIHVoX0jqUWALQhfY= +golang.org/x/text v0.36.0 h1:JfKh3XmcRPqZPKevfXVpI1wXPTqbkE5f7JA92a55Yxg= +golang.org/x/text v0.36.0/go.mod h1:NIdBknypM8iqVmPiuco0Dh6P5Jcdk8lJL0CUebqK164= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE= diff --git a/internal/ailego/vector_math.go b/internal/ailego/vector_math.go index 390f7a0..c4e8162 100644 --- a/internal/ailego/vector_math.go +++ b/internal/ailego/vector_math.go @@ -17,6 +17,8 @@ package ailego import ( "errors" "math" + + "github.com/gorse-io/xvec/internal/floats" ) var ( @@ -26,22 +28,15 @@ var ( ErrInvalidSparseOrder = errors.New("ailego: sparse indices must be strictly increasing") ) -// L2Squared computes squared Euclidean distance using float64 accumulation and -// returns the float32 score used by indexes. +// L2Squared computes squared Euclidean distance using float32 accumulation. func L2Squared(left, right []float32) (float32, error) { if err := validateDensePair(left, right); err != nil { return 0, err } - var sum float64 - for index, leftValue := range left { - rightValue := right[index] - if !finite32(leftValue) || !finite32(rightValue) { - return 0, ErrNonFiniteVector - } - difference := float64(leftValue) - float64(rightValue) - sum += difference * difference + if err := validateDenseFinite(left, right); err != nil { + return 0, err } - return finiteScore(sum) + return finiteScore(float64(floats.L2Squared(left, right))) } // InnerProduct computes the dot-product similarity. Higher scores are better. @@ -49,15 +44,10 @@ func InnerProduct(left, right []float32) (float32, error) { if err := validateDensePair(left, right); err != nil { return 0, err } - var sum float64 - for index, leftValue := range left { - rightValue := right[index] - if !finite32(leftValue) || !finite32(rightValue) { - return 0, ErrNonFiniteVector - } - sum += float64(leftValue) * float64(rightValue) + if err := validateDenseFinite(left, right); err != nil { + return 0, err } - return finiteScore(sum) + return finiteScore(float64(floats.InnerProduct(left, right))) } // CosineDistance computes 1-cos(left,right). Lower scores are better. Two zero @@ -66,19 +56,25 @@ func CosineDistance(left, right []float32) (float32, error) { if err := validateDensePair(left, right); err != nil { return 0, err } - inner, leftNorm, rightNorm, err := denseProducts(left, right) - if err != nil { + if err := validateDenseFinite(left, right); err != nil { return 0, err } + return finiteScore(float64(cosineDistance(left, right))) +} + +func cosineDistance(left, right []float32) float32 { + inner, leftNorm, rightNorm := floats.DotNorms(left, right) if leftNorm == 0 && rightNorm == 0 { - return 0, nil + return 0 } if leftNorm == 0 || rightNorm == 0 { - return 1, nil + return 1 } - cosine := inner / math.Sqrt(leftNorm*rightNorm) + leftMagnitude := float32(math.Sqrt(float64(leftNorm))) + rightMagnitude := float32(math.Sqrt(float64(rightNorm))) + cosine := inner / (leftMagnitude * rightMagnitude) cosine = min(1, max(-1, cosine)) - return finiteScore(1 - cosine) + return 1 - cosine } // MIPSL2Squared computes the baseline localized-spherical MIPS-to-L2 @@ -88,15 +84,19 @@ func MIPSL2Squared(left, right []float32) (float32, error) { if err := validateDensePair(left, right); err != nil { return 0, err } - inner, leftNorm, rightNorm, err := denseProducts(left, right) - if err != nil { + if err := validateDenseFinite(left, right); err != nil { return 0, err } + return finiteScore(float64(mipsL2Squared(left, right))) +} + +func mipsL2Squared(left, right []float32) float32 { + inner, leftNorm, rightNorm := floats.DotNorms(left, right) denominator := max(leftNorm, rightNorm) if denominator == 0 { - return 0, nil + return 0 } - return finiteScore(2 - 2*inner/denominator) + return 2 - 2*inner/denominator } // SparseInnerProduct computes the dot product of canonical sparse vectors. @@ -139,19 +139,14 @@ func validateDensePair(left, right []float32) error { return nil } -func denseProducts(left, right []float32) (inner, leftNorm, rightNorm float64, err error) { +func validateDenseFinite(left, right []float32) error { for index, leftValue := range left { rightValue := right[index] if !finite32(leftValue) || !finite32(rightValue) { - return 0, 0, 0, ErrNonFiniteVector + return ErrNonFiniteVector } - leftFloat := float64(leftValue) - rightFloat := float64(rightValue) - inner += leftFloat * rightFloat - leftNorm += leftFloat * leftFloat - rightNorm += rightFloat * rightFloat } - return inner, leftNorm, rightNorm, nil + return nil } func validateSparse(indices []uint32, values []float32) error { diff --git a/internal/ailego/vector_math_prevalidated.go b/internal/ailego/vector_math_prevalidated.go new file mode 100644 index 0000000..b1be96d --- /dev/null +++ b/internal/ailego/vector_math_prevalidated.go @@ -0,0 +1,43 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package ailego + +import "github.com/gorse-io/xvec/internal/floats" + +// DenseDistance computes a score for two already validated dense vectors. +// Callers must guarantee equal, non-zero dimensions and finite components. +type DenseDistance func(left, right []float32) (float32, error) + +// L2SquaredPrevalidated computes squared Euclidean distance without validating +// its inputs. It is intended for index hot paths whose storage boundary has +// already validated every vector. +func L2SquaredPrevalidated(left, right []float32) (float32, error) { + return finiteScore(float64(floats.L2Squared(left, right))) +} + +// InnerProductPrevalidated computes inner product without validating inputs. +func InnerProductPrevalidated(left, right []float32) (float32, error) { + return finiteScore(float64(floats.InnerProduct(left, right))) +} + +// CosineDistancePrevalidated computes cosine distance without validating inputs. +func CosineDistancePrevalidated(left, right []float32) (float32, error) { + return finiteScore(float64(cosineDistance(left, right))) +} + +// MIPSL2SquaredPrevalidated computes MIPS-to-L2 distance without validating inputs. +func MIPSL2SquaredPrevalidated(left, right []float32) (float32, error) { + return finiteScore(float64(mipsL2Squared(left, right))) +} diff --git a/internal/ailego/vector_math_scalar.go b/internal/ailego/vector_math_scalar.go deleted file mode 100644 index c23d30d..0000000 --- a/internal/ailego/vector_math_scalar.go +++ /dev/null @@ -1,110 +0,0 @@ -// Copyright 2026-present the xvec project -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -package ailego - -import "math" - -// DenseDistance computes a score for two already validated dense vectors. -// Callers must guarantee equal, non-zero dimensions and finite components. -type DenseDistance func(left, right []float32) (float32, error) - -type denseDistanceKernel func(left, right []float32) float32 - -// Keep the dispatch table separate from the checked API so architecture files -// can select SIMD kernels later without changing index code. The baseline uses -// portable scalar implementations on every platform. -var denseDistanceKernels = struct { - l2 denseDistanceKernel - ip denseDistanceKernel - cosine denseDistanceKernel - mipsL2 denseDistanceKernel -}{ - l2: l2SquaredScalar, - ip: innerProductScalar, - cosine: cosineDistanceScalar, - mipsL2: mipsL2SquaredScalar, -} - -// L2SquaredPrevalidated computes squared Euclidean distance without validating -// its inputs. It is intended for index hot paths whose storage boundary has -// already validated every vector. -func L2SquaredPrevalidated(left, right []float32) (float32, error) { - return finiteScore(float64(denseDistanceKernels.l2(left, right))) -} - -// InnerProductPrevalidated computes inner product without validating inputs. -func InnerProductPrevalidated(left, right []float32) (float32, error) { - return finiteScore(float64(denseDistanceKernels.ip(left, right))) -} - -// CosineDistancePrevalidated computes cosine distance without validating inputs. -func CosineDistancePrevalidated(left, right []float32) (float32, error) { - return finiteScore(float64(denseDistanceKernels.cosine(left, right))) -} - -// MIPSL2SquaredPrevalidated computes MIPS-to-L2 distance without validating inputs. -func MIPSL2SquaredPrevalidated(left, right []float32) (float32, error) { - return finiteScore(float64(denseDistanceKernels.mipsL2(left, right))) -} - -func l2SquaredScalar(left, right []float32) float32 { - var sum float64 - for index, leftValue := range left { - difference := float64(leftValue) - float64(right[index]) - sum += difference * difference - } - return float32(sum) -} - -func innerProductScalar(left, right []float32) float32 { - var sum float64 - for index, leftValue := range left { - sum += float64(leftValue) * float64(right[index]) - } - return float32(sum) -} - -func cosineDistanceScalar(left, right []float32) float32 { - inner, leftNorm, rightNorm := denseProductsScalar(left, right) - if leftNorm == 0 && rightNorm == 0 { - return 0 - } - if leftNorm == 0 || rightNorm == 0 { - return 1 - } - cosine := inner / math.Sqrt(leftNorm*rightNorm) - cosine = min(1, max(-1, cosine)) - return float32(1 - cosine) -} - -func mipsL2SquaredScalar(left, right []float32) float32 { - inner, leftNorm, rightNorm := denseProductsScalar(left, right) - denominator := max(leftNorm, rightNorm) - if denominator == 0 { - return 0 - } - return float32(2 - 2*inner/denominator) -} - -func denseProductsScalar(left, right []float32) (inner, leftNorm, rightNorm float64) { - for index, leftValue := range left { - leftFloat := float64(leftValue) - rightFloat := float64(right[index]) - inner += leftFloat * rightFloat - leftNorm += leftFloat * leftFloat - rightNorm += rightFloat * rightFloat - } - return inner, leftNorm, rightNorm -} diff --git a/internal/ailego/vector_math_test.go b/internal/ailego/vector_math_test.go index fb2d738..26541ca 100644 --- a/internal/ailego/vector_math_test.go +++ b/internal/ailego/vector_math_test.go @@ -117,6 +117,53 @@ func TestDenseMetricValidation(t *testing.T) { } } +func TestCosineDistanceAvoidsNormProductOverflowAndUnderflow(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + left []float32 + right []float32 + expected float32 + }{ + { + name: "large identical vectors", + left: []float32{1e10, 1e10}, + right: []float32{1e10, 1e10}, + expected: 0, + }, + { + name: "tiny orthogonal vectors", + left: []float32{1e-20, 0}, + right: []float32{0, 1e-20}, + expected: 1, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + checked, err := CosineDistance(test.left, test.right) + require.NoError(t, err) + require.InDelta(t, test.expected, checked, 1e-6) + prevalidated, err := CosineDistancePrevalidated(test.left, test.right) + require.NoError(t, err) + require.InDelta(t, test.expected, prevalidated, 1e-6) + }) + } +} + +func TestDenseMetricsUseFloat32Accumulation(t *testing.T) { + t.Parallel() + + left := []float32{1e8, 1, -1e8} + right := []float32{1, 1, 1} + checked, err := InnerProduct(left, right) + require.NoError(t, err) + require.Zero(t, checked) + prevalidated, err := InnerProductPrevalidated(left, right) + require.NoError(t, err) + require.Zero(t, prevalidated) +} + func TestPrevalidatedDenseDistanceKernelsMatchCheckedMetrics(t *testing.T) { left := []float32{0.2, 0.9, -0.4, 0.7} right := []float32{0.3, 0.5, 0.8, -0.1} diff --git a/internal/floats/Makefile b/internal/floats/Makefile new file mode 100644 index 0000000..5f7cdd1 --- /dev/null +++ b/internal/floats/Makefile @@ -0,0 +1,18 @@ +GO ?= go +GOAT ?= $(GO) tool goat + +.PHONY: generate avx avx512 neon clean + +generate: avx avx512 neon clean + +avx: + $(GOAT) src/floats_avx.c --target amd64 -O3 -mavx + +avx512: + $(GOAT) src/floats_avx512.c --target amd64 -O3 -mavx -mfma -mavx512f + +neon: + $(GOAT) src/floats_neon.c --target arm64 -O3 + +clean: + $(RM) src/*.o src/*.s diff --git a/internal/floats/floats.go b/internal/floats/floats.go new file mode 100644 index 0000000..61791e9 --- /dev/null +++ b/internal/floats/floats.go @@ -0,0 +1,70 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Package floats provides allocation-free float32 vector kernels. Callers must +// pass equal, non-empty slices; validation belongs at the API boundary. +package floats + +type binaryKernel func(left, right []float32) float32 +type productsKernel func(left, right []float32) (dot, leftNorm, rightNorm float32) + +var kernels = struct { + l2 binaryKernel + dot binaryKernel + products productsKernel +}{ + l2: l2SquaredScalar, + dot: innerProductScalar, + products: dotNormsScalar, +} + +// L2Squared returns the squared Euclidean distance between left and right. +func L2Squared(left, right []float32) float32 { + return kernels.l2(left, right) +} + +// InnerProduct returns the dot product of left and right. +func InnerProduct(left, right []float32) float32 { + return kernels.dot(left, right) +} + +// DotNorms computes the dot product and both squared norms in one pass. +func DotNorms(left, right []float32) (dot, leftNorm, rightNorm float32) { + return kernels.products(left, right) +} + +func l2SquaredScalar(left, right []float32) (sum float32) { + for index, leftValue := range left { + difference := leftValue - right[index] + sum += difference * difference + } + return +} + +func innerProductScalar(left, right []float32) (sum float32) { + for index, leftValue := range left { + sum += leftValue * right[index] + } + return +} + +func dotNormsScalar(left, right []float32) (dot, leftNorm, rightNorm float32) { + for index, leftValue := range left { + rightValue := right[index] + dot += leftValue * rightValue + leftNorm += leftValue * leftValue + rightNorm += rightValue * rightValue + } + return +} diff --git a/internal/floats/floats_amd64.go b/internal/floats/floats_amd64.go new file mode 100644 index 0000000..e42ad7c --- /dev/null +++ b/internal/floats/floats_amd64.go @@ -0,0 +1,97 @@ +//go:build !noasm && amd64 + +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package floats + +import ( + "unsafe" + + "golang.org/x/sys/cpu" +) + +//go:generate go tool goat src/floats_avx.c -O3 -mavx +//go:generate go tool goat src/floats_avx512.c -O3 -mavx -mfma -mavx512f + +func init() { + switch { + case cpu.X86.HasAVX && cpu.X86.HasFMA && cpu.X86.HasAVX512F: + kernels.l2 = l2SquaredAVX512 + kernels.dot = innerProductAVX512 + kernels.products = dotNormsAVX512 + case cpu.X86.HasAVX: + kernels.l2 = l2SquaredAVX + kernels.dot = innerProductAVX + kernels.products = dotNormsAVX + } +} + +func l2SquaredAVX(left, right []float32) float32 { + if len(left) < 8 { + return l2SquaredScalar(left, right) + } + var result float32 + xvec_avx_l2_squared(unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), unsafe.Pointer(&result)) + return result +} + +func innerProductAVX(left, right []float32) float32 { + if len(left) < 8 { + return innerProductScalar(left, right) + } + var result float32 + xvec_avx_inner_product(unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), unsafe.Pointer(&result)) + return result +} + +func dotNormsAVX(left, right []float32) (dot, leftNorm, rightNorm float32) { + if len(left) < 8 { + return dotNormsScalar(left, right) + } + xvec_avx_dot_norms( + unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), + unsafe.Pointer(&dot), unsafe.Pointer(&leftNorm), unsafe.Pointer(&rightNorm), + ) + return +} + +func l2SquaredAVX512(left, right []float32) float32 { + if len(left) < 16 { + return l2SquaredAVX(left, right) + } + var result float32 + xvec_avx512_l2_squared(unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), unsafe.Pointer(&result)) + return result +} + +func innerProductAVX512(left, right []float32) float32 { + if len(left) < 16 { + return innerProductAVX(left, right) + } + var result float32 + xvec_avx512_inner_product(unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), unsafe.Pointer(&result)) + return result +} + +func dotNormsAVX512(left, right []float32) (dot, leftNorm, rightNorm float32) { + if len(left) < 16 { + return dotNormsAVX(left, right) + } + xvec_avx512_dot_norms( + unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), + unsafe.Pointer(&dot), unsafe.Pointer(&leftNorm), unsafe.Pointer(&rightNorm), + ) + return +} diff --git a/internal/floats/floats_amd64_test.go b/internal/floats/floats_amd64_test.go new file mode 100644 index 0000000..c0f3f43 --- /dev/null +++ b/internal/floats/floats_amd64_test.go @@ -0,0 +1,37 @@ +//go:build !noasm && amd64 + +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package floats + +import ( + "testing" + + "golang.org/x/sys/cpu" +) + +func TestAVXDistanceKernels(t *testing.T) { + if !cpu.X86.HasAVX { + t.Skip("AVX is not supported by this CPU") + } + testArchitectureKernels(t, l2SquaredAVX, innerProductAVX, dotNormsAVX) +} + +func TestAVX512DistanceKernels(t *testing.T) { + if !cpu.X86.HasAVX || !cpu.X86.HasFMA || !cpu.X86.HasAVX512F { + t.Skip("AVX-512/FMA is not supported by this CPU") + } + testArchitectureKernels(t, l2SquaredAVX512, innerProductAVX512, dotNormsAVX512) +} diff --git a/internal/floats/floats_arm64.go b/internal/floats/floats_arm64.go new file mode 100644 index 0000000..15fd6f2 --- /dev/null +++ b/internal/floats/floats_arm64.go @@ -0,0 +1,62 @@ +//go:build !noasm && arm64 + +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package floats + +import ( + "unsafe" + + "golang.org/x/sys/cpu" +) + +//go:generate go tool goat src/floats_neon.c -O3 + +func init() { + if cpu.ARM64.HasASIMD { + kernels.l2 = l2SquaredNEON + kernels.dot = innerProductNEON + kernels.products = dotNormsNEON + } +} + +func l2SquaredNEON(left, right []float32) float32 { + if len(left) < 4 { + return l2SquaredScalar(left, right) + } + var result float32 + xvec_neon_l2_squared(unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), unsafe.Pointer(&result)) + return result +} + +func innerProductNEON(left, right []float32) float32 { + if len(left) < 4 { + return innerProductScalar(left, right) + } + var result float32 + xvec_neon_inner_product(unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), unsafe.Pointer(&result)) + return result +} + +func dotNormsNEON(left, right []float32) (dot, leftNorm, rightNorm float32) { + if len(left) < 4 { + return dotNormsScalar(left, right) + } + xvec_neon_dot_norms( + unsafe.Pointer(&left[0]), unsafe.Pointer(&right[0]), int64(len(left)), + unsafe.Pointer(&dot), unsafe.Pointer(&leftNorm), unsafe.Pointer(&rightNorm), + ) + return +} diff --git a/internal/floats/floats_arm64_test.go b/internal/floats/floats_arm64_test.go new file mode 100644 index 0000000..c886ba1 --- /dev/null +++ b/internal/floats/floats_arm64_test.go @@ -0,0 +1,30 @@ +//go:build !noasm && arm64 + +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package floats + +import ( + "testing" + + "golang.org/x/sys/cpu" +) + +func TestNEONDistanceKernels(t *testing.T) { + if !cpu.ARM64.HasASIMD { + t.Skip("NEON/ASIMD is not supported by this CPU") + } + testArchitectureKernels(t, l2SquaredNEON, innerProductNEON, dotNormsNEON) +} diff --git a/internal/floats/floats_avx.go b/internal/floats/floats_avx.go new file mode 100644 index 0000000..42b33d0 --- /dev/null +++ b/internal/floats/floats_avx.go @@ -0,0 +1,20 @@ +//go:build !noasm && amd64 +// Code generated by GoAT. DO NOT EDIT. +// versions: +// clang 21.1.8 (6ubuntu1) +// objdump 2.46 +// flags: -mavx -O3 +// source: src/floats_avx.c + +package floats + +import "unsafe" + +//go:noescape +func xvec_avx_l2_squared(left, right unsafe.Pointer, size int64, output unsafe.Pointer) + +//go:noescape +func xvec_avx_inner_product(left, right unsafe.Pointer, size int64, output unsafe.Pointer) + +//go:noescape +func xvec_avx_dot_norms(left, right unsafe.Pointer, size int64, dot, left_norm, right_norm unsafe.Pointer) diff --git a/internal/floats/floats_avx.s b/internal/floats/floats_avx.s new file mode 100644 index 0000000..229e4f9 --- /dev/null +++ b/internal/floats/floats_avx.s @@ -0,0 +1,517 @@ +//go:build !noasm && amd64 +// Code generated by GoAT. DO NOT EDIT. +// versions: +// clang 21.1.8 (6ubuntu1) +// objdump 2.46 +// flags: -mavx -O3 +// source: src/floats_avx.c + +TEXT ·xvec_avx_l2_squared(SB), $0-32 + MOVQ left+0(FP), DI + MOVQ right+8(FP), SI + MOVQ size+16(FP), DX + MOVQ output+24(FP), CX + BYTE $0x55 // pushq %rbp + WORD $0x8948; BYTE $0xe5 // movq %rsp, %rbp + LONG $0xf8e48348 // andq $-8, %rsp + LONG $0x07428d4c // leaq 7(%rdx), %r8 + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + LONG $0xc2490f4c // cmovnsq %rdx, %r8 + WORD $0x894d; BYTE $0xc1 // movq %r8, %r9 + LONG $0xf8e18349 // andq $-8, %r9 + WORD $0x8948; BYTE $0xd0 // movq %rdx, %rax + WORD $0x294c; BYTE $0xc8 // subq %r9, %rax + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + LONG $0x08fa8348 // cmpq $8, %rdx + JL LBB0_7 + LONG $0x03f8c149 // sarq $3, %r8 + LONG $0xff488d4d // leaq -1(%r8), %r9 + WORD $0x8944; BYTE $0xc2 // movl %r8d, %edx + WORD $0xe283; BYTE $0x03 // andl $3, %edx + LONG $0x03f98349 // cmpq $3, %r9 + JAE LBB0_16 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + JMP LBB0_3 + +LBB0_16: + QUAD $0xfffffffffffcbb49; WORD $0x0fff // movabsq $1152921504606846972, %r11 # imm = 0xFFFFFFFFFFFFFFC + WORD $0x214d; BYTE $0xc3 // andq %r8, %r11 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + +LBB0_17: + LONG $0x107cc1c4; BYTE $0x09 // vmovups (%r9), %ymm1 + LONG $0x107cc1c4; WORD $0x2051 // vmovups 32(%r9), %ymm2 + LONG $0x107cc1c4; WORD $0x4059 // vmovups 64(%r9), %ymm3 + LONG $0x107cc1c4; WORD $0x6061 // vmovups 96(%r9), %ymm4 + LONG $0x5c74c1c4; BYTE $0x0a // vsubps (%r10), %ymm1, %ymm1 + LONG $0xc959f4c5 // vmulps %ymm1, %ymm1, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x5c6cc1c4; WORD $0x204a // vsubps 32(%r10), %ymm2, %ymm1 + LONG $0xc959f4c5 // vmulps %ymm1, %ymm1, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x5c64c1c4; WORD $0x404a // vsubps 64(%r10), %ymm3, %ymm1 + LONG $0xc959f4c5 // vmulps %ymm1, %ymm1, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x5c5cc1c4; WORD $0x604a // vsubps 96(%r10), %ymm4, %ymm1 + LONG $0xc959f4c5 // vmulps %ymm1, %ymm1, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x80e98349 // subq $-128, %r9 + LONG $0x80ea8349 // subq $-128, %r10 + LONG $0xfcc38349 // addq $-4, %r11 + JNE LBB0_17 + +LBB0_3: + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + JE LBB0_6 + WORD $0xe2c1; BYTE $0x05 // shll $5, %edx + WORD $0x3145; BYTE $0xdb // xorl %r11d, %r11d + +LBB0_5: + LONG $0x107c81c4; WORD $0x190c // vmovups (%r9,%r11), %ymm1 + LONG $0x5c7481c4; WORD $0x1a0c // vsubps (%r10,%r11), %ymm1, %ymm1 + LONG $0xc959f4c5 // vmulps %ymm1, %ymm1, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x20c38349 // addq $32, %r11 + WORD $0x394c; BYTE $0xda // cmpq %r11, %rdx + JNE LBB0_5 + +LBB0_6: + LONG $0x05e0c149 // shlq $5, %r8 + WORD $0x014c; BYTE $0xc7 // addq %r8, %rdi + WORD $0x014c; BYTE $0xc6 // addq %r8, %rsi + +LBB0_7: + LONG $0xc816fac5 // vmovshdup %xmm0, %xmm1 # xmm1 = xmm0[1,1,3,3] + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x197de3c4; WORD $0x01c0 // vextractf128 $1, %ymm0, %xmm0 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xc0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm0 # xmm0 = xmm0[3,3,3,3] + LONG $0xc158fac5 // vaddss %xmm1, %xmm0, %xmm0 + WORD $0x8548; BYTE $0xc0 // testq %rax, %rax + JLE LBB0_15 + LONG $0x0f10fac5 // vmovss (%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x0e5cf2c5 // vsubss (%rsi), %xmm1, %xmm1 + LONG $0xc959f2c5 // vmulss %xmm1, %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x01f88348 // cmpq $1, %rax + JE LBB0_15 + LONG $0x4f10fac5; BYTE $0x04 // vmovss 4(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e5cf2c5; BYTE $0x04 // vsubss 4(%rsi), %xmm1, %xmm1 + LONG $0xc959f2c5 // vmulss %xmm1, %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x02f88348 // cmpq $2, %rax + JE LBB0_15 + LONG $0x4f10fac5; BYTE $0x08 // vmovss 8(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e5cf2c5; BYTE $0x08 // vsubss 8(%rsi), %xmm1, %xmm1 + LONG $0xc959f2c5 // vmulss %xmm1, %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x03f88348 // cmpq $3, %rax + JE LBB0_15 + LONG $0x4f10fac5; BYTE $0x0c // vmovss 12(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e5cf2c5; BYTE $0x0c // vsubss 12(%rsi), %xmm1, %xmm1 + LONG $0xc959f2c5 // vmulss %xmm1, %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x04f88348 // cmpq $4, %rax + JE LBB0_15 + LONG $0x4f10fac5; BYTE $0x10 // vmovss 16(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e5cf2c5; BYTE $0x10 // vsubss 16(%rsi), %xmm1, %xmm1 + LONG $0xc959f2c5 // vmulss %xmm1, %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x05f88348 // cmpq $5, %rax + JE LBB0_15 + LONG $0x4f10fac5; BYTE $0x14 // vmovss 20(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e5cf2c5; BYTE $0x14 // vsubss 20(%rsi), %xmm1, %xmm1 + LONG $0xc959f2c5 // vmulss %xmm1, %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x06f88348 // cmpq $6, %rax + JE LBB0_15 + LONG $0x4f10fac5; BYTE $0x18 // vmovss 24(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e5cf2c5; BYTE $0x18 // vsubss 24(%rsi), %xmm1, %xmm1 + LONG $0xc959f2c5 // vmulss %xmm1, %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + +LBB0_15: + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + WORD $0x8948; BYTE $0xec // movq %rbp, %rsp + BYTE $0x5d // popq %rbp + WORD $0xf8c5; BYTE $0x77 // vzeroupper + RET + +TEXT ·xvec_avx_inner_product(SB), $0-32 + MOVQ left+0(FP), DI + MOVQ right+8(FP), SI + MOVQ size+16(FP), DX + MOVQ output+24(FP), CX + BYTE $0x55 // pushq %rbp + WORD $0x8948; BYTE $0xe5 // movq %rsp, %rbp + LONG $0xf8e48348 // andq $-8, %rsp + LONG $0x07428d4c // leaq 7(%rdx), %r8 + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + LONG $0xc2490f4c // cmovnsq %rdx, %r8 + WORD $0x894d; BYTE $0xc1 // movq %r8, %r9 + LONG $0xf8e18349 // andq $-8, %r9 + WORD $0x8948; BYTE $0xd0 // movq %rdx, %rax + WORD $0x294c; BYTE $0xc8 // subq %r9, %rax + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + LONG $0x08fa8348 // cmpq $8, %rdx + JL LBB1_7 + LONG $0x03f8c149 // sarq $3, %r8 + LONG $0xff488d4d // leaq -1(%r8), %r9 + WORD $0x8944; BYTE $0xc2 // movl %r8d, %edx + WORD $0xe283; BYTE $0x03 // andl $3, %edx + LONG $0x03f98349 // cmpq $3, %r9 + JAE LBB1_16 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + JMP LBB1_3 + +LBB1_16: + QUAD $0xfffffffffffcbb49; WORD $0x0fff // movabsq $1152921504606846972, %r11 # imm = 0xFFFFFFFFFFFFFFC + WORD $0x214d; BYTE $0xc3 // andq %r8, %r11 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + +LBB1_17: + LONG $0x107cc1c4; BYTE $0x09 // vmovups (%r9), %ymm1 + LONG $0x107cc1c4; WORD $0x2051 // vmovups 32(%r9), %ymm2 + LONG $0x107cc1c4; WORD $0x4059 // vmovups 64(%r9), %ymm3 + LONG $0x107cc1c4; WORD $0x6061 // vmovups 96(%r9), %ymm4 + LONG $0x5974c1c4; BYTE $0x0a // vmulps (%r10), %ymm1, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x596cc1c4; WORD $0x204a // vmulps 32(%r10), %ymm2, %ymm1 + LONG $0x5964c1c4; WORD $0x4052 // vmulps 64(%r10), %ymm3, %ymm2 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0xc258fcc5 // vaddps %ymm2, %ymm0, %ymm0 + LONG $0x595cc1c4; WORD $0x604a // vmulps 96(%r10), %ymm4, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x80e98349 // subq $-128, %r9 + LONG $0x80ea8349 // subq $-128, %r10 + LONG $0xfcc38349 // addq $-4, %r11 + JNE LBB1_17 + +LBB1_3: + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + JE LBB1_6 + WORD $0xe2c1; BYTE $0x05 // shll $5, %edx + WORD $0x3145; BYTE $0xdb // xorl %r11d, %r11d + +LBB1_5: + LONG $0x107c81c4; WORD $0x190c // vmovups (%r9,%r11), %ymm1 + LONG $0x597481c4; WORD $0x1a0c // vmulps (%r10,%r11), %ymm1, %ymm1 + LONG $0xc158fcc5 // vaddps %ymm1, %ymm0, %ymm0 + LONG $0x20c38349 // addq $32, %r11 + WORD $0x394c; BYTE $0xda // cmpq %r11, %rdx + JNE LBB1_5 + +LBB1_6: + LONG $0x05e0c149 // shlq $5, %r8 + WORD $0x014c; BYTE $0xc7 // addq %r8, %rdi + WORD $0x014c; BYTE $0xc6 // addq %r8, %rsi + +LBB1_7: + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x197de3c4; WORD $0x01c0 // vextractf128 $1, %ymm0, %xmm0 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xc0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm0 # xmm0 = xmm0[3,3,3,3] + LONG $0xc158fac5 // vaddss %xmm1, %xmm0, %xmm0 + WORD $0x8548; BYTE $0xc0 // testq %rax, %rax + JLE LBB1_15 + LONG $0x0f10fac5 // vmovss (%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x0e59f2c5 // vmulss (%rsi), %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x01f88348 // cmpq $1, %rax + JE LBB1_15 + LONG $0x4f10fac5; BYTE $0x04 // vmovss 4(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e59f2c5; BYTE $0x04 // vmulss 4(%rsi), %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x02f88348 // cmpq $2, %rax + JE LBB1_15 + LONG $0x4f10fac5; BYTE $0x08 // vmovss 8(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e59f2c5; BYTE $0x08 // vmulss 8(%rsi), %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x03f88348 // cmpq $3, %rax + JE LBB1_15 + LONG $0x4f10fac5; BYTE $0x0c // vmovss 12(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e59f2c5; BYTE $0x0c // vmulss 12(%rsi), %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x04f88348 // cmpq $4, %rax + JE LBB1_15 + LONG $0x4f10fac5; BYTE $0x10 // vmovss 16(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e59f2c5; BYTE $0x10 // vmulss 16(%rsi), %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x05f88348 // cmpq $5, %rax + JE LBB1_15 + LONG $0x4f10fac5; BYTE $0x14 // vmovss 20(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e59f2c5; BYTE $0x14 // vmulss 20(%rsi), %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + LONG $0x06f88348 // cmpq $6, %rax + JE LBB1_15 + LONG $0x4f10fac5; BYTE $0x18 // vmovss 24(%rdi), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x4e59f2c5; BYTE $0x18 // vmulss 24(%rsi), %xmm1, %xmm1 + LONG $0xc058f2c5 // vaddss %xmm0, %xmm1, %xmm0 + +LBB1_15: + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + WORD $0x8948; BYTE $0xec // movq %rbp, %rsp + BYTE $0x5d // popq %rbp + WORD $0xf8c5; BYTE $0x77 // vzeroupper + RET + +TEXT ·xvec_avx_dot_norms(SB), $0-48 + MOVQ left+0(FP), DI + MOVQ right+8(FP), SI + MOVQ size+16(FP), DX + MOVQ dot+24(FP), CX + MOVQ left_norm+32(FP), R8 + MOVQ right_norm+40(FP), R9 + BYTE $0x55 // pushq %rbp + WORD $0x8948; BYTE $0xe5 // movq %rsp, %rbp + BYTE $0x53 // pushq %rbx + LONG $0xf8e48348 // andq $-8, %rsp + LONG $0x07528d4c // leaq 7(%rdx), %r10 + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + LONG $0xd2490f4c // cmovnsq %rdx, %r10 + WORD $0x894d; BYTE $0xd3 // movq %r10, %r11 + LONG $0xf8e38349 // andq $-8, %r11 + WORD $0x8948; BYTE $0xd0 // movq %rdx, %rax + WORD $0x294c; BYTE $0xd8 // subq %r11, %rax + LONG $0xd257e8c5 // vxorps %xmm2, %xmm2, %xmm2 + LONG $0x08fa8348 // cmpq $8, %rdx + JL LBB2_1 + LONG $0x03fac149 // sarq $3, %r10 + QUAD $0xfffffffffff8bb49; WORD $0x7fff // movabsq $9223372036854775800, %r11 # imm = 0x7FFFFFFFFFFFFFF8 + WORD $0x214c; BYTE $0xda // andq %r11, %rdx + LONG $0x08fa8348 // cmpq $8, %rdx + JNE LBB2_9 + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + WORD $0x8948; BYTE $0xfa // movq %rdi, %rdx + WORD $0x8949; BYTE $0xf3 // movq %rsi, %r11 + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + JMP LBB2_4 + +LBB2_1: + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + JMP LBB2_7 + +LBB2_9: + QUAD $0xfffffffffffebb48; WORD $0x0fff // movabsq $1152921504606846974, %rbx # imm = 0xFFFFFFFFFFFFFFE + WORD $0x214c; BYTE $0xd3 // andq %r10, %rbx + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + WORD $0x8948; BYTE $0xfa // movq %rdi, %rdx + WORD $0x8949; BYTE $0xf3 // movq %rsi, %r11 + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + +LBB2_10: + LONG $0x1a10fcc5 // vmovups (%rdx), %ymm3 + LONG $0x6210fcc5; BYTE $0x20 // vmovups 32(%rdx), %ymm4 + LONG $0x107cc1c4; BYTE $0x2b // vmovups (%r11), %ymm5 + LONG $0x107cc1c4; WORD $0x2073 // vmovups 32(%r11), %ymm6 + LONG $0xfd59e4c5 // vmulps %ymm5, %ymm3, %ymm7 + LONG $0xd758ecc5 // vaddps %ymm7, %ymm2, %ymm2 + LONG $0xdb59e4c5 // vmulps %ymm3, %ymm3, %ymm3 + LONG $0xcb58f4c5 // vaddps %ymm3, %ymm1, %ymm1 + LONG $0xdd59d4c5 // vmulps %ymm5, %ymm5, %ymm3 + LONG $0xc358fcc5 // vaddps %ymm3, %ymm0, %ymm0 + LONG $0xde59dcc5 // vmulps %ymm6, %ymm4, %ymm3 + LONG $0xd358ecc5 // vaddps %ymm3, %ymm2, %ymm2 + LONG $0xdc59dcc5 // vmulps %ymm4, %ymm4, %ymm3 + LONG $0xcb58f4c5 // vaddps %ymm3, %ymm1, %ymm1 + LONG $0xde59ccc5 // vmulps %ymm6, %ymm6, %ymm3 + LONG $0xc358fcc5 // vaddps %ymm3, %ymm0, %ymm0 + LONG $0x40c28348 // addq $64, %rdx + LONG $0x40c38349 // addq $64, %r11 + LONG $0xfec38348 // addq $-2, %rbx + JNE LBB2_10 + +LBB2_4: + LONG $0x01c2f641 // testb $1, %r10b + JE LBB2_6 + LONG $0x1a10fcc5 // vmovups (%rdx), %ymm3 + LONG $0x107cc1c4; BYTE $0x23 // vmovups (%r11), %ymm4 + LONG $0xec59e4c5 // vmulps %ymm4, %ymm3, %ymm5 + LONG $0xd558ecc5 // vaddps %ymm5, %ymm2, %ymm2 + LONG $0xdb59e4c5 // vmulps %ymm3, %ymm3, %ymm3 + LONG $0xcb58f4c5 // vaddps %ymm3, %ymm1, %ymm1 + LONG $0xdc59dcc5 // vmulps %ymm4, %ymm4, %ymm3 + LONG $0xc358fcc5 // vaddps %ymm3, %ymm0, %ymm0 + +LBB2_6: + LONG $0x05e2c149 // shlq $5, %r10 + WORD $0x014c; BYTE $0xd7 // addq %r10, %rdi + WORD $0x014c; BYTE $0xd6 // addq %r10, %rsi + +LBB2_7: + LONG $0xdb57e0c5 // vxorps %xmm3, %xmm3, %xmm3 + LONG $0xdb58eac5 // vaddss %xmm3, %xmm2, %xmm3 + LONG $0xe216fac5 // vmovshdup %xmm2, %xmm4 # xmm4 = xmm2[1,1,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xe2c6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm4 # xmm4 = xmm2[1,0] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xe2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm4 # xmm4 = xmm2[3,3,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0x197de3c4; WORD $0x01d2 // vextractf128 $1, %ymm2, %xmm2 + LONG $0xdb58eac5 // vaddss %xmm3, %xmm2, %xmm3 + LONG $0xe216fac5 // vmovshdup %xmm2, %xmm4 # xmm4 = xmm2[1,1,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xe2c6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm4 # xmm4 = xmm2[1,0] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xd358eac5 // vaddss %xmm3, %xmm2, %xmm2 + LONG $0x1111fac5 // vmovss %xmm2, (%rcx) + LONG $0xd116fac5 // vmovshdup %xmm1, %xmm2 # xmm2 = xmm1[1,1,3,3] + LONG $0xd258f2c5 // vaddss %xmm2, %xmm1, %xmm2 + LONG $0xd9c6f1c5; BYTE $0x01 // vshufpd $1, %xmm1, %xmm1, %xmm3 # xmm3 = xmm1[1,0] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xd9c6f0c5; BYTE $0xff // vshufps $255, %xmm1, %xmm1, %xmm3 # xmm3 = xmm1[3,3,3,3] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0x197de3c4; WORD $0x01c9 // vextractf128 $1, %ymm1, %xmm1 + LONG $0xd258f2c5 // vaddss %xmm2, %xmm1, %xmm2 + LONG $0xd916fac5 // vmovshdup %xmm1, %xmm3 # xmm3 = xmm1[1,1,3,3] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xd9c6f1c5; BYTE $0x01 // vshufpd $1, %xmm1, %xmm1, %xmm3 # xmm3 = xmm1[1,0] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xc9c6f0c5; BYTE $0xff // vshufps $255, %xmm1, %xmm1, %xmm1 # xmm1 = xmm1[3,3,3,3] + LONG $0xca58f2c5 // vaddss %xmm2, %xmm1, %xmm1 + LONG $0x117ac1c4; BYTE $0x08 // vmovss %xmm1, (%r8) + LONG $0xc816fac5 // vmovshdup %xmm0, %xmm1 # xmm1 = xmm0[1,1,3,3] + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x197de3c4; WORD $0x01c0 // vextractf128 $1, %ymm0, %xmm0 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xc0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm0 # xmm0 = xmm0[3,3,3,3] + LONG $0xc158fac5 // vaddss %xmm1, %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + WORD $0x8548; BYTE $0xc0 // testq %rax, %rax + JLE LBB2_8 + LONG $0x0710fac5 // vmovss (%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x0659fac5 // vmulss (%rsi), %xmm0, %xmm0 + LONG $0x0158fac5 // vaddss (%rcx), %xmm0, %xmm0 + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + LONG $0x0710fac5 // vmovss (%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x00 // vaddss (%r8), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x0610fac5 // vmovss (%rsi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x01 // vaddss (%r9), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + LONG $0x01f88348 // cmpq $1, %rax + JE LBB2_8 + LONG $0x4710fac5; BYTE $0x04 // vmovss 4(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x4659fac5; BYTE $0x04 // vmulss 4(%rsi), %xmm0, %xmm0 + LONG $0x0158fac5 // vaddss (%rcx), %xmm0, %xmm0 + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + LONG $0x4710fac5; BYTE $0x04 // vmovss 4(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x00 // vaddss (%r8), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x4610fac5; BYTE $0x04 // vmovss 4(%rsi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x01 // vaddss (%r9), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + LONG $0x02f88348 // cmpq $2, %rax + JE LBB2_8 + LONG $0x4710fac5; BYTE $0x08 // vmovss 8(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x4659fac5; BYTE $0x08 // vmulss 8(%rsi), %xmm0, %xmm0 + LONG $0x0158fac5 // vaddss (%rcx), %xmm0, %xmm0 + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + LONG $0x4710fac5; BYTE $0x08 // vmovss 8(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x00 // vaddss (%r8), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x4610fac5; BYTE $0x08 // vmovss 8(%rsi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x01 // vaddss (%r9), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + LONG $0x03f88348 // cmpq $3, %rax + JE LBB2_8 + LONG $0x4710fac5; BYTE $0x0c // vmovss 12(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x4659fac5; BYTE $0x0c // vmulss 12(%rsi), %xmm0, %xmm0 + LONG $0x0158fac5 // vaddss (%rcx), %xmm0, %xmm0 + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + LONG $0x4710fac5; BYTE $0x0c // vmovss 12(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x00 // vaddss (%r8), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x4610fac5; BYTE $0x0c // vmovss 12(%rsi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x01 // vaddss (%r9), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + LONG $0x04f88348 // cmpq $4, %rax + JE LBB2_8 + LONG $0x4710fac5; BYTE $0x10 // vmovss 16(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x4659fac5; BYTE $0x10 // vmulss 16(%rsi), %xmm0, %xmm0 + LONG $0x0158fac5 // vaddss (%rcx), %xmm0, %xmm0 + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + LONG $0x4710fac5; BYTE $0x10 // vmovss 16(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x00 // vaddss (%r8), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x4610fac5; BYTE $0x10 // vmovss 16(%rsi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x01 // vaddss (%r9), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + LONG $0x05f88348 // cmpq $5, %rax + JE LBB2_8 + LONG $0x4710fac5; BYTE $0x14 // vmovss 20(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x4659fac5; BYTE $0x14 // vmulss 20(%rsi), %xmm0, %xmm0 + LONG $0x0158fac5 // vaddss (%rcx), %xmm0, %xmm0 + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + LONG $0x4710fac5; BYTE $0x14 // vmovss 20(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x00 // vaddss (%r8), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x4610fac5; BYTE $0x14 // vmovss 20(%rsi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x01 // vaddss (%r9), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + LONG $0x06f88348 // cmpq $6, %rax + JE LBB2_8 + LONG $0x4710fac5; BYTE $0x18 // vmovss 24(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x4659fac5; BYTE $0x18 // vmulss 24(%rsi), %xmm0, %xmm0 + LONG $0x0158fac5 // vaddss (%rcx), %xmm0, %xmm0 + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + LONG $0x4710fac5; BYTE $0x18 // vmovss 24(%rdi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x00 // vaddss (%r8), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x4610fac5; BYTE $0x18 // vmovss 24(%rsi), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xc059fac5 // vmulss %xmm0, %xmm0, %xmm0 + LONG $0x587ac1c4; BYTE $0x01 // vaddss (%r9), %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + +LBB2_8: + LONG $0xf8658d48 // leaq -8(%rbp), %rsp + BYTE $0x5b // popq %rbx + BYTE $0x5d // popq %rbp + WORD $0xf8c5; BYTE $0x77 // vzeroupper + RET diff --git a/internal/floats/floats_avx512.go b/internal/floats/floats_avx512.go new file mode 100644 index 0000000..85d7fe6 --- /dev/null +++ b/internal/floats/floats_avx512.go @@ -0,0 +1,20 @@ +//go:build !noasm && amd64 +// Code generated by GoAT. DO NOT EDIT. +// versions: +// clang 21.1.8 (6ubuntu1) +// objdump 2.46 +// flags: -mavx -mfma -mavx512f -O3 +// source: src/floats_avx512.c + +package floats + +import "unsafe" + +//go:noescape +func xvec_avx512_l2_squared(left, right unsafe.Pointer, size int64, output unsafe.Pointer) + +//go:noescape +func xvec_avx512_inner_product(left, right unsafe.Pointer, size int64, output unsafe.Pointer) + +//go:noescape +func xvec_avx512_dot_norms(left, right unsafe.Pointer, size int64, dot, left_norm, right_norm unsafe.Pointer) diff --git a/internal/floats/floats_avx512.s b/internal/floats/floats_avx512.s new file mode 100644 index 0000000..aff450b --- /dev/null +++ b/internal/floats/floats_avx512.s @@ -0,0 +1,528 @@ +//go:build !noasm && amd64 +// Code generated by GoAT. DO NOT EDIT. +// versions: +// clang 21.1.8 (6ubuntu1) +// objdump 2.46 +// flags: -mavx -mfma -mavx512f -O3 +// source: src/floats_avx512.c + +TEXT ·xvec_avx512_l2_squared(SB), $0-32 + MOVQ left+0(FP), DI + MOVQ right+8(FP), SI + MOVQ size+16(FP), DX + MOVQ output+24(FP), CX + BYTE $0x55 // pushq %rbp + WORD $0x8948; BYTE $0xe5 // movq %rsp, %rbp + LONG $0xf8e48348 // andq $-8, %rsp + LONG $0x0f428d4c // leaq 15(%rdx), %r8 + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + LONG $0xc2490f4c // cmovnsq %rdx, %r8 + WORD $0x894d; BYTE $0xc1 // movq %r8, %r9 + LONG $0xf0e18349 // andq $-16, %r9 + WORD $0x8948; BYTE $0xd0 // movq %rdx, %rax + WORD $0x294c; BYTE $0xc8 // subq %r9, %rax + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + LONG $0x10fa8348 // cmpq $16, %rdx + JL LBB0_7 + LONG $0x04f8c149 // sarq $4, %r8 + LONG $0xff488d4d // leaq -1(%r8), %r9 + WORD $0x8944; BYTE $0xc2 // movl %r8d, %edx + WORD $0xe283; BYTE $0x03 // andl $3, %edx + LONG $0x03f98349 // cmpq $3, %r9 + JAE LBB0_14 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + JMP LBB0_3 + +LBB0_14: + QUAD $0xfffffffffffcbb49; WORD $0x07ff // movabsq $576460752303423484, %r11 # imm = 0x7FFFFFFFFFFFFFC + WORD $0x214d; BYTE $0xc3 // andq %r8, %r11 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + +LBB0_15: + LONG $0x487cd162; WORD $0x0910 // vmovups (%r9), %zmm1 + LONG $0x487cd162; WORD $0x5110; BYTE $0x01 // vmovups 64(%r9), %zmm2 + LONG $0x487cd162; WORD $0x5910; BYTE $0x02 // vmovups 128(%r9), %zmm3 + LONG $0x487cd162; WORD $0x6110; BYTE $0x03 // vmovups 192(%r9), %zmm4 + LONG $0x4874d162; WORD $0x0a5c // vsubps (%r10), %zmm1, %zmm1 + LONG $0x4875f262; WORD $0xc8a8 // vfmadd213ps %zmm0, %zmm1, %zmm1 # zmm1 = (zmm1 * zmm1) + zmm0 + LONG $0x486cd162; WORD $0x425c; BYTE $0x01 // vsubps 64(%r10), %zmm2, %zmm0 + LONG $0x4864d162; WORD $0x525c; BYTE $0x02 // vsubps 128(%r10), %zmm3, %zmm2 + LONG $0x487df262; WORD $0xc1a8 // vfmadd213ps %zmm1, %zmm0, %zmm0 # zmm0 = (zmm0 * zmm0) + zmm1 + LONG $0x486df262; WORD $0xd0a8 // vfmadd213ps %zmm0, %zmm2, %zmm2 # zmm2 = (zmm2 * zmm2) + zmm0 + LONG $0x485cd162; WORD $0x425c; BYTE $0x03 // vsubps 192(%r10), %zmm4, %zmm0 + LONG $0x487df262; WORD $0xc2a8 // vfmadd213ps %zmm2, %zmm0, %zmm0 # zmm0 = (zmm0 * zmm0) + zmm2 + LONG $0x00c18149; WORD $0x0001; BYTE $0x00 // addq $256, %r9 # imm = 0x100 + LONG $0x00c28149; WORD $0x0001; BYTE $0x00 // addq $256, %r10 # imm = 0x100 + LONG $0xfcc38349 // addq $-4, %r11 + JNE LBB0_15 + +LBB0_3: + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + JE LBB0_6 + WORD $0xe2c1; BYTE $0x06 // shll $6, %edx + WORD $0x3145; BYTE $0xdb // xorl %r11d, %r11d + +LBB0_5: + LONG $0x487c9162; WORD $0x0c10; BYTE $0x19 // vmovups (%r9,%r11), %zmm1 + LONG $0x48749162; WORD $0x0c5c; BYTE $0x1a // vsubps (%r10,%r11), %zmm1, %zmm1 + LONG $0x4875f262; WORD $0xc1b8 // vfmadd231ps %zmm1, %zmm1, %zmm0 # zmm0 = (zmm1 * zmm1) + zmm0 + LONG $0x40c38349 // addq $64, %r11 + WORD $0x394c; BYTE $0xda // cmpq %r11, %rdx + JNE LBB0_5 + +LBB0_6: + LONG $0x06e0c149 // shlq $6, %r8 + WORD $0x014c; BYTE $0xc7 // addq %r8, %rdi + WORD $0x014c; BYTE $0xc6 // addq %r8, %rsi + +LBB0_7: + LONG $0xc816fac5 // vmovshdup %xmm0, %xmm1 # xmm1 = xmm0[1,1,3,3] + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x197de3c4; WORD $0x01c2 // vextractf128 $1, %ymm0, %xmm2 + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xda16fac5 // vmovshdup %xmm2, %xmm3 # xmm3 = xmm2[1,1,3,3] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xdac6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm3 # xmm3 = xmm2[1,0] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x487df362; WORD $0xc219; BYTE $0x02 // vextractf32x4 $2, %zmm0, %xmm2 + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xda16fac5 // vmovshdup %xmm2, %xmm3 # xmm3 = xmm2[1,1,3,3] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xdac6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm3 # xmm3 = xmm2[1,0] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x487df362; WORD $0xc019; BYTE $0x03 // vextractf32x4 $3, %zmm0, %xmm0 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xc0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm0 # xmm0 = xmm0[3,3,3,3] + LONG $0xc158fac5 // vaddss %xmm1, %xmm0, %xmm0 + WORD $0x8548; BYTE $0xc0 // testq %rax, %rax + JLE LBB0_13 + WORD $0xc289 // movl %eax, %edx + WORD $0xe283; BYTE $0x03 // andl $3, %edx + LONG $0x04f88348 // cmpq $4, %rax + JAE LBB0_16 + WORD $0x3145; BYTE $0xc0 // xorl %r8d, %r8d + JMP LBB0_10 + +LBB0_16: + QUAD $0xfffffffffffcb849; WORD $0x7fff // movabsq $9223372036854775804, %r8 # imm = 0x7FFFFFFFFFFFFFFC + WORD $0x214c; BYTE $0xc0 // andq %r8, %rax + WORD $0x3145; BYTE $0xc0 // xorl %r8d, %r8d + +LBB0_17: + LONG $0x107aa1c4; WORD $0x870c // vmovss (%rdi,%r8,4), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x107aa1c4; WORD $0x8754; BYTE $0x04 // vmovss 4(%rdi,%r8,4), %xmm2 # xmm2 = mem[0],zero,zero,zero + LONG $0x5c72a1c4; WORD $0x860c // vsubss (%rsi,%r8,4), %xmm1, %xmm1 + LONG $0x5c6aa1c4; WORD $0x8654; BYTE $0x04 // vsubss 4(%rsi,%r8,4), %xmm2, %xmm2 + LONG $0xa971e2c4; BYTE $0xc8 // vfmadd213ss %xmm0, %xmm1, %xmm1 # xmm1 = (xmm1 * xmm1) + xmm0 + LONG $0x107aa1c4; WORD $0x8744; BYTE $0x08 // vmovss 8(%rdi,%r8,4), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x5c7aa1c4; WORD $0x865c; BYTE $0x08 // vsubss 8(%rsi,%r8,4), %xmm0, %xmm3 + LONG $0xa969e2c4; BYTE $0xd1 // vfmadd213ss %xmm1, %xmm2, %xmm2 # xmm2 = (xmm2 * xmm2) + xmm1 + LONG $0x107aa1c4; WORD $0x8744; BYTE $0x0c // vmovss 12(%rdi,%r8,4), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x5c7aa1c4; WORD $0x8644; BYTE $0x0c // vsubss 12(%rsi,%r8,4), %xmm0, %xmm0 + LONG $0xa961e2c4; BYTE $0xda // vfmadd213ss %xmm2, %xmm3, %xmm3 # xmm3 = (xmm3 * xmm3) + xmm2 + LONG $0xa979e2c4; BYTE $0xc3 // vfmadd213ss %xmm3, %xmm0, %xmm0 # xmm0 = (xmm0 * xmm0) + xmm3 + LONG $0x04c08349 // addq $4, %r8 + WORD $0x394c; BYTE $0xc0 // cmpq %r8, %rax + JNE LBB0_17 + +LBB0_10: + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + JE LBB0_13 + LONG $0x86048d4a // leaq (%rsi,%r8,4), %rax + LONG $0x87348d4a // leaq (%rdi,%r8,4), %rsi + WORD $0xff31 // xorl %edi, %edi + +LBB0_12: + LONG $0x0c10fac5; BYTE $0xbe // vmovss (%rsi,%rdi,4), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x0c5cf2c5; BYTE $0xb8 // vsubss (%rax,%rdi,4), %xmm1, %xmm1 + LONG $0xb971e2c4; BYTE $0xc1 // vfmadd231ss %xmm1, %xmm1, %xmm0 # xmm0 = (xmm1 * xmm1) + xmm0 + WORD $0xff48; BYTE $0xc7 // incq %rdi + WORD $0x3948; BYTE $0xfa // cmpq %rdi, %rdx + JNE LBB0_12 + +LBB0_13: + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + WORD $0x8948; BYTE $0xec // movq %rbp, %rsp + BYTE $0x5d // popq %rbp + WORD $0xf8c5; BYTE $0x77 // vzeroupper + RET + +TEXT ·xvec_avx512_inner_product(SB), $0-32 + MOVQ left+0(FP), DI + MOVQ right+8(FP), SI + MOVQ size+16(FP), DX + MOVQ output+24(FP), CX + BYTE $0x55 // pushq %rbp + WORD $0x8948; BYTE $0xe5 // movq %rsp, %rbp + LONG $0xf8e48348 // andq $-8, %rsp + LONG $0x0f428d4c // leaq 15(%rdx), %r8 + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + LONG $0xc2490f4c // cmovnsq %rdx, %r8 + WORD $0x894d; BYTE $0xc1 // movq %r8, %r9 + LONG $0xf0e18349 // andq $-16, %r9 + WORD $0x8948; BYTE $0xd0 // movq %rdx, %rax + WORD $0x294c; BYTE $0xc8 // subq %r9, %rax + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + LONG $0x10fa8348 // cmpq $16, %rdx + JL LBB1_7 + LONG $0x04f8c149 // sarq $4, %r8 + LONG $0xff488d4d // leaq -1(%r8), %r9 + WORD $0x8944; BYTE $0xc2 // movl %r8d, %edx + WORD $0xe283; BYTE $0x03 // andl $3, %edx + LONG $0x03f98349 // cmpq $3, %r9 + JAE LBB1_14 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + JMP LBB1_3 + +LBB1_14: + QUAD $0xfffffffffffcbb49; WORD $0x07ff // movabsq $576460752303423484, %r11 # imm = 0x7FFFFFFFFFFFFFC + WORD $0x214d; BYTE $0xc3 // andq %r8, %r11 + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + WORD $0x8949; BYTE $0xf9 // movq %rdi, %r9 + WORD $0x8949; BYTE $0xf2 // movq %rsi, %r10 + +LBB1_15: + LONG $0x487cd162; WORD $0x0110 // vmovups (%r9), %zmm0 + LONG $0x487cd162; WORD $0x5110; BYTE $0x01 // vmovups 64(%r9), %zmm2 + LONG $0x487cd162; WORD $0x5910; BYTE $0x02 // vmovups 128(%r9), %zmm3 + LONG $0x487cd162; WORD $0x6110; BYTE $0x03 // vmovups 192(%r9), %zmm4 + LONG $0x4875d262; WORD $0x0298 // vfmadd132ps (%r10), %zmm1, %zmm0 # zmm0 = (zmm0 * mem) + zmm1 + LONG $0x486dd262; WORD $0x42b8; BYTE $0x01 // vfmadd231ps 64(%r10), %zmm2, %zmm0 # zmm0 = (zmm2 * mem) + zmm0 + LONG $0x4865d262; WORD $0x42b8; BYTE $0x02 // vfmadd231ps 128(%r10), %zmm3, %zmm0 # zmm0 = (zmm3 * mem) + zmm0 + LONG $0x485dd262; WORD $0x42b8; BYTE $0x03 // vfmadd231ps 192(%r10), %zmm4, %zmm0 # zmm0 = (zmm4 * mem) + zmm0 + LONG $0x00c18149; WORD $0x0001; BYTE $0x00 // addq $256, %r9 # imm = 0x100 + LONG $0x00c28149; WORD $0x0001; BYTE $0x00 // addq $256, %r10 # imm = 0x100 + LONG $0x487cf162; WORD $0xc828 // vmovaps %zmm0, %zmm1 + LONG $0xfcc38349 // addq $-4, %r11 + JNE LBB1_15 + +LBB1_3: + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + JE LBB1_6 + WORD $0xe2c1; BYTE $0x06 // shll $6, %edx + WORD $0x3145; BYTE $0xdb // xorl %r11d, %r11d + +LBB1_5: + LONG $0x487c9162; WORD $0x0c10; BYTE $0x19 // vmovups (%r9,%r11), %zmm1 + LONG $0x48759262; WORD $0x04b8; BYTE $0x1a // vfmadd231ps (%r10,%r11), %zmm1, %zmm0 # zmm0 = (zmm1 * mem) + zmm0 + LONG $0x40c38349 // addq $64, %r11 + WORD $0x394c; BYTE $0xda // cmpq %r11, %rdx + JNE LBB1_5 + +LBB1_6: + LONG $0x06e0c149 // shlq $6, %r8 + WORD $0x014c; BYTE $0xc7 // addq %r8, %rdi + WORD $0x014c; BYTE $0xc6 // addq %r8, %rsi + +LBB1_7: + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x197de3c4; WORD $0x01c2 // vextractf128 $1, %ymm0, %xmm2 + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xda16fac5 // vmovshdup %xmm2, %xmm3 # xmm3 = xmm2[1,1,3,3] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xdac6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm3 # xmm3 = xmm2[1,0] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x487df362; WORD $0xc219; BYTE $0x02 // vextractf32x4 $2, %zmm0, %xmm2 + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xda16fac5 // vmovshdup %xmm2, %xmm3 # xmm3 = xmm2[1,1,3,3] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xdac6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm3 # xmm3 = xmm2[1,0] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x487df362; WORD $0xc019; BYTE $0x03 // vextractf32x4 $3, %zmm0, %xmm0 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xc0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm0 # xmm0 = xmm0[3,3,3,3] + LONG $0xc158fac5 // vaddss %xmm1, %xmm0, %xmm0 + WORD $0x8548; BYTE $0xc0 // testq %rax, %rax + JLE LBB1_13 + WORD $0xc289 // movl %eax, %edx + WORD $0xe283; BYTE $0x03 // andl $3, %edx + LONG $0x04f88348 // cmpq $4, %rax + JAE LBB1_16 + WORD $0x3145; BYTE $0xc0 // xorl %r8d, %r8d + JMP LBB1_10 + +LBB1_16: + QUAD $0xfffffffffffcb849; WORD $0x7fff // movabsq $9223372036854775804, %r8 # imm = 0x7FFFFFFFFFFFFFFC + WORD $0x214c; BYTE $0xc0 // andq %r8, %rax + WORD $0x3145; BYTE $0xc0 // xorl %r8d, %r8d + +LBB1_17: + LONG $0x107aa1c4; WORD $0x870c // vmovss (%rdi,%r8,4), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0x107aa1c4; WORD $0x8754; BYTE $0x04 // vmovss 4(%rdi,%r8,4), %xmm2 # xmm2 = mem[0],zero,zero,zero + LONG $0x9979a2c4; WORD $0x860c // vfmadd132ss (%rsi,%r8,4), %xmm0, %xmm1 # xmm1 = (xmm1 * mem) + xmm0 + LONG $0xb969a2c4; WORD $0x864c; BYTE $0x04 // vfmadd231ss 4(%rsi,%r8,4), %xmm2, %xmm1 # xmm1 = (xmm2 * mem) + xmm1 + LONG $0x107aa1c4; WORD $0x8754; BYTE $0x08 // vmovss 8(%rdi,%r8,4), %xmm2 # xmm2 = mem[0],zero,zero,zero + LONG $0x9971a2c4; WORD $0x8654; BYTE $0x08 // vfmadd132ss 8(%rsi,%r8,4), %xmm1, %xmm2 # xmm2 = (xmm2 * mem) + xmm1 + LONG $0x107aa1c4; WORD $0x8744; BYTE $0x0c // vmovss 12(%rdi,%r8,4), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x9969a2c4; WORD $0x8644; BYTE $0x0c // vfmadd132ss 12(%rsi,%r8,4), %xmm2, %xmm0 # xmm0 = (xmm0 * mem) + xmm2 + LONG $0x04c08349 // addq $4, %r8 + WORD $0x394c; BYTE $0xc0 // cmpq %r8, %rax + JNE LBB1_17 + +LBB1_10: + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + JE LBB1_13 + LONG $0x86048d4a // leaq (%rsi,%r8,4), %rax + LONG $0x87348d4a // leaq (%rdi,%r8,4), %rsi + WORD $0xff31 // xorl %edi, %edi + +LBB1_12: + LONG $0x0c10fac5; BYTE $0xbe // vmovss (%rsi,%rdi,4), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0xb971e2c4; WORD $0xb804 // vfmadd231ss (%rax,%rdi,4), %xmm1, %xmm0 # xmm0 = (xmm1 * mem) + xmm0 + WORD $0xff48; BYTE $0xc7 // incq %rdi + WORD $0x3948; BYTE $0xfa // cmpq %rdi, %rdx + JNE LBB1_12 + +LBB1_13: + LONG $0x0111fac5 // vmovss %xmm0, (%rcx) + WORD $0x8948; BYTE $0xec // movq %rbp, %rsp + BYTE $0x5d // popq %rbp + WORD $0xf8c5; BYTE $0x77 // vzeroupper + RET + +TEXT ·xvec_avx512_dot_norms(SB), $0-48 + MOVQ left+0(FP), DI + MOVQ right+8(FP), SI + MOVQ size+16(FP), DX + MOVQ dot+24(FP), CX + MOVQ left_norm+32(FP), R8 + MOVQ right_norm+40(FP), R9 + BYTE $0x55 // pushq %rbp + WORD $0x8948; BYTE $0xe5 // movq %rsp, %rbp + WORD $0x5641 // pushq %r14 + BYTE $0x53 // pushq %rbx + LONG $0xf8e48348 // andq $-8, %rsp + LONG $0x0f528d4c // leaq 15(%rdx), %r10 + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + LONG $0xd2490f4c // cmovnsq %rdx, %r10 + WORD $0x894d; BYTE $0xd3 // movq %r10, %r11 + LONG $0xf0e38349 // andq $-16, %r11 + WORD $0x8948; BYTE $0xd0 // movq %rdx, %rax + WORD $0x294c; BYTE $0xd8 // subq %r11, %rax + LONG $0xd257e8c5 // vxorps %xmm2, %xmm2, %xmm2 + LONG $0x10fa8348 // cmpq $16, %rdx + JL LBB2_1 + LONG $0x04fac149 // sarq $4, %r10 + LONG $0xff5a8d4d // leaq -1(%r10), %r11 + WORD $0x8944; BYTE $0xd2 // movl %r10d, %edx + WORD $0xe283; BYTE $0x03 // andl $3, %edx + LONG $0x03fb8349 // cmpq $3, %r11 + JAE LBB2_12 + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + WORD $0x8949; BYTE $0xfb // movq %rdi, %r11 + WORD $0x8948; BYTE $0xf3 // movq %rsi, %rbx + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + JMP LBB2_4 + +LBB2_1: + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + JMP LBB2_8 + +LBB2_12: + QUAD $0xfffffffffffcbe49; WORD $0x07ff // movabsq $576460752303423484, %r14 # imm = 0x7FFFFFFFFFFFFFC + WORD $0x214d; BYTE $0xd6 // andq %r10, %r14 + LONG $0xc057f8c5 // vxorps %xmm0, %xmm0, %xmm0 + WORD $0x8949; BYTE $0xfb // movq %rdi, %r11 + WORD $0x8948; BYTE $0xf3 // movq %rsi, %rbx + LONG $0xc957f0c5 // vxorps %xmm1, %xmm1, %xmm1 + +LBB2_13: + LONG $0x487cd162; WORD $0x1b10 // vmovups (%r11), %zmm3 + LONG $0x487cd162; WORD $0x6310; BYTE $0x01 // vmovups 64(%r11), %zmm4 + LONG $0x487cd162; WORD $0x6b10; BYTE $0x02 // vmovups 128(%r11), %zmm5 + LONG $0x487cd162; WORD $0x7310; BYTE $0x03 // vmovups 192(%r11), %zmm6 + LONG $0x487cf162; WORD $0x3b10 // vmovups (%rbx), %zmm7 + LONG $0x487c7162; WORD $0x4310; BYTE $0x01 // vmovups 64(%rbx), %zmm8 + LONG $0x487c7162; WORD $0x4b10; BYTE $0x02 // vmovups 128(%rbx), %zmm9 + LONG $0x487c7162; WORD $0x5310; BYTE $0x03 // vmovups 192(%rbx), %zmm10 + LONG $0x4865f262; WORD $0xd7b8 // vfmadd231ps %zmm7, %zmm3, %zmm2 # zmm2 = (zmm3 * zmm7) + zmm2 + LONG $0x4865f262; WORD $0xcbb8 // vfmadd231ps %zmm3, %zmm3, %zmm1 # zmm1 = (zmm3 * zmm3) + zmm1 + LONG $0x4845f262; WORD $0xc7b8 // vfmadd231ps %zmm7, %zmm7, %zmm0 # zmm0 = (zmm7 * zmm7) + zmm0 + LONG $0x485dd262; WORD $0xd0b8 // vfmadd231ps %zmm8, %zmm4, %zmm2 # zmm2 = (zmm4 * zmm8) + zmm2 + LONG $0x485df262; WORD $0xccb8 // vfmadd231ps %zmm4, %zmm4, %zmm1 # zmm1 = (zmm4 * zmm4) + zmm1 + LONG $0x483dd262; WORD $0xc0b8 // vfmadd231ps %zmm8, %zmm8, %zmm0 # zmm0 = (zmm8 * zmm8) + zmm0 + LONG $0x4855d262; WORD $0xd1b8 // vfmadd231ps %zmm9, %zmm5, %zmm2 # zmm2 = (zmm5 * zmm9) + zmm2 + LONG $0x4855f262; WORD $0xcdb8 // vfmadd231ps %zmm5, %zmm5, %zmm1 # zmm1 = (zmm5 * zmm5) + zmm1 + LONG $0x4835d262; WORD $0xc1b8 // vfmadd231ps %zmm9, %zmm9, %zmm0 # zmm0 = (zmm9 * zmm9) + zmm0 + LONG $0x484dd262; WORD $0xd2b8 // vfmadd231ps %zmm10, %zmm6, %zmm2 # zmm2 = (zmm6 * zmm10) + zmm2 + LONG $0x484df262; WORD $0xceb8 // vfmadd231ps %zmm6, %zmm6, %zmm1 # zmm1 = (zmm6 * zmm6) + zmm1 + LONG $0x482dd262; WORD $0xc2b8 // vfmadd231ps %zmm10, %zmm10, %zmm0 # zmm0 = (zmm10 * zmm10) + zmm0 + LONG $0x00c38149; WORD $0x0001; BYTE $0x00 // addq $256, %r11 # imm = 0x100 + LONG $0x00c38148; WORD $0x0001; BYTE $0x00 // addq $256, %rbx # imm = 0x100 + LONG $0xfcc68349 // addq $-4, %r14 + JNE LBB2_13 + +LBB2_4: + WORD $0x8548; BYTE $0xd2 // testq %rdx, %rdx + JE LBB2_7 + WORD $0xe2c1; BYTE $0x06 // shll $6, %edx + WORD $0x3145; BYTE $0xf6 // xorl %r14d, %r14d + +LBB2_6: + LONG $0x487c9162; WORD $0x1c10; BYTE $0x33 // vmovups (%r11,%r14), %zmm3 + LONG $0x487cb162; WORD $0x2410; BYTE $0x33 // vmovups (%rbx,%r14), %zmm4 + LONG $0x4865f262; WORD $0xd4b8 // vfmadd231ps %zmm4, %zmm3, %zmm2 # zmm2 = (zmm3 * zmm4) + zmm2 + LONG $0x4865f262; WORD $0xcbb8 // vfmadd231ps %zmm3, %zmm3, %zmm1 # zmm1 = (zmm3 * zmm3) + zmm1 + LONG $0x485df262; WORD $0xc4b8 // vfmadd231ps %zmm4, %zmm4, %zmm0 # zmm0 = (zmm4 * zmm4) + zmm0 + LONG $0x40c68349 // addq $64, %r14 + WORD $0x394c; BYTE $0xf2 // cmpq %r14, %rdx + JNE LBB2_6 + +LBB2_7: + LONG $0x06e2c149 // shlq $6, %r10 + WORD $0x014c; BYTE $0xd7 // addq %r10, %rdi + WORD $0x014c; BYTE $0xd6 // addq %r10, %rsi + +LBB2_8: + LONG $0xdb57e0c5 // vxorps %xmm3, %xmm3, %xmm3 + LONG $0xdb58eac5 // vaddss %xmm3, %xmm2, %xmm3 + LONG $0xe216fac5 // vmovshdup %xmm2, %xmm4 # xmm4 = xmm2[1,1,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xe2c6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm4 # xmm4 = xmm2[1,0] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xe2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm4 # xmm4 = xmm2[3,3,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0x197de3c4; WORD $0x01d4 // vextractf128 $1, %ymm2, %xmm4 + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xec16fac5 // vmovshdup %xmm4, %xmm5 # xmm5 = xmm4[1,1,3,3] + LONG $0xdb58d2c5 // vaddss %xmm3, %xmm5, %xmm3 + LONG $0xecc6d9c5; BYTE $0x01 // vshufpd $1, %xmm4, %xmm4, %xmm5 # xmm5 = xmm4[1,0] + LONG $0xdb58d2c5 // vaddss %xmm3, %xmm5, %xmm3 + LONG $0xe4c6d8c5; BYTE $0xff // vshufps $255, %xmm4, %xmm4, %xmm4 # xmm4 = xmm4[3,3,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0x487df362; WORD $0xd419; BYTE $0x02 // vextractf32x4 $2, %zmm2, %xmm4 + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xec16fac5 // vmovshdup %xmm4, %xmm5 # xmm5 = xmm4[1,1,3,3] + LONG $0xdb58d2c5 // vaddss %xmm3, %xmm5, %xmm3 + LONG $0xecc6d9c5; BYTE $0x01 // vshufpd $1, %xmm4, %xmm4, %xmm5 # xmm5 = xmm4[1,0] + LONG $0xdb58d2c5 // vaddss %xmm3, %xmm5, %xmm3 + LONG $0xe4c6d8c5; BYTE $0xff // vshufps $255, %xmm4, %xmm4, %xmm4 # xmm4 = xmm4[3,3,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0x487df362; WORD $0xd219; BYTE $0x03 // vextractf32x4 $3, %zmm2, %xmm2 + LONG $0xdb58eac5 // vaddss %xmm3, %xmm2, %xmm3 + LONG $0xe216fac5 // vmovshdup %xmm2, %xmm4 # xmm4 = xmm2[1,1,3,3] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xe2c6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm4 # xmm4 = xmm2[1,0] + LONG $0xdb58dac5 // vaddss %xmm3, %xmm4, %xmm3 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xd358eac5 // vaddss %xmm3, %xmm2, %xmm2 + LONG $0x1111fac5 // vmovss %xmm2, (%rcx) + LONG $0xd116fac5 // vmovshdup %xmm1, %xmm2 # xmm2 = xmm1[1,1,3,3] + LONG $0xd258f2c5 // vaddss %xmm2, %xmm1, %xmm2 + LONG $0xd9c6f1c5; BYTE $0x01 // vshufpd $1, %xmm1, %xmm1, %xmm3 # xmm3 = xmm1[1,0] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xd9c6f0c5; BYTE $0xff // vshufps $255, %xmm1, %xmm1, %xmm3 # xmm3 = xmm1[3,3,3,3] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0x197de3c4; WORD $0x01cb // vextractf128 $1, %ymm1, %xmm3 + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xe316fac5 // vmovshdup %xmm3, %xmm4 # xmm4 = xmm3[1,1,3,3] + LONG $0xd258dac5 // vaddss %xmm2, %xmm4, %xmm2 + LONG $0xe3c6e1c5; BYTE $0x01 // vshufpd $1, %xmm3, %xmm3, %xmm4 # xmm4 = xmm3[1,0] + LONG $0xd258dac5 // vaddss %xmm2, %xmm4, %xmm2 + LONG $0xdbc6e0c5; BYTE $0xff // vshufps $255, %xmm3, %xmm3, %xmm3 # xmm3 = xmm3[3,3,3,3] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0x487df362; WORD $0xcb19; BYTE $0x02 // vextractf32x4 $2, %zmm1, %xmm3 + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xe316fac5 // vmovshdup %xmm3, %xmm4 # xmm4 = xmm3[1,1,3,3] + LONG $0xd258dac5 // vaddss %xmm2, %xmm4, %xmm2 + LONG $0xe3c6e1c5; BYTE $0x01 // vshufpd $1, %xmm3, %xmm3, %xmm4 # xmm4 = xmm3[1,0] + LONG $0xd258dac5 // vaddss %xmm2, %xmm4, %xmm2 + LONG $0xdbc6e0c5; BYTE $0xff // vshufps $255, %xmm3, %xmm3, %xmm3 # xmm3 = xmm3[3,3,3,3] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0x487df362; WORD $0xc919; BYTE $0x03 // vextractf32x4 $3, %zmm1, %xmm1 + LONG $0xd258f2c5 // vaddss %xmm2, %xmm1, %xmm2 + LONG $0xd916fac5 // vmovshdup %xmm1, %xmm3 # xmm3 = xmm1[1,1,3,3] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xd9c6f1c5; BYTE $0x01 // vshufpd $1, %xmm1, %xmm1, %xmm3 # xmm3 = xmm1[1,0] + LONG $0xd258e2c5 // vaddss %xmm2, %xmm3, %xmm2 + LONG $0xc9c6f0c5; BYTE $0xff // vshufps $255, %xmm1, %xmm1, %xmm1 # xmm1 = xmm1[3,3,3,3] + LONG $0xca58f2c5 // vaddss %xmm2, %xmm1, %xmm1 + LONG $0x117ac1c4; BYTE $0x08 // vmovss %xmm1, (%r8) + LONG $0xc816fac5 // vmovshdup %xmm0, %xmm1 # xmm1 = xmm0[1,1,3,3] + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x197de3c4; WORD $0x01c2 // vextractf128 $1, %ymm0, %xmm2 + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xda16fac5 // vmovshdup %xmm2, %xmm3 # xmm3 = xmm2[1,1,3,3] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xdac6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm3 # xmm3 = xmm2[1,0] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x487df362; WORD $0xc219; BYTE $0x02 // vextractf32x4 $2, %zmm0, %xmm2 + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xda16fac5 // vmovshdup %xmm2, %xmm3 # xmm3 = xmm2[1,1,3,3] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xdac6e9c5; BYTE $0x01 // vshufpd $1, %xmm2, %xmm2, %xmm3 # xmm3 = xmm2[1,0] + LONG $0xc958e2c5 // vaddss %xmm1, %xmm3, %xmm1 + LONG $0xd2c6e8c5; BYTE $0xff // vshufps $255, %xmm2, %xmm2, %xmm2 # xmm2 = xmm2[3,3,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0x487df362; WORD $0xc019; BYTE $0x03 // vextractf32x4 $3, %zmm0, %xmm0 + LONG $0xc958fac5 // vaddss %xmm1, %xmm0, %xmm1 + LONG $0xd016fac5 // vmovshdup %xmm0, %xmm2 # xmm2 = xmm0[1,1,3,3] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xd0c6f9c5; BYTE $0x01 // vshufpd $1, %xmm0, %xmm0, %xmm2 # xmm2 = xmm0[1,0] + LONG $0xc958eac5 // vaddss %xmm1, %xmm2, %xmm1 + LONG $0xc0c6f8c5; BYTE $0xff // vshufps $255, %xmm0, %xmm0, %xmm0 # xmm0 = xmm0[3,3,3,3] + LONG $0xc158fac5 // vaddss %xmm1, %xmm0, %xmm0 + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + WORD $0x8548; BYTE $0xc0 // testq %rax, %rax + JLE LBB2_11 + WORD $0xd231 // xorl %edx, %edx + +LBB2_10: + LONG $0x0410fac5; BYTE $0x97 // vmovss (%rdi,%rdx,4), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0x0c10fac5; BYTE $0x96 // vmovss (%rsi,%rdx,4), %xmm1 # xmm1 = mem[0],zero,zero,zero + LONG $0xa979e2c4; BYTE $0x09 // vfmadd213ss (%rcx), %xmm0, %xmm1 # xmm1 = (xmm0 * xmm1) + mem + LONG $0x0911fac5 // vmovss %xmm1, (%rcx) + LONG $0x0410fac5; BYTE $0x97 // vmovss (%rdi,%rdx,4), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xa979c2c4; BYTE $0x00 // vfmadd213ss (%r8), %xmm0, %xmm0 # xmm0 = (xmm0 * xmm0) + mem + LONG $0x117ac1c4; BYTE $0x00 // vmovss %xmm0, (%r8) + LONG $0x0410fac5; BYTE $0x96 // vmovss (%rsi,%rdx,4), %xmm0 # xmm0 = mem[0],zero,zero,zero + LONG $0xa979c2c4; BYTE $0x01 // vfmadd213ss (%r9), %xmm0, %xmm0 # xmm0 = (xmm0 * xmm0) + mem + LONG $0x117ac1c4; BYTE $0x01 // vmovss %xmm0, (%r9) + WORD $0xff48; BYTE $0xc2 // incq %rdx + WORD $0x3948; BYTE $0xd0 // cmpq %rdx, %rax + JNE LBB2_10 + +LBB2_11: + LONG $0xf0658d48 // leaq -16(%rbp), %rsp + BYTE $0x5b // popq %rbx + WORD $0x5e41 // popq %r14 + BYTE $0x5d // popq %rbp + WORD $0xf8c5; BYTE $0x77 // vzeroupper + RET diff --git a/internal/floats/floats_neon.go b/internal/floats/floats_neon.go new file mode 100644 index 0000000..4e7a1ce --- /dev/null +++ b/internal/floats/floats_neon.go @@ -0,0 +1,20 @@ +//go:build !noasm && arm64 +// Code generated by GoAT. DO NOT EDIT. +// versions: +// clang 21.1.8 (6ubuntu1) +// objdump 2.46 +// flags: -O3 +// source: src/floats_neon.c + +package floats + +import "unsafe" + +//go:noescape +func xvec_neon_l2_squared(left, right unsafe.Pointer, size int64, output unsafe.Pointer) + +//go:noescape +func xvec_neon_inner_product(left, right unsafe.Pointer, size int64, output unsafe.Pointer) + +//go:noescape +func xvec_neon_dot_norms(left, right unsafe.Pointer, size int64, dot, left_norm, right_norm unsafe.Pointer) diff --git a/internal/floats/floats_neon.s b/internal/floats/floats_neon.s new file mode 100644 index 0000000..b99617b --- /dev/null +++ b/internal/floats/floats_neon.s @@ -0,0 +1,234 @@ +//go:build !noasm && arm64 +// Code generated by GoAT. DO NOT EDIT. +// versions: +// clang 21.1.8 (6ubuntu1) +// objdump 2.46 +// flags: -O3 +// source: src/floats_neon.c + +TEXT ·xvec_neon_l2_squared(SB), $0-32 + MOVD left+0(FP), R0 + MOVD right+8(FP), R1 + MOVD size+16(FP), R2 + MOVD output+24(FP), R3 + WORD $0x91000c48 // add x8, x2, #3 + WORD $0xf100005f // cmp x2, #0 + WORD $0x6f00e400 // movi v0.2d, #0000000000000000 + WORD $0x9a82b109 // csel x9, x8, x2, lt + WORD $0xf100105f // cmp x2, #4 + WORD $0x927ef528 // and x8, x9, #0xfffffffffffffffc + WORD $0xcb080048 // sub x8, x2, x8 + BLT LBB0_4 + WORD $0x9342fd29 // asr x9, x9, #2 + WORD $0xaa0003eb // mov x11, x0 + WORD $0xaa0103ec // mov x12, x1 + WORD $0xaa0903ea // mov x10, x9 + +LBB0_2: + WORD $0x3cc10561 // ldr q1, [x11], #16 + WORD $0xf100054a // subs x10, x10, #1 + WORD $0x3cc10582 // ldr q2, [x12], #16 + WORD $0x4ea2d421 // fsub v1.4s, v1.4s, v2.4s + WORD $0x4e21cc20 // fmla v0.4s, v1.4s, v1.4s + BNE LBB0_2 + WORD $0xd37ced29 // lsl x9, x9, #4 + WORD $0x8b090000 // add x0, x0, x9 + WORD $0x8b090021 // add x1, x1, x9 + +LBB0_4: + WORD $0x4e0c0401 // dup v1.4s, v0.s[1] + WORD $0x4e140402 // dup v2.4s, v0.s[2] + WORD $0xf100011f // cmp x8, #0 + WORD $0x4e21d401 // fadd v1.4s, v0.4s, v1.4s + WORD $0x4e1c0400 // dup v0.4s, v0.s[3] + WORD $0x4e21d441 // fadd v1.4s, v2.4s, v1.4s + WORD $0x4e21d400 // fadd v0.4s, v0.4s, v1.4s + BLE LBB0_8 + WORD $0xbd400001 // ldr s1, [x0] + WORD $0xbd400022 // ldr s2, [x1] + WORD $0xf100051f // cmp x8, #1 + WORD $0x1e223821 // fsub s1, s1, s2 + WORD $0x1f010020 // fmadd s0, s1, s1, s0 + BEQ LBB0_8 + WORD $0xbd400401 // ldr s1, [x0, #4] + WORD $0xbd400422 // ldr s2, [x1, #4] + WORD $0xf100091f // cmp x8, #2 + WORD $0x1e223821 // fsub s1, s1, s2 + WORD $0x1f010020 // fmadd s0, s1, s1, s0 + BEQ LBB0_8 + WORD $0xbd400801 // ldr s1, [x0, #8] + WORD $0xbd400822 // ldr s2, [x1, #8] + WORD $0x1e223821 // fsub s1, s1, s2 + WORD $0x1f010020 // fmadd s0, s1, s1, s0 + +LBB0_8: + WORD $0xbd000060 // str s0, [x3] + RET + +TEXT ·xvec_neon_inner_product(SB), $0-32 + MOVD left+0(FP), R0 + MOVD right+8(FP), R1 + MOVD size+16(FP), R2 + MOVD output+24(FP), R3 + WORD $0x91000c48 // add x8, x2, #3 + WORD $0xf100005f // cmp x2, #0 + WORD $0x6f00e400 // movi v0.2d, #0000000000000000 + WORD $0x9a82b109 // csel x9, x8, x2, lt + WORD $0xf100105f // cmp x2, #4 + WORD $0x927ef528 // and x8, x9, #0xfffffffffffffffc + WORD $0xcb080048 // sub x8, x2, x8 + BLT LBB1_4 + WORD $0x9342fd29 // asr x9, x9, #2 + WORD $0xaa0003eb // mov x11, x0 + WORD $0xaa0103ec // mov x12, x1 + WORD $0xaa0903ea // mov x10, x9 + +LBB1_2: + WORD $0x3cc10561 // ldr q1, [x11], #16 + WORD $0xf100054a // subs x10, x10, #1 + WORD $0x3cc10582 // ldr q2, [x12], #16 + WORD $0x4e21cc40 // fmla v0.4s, v2.4s, v1.4s + BNE LBB1_2 + WORD $0xd37ced29 // lsl x9, x9, #4 + WORD $0x8b090000 // add x0, x0, x9 + WORD $0x8b090021 // add x1, x1, x9 + +LBB1_4: + WORD $0x4e0c0401 // dup v1.4s, v0.s[1] + WORD $0x4e140402 // dup v2.4s, v0.s[2] + WORD $0xf100011f // cmp x8, #0 + WORD $0x4e21d401 // fadd v1.4s, v0.4s, v1.4s + WORD $0x4e1c0400 // dup v0.4s, v0.s[3] + WORD $0x4e21d441 // fadd v1.4s, v2.4s, v1.4s + WORD $0x4e21d400 // fadd v0.4s, v0.4s, v1.4s + BLE LBB1_8 + WORD $0xbd400001 // ldr s1, [x0] + WORD $0xbd400022 // ldr s2, [x1] + WORD $0xf100051f // cmp x8, #1 + WORD $0x1f020020 // fmadd s0, s1, s2, s0 + BEQ LBB1_8 + WORD $0xbd400401 // ldr s1, [x0, #4] + WORD $0xbd400422 // ldr s2, [x1, #4] + WORD $0xf100091f // cmp x8, #2 + WORD $0x1f020020 // fmadd s0, s1, s2, s0 + BEQ LBB1_8 + WORD $0xbd400801 // ldr s1, [x0, #8] + WORD $0xbd400822 // ldr s2, [x1, #8] + WORD $0x1f020020 // fmadd s0, s1, s2, s0 + +LBB1_8: + WORD $0xbd000060 // str s0, [x3] + RET + +TEXT ·xvec_neon_dot_norms(SB), $0-48 + MOVD left+0(FP), R0 + MOVD right+8(FP), R1 + MOVD size+16(FP), R2 + MOVD dot+24(FP), R3 + MOVD left_norm+32(FP), R4 + MOVD right_norm+40(FP), R5 + WORD $0x91000c48 // add x8, x2, #3 + WORD $0xf100005f // cmp x2, #0 + WORD $0x6f00e400 // movi v0.2d, #0000000000000000 + WORD $0x9a82b109 // csel x9, x8, x2, lt + WORD $0xf100105f // cmp x2, #4 + WORD $0x927ef528 // and x8, x9, #0xfffffffffffffffc + WORD $0xcb080048 // sub x8, x2, x8 + BLT LBB2_4 + WORD $0x6f00e401 // movi v1.2d, #0000000000000000 + WORD $0x6f00e402 // movi v2.2d, #0000000000000000 + WORD $0x9342fd29 // asr x9, x9, #2 + WORD $0xaa0003eb // mov x11, x0 + WORD $0xaa0103ec // mov x12, x1 + WORD $0xaa0903ea // mov x10, x9 + +LBB2_2: + WORD $0x3cc10563 // ldr q3, [x11], #16 + WORD $0xf100054a // subs x10, x10, #1 + WORD $0x3cc10584 // ldr q4, [x12], #16 + WORD $0x4e23cc62 // fmla v2.4s, v3.4s, v3.4s + WORD $0x4e23cc81 // fmla v1.4s, v4.4s, v3.4s + WORD $0x4e24cc80 // fmla v0.4s, v4.4s, v4.4s + BNE LBB2_2 + WORD $0xd37ced29 // lsl x9, x9, #4 + WORD $0x8b090000 // add x0, x0, x9 + WORD $0x8b090021 // add x1, x1, x9 + B LBB2_5 + +LBB2_4: + WORD $0x6f00e402 // movi v2.2d, #0000000000000000 + WORD $0x6f00e401 // movi v1.2d, #0000000000000000 + +LBB2_5: + WORD $0x4e0c0423 // dup v3.4s, v1.s[1] + WORD $0x4e0c0444 // dup v4.4s, v2.s[1] + WORD $0xf100011f // cmp x8, #0 + WORD $0x4e0c0405 // dup v5.4s, v0.s[1] + WORD $0x4e140426 // dup v6.4s, v1.s[2] + WORD $0x4e140447 // dup v7.4s, v2.s[2] + WORD $0x4e140410 // dup v16.4s, v0.s[2] + WORD $0x4e23d423 // fadd v3.4s, v1.4s, v3.4s + WORD $0x4e24d444 // fadd v4.4s, v2.4s, v4.4s + WORD $0x4e1c0421 // dup v1.4s, v1.s[3] + WORD $0x4e25d405 // fadd v5.4s, v0.4s, v5.4s + WORD $0x4e1c0442 // dup v2.4s, v2.s[3] + WORD $0x4e1c0400 // dup v0.4s, v0.s[3] + WORD $0x4e23d4c3 // fadd v3.4s, v6.4s, v3.4s + WORD $0x4e24d4e4 // fadd v4.4s, v7.4s, v4.4s + WORD $0x4e25d605 // fadd v5.4s, v16.4s, v5.4s + WORD $0x4e23d421 // fadd v1.4s, v1.4s, v3.4s + WORD $0x4e24d442 // fadd v2.4s, v2.4s, v4.4s + WORD $0x4e25d400 // fadd v0.4s, v0.4s, v5.4s + WORD $0xbd000061 // str s1, [x3] + WORD $0xbd000082 // str s2, [x4] + WORD $0xbd0000a0 // str s0, [x5] + BLE LBB2_8 + WORD $0xbd400000 // ldr s0, [x0] + WORD $0xbd400021 // ldr s1, [x1] + WORD $0xf100051f // cmp x8, #1 + WORD $0xbd400062 // ldr s2, [x3] + WORD $0x1f010800 // fmadd s0, s0, s1, s2 + WORD $0xbd000060 // str s0, [x3] + WORD $0xbd400000 // ldr s0, [x0] + WORD $0xbd400081 // ldr s1, [x4] + WORD $0x1f000400 // fmadd s0, s0, s0, s1 + WORD $0xbd000080 // str s0, [x4] + WORD $0xbd400020 // ldr s0, [x1] + WORD $0xbd4000a1 // ldr s1, [x5] + WORD $0x1f000400 // fmadd s0, s0, s0, s1 + WORD $0xbd0000a0 // str s0, [x5] + BEQ LBB2_8 + WORD $0xbd400400 // ldr s0, [x0, #4] + WORD $0xbd400421 // ldr s1, [x1, #4] + WORD $0xf100091f // cmp x8, #2 + WORD $0xbd400062 // ldr s2, [x3] + WORD $0x1f010800 // fmadd s0, s0, s1, s2 + WORD $0xbd000060 // str s0, [x3] + WORD $0xbd400400 // ldr s0, [x0, #4] + WORD $0xbd400081 // ldr s1, [x4] + WORD $0x1f000400 // fmadd s0, s0, s0, s1 + WORD $0xbd000080 // str s0, [x4] + WORD $0xbd400420 // ldr s0, [x1, #4] + WORD $0xbd4000a1 // ldr s1, [x5] + WORD $0x1f000400 // fmadd s0, s0, s0, s1 + WORD $0xbd0000a0 // str s0, [x5] + BNE LBB2_9 + +LBB2_8: + RET + +LBB2_9: + WORD $0xbd400800 // ldr s0, [x0, #8] + WORD $0xbd400821 // ldr s1, [x1, #8] + WORD $0xbd400062 // ldr s2, [x3] + WORD $0x1f010800 // fmadd s0, s0, s1, s2 + WORD $0xbd000060 // str s0, [x3] + WORD $0xbd400800 // ldr s0, [x0, #8] + WORD $0xbd400081 // ldr s1, [x4] + WORD $0x1f000400 // fmadd s0, s0, s0, s1 + WORD $0xbd000080 // str s0, [x4] + WORD $0xbd400820 // ldr s0, [x1, #8] + WORD $0xbd4000a1 // ldr s1, [x5] + WORD $0x1f000400 // fmadd s0, s0, s0, s1 + WORD $0xbd0000a0 // str s0, [x5] + RET diff --git a/internal/floats/floats_test.go b/internal/floats/floats_test.go new file mode 100644 index 0000000..d96643d --- /dev/null +++ b/internal/floats/floats_test.go @@ -0,0 +1,156 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package floats + +import ( + "fmt" + "math" + "math/rand" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestDistanceKernelsMatchFloat32Oracle(t *testing.T) { + t.Parallel() + + for _, dimension := range []int{1, 3, 4, 7, 8, 15, 16, 17, 127, 128, 129, 768, 1536} { + t.Run(fmt.Sprintf("dimension_%d", dimension), func(t *testing.T) { + random := rand.New(rand.NewSource(int64(dimension))) + left := make([]float32, dimension+1) + right := make([]float32, dimension+1) + for index := 1; index <= dimension; index++ { + left[index] = random.Float32()*2 - 1 + right[index] = random.Float32()*2 - 1 + } + // Offset both slices to exercise unaligned addresses. + left = left[1:] + right = right[1:] + + wantL2, wantDot, wantLeftNorm, wantRightNorm := distanceOracle(left, right) + requireFloat32Close(t, wantL2, L2Squared(left, right)) + requireFloat32Close(t, wantDot, InnerProduct(left, right)) + dot, leftNorm, rightNorm := DotNorms(left, right) + requireFloat32Close(t, wantDot, dot) + requireFloat32Close(t, wantLeftNorm, leftNorm) + requireFloat32Close(t, wantRightNorm, rightNorm) + }) + } +} + +func TestDistanceKernelsUseFloat32Accumulation(t *testing.T) { + t.Parallel() + + // At float32 precision, adding one after 2^24 no longer changes the sum. + left := []float32{4096, 1} + right := []float32{4096, 1} + require.Equal(t, float32(1<<24), InnerProduct(left, right)) + dot, leftNorm, rightNorm := DotNorms(left, right) + require.Equal(t, float32(1<<24), dot) + require.Equal(t, float32(1<<24), leftNorm) + require.Equal(t, float32(1<<24), rightNorm) +} + +func TestDistanceKernelsDoNotAllocateOrMutate(t *testing.T) { + left := []float32{0.2, 0.9, -0.4, 0.7} + right := []float32{0.3, 0.5, 0.8, -0.1} + leftCopy := append([]float32(nil), left...) + rightCopy := append([]float32(nil), right...) + + require.Zero(t, testing.AllocsPerRun(100, func() { + benchmarkL2 = L2Squared(left, right) + benchmarkInnerProduct = InnerProduct(left, right) + benchmarkDot, benchmarkLeftNorm, benchmarkRightNorm = DotNorms(left, right) + })) + require.Equal(t, leftCopy, left) + require.Equal(t, rightCopy, right) +} + +func BenchmarkDistanceKernels(b *testing.B) { + for _, dimension := range []int{128, 768, 1536} { + left := make([]float32, dimension) + right := make([]float32, dimension) + for index := range left { + left[index] = float32(index%17)/17 - 0.5 + right[index] = float32(index%23)/23 - 0.5 + } + b.Run(fmt.Sprintf("L2/%d", dimension), func(b *testing.B) { + for b.Loop() { + benchmarkL2 = L2Squared(left, right) + } + }) + b.Run(fmt.Sprintf("InnerProduct/%d", dimension), func(b *testing.B) { + for b.Loop() { + benchmarkDot = InnerProduct(left, right) + } + }) + b.Run(fmt.Sprintf("DotNorms/%d", dimension), func(b *testing.B) { + for b.Loop() { + benchmarkDot, benchmarkLeftNorm, benchmarkRightNorm = DotNorms(left, right) + } + }) + } +} + +func testArchitectureKernels( + t *testing.T, + l2Kernel binaryKernel, + dotKernel binaryKernel, + productsKernel productsKernel, +) { + t.Helper() + for _, dimension := range []int{1, 3, 4, 7, 8, 15, 16, 17, 127, 128, 129} { + left := make([]float32, dimension+1) + right := make([]float32, dimension+1) + for index := 1; index <= dimension; index++ { + left[index] = float32((index*7)%19)/19 - 0.5 + right[index] = float32((index*11)%23)/23 - 0.5 + } + left, right = left[1:], right[1:] + wantL2, wantDot, wantLeftNorm, wantRightNorm := distanceOracle(left, right) + requireFloat32Close(t, wantL2, l2Kernel(left, right)) + requireFloat32Close(t, wantDot, dotKernel(left, right)) + dot, leftNorm, rightNorm := productsKernel(left, right) + requireFloat32Close(t, wantDot, dot) + requireFloat32Close(t, wantLeftNorm, leftNorm) + requireFloat32Close(t, wantRightNorm, rightNorm) + } +} + +func distanceOracle(left, right []float32) (l2, dot, leftNorm, rightNorm float32) { + for index, leftValue := range left { + rightValue := right[index] + difference := leftValue - rightValue + l2 += difference * difference + dot += leftValue * rightValue + leftNorm += leftValue * leftValue + rightNorm += rightValue * rightValue + } + return +} + +func requireFloat32Close(t *testing.T, expected, actual float32) { + t.Helper() + tolerance := float32(1e-5) * max(1, float32(math.Abs(float64(expected)))) + require.InDelta(t, expected, actual, float64(tolerance)) +} + +var ( + benchmarkL2 float32 + benchmarkInnerProduct float32 + benchmarkDot float32 + benchmarkLeftNorm float32 + benchmarkRightNorm float32 +) diff --git a/internal/floats/src/.gitignore b/internal/floats/src/.gitignore new file mode 100644 index 0000000..d897b56 --- /dev/null +++ b/internal/floats/src/.gitignore @@ -0,0 +1,2 @@ +*.o +*.s diff --git a/internal/floats/src/floats_avx.c b/internal/floats/src/floats_avx.c new file mode 100644 index 0000000..e3f87ce --- /dev/null +++ b/internal/floats/src/floats_avx.c @@ -0,0 +1,90 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include + +static inline float reduce256(__m256 value) { + float partial[8]; + _mm256_storeu_ps(partial, value); + float sum = 0; + for (int index = 0; index < 8; index++) { + sum += partial[index]; + } + return sum; +} + +void xvec_avx_l2_squared(float *left, float *right, int64_t size, float *output) { + int64_t vectors = size / 8; + int64_t remain = size % 8; + __m256 sum = _mm256_setzero_ps(); + for (int64_t index = 0; index < vectors; index++) { + __m256 left_value = _mm256_loadu_ps(left); + __m256 right_value = _mm256_loadu_ps(right); + __m256 difference = _mm256_sub_ps(left_value, right_value); + sum = _mm256_add_ps(sum, _mm256_mul_ps(difference, difference)); + left += 8; + right += 8; + } + float result = reduce256(sum); + for (int64_t index = 0; index < remain; index++) { + float difference = left[index] - right[index]; + result += difference * difference; + } + *output = result; +} + +void xvec_avx_inner_product(float *left, float *right, int64_t size, float *output) { + int64_t vectors = size / 8; + int64_t remain = size % 8; + __m256 sum = _mm256_setzero_ps(); + for (int64_t index = 0; index < vectors; index++) { + __m256 left_value = _mm256_loadu_ps(left); + __m256 right_value = _mm256_loadu_ps(right); + sum = _mm256_add_ps(sum, _mm256_mul_ps(left_value, right_value)); + left += 8; + right += 8; + } + float result = reduce256(sum); + for (int64_t index = 0; index < remain; index++) { + result += left[index] * right[index]; + } + *output = result; +} + +void xvec_avx_dot_norms(float *left, float *right, int64_t size, + float *dot, float *left_norm, float *right_norm) { + int64_t vectors = size / 8; + int64_t remain = size % 8; + __m256 dot_sum = _mm256_setzero_ps(); + __m256 left_sum = _mm256_setzero_ps(); + __m256 right_sum = _mm256_setzero_ps(); + for (int64_t index = 0; index < vectors; index++) { + __m256 left_value = _mm256_loadu_ps(left); + __m256 right_value = _mm256_loadu_ps(right); + dot_sum = _mm256_add_ps(dot_sum, _mm256_mul_ps(left_value, right_value)); + left_sum = _mm256_add_ps(left_sum, _mm256_mul_ps(left_value, left_value)); + right_sum = _mm256_add_ps(right_sum, _mm256_mul_ps(right_value, right_value)); + left += 8; + right += 8; + } + *dot = reduce256(dot_sum); + *left_norm = reduce256(left_sum); + *right_norm = reduce256(right_sum); + for (int64_t index = 0; index < remain; index++) { + *dot += left[index] * right[index]; + *left_norm += left[index] * left[index]; + *right_norm += right[index] * right[index]; + } +} diff --git a/internal/floats/src/floats_avx512.c b/internal/floats/src/floats_avx512.c new file mode 100644 index 0000000..f3e26d9 --- /dev/null +++ b/internal/floats/src/floats_avx512.c @@ -0,0 +1,90 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include + +static inline float reduce512(__m512 value) { + float partial[16]; + _mm512_storeu_ps(partial, value); + float sum = 0; + for (int index = 0; index < 16; index++) { + sum += partial[index]; + } + return sum; +} + +void xvec_avx512_l2_squared(float *left, float *right, int64_t size, float *output) { + int64_t vectors = size / 16; + int64_t remain = size % 16; + __m512 sum = _mm512_setzero_ps(); + for (int64_t index = 0; index < vectors; index++) { + __m512 left_value = _mm512_loadu_ps(left); + __m512 right_value = _mm512_loadu_ps(right); + __m512 difference = _mm512_sub_ps(left_value, right_value); + sum = _mm512_fmadd_ps(difference, difference, sum); + left += 16; + right += 16; + } + float result = reduce512(sum); + for (int64_t index = 0; index < remain; index++) { + float difference = left[index] - right[index]; + result += difference * difference; + } + *output = result; +} + +void xvec_avx512_inner_product(float *left, float *right, int64_t size, float *output) { + int64_t vectors = size / 16; + int64_t remain = size % 16; + __m512 sum = _mm512_setzero_ps(); + for (int64_t index = 0; index < vectors; index++) { + __m512 left_value = _mm512_loadu_ps(left); + __m512 right_value = _mm512_loadu_ps(right); + sum = _mm512_fmadd_ps(left_value, right_value, sum); + left += 16; + right += 16; + } + float result = reduce512(sum); + for (int64_t index = 0; index < remain; index++) { + result += left[index] * right[index]; + } + *output = result; +} + +void xvec_avx512_dot_norms(float *left, float *right, int64_t size, + float *dot, float *left_norm, float *right_norm) { + int64_t vectors = size / 16; + int64_t remain = size % 16; + __m512 dot_sum = _mm512_setzero_ps(); + __m512 left_sum = _mm512_setzero_ps(); + __m512 right_sum = _mm512_setzero_ps(); + for (int64_t index = 0; index < vectors; index++) { + __m512 left_value = _mm512_loadu_ps(left); + __m512 right_value = _mm512_loadu_ps(right); + dot_sum = _mm512_fmadd_ps(left_value, right_value, dot_sum); + left_sum = _mm512_fmadd_ps(left_value, left_value, left_sum); + right_sum = _mm512_fmadd_ps(right_value, right_value, right_sum); + left += 16; + right += 16; + } + *dot = reduce512(dot_sum); + *left_norm = reduce512(left_sum); + *right_norm = reduce512(right_sum); + for (int64_t index = 0; index < remain; index++) { + *dot += left[index] * right[index]; + *left_norm += left[index] * left[index]; + *right_norm += right[index] * right[index]; + } +} diff --git a/internal/floats/src/floats_neon.c b/internal/floats/src/floats_neon.c new file mode 100644 index 0000000..dd60bf5 --- /dev/null +++ b/internal/floats/src/floats_neon.c @@ -0,0 +1,86 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include + +static inline float reduce_neon(float32x4_t value) { + float partial[4]; + vst1q_f32(partial, value); + return partial[0] + partial[1] + partial[2] + partial[3]; +} + +void xvec_neon_l2_squared(float *left, float *right, int64_t size, float *output) { + int64_t vectors = size / 4; + int64_t remain = size % 4; + float32x4_t sum = vdupq_n_f32(0); + for (int64_t index = 0; index < vectors; index++) { + float32x4_t left_value = vld1q_f32(left); + float32x4_t right_value = vld1q_f32(right); + float32x4_t difference = vsubq_f32(left_value, right_value); + sum = vmlaq_f32(sum, difference, difference); + left += 4; + right += 4; + } + float result = reduce_neon(sum); + for (int64_t index = 0; index < remain; index++) { + float difference = left[index] - right[index]; + result += difference * difference; + } + *output = result; +} + +void xvec_neon_inner_product(float *left, float *right, int64_t size, float *output) { + int64_t vectors = size / 4; + int64_t remain = size % 4; + float32x4_t sum = vdupq_n_f32(0); + for (int64_t index = 0; index < vectors; index++) { + float32x4_t left_value = vld1q_f32(left); + float32x4_t right_value = vld1q_f32(right); + sum = vmlaq_f32(sum, left_value, right_value); + left += 4; + right += 4; + } + float result = reduce_neon(sum); + for (int64_t index = 0; index < remain; index++) { + result += left[index] * right[index]; + } + *output = result; +} + +void xvec_neon_dot_norms(float *left, float *right, int64_t size, + float *dot, float *left_norm, float *right_norm) { + int64_t vectors = size / 4; + int64_t remain = size % 4; + float32x4_t dot_sum = vdupq_n_f32(0); + float32x4_t left_sum = vdupq_n_f32(0); + float32x4_t right_sum = vdupq_n_f32(0); + for (int64_t index = 0; index < vectors; index++) { + float32x4_t left_value = vld1q_f32(left); + float32x4_t right_value = vld1q_f32(right); + dot_sum = vmlaq_f32(dot_sum, left_value, right_value); + left_sum = vmlaq_f32(left_sum, left_value, left_value); + right_sum = vmlaq_f32(right_sum, right_value, right_value); + left += 4; + right += 4; + } + *dot = reduce_neon(dot_sum); + *left_norm = reduce_neon(left_sum); + *right_norm = reduce_neon(right_sum); + for (int64_t index = 0; index < remain; index++) { + *dot += left[index] * right[index]; + *left_norm += left[index] * left[index]; + *right_norm += right[index] * right[index]; + } +}