FluxML / FluxML/Optimisers.jl

Restructure makes a copy

Open
#146 4 comments 0 reactions 0 assignees View on GitHub
enhancement
Dominant language
Julia
Stars
96
Forks
30
PR merge metrics
No merged PRs in 30d

Description

In some situations, you have to restructure a lot if you use Flux, for instance if you want to run your batches as seperate solves in DiffEqFlux using an EnsembleProblem. You have to use something like a ComponentArray to pass the parameters through the solver and to let the adjoint methods do their work in differentiating the solve. But restructuring using a ComponentArray is unreasonably(?) slow in Flux. Switching to Lux eliminates those problems, but it seems like something that could be implemented better in Flux or ComponentArrays.

Example code:
```
using Flux, Profile, PProf, Random, ComponentArrays
layer_size = 256
layer = Flux.Dense(layer_size => layer_size)
params_, re = Flux.destructure(layer)
params = ComponentArray(params_)
function eval(steps, input)
vec = input
for i in 1:steps
vec = re(params)(vec)
end
end
Profile.clear()
@profile eval(100000, rand(Float32, layer_size))
pprof(; web=true)
```

Here, we spend 10% of the time in sgemv matrix multiplication, another 10% in the rest of the Dense call and about 75% in the Restructure. This gets worse if the networks are smaller. As far as I can read the flame graph, the restructure seems to spend a lot of time in the GC:

![0d19ec3a732aaf409e85b06ebb4ebe6392198e20_2_1044x1000](https://user-images.githubusercontent.com/8374054/234003126-cb3bd664-8a9c-42b2-bf81-04e92df3adaf.png)

Could there be a way to mitigate this specific problem? In particular if you use the same parameters. I think this would make some example code a lot faster too.

I also opened a discourse about this because I'm not sure if it's an issue with Flux specifically: https://discourse.julialang.org/t/flux-restructure-for-componentarrays-jl-unreasonably-slow/97849/3

Contributor guide

No contributing guide indexed for this repository

Research direction

Start with the provided Julia example using Flux.destructure, ComponentArray, and the eval loop, then inspect the profile showing time in Restructure and garbage collection. Compare the behavior with the same parameters across Flux and ComponentArrays, using the linked Discourse discussion for context. Done means identifying and addressing the repeated restructuring overhead without changing the solve behavior.

Written by the indexing model from the issue text.

Assessment

Tech stack
julia
Domain
machine-learning, performance
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.