From 45b5779d4e1b83eb7f28e616acefa49c912ae9ac Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Sun, 30 Aug 2026 08:29:03 -0400 Subject: [PATCH 1/2] Preserve Tracker gradients in GPUArraysCore restructure `restructure` on a GPUArraysCore target adapted the TrackedArray source through `Tracker.adapt_structure`, which calls `param()` and starts a new tape leaf. The forward value stayed tracked but the source gradient was silently zero. Reshape the TrackedArray instead, matching the existing Array/TrackedArray methods. Fixes https://github.com/JuliaArrays/ArrayInterface.jl/issues/504 Co-Authored-By: Chris Rackauckas Co-Authored-By: Grok Agent-Harness: Grok CLI 1.0.13 Agent-Model: grok-4.6 Agent-Session: 01a05298-11d9-7a71-89a3-e977fe6a9a3e --- Project.toml | 7 +++++-- ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl | 13 +++++++++++++ test/ad.jl | 19 ++++++++++++++++++- 3 files changed, 36 insertions(+), 3 deletions(-) create mode 100644 ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl diff --git a/Project.toml b/Project.toml index ac60df81..481b12b5 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "ArrayInterface" uuid = "4fba245c-0d91-5ea0-9b3e-6abc04ee57a9" -version = "7.30.0" +version = "7.30.1" [deps] Adapt = "79e6a3ab-5dfb-504d-930d-738a2a938a0e" @@ -32,6 +32,7 @@ ArrayInterfaceChainRulesCoreExt = "ChainRulesCore" ArrayInterfaceChainRulesExt = "ChainRules" ArrayInterfaceFillArraysExt = "FillArrays" ArrayInterfaceGPUArraysCoreExt = "GPUArraysCore" +ArrayInterfaceGPUArraysCoreTrackerExt = ["GPUArraysCore", "Tracker"] ArrayInterfaceMetalExt = "Metal" ArrayInterfaceReverseDiffExt = "ReverseDiff" ArrayInterfaceSparseArraysExt = "SparseArrays" @@ -50,6 +51,7 @@ ChainRulesCore = "1" ChainRulesTestUtils = "1" FillArrays = "1" GPUArraysCore = "0.1, 0.2" +JLArrays = "0.3" LinearAlgebra = "1.10" Metal = "1" ReverseDiff = "1" @@ -68,6 +70,7 @@ ChainRulesTestUtils = "cdddcdb0-9152-4a09-a978-84456f9df70a" ComponentArrays = "b0b7db55-cfe3-40fc-9ded-d10e2dbeff66" FillArrays = "1a297f60-69ca-5386-bcde-b61e274b549b" GPUArraysCore = "46192b85-c4d5-4398-a991-12ede77f4527" +JLArrays = "27aeb0d3-9eb9-45fb-866b-73c2ecf80fcb" JuliaFormatter = "98e50ef6-434e-11e9-1051-2b60c6c9e899" Metal = "dde4c033-4e86-420c-a63e-0dd931031962" Pkg = "44cfe95a-1eb2-52ea-b672-e2afdf69b78f" @@ -82,4 +85,4 @@ Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" Tracker = "9f7883ad-71c0-57eb-9f7f-b5c9e6d3789c" [targets] -test = ["SafeTestsets", "Pkg", "Test", "Aqua", "Random", "SparseArrays", "SuiteSparse", "BandedMatrices", "BlockBandedMatrices", "GPUArraysCore", "StaticArrays", "Tracker", "ReverseDiff", "ChainRules", "FillArrays", "ComponentArrays", "ChainRulesTestUtils"] +test = ["SafeTestsets", "Pkg", "Test", "Aqua", "Random", "SparseArrays", "SuiteSparse", "BandedMatrices", "BlockBandedMatrices", "GPUArraysCore", "JLArrays", "StaticArrays", "Tracker", "ReverseDiff", "ChainRules", "FillArrays", "ComponentArrays", "ChainRulesTestUtils"] diff --git a/ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl b/ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl new file mode 100644 index 00000000..5f8fef59 --- /dev/null +++ b/ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl @@ -0,0 +1,13 @@ +module ArrayInterfaceGPUArraysCoreTrackerExt + +using ArrayInterface +import GPUArraysCore +import Tracker + +# Tracker.adapt_structure uses `param(adapt(T, data(xs)))`, which severs the tape. +function ArrayInterface.restructure( + x::GPUArraysCore.AbstractGPUArray, y::Tracker.TrackedArray) + reshape(y, Base.size(x)...) +end + +end diff --git a/test/ad.jl b/test/ad.jl index aa2bb03a..95102ce6 100644 --- a/test/ad.jl +++ b/test/ad.jl @@ -1,4 +1,4 @@ -using ArrayInterface, ReverseDiff, Tracker, Test +using ArrayInterface, ReverseDiff, Tracker, Test, JLArrays x = ReverseDiff.track([4.0]) @test ArrayInterface.aos_to_soa(x) isa ReverseDiff.TrackedArray x = reshape([ReverseDiff.track(rand(1, 1, 1))[1]], 1, 1, 1) @@ -51,3 +51,20 @@ x = rand(4) @test ArrayInterface.restructure(x, y) isa Array @test eltype(ArrayInterface.restructure(x, y)) <: ReverseDiff.TrackedReal @test size(ArrayInterface.restructure(x, y)) == (4,) + +@testset "restructure GPUArraysCore + Tracker" begin + target = JLArray(reshape(Float32.(1:6), 2, 3)) + src = JLArray(copy(vec(Array(target)))) + + yr = ArrayInterface.restructure(target, src) + @test yr isa JLArray + @test size(yr) == (2, 3) + @test Array(yr) == reshape(Array(src), 2, 3) + + y, back = Tracker.forward(src) do t + sum(ArrayInterface.restructure(target, t)) + end + dx = only(back(1.0f0)) + @test Tracker.data(y) == 21.0f0 + @test Array(Tracker.data(dx)) == ones(Float32, 6) +end From a5bc7256ba154a9644f89cef7ec00271af38cb1a Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Sun, 30 Aug 2026 08:56:50 -0400 Subject: [PATCH 2/2] Adapt Tracker storage onto the GPU restructure target A reshape-only path preserved the tape but left a CPU TrackedArray on the host when the source was not already a GPUArraysCore array. Adapt through a tracked primitive when the data is not already the target type; skip Adapt when it is, since Tracker.adapt_structure would start a new tape leaf. Co-Authored-By: Chris Rackauckas Co-Authored-By: Grok Agent-Harness: Grok CLI 1.0.13 Agent-Model: grok-4.6 Agent-Session: 01a05298-11d9-7a71-89a3-e977fe6a9a3e --- ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl | 17 ++++++++++++++++- test/ad.jl | 16 +++++++++++++++- 2 files changed, 31 insertions(+), 2 deletions(-) diff --git a/ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl b/ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl index 5f8fef59..d7438cf6 100644 --- a/ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl +++ b/ext/ArrayInterfaceGPUArraysCoreTrackerExt.jl @@ -1,13 +1,28 @@ module ArrayInterfaceGPUArraysCoreTrackerExt +using Adapt using ArrayInterface import GPUArraysCore import Tracker # Tracker.adapt_structure uses `param(adapt(T, data(xs)))`, which severs the tape. +function _adapt_tracked_storage(T, y) + Adapt.adapt(T, y) +end +function _adapt_tracked_storage(T, y::Tracker.TrackedArray) + Tracker.track(_adapt_tracked_storage, T, y) +end +Tracker.@grad function _adapt_tracked_storage(T, y) + ydata = Tracker.data(y) + Adapt.adapt(T, ydata), + Δ -> (nothing, Adapt.adapt(ArrayInterface.parameterless_type(ydata), Tracker.data(Δ))) +end + function ArrayInterface.restructure( x::GPUArraysCore.AbstractGPUArray, y::Tracker.TrackedArray) - reshape(y, Base.size(x)...) + T = ArrayInterface.parameterless_type(x) + yT = Tracker.data(y) isa T ? y : _adapt_tracked_storage(T, y) + reshape(yT, Base.size(x)...) end end diff --git a/test/ad.jl b/test/ad.jl index 95102ce6..e08507db 100644 --- a/test/ad.jl +++ b/test/ad.jl @@ -55,6 +55,7 @@ x = rand(4) @testset "restructure GPUArraysCore + Tracker" begin target = JLArray(reshape(Float32.(1:6), 2, 3)) src = JLArray(copy(vec(Array(target)))) + src_cpu = Array(src) yr = ArrayInterface.restructure(target, src) @test yr isa JLArray @@ -62,9 +63,22 @@ x = rand(4) @test Array(yr) == reshape(Array(src), 2, 3) y, back = Tracker.forward(src) do t - sum(ArrayInterface.restructure(target, t)) + r = ArrayInterface.restructure(target, t) + @test Tracker.data(r) isa JLArray + @test size(r) == (2, 3) + sum(r) end dx = only(back(1.0f0)) @test Tracker.data(y) == 21.0f0 @test Array(Tracker.data(dx)) == ones(Float32, 6) + + y_cpu, back_cpu = Tracker.forward(src_cpu) do t + r = ArrayInterface.restructure(target, t) + @test Tracker.data(r) isa JLArray + @test size(r) == (2, 3) + sum(r) + end + dx_cpu = only(back_cpu(1.0f0)) + @test Tracker.data(y_cpu) == 21.0f0 + @test dx_cpu == ones(Float32, 6) end