From 40ed845bd77899fad5af573da444015383fe9c47 Mon Sep 17 00:00:00 2001 From: Jeffrey Ward Date: Sat, 11 Jul 2026 14:28:06 -0400 Subject: [PATCH] =?UTF-8?q?feat:=20ScalarFunction=20=C2=B1=20Number=20adds?= =?UTF-8?q?=20the=20constant=20function?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A scalar added to a function means "add the constant function of that value", projected onto the basis: f + c == f + dg_function(dg, fill(c, n)). Only functions absorb a scalar this way — the other tensors have no canonical constant element, so +/- against a Number stays a MethodError there, matching upstream. Flips the `function plus scalar` @test_broken in the ported Python suite. Closes #11 --- src/tensors/base_tensor/base_tensor.jl | 11 ++++++ test/test_pyclasses.jl | 50 ++++++++++++++++++-------- 2 files changed, 47 insertions(+), 14 deletions(-) diff --git a/src/tensors/base_tensor/base_tensor.jl b/src/tensors/base_tensor/base_tensor.jl index 42bf17d..91884d0 100644 --- a/src/tensors/base_tensor/base_tensor.jl +++ b/src/tensors/base_tensor/base_tensor.jl @@ -24,6 +24,17 @@ Base.:*(s::Number, t::AbstractTensor) = wrap(t.space, s .* t.coeffs) Base.:*(t::AbstractTensor, s::Number) = s * t Base.:/(t::AbstractTensor, s::Number) = wrap(t.space, t.coeffs ./ s) +# A scalar added to a function is the constant function of that value, projected +# onto the basis. Only functions absorb a scalar this way — the other tensors have +# no canonical constant element, and `+`/`-` against a Number stays a MethodError. +_constant_function(f::ScalarFunction, c::Number) = + dg_function(geometry(f), fill(c, npoints(geometry(f)))) + +Base.:+(f::ScalarFunction, c::Number) = f + _constant_function(f, c) +Base.:+(c::Number, f::ScalarFunction) = f + c +Base.:-(f::ScalarFunction, c::Number) = f + (-c) +Base.:-(c::Number, f::ScalarFunction) = (-f) + c + # ── Pointwise products with a function ───────────────────────────────────────── function _pointwise_product(t::AbstractTensor, f) dg = geometry(t) diff --git a/test/test_pyclasses.jl b/test/test_pyclasses.jl index 9d7a350..7c4b071 100644 --- a/test/test_pyclasses.jl +++ b/test/test_pyclasses.jl @@ -4,14 +4,10 @@ # wrappers, batching/broadcasting, direct sums, spectral operators, transposes and # error paths. It runs over the same d = 1..4 config sweep (see pysuite.jl). # -# Two upstream behaviours have no Julia implementation and are recorded as -# @test_broken rather than dropped — they will flip to a failure the day someone -# adds them: -# -# 1. `ScalarFunction ± Number` (Python's `f + 5`, meaning "add the constant -# function"). Julia defines `*`/`/` against a Number but not `+`/`-`. -# 2. Scaling a batched tensor by a *vector* of per-batch scalars -# (Python's `weights * omega`, weights of length B). +# One upstream behaviour exercised here has no Julia implementation, and is +# recorded as @test_broken rather than dropped — it will flip to a failure the day +# someone adds it: scaling a batched tensor by a *vector* of per-batch scalars +# (Python's `weights * omega`, weights of length B). # # Purely Python-shaped tests are skipped with a note where they appear: numpy ufunc # dispatch (`np.multiply(..., where=...)`), `repr` string contents (Julia defines @@ -1025,22 +1021,48 @@ end @test_throws AssertionError op(f, T) @test_throws AssertionError op(v, f) - # Non-Function tensors cannot absorb a scalar. + # Non-Function tensors cannot absorb a scalar, in either order. @test_throws MethodError op(v, 3.0) @test_throws MethodError op(ω, 3.0) @test_throws MethodError op(T, 3.0) + @test_throws MethodError op(3.0, v) + @test_throws MethodError op(3.0, ω) + @test_throws MethodError op(3.0, T) end @test coeffs(-v) ≈ -coeffs(v) end - # GAP: Python defines `Function ± scalar` as adding the constant function. + # `Function ± scalar` adds/subtracts the constant function. @testset "function plus scalar" begin f = wrap(fs, rand(n0)) - @test_broken (f + 3.5) isa ScalarFunction - @test_broken (3.5 + f) isa ScalarFunction - @test_broken (f - 3.5) isa ScalarFunction - @test_broken (3.5 - f) isa ScalarFunction + c = 3.5 + const_c = dg_function(dg, fill(c, n)) + + @test (f + c) isa ScalarFunction + @test (c + f) isa ScalarFunction + @test (f - c) isa ScalarFunction + @test (c - f) isa ScalarFunction + + @test coeffs(f + c) ≈ coeffs(f) .+ coeffs(const_c) + @test coeffs(c + f) ≈ coeffs(f + c) + @test coeffs(f - c) ≈ coeffs(f) .- coeffs(const_c) + @test coeffs(c - f) ≈ coeffs(const_c) .- coeffs(f) + + # φ₀ is the constant Perron eigenfunction, so the constant function + # lands entirely in the first coefficient. + diff = coeffs(f + c) .- coeffs(f) + @test diff[1] ≈ c * sqrt(sum(measure(dg))) rtol = 1e-6 + @test all(isapprox.(diff[2:end], 0.0; atol=1e-8)) + + # Batched functions absorb a scalar too: the unbatched constant + # broadcasts across every batch row, in both operand orders. + fb = wrap(fs, randn(Xoshiro(0), 3, n0)) + cb = reshape(coeffs(const_c), 1, :) + @test coeffs(fb + c) ≈ coeffs(fb) .+ cb + @test coeffs(c + fb) ≈ coeffs(fb) .+ cb + @test coeffs(fb - c) ≈ coeffs(fb) .- cb + @test coeffs(c - fb) ≈ cb .- coeffs(fb) end @testset "exhaustive products" begin