JuliaDiff / JuliaDiff/ChainRules.jl
stack overflow issue
Nobody has claimed this yet.
- Dominant language
- Julia
- Stars
- 475
- Forks
- 98
- PR merge metrics
- No merged PRs in 30d
Description
Hello
I have a memory-constrained problem with a Lux.jl model that uses Zygote for most of the backpropagation.
I tried to approach this from chainrules perspective I need to checkpoint each Lux.jl layer in neural network. So I tried to achieve it like that :
function ChainRulesCore.rrule(::typeof(Lux.apply), l::Lux.AbstractExplicitLayer, x, ps, st)
y = Lux.apply(l, x, ps, st)
function pullback_checkpointed(Δy)
y, pb =Zygote.pullback(Lux.apply,l, x, ps, st)
return NoTangent(), pb(Δy)
end
y, pullback_checkpointed
end
Rule gets invoked in backpropagation Hovewer the issue is that for some reason it try recursively to do backpropagation of the first line
y = Lux.apply(l, x, ps, st)
so I get stack overflow error; how to correct it?
I had also posted this issue in https://discourse.julialang.org/t/avoid-storing-intermediate-results-from-the-forward-pass-by-default/119694/4?u=jakub_mitura
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with the custom ChainRulesCore.rrule for Lux.apply and trace how its Lux.apply forward call interacts with Zygote.pullback. Use the linked Discourse discussion for context; done means the rule avoids the recursive stack overflow while supporting checkpointed backpropagation for the Lux model.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- julia
- Domain
- machine-learning, tooling
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 30/100