From 19c81379b0194fc9eb0d1f90368d8e71bf827698 Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Thu, 27 Aug 2026 14:55:40 -0400 Subject: [PATCH] Fix batched ForwardWithPrimal sret width Co-Authored-By: Chris Rackauckas Co-Authored-By: Claude Claude-Session: https://chatgpt.com/codex/tasks/01a03f84-af17-7a90-955c-10b5d980e4da --- src/compiler.jl | 3 ++- test/typeunstable.jl | 23 +++++++++++++++++++---- 2 files changed, 21 insertions(+), 5 deletions(-) diff --git a/src/compiler.jl b/src/compiler.jl index 12f0c47ab4..12d01dc233 100644 --- a/src/compiler.jl +++ b/src/compiler.jl @@ -3866,7 +3866,8 @@ function create_abi_wrapper( twidth = if width == 1 1 else - if (rettype <: Const) && returnNum == 0 + if ((rettype <: Const) && returnNum == 0) || + (returnPrimal && returnNum == count_Sret - 1) 1 else width diff --git a/test/typeunstable.jl b/test/typeunstable.jl index 31c8c5e8a7..68e6c33d00 100644 --- a/test/typeunstable.jl +++ b/test/typeunstable.jl @@ -334,15 +334,30 @@ end batched_uninferred_target(x) = (2 .* x,) batched_uninferred_return(x) = Base.invokelatest(batched_uninferred_target, x) -@testset "Batched forward concrete annotation with uninferred return" begin +function batched_uninferred_setup() x = [1.0, 2.0] - dx1 = ones(2) - dx2 = fill(3.0, 2) + tangents = (ones(2), fill(3.0, 2)) return_activity = BatchDuplicated{Tuple{Vector{Float64}}, 2} + return x, tangents, return_activity +end + +@testset "Batched forward concrete annotation with uninferred return" begin + x, tangents, return_activity = batched_uninferred_setup() result = Enzyme.autodiff( Enzyme.set_runtime_activity(Forward), Const(batched_uninferred_return), - return_activity, BatchDuplicated(x, (dx1, dx2)) + return_activity, BatchDuplicated(x, tangents) + ) + @test result[1][1][1] == [2.0, 2.0] + @test result[1][2][1] == [6.0, 6.0] +end + +@testset "Batched forward-with-primal concrete annotation with uninferred return" begin + x, tangents, return_activity = batched_uninferred_setup() + result = Enzyme.autodiff( + Enzyme.set_runtime_activity(ForwardWithPrimal), Const(batched_uninferred_return), + return_activity, BatchDuplicated(x, tangents) ) @test result[1][1][1] == [2.0, 2.0] @test result[1][2][1] == [6.0, 6.0] + @test result[2][1] == [2.0, 4.0] end