Skip to content

Don't copy or retain C in the ChainRules rules when β = Zero() - #308

Merged
lkdvos merged 3 commits into
mainfrom
lb/zero-beta-rrules
Sep 25, 2026
Merged

lkdvos merged 3 commits into
mainfrom
lb/zero-beta-rrules

Conversation

@leburgel

@leburgel leburgel commented Sep 24, 2026 •

Copy link
Copy Markdown
Member

A next thing I ran into while trying to reduce AD tape sizes used in implicit differentiation: the Chainrules implementations quite often track things they don't need to, leading to chained pullbacks carrying around more memory than they need to (at least for my specific use cases).

I'm not sure if the proposed fix here is the right way to go, but I thought I might as well attach it here as a suggested fix for the actual issue I'm reporting here.

Description

@tensor fills every new or temporary array with tensorcontract!(C, …, α, Zero()) (likewise
tensoradd!, tensortrace!). The rules then copy C, although with β = Zero() its contents
are never read, and the pullback keeps the original C alive until the reverse pass, since it
passes C to helpers that don't use it. Every @tensor contraction on the tape thus holds a
dead array of the output's size.

Change

  • With β === Zero(), write into similar(C) instead of copy(C). The kernels never read the
    destination then, but the default and StridedNative backends touch it, which fails for an
    uninitialized non-isbits array, hence zerovector!!(similar(C)) in that case.
  • Capture C only if dβ needs it.
  • Return NoTangent() rather than ZeroTangent() for a constant β, so that test_rrule
    accepts β = Zero().

Effect

One enlarged-corner (issue surfaced in the context of PEPSKit.jl) @tensor contraction on Array{Float64}, differentiated with Zygote:

χ, D memory held between forward and reverse pass forward allocations gradient
24, 3 2.9 → 1.5 MiB 4.5 MiB, unchanged bit-identical
48, 4 36.6 → 18.3 MiB 54.6 MiB, unchanged bit-identical
64, 5 157.8 → 78.9 MiB 236.0 MiB, unchanged bit-identical

"Memory held" is the live heap after a full GC with the pullback alive, minus that before the
forward pass. This lowers peak memory under AD; allocations and GC work are unchanged.

Reproducer
using TensorOperations, ChainRulesCore, Zygote, LinearAlgebra, Random
using VectorInterface: One, Zero

# 1. The pullback closure holds on to `C`
A = randn(30, 30); B = randn(30, 30); C = similar(A)
_, back = rrule(
    tensorcontract!, C, A, ((1,), (2,)), false, B, ((1,), (2,)), false, ((1, 2), ()),
    One(), Zero()
)
println("pullback captures C: ", any(f -> getfield(back, f) === C, fieldnames(typeof(back))))

# 2. What that costs for an ordinary @tensor expression on arrays
function f(C, EW, EN, A)
    @tensor Q[χS, DSt, DSb, χE, DEt, DEb] :=
        EW[χS, DWt, DWb, χ1] * C[χ1, χ2] * EN[χ2, DNt, DNb, χE] *
        A[d, DNt, DEt, DSt, DWt] * conj(A[d, DNb, DEb, DSb, DWb])
    return real(dot(Q, Q))
end
for (χ, D) in ((24, 3), (48, 4), (64, 5))
    Random.seed!(1234)
    args = (randn(χ, χ), randn(χ, D, D, χ), randn(χ, D, D, χ), randn(2, D, D, D, D))
    Zygote.gradient(f, args...)                        # compile
    GC.gc(true)
    live0 = Base.gc_live_bytes()
    bytes = @allocated ((y, pb) = Zygote.pullback(f, args...))
    tape = Base.summarysize(pb)
    GC.gc(true)
    live = Base.gc_live_bytes() - live0     # what stays alive between forward and reverse
    g = pb(one(y))
    MiB(b) = round(b / 2^20; digits = 1)
    println(
        "χ=$χ D=$D: forward allocates $(MiB(bytes)) MiB, tape holds $(MiB(tape)) MiB, ",
        "live heap +$(MiB(live)) MiB, gradient hash $(hash(g))"
    )
end

Tests

test/ad.jl gains a β = Zero() test_rrule case for each rule, and a check that an
uninitialized BigFloat C works with β = Zero().

@Jutho

Jutho commented Sep 24, 2026

Copy link
Copy Markdown
Member

The main gain here comes from the _needs_tangent(β) cases, right? Does copy(C) vs similar(C) make that much of a difference? It's the same memory I guess; does it have a significant runtime effect?

@leburgel

Copy link
Copy Markdown
Member Author

Yes, the gain comes from C_β = _needs_tangent(β) ? C : nothing, specifically the case where it's false, so we don't hold onto a dead copy of C when β doesn't need a tangent. copy(C) vs similar(C) doesn't matter for memory at all, and there's no real runtime effect (the timings I did to motivate the difference are essentially just noise now that I look at them better). I'll drop the distinction and just use copy.

@codecov

codecov Bot commented Sep 25, 2026

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

Files with missing lines Coverage Δ
ext/TensorOperationsChainRulesCoreExt.jl 83.52% <100.00%> (+4.26%) ⬆️

... and 7 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@lkdvos
lkdvos merged commit b5e6e76 into main Sep 25, 2026
13 checks passed
@lkdvos
lkdvos deleted the lb/zero-beta-rrules branch September 25, 2026 17:49
@lkdvos lkdvos mentioned this pull request Sep 30, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants