Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
name = "VectorInterface"
uuid = "409d34a3-91d5-4945-b6ec-7529ddf182d8"
authors = ["Jutho Haegeman <jutho.haegeman@ugent.be> and contributors"]
version = "0.6.0"
version = "0.6.1"

[deps]
LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e"
Expand Down
2 changes: 2 additions & 0 deletions src/abstractarray.jl
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,8 @@ function add!(
) where {T <: BlasFloat}
if β === One()
LinearAlgebra.axpy!(convert(T, α), x, y)
elseif β === Zero() || β === false

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'm slightly confused you also added the === false branch here, does that effectively mean that we want to interpret false as a strong zero as well?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In Julia false is always strong zero, I thought?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

julia> 0.0 * NaN
NaN

julia> false * NaN
0.0

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I did not know this 😆 I thought that was a BLAS-specific thing... In that case, ignore my comments :p

A different thing that just popped in my brain is whether it would make sense to replace these checks everywhere with isstrongzero and isstrongone, to enforce that this is handled consistently?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think the strong zero check is only made here, so I don't know if it makes sense to create a whole new function for it

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The reason I was thinking about this is because we also use this downstream every now and again, e.g. https://github.com/QuantumKitHub/TensorOperations.jl/blob/ef937adcfa4e6bc4476f8197a62a2da45cf3b1d9/src/implementation/strided.jl#L142 where we actually turn every zero into a strong one, and I do like the idea of having this a bit more formalized/consistent. Does not have to be in this PR though!

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I might just merge this and tag then to keep stuff moving for the other PR

scale!(y, x, convert(T, α))
else
LinearAlgebra.axpby!(convert(T, α), x, convert(T, β), y)
end
Expand Down
1 change: 1 addition & 0 deletions test/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
24 changes: 24 additions & 0 deletions test/complicated.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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))
Comment thread
kshyatt marked this conversation as resolved.

α = 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
Expand Down
129 changes: 129 additions & 0 deletions test/jlarray.jl
Original file line number Diff line number Diff line change
@@ -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
22 changes: 22 additions & 0 deletions test/simple.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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))
Comment thread
kshyatt marked this conversation as resolved.

α = 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, α, β)
Expand Down
Loading