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" diff --git a/src/abstractarray.jl b/src/abstractarray.jl index b77514e..e1a6d67 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() || β === false + scale!(y, x, convert(T, α)) else LinearAlgebra.axpby!(convert(T, α), x, convert(T, β), y) end 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/complicated.jl b/test/complicated.jl index ee7cc17..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) @@ -146,6 +150,26 @@ 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)) + + α = 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 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 diff --git a/test/simple.jl b/test/simple.jl index 9247e40..466e7e2 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) @@ -143,6 +145,26 @@ 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(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) # the BLAS call actually also forces strong zero even for 0.0 + α, β = randn(ComplexF64, 2) @test_throws InexactError add!(deepcopy(y), xcopy, α) @test_throws InexactError add!(deepcopy(y), xcopy, α, β)