From 894696b9cd01be7069f5db099b493760279a8538 Mon Sep 17 00:00:00 2001 From: Michael Goerz Date: Mon, 4 May 2026 15:09:58 -0400 Subject: [PATCH] Allow `J_a_fluence` on non-uniform time grids --- src/functionals.jl | 40 ++++++++++++++++++++++++++++------------ test/test_functionals.jl | 23 +++++++++++++++++++++++ 2 files changed, 51 insertions(+), 12 deletions(-) diff --git a/src/functionals.jl b/src/functionals.jl index 7725b43..b201647 100644 --- a/src/functionals.jl +++ b/src/functionals.jl @@ -379,6 +379,12 @@ vector `∇J_a` containing the vectorized elements ``∂J_a/∂ϵ_{ln}``. The function `J_a` must have the interface `J_a(pulsevals, tlist)`, see, e.g., [`J_a_fluence`](@ref). +In `pulsevals`, the values ``ϵ_{nl}`` are vectorized with `n` (time interval +index) varying faster than `l` (control index), i.e., +`pulsevals = [ϵ₁₁, ϵ₂₁, …, ϵ_{N_T,1}, ϵ₁₂, ϵ₂₂, …, ϵ_{N_T,2}, …]`, +where ``N_T = `` `length(tlist) - 1`. The pulse values for each control are +contiguous. + The parameters `mode` and `automatic` are handled as in [`make_chi`](@ref), where `mode` is one of `:any`, `:analytic`, `:automatic`, and `automatic` is he loaded module of an automatic differentiation framework, where `:default` @@ -924,16 +930,24 @@ J_a = J_a_fluence(pulsevals, tlist) calculates ```math -J_a = \sum_l \int_0^T |ϵ_l(t)|^2 dt = \left(\sum_{nl} |ϵ_{nl}|^2 \right) dt +J_a = \sum_l \int_0^T |ϵ_l(t)|^2 dt \approx \sum_{nl} |ϵ_{nl}|^2 \, dt_n ``` -where ``ϵ_{nl}`` are the values in the (vectorized) `pulsevals`, `n` is the -index of the intervals of the time grid, and ``dt`` is the time step, taken -from the first time interval of `tlist` and assumed to be uniform. +where ``ϵ_{nl}`` are the values in `pulsevals`, with `n` +the index of the time interval and `l` the index of the control, and +``dt_n = `` `tlist[n+1] - tlist[n]` is the duration of interval `n`. +The `pulsevals` are vectorized as ``[ϵ₁₁, ϵ₂₁, …, ϵ_{N_T,1}, ϵ₁₂, ϵ₂₂, …]``, +where `N_T = length(tlist) - 1`. Supports non-uniform time grids. + +# See also + +* [`grad_J_a_fluence`](@ref) — analytic (automatic) gradient """ function J_a_fluence(pulsevals, tlist) - dt = tlist[begin+1] - tlist[begin] - return sum(abs2.(pulsevals)) * dt + N_T = length(tlist) - 1 + dt = reshape(diff(tlist), :, 1) # (N_T, 1) for column broadcasting + pv = reshape(pulsevals, N_T, :) # (N_T, N_L), no copy + return sum(abs2.(pv) .* dt) end @@ -943,14 +957,16 @@ end ∇J_a = grad_J_a_fluence(pulsevals, tlist) ``` -returns the `∇J_a`, which contains the (vectorized) elements ``2 ϵ_{nl} dt``, -where ``ϵ_{nl}`` are the (vectorized) elements of `pulsevals` and ``dt`` is the -time step, taken from the first time interval of `tlist` and assumed to be -uniform. +returns `∇J_a`, which contains the (vectorized) elements ``2 ϵ_{nl} dt_n``, +where ``ϵ_{nl}`` are the (vectorized) elements of `pulsevals` and +``dt_n = `` `tlist[n+1] - tlist[n]` is the duration of interval `n`. +Supports non-uniform time grids. """ function grad_J_a_fluence(pulsevals, tlist) - dt = tlist[begin+1] - tlist[begin] - return (2 * dt) * pulsevals + N_T = length(tlist) - 1 + dt = reshape(diff(tlist), :, 1) # (N_T, 1) for column broadcasting + pv = reshape(pulsevals, N_T, :) # (N_T, N_L), no copy + return vec(2 .* dt .* pv) end diff --git a/test/test_functionals.jl b/test/test_functionals.jl index 8eef4bb..1d3db3b 100644 --- a/test/test_functionals.jl +++ b/test/test_functionals.jl @@ -125,6 +125,29 @@ end end +@testset "J_a_fluence non-uniform grid" begin + + # Non-uniform tlist with 4 intervals, 2 controls + # pulsevals layout: [ϵ₁₁, ϵ₂₁, ϵ₃₁, ϵ₄₁, ϵ₁₂, ϵ₂₂, ϵ₃₂, ϵ₄₂] + tlist_nu = [0.0, 0.1, 0.3, 0.6, 1.0] + dt_nu = [0.1, 0.2, 0.3, 0.4] + pv1 = [1.0, 2.0, 3.0, 4.0] + pv2 = [0.5, 1.5, 2.5, 3.5] + pulsevals_nu = vcat(pv1, pv2) + + J_expected = sum(abs2.(pv1) .* dt_nu) + sum(abs2.(pv2) .* dt_nu) + @test J_a_fluence(pulsevals_nu, tlist_nu) ≈ J_expected + + G_expected = vcat(2 .* pv1 .* dt_nu, 2 .* pv2 .* dt_nu) + @test grad_J_a_fluence(pulsevals_nu, tlist_nu) ≈ G_expected + + grad_J_a_zygote_nu = + make_grad_J_a(J_a_fluence, tlist_nu; mode = :automatic, automatic = Zygote) + @test norm(grad_J_a_zygote_nu(pulsevals_nu, tlist_nu) - G_expected) < 1e-12 + +end + + @testset "J_T without analytic derivative" begin QuantumControl.set_default_ad_framework(nothing; quiet = true)