From eaa043879155e30555fe8af9328de61e2498f79c Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sat, 19 Sep 2026 10:56:18 +0200 Subject: [PATCH 1/9] Respect strong zero in axpby --- src/abstractarray.jl | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/abstractarray.jl b/src/abstractarray.jl index b77514e..57a6e4e 100644 --- a/src/abstractarray.jl +++ b/src/abstractarray.jl @@ -77,6 +77,8 @@ function add!( ) where {T <: BlasFloat} if β === One() LinearAlgebra.axpy!(convert(T, α), x, y) + elseif β === Zero() + scale!(y, x, convert(T, α)) else LinearAlgebra.axpby!(convert(T, α), x, convert(T, β), y) end From b18b799c95eb847585a1aad1a1484ded9a601fb0 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sat, 19 Sep 2026 12:10:54 +0200 Subject: [PATCH 2/9] false is also a strong zero --- src/abstractarray.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/abstractarray.jl b/src/abstractarray.jl index 57a6e4e..e1a6d67 100644 --- a/src/abstractarray.jl +++ b/src/abstractarray.jl @@ -77,7 +77,7 @@ function add!( ) where {T <: BlasFloat} if β === One() LinearAlgebra.axpy!(convert(T, α), x, y) - elseif β === Zero() + elseif β === Zero() || β === false scale!(y, x, convert(T, α)) else LinearAlgebra.axpby!(convert(T, α), x, convert(T, β), y) From cd2d212b8533c9d737ef730e5678cbf3dccebb82 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sat, 19 Sep 2026 16:31:04 +0200 Subject: [PATCH 3/9] Add test --- test/simple.jl | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/test/simple.jl b/test/simple.jl index 9247e40..e7dec9f 100644 --- a/test/simple.jl +++ b/test/simple.jl @@ -143,6 +143,13 @@ end @test deepcollect(z5) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* β)) @test all(deepcollect(xcopy) .== deepcollect(x)) + # test strong zero + α = randn(ComplexF64) + z6 = @constinferred add(y, x, α, Zero()) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* Zero())) + z6 = @constinferred add(y, x, α, false) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* false)) + α, β = randn(ComplexF64, 2) @test_throws InexactError add!(deepcopy(y), xcopy, α) @test_throws InexactError add!(deepcopy(y), xcopy, α, β) From 3f247681ac07e647ba58a9d1a31bc8e07b4c6d7d Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sat, 19 Sep 2026 17:22:07 +0200 Subject: [PATCH 4/9] Add test to complicated as well --- test/complicated.jl | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/test/complicated.jl b/test/complicated.jl index ee7cc17..472c914 100644 --- a/test/complicated.jl +++ b/test/complicated.jl @@ -146,6 +146,13 @@ end @test deepcollect(z5) ≈ (muladd.(deepcollect(x), α, deepcollect(y))) z5 = @constinferred add!!(deepcopy(y), deepcopy(x), α, β) @test deepcollect(z5) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* β)) + + # test strong zero + α = randn(ComplexF64) + z6 = @constinferred add(y, x, α, Zero()) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* Zero())) + z6 = @constinferred add(y, x, α, false) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* false)) end @testset "inner" begin From a3dbd28b6acc3422f46d57a54dbd387786421994 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 20 Sep 2026 08:46:17 -0400 Subject: [PATCH 5/9] Add some NaN tests --- test/Project.toml | 1 + test/jlarray.jl | 129 ++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 130 insertions(+) create mode 100644 test/jlarray.jl diff --git a/test/Project.toml b/test/Project.toml index 89ea21a..94fe686 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -6,6 +6,7 @@ ChainRulesCore = "d360d2e6-b24c-11e9-a2a3-2a2ae2dbcce4" ChainRulesTestUtils = "cdddcdb0-9152-4a09-a978-84456f9df70a" Enzyme = "7da242da-08ed-463a-9acd-ee780be4f1d9" EnzymeTestUtils = "12d8515a-0907-448a-8884-5fe00fdf1c5a" +JLArrays = "27aeb0d3-9eb9-45fb-866b-73c2ecf80fcb" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6" ParallelTestRunner = "d3525ed8-44d0-4b2c-a655-542cee43accc" diff --git a/test/jlarray.jl b/test/jlarray.jl new file mode 100644 index 0000000..d423495 --- /dev/null +++ b/test/jlarray.jl @@ -0,0 +1,129 @@ +module StaticSVec +using VectorInterface +using JLArrays +using Test +using TestExtras + +deepcollect(x::JLArray) = collect(x) +deepcollect(x::Number) = x + +x = JLVector(randn(3)) +y = JLVector(randn(3)) +nan_y = JLVector(vcat(randn(2), NaN)) + +@testset "scalartype" begin + s = @constinferred scalartype(x) + @test s == Float64 +end + +@testset "zerovector" begin + z = @constinferred zerovector(x) + @test z isa JLVector{Float64} + @test all(iszero, deepcollect(z)) + @test all(deepcollect(z) .=== zero(scalartype(x))) + z1 = @constinferred zerovector!!(x) + @test z1 isa JLVector{Float64} + @test all(deepcollect(z1) .=== zero(scalartype(x))) + + z3 = @constinferred zerovector(x, ComplexF64) + @test z3 isa JLVector{ComplexF64} + @test all(deepcollect(z3) .=== zero(ComplexF64)) + z4 = @constinferred zerovector!!(x, ComplexF64) + @test z4 isa JLVector{ComplexF64} + @test all(deepcollect(z4) .=== zero(ComplexF64)) +end + +@testset "scale" begin + α = randn() + z = @constinferred scale(x, α) + @test z isa JLVector{Float64} + @test all(deepcollect(z) .== α .* deepcollect(x)) + + z2 = @constinferred scale!!(x, α) + @test z2 isa JLVector{Float64} + @test deepcollect(z2) ≈ (α .* deepcollect(x)) + z2 = @constinferred scale!!(y, x, α) + @test z2 isa JLVector{Float64} + @test deepcollect(z2) ≈ (α .* deepcollect(x)) + + α = randn(ComplexF64) + z4 = @constinferred scale(x, α) + @test z4 isa JLVector{ComplexF64} + @test deepcollect(z4) ≈ (α .* deepcollect(x)) + z5 = @constinferred scale!!(x, α) + @test z5 isa JLVector{ComplexF64} + @test deepcollect(z5) ≈ (α .* deepcollect(x)) + + z6 = @constinferred scale!!(zerovector(x), x, α) + @test z6 isa JLVector{ComplexF64} + @test deepcollect(z6) ≈ (α .* deepcollect(x)) + + ycomplex = zerovector(y, ComplexF64) + α = randn(Float64) + z8 = @constinferred scale!!(ycomplex, x, α) + @test scalartype(z8) == ComplexF64 + @test all(deepcollect(z8) .== α .* deepcollect(x)) +end + +@testset "add" begin + α, β = randn(2) + z = add(y, x) + @test z isa JLVector{Float64} + @test all(deepcollect(z) .== deepcollect(x) .+ deepcollect(y)) + z = add(y, x, α) + @test deepcollect(z) ≈ muladd.(deepcollect(x), α, deepcollect(y)) + z = add(y, x, α, β) + @test deepcollect(z) ≈ muladd.(deepcollect(x), α, deepcollect(y) .* β) + + z2 = @constinferred add!!(y, x) + @test z2 isa JLVector{Float64} + @test deepcollect(z2) ≈ (deepcollect(x) .+ deepcollect(y)) + z2 = @constinferred add!!(y, x, α) + @test deepcollect(z2) ≈ (muladd.(deepcollect(x), α, deepcollect(y))) + z2 = @constinferred add!!(y, x, α, β) + @test deepcollect(z2) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* β)) + + α, β = randn(ComplexF64, 2) + z4 = add(y, x, α) + @test z4 isa JLVector{ComplexF64} + @test deepcollect(z4) ≈ (muladd.(deepcollect(x), α, deepcollect(y))) + z4 = add(y, x, α, β) + @test deepcollect(z4) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* β)) + + z5 = @constinferred add!!(y, x, α) + @test z5 isa JLVector{ComplexF64} + @test deepcollect(z5) ≈ (muladd.(deepcollect(x), α, deepcollect(y))) + z5 = @constinferred add!!(y, x, α, β) + @test deepcollect(z5) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* β)) + + # test strong zero + α = randn(ComplexF64) + z6 = add(y, x, α, Zero()) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* Zero())) + z6 = add(y, x, α, false) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* false)) + + α = randn(scalartype(x)) + z6 = deepcopy(nan_y) + z6 = @constinferred add!(z6, x, α, Zero()) + @test !any(isnan, z6) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(nan_y) .* Zero())) + z6 = deepcopy(nan_y) + z6 = @constinferred add!(z6, x, α, false) + @test !any(isnan, z6) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(nan_y) .* false)) + z6 = deepcopy(nan_y) + z6 = @constinferred add!(z6, x, α, 0.0) + @test any(isnan, z6) +end + +@testset "inner" begin + s = @constinferred inner(x, y) + @test s ≈ inner(deepcollect(x), deepcollect(y)) + + α, β = randn(ComplexF64, 2) + s2 = @constinferred inner(scale(x, α), scale(y, β)) + @test s2 ≈ inner(α * deepcollect(x), β * deepcollect(y)) +end + +end From dd6f58b8056293d92b245d194bd43565c4c0c816 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Sun, 20 Sep 2026 14:43:48 -0400 Subject: [PATCH 6/9] Check the strong zero in the non-JLArrays simple case too --- test/simple.jl | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/test/simple.jl b/test/simple.jl index e7dec9f..fb7bc20 100644 --- a/test/simple.jl +++ b/test/simple.jl @@ -7,6 +7,8 @@ deepcollect(x::Number) = x x = randn(3, 3, 3) y = randn(3, 3, 3) +nan_y = randn(3, 3, 3) +nan_y[end] = NaN @testset "scalartype" begin s = @constinferred scalartype(x) @@ -150,6 +152,19 @@ end z6 = @constinferred add(y, x, α, false) @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* false)) + α = randn(scalartype(x)) + z6 = deepcopy(nan_y) + z6 = @constinferred add!(z6, x, α, Zero()) + @test !any(isnan, z6) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(nan_y) .* Zero())) + z6 = deepcopy(nan_y) + z6 = @constinferred add!(z6, x, α, false) + @test !any(isnan, z6) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(nan_y) .* false)) + z6 = deepcopy(nan_y) + z6 = @constinferred add!(z6, x, α, 0.0) + @test any(isnan, z6) + α, β = randn(ComplexF64, 2) @test_throws InexactError add!(deepcopy(y), xcopy, α) @test_throws InexactError add!(deepcopy(y), xcopy, α, β) From f49b8ef6190efde11b6b1f161b4b3bb2cdbde7b3 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 21 Sep 2026 01:24:16 -0400 Subject: [PATCH 7/9] One last fix --- test/simple.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/simple.jl b/test/simple.jl index fb7bc20..466e7e2 100644 --- a/test/simple.jl +++ b/test/simple.jl @@ -163,7 +163,7 @@ end @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(nan_y) .* false)) z6 = deepcopy(nan_y) z6 = @constinferred add!(z6, x, α, 0.0) - @test any(isnan, z6) + @test !any(isnan, z6) # the BLAS call actually also forces strong zero even for 0.0 α, β = randn(ComplexF64, 2) @test_throws InexactError add!(deepcopy(y), xcopy, α) From 3c47539ce437919e58af99c6ccdb718d40dc52c4 Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 21 Sep 2026 01:35:37 -0400 Subject: [PATCH 8/9] NaN tests for complicated too --- test/complicated.jl | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/test/complicated.jl b/test/complicated.jl index 472c914..cde3854 100644 --- a/test/complicated.jl +++ b/test/complicated.jl @@ -15,6 +15,10 @@ y = ( NamedTuple{(:x, :y)}.(collect(zip(randn(2, 2), rand(2, 2)))), (randn(), randn(3), randn(2, 2)'), randn(), (view(randn(4, 4), 1:2, [1, 3, 4]),), ) +nan_y = ( + NamedTuple{(:x, :y)}.(collect(zip(randn(2, 2), rand(2, 2)))), + (NaN, randn(3), randn(2, 2)'), randn(), (view(randn(4, 4), 1:2, [1, 3, 4]),), +) @testset "scalartype" begin s = @constinferred scalartype(x) @@ -153,6 +157,19 @@ end @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* Zero())) z6 = @constinferred add(y, x, α, false) @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(y) .* false)) + + α = randn(scalartype(x)) + z6 = deepcopy(nan_y) + z6 = @constinferred add(z6, x, α, Zero()) + @test !any(isnan, deepcollect(z6)) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(nan_y) .* Zero())) + z6 = deepcopy(nan_y) + z6 = @constinferred add(z6, x, α, false) + @test !any(isnan, deepcollect(z6)) + @test deepcollect(z6) ≈ (muladd.(deepcollect(x), α, deepcollect(nan_y) .* false)) + z6 = deepcopy(nan_y) + z6 = @constinferred add(z6, x, α, 0.0) + @test !any(isnan, deepcollect(z6)) # underlying scale! call forces strong zero end @testset "inner" begin From 3d19abe96e3e4db461780d129cde7b21ef1317ed Mon Sep 17 00:00:00 2001 From: Katharine Hyatt Date: Mon, 21 Sep 2026 15:52:15 +0200 Subject: [PATCH 9/9] Bump patch version --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index 47fee58..901f19c 100644 --- a/Project.toml +++ b/Project.toml @@ -1,7 +1,7 @@ name = "VectorInterface" uuid = "409d34a3-91d5-4945-b6ec-7529ddf182d8" authors = ["Jutho Haegeman and contributors"] -version = "0.6.0" +version = "0.6.1" [deps] LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e"