# Passing whole sequence to next layer from RNNCell in Lux

**URL:** <https://discourse.julialang.org/t/passing-whole-sequence-to-next-layer-from-rnncell-in-lux/136779>\
**Category:** Performance\
**Tags:** lux\
**Created:** [April 19, 2026, 4:42pm UTC](https://discourse.julialang.org/t/passing-whole-sequence-to-next-layer-from-rnncell-in-lux/136779 "2026-04-19T16:42:39Z")\
**Posts on this page:** 1\
**Page:** 1

<div class="post-metadata">

**Author:** ![alequa](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/alequa/32/12338_2.png) [@alequa](https://discourse.julialang.org/u/alequa)\
**Post date:** [April 19, 2026, 4:42pm UTC](https://discourse.julialang.org/t/passing-whole-sequence-to-next-layer-from-rnncell-in-lux/136779/1 "2026-04-19T16:42:39Z")

</div>

Hello,

I am trying to reproduce some basic code for training recurrent spiking networks. Here you can find an old porting for Flux from spytorch.

I already adapted to the current Flux version, but I would really like to move to Lux as I believe the support for Neural ODE could come in very handy in later stages of the project.

The idea is that there is a hidden spiking layer that receives spike inputs and yields spikes to an output layer. As such, the input is a `InputNeurons x Time x Batchsize` matrix, and the output should be an `OutputNeuron x Time x Batchsize` matrix.

On my way there, I tried to do this using a RNNCell and I can do it as below:

```julia-auto
using Lux, Random, Optimisers

function RNNClassifierCompact(in_dims, hidden_dims, out_dims)
    return @compact(;
        input=Dense(in_dims=>hidden_dims, sigmoid),
        rnn_cell=RNNCell(hidden_dims => hidden_dims) |> x-> Lux.Recurrence(x, return_sequence=true),
        classifier=Dense(hidden_dims => out_dims, sigmoid)
    ) do x::AbstractArray{T,3} where {T}
        out = map(rnn_cell(input(x))) do x
                classifier(x)
        end
        @return cat(out..., dims=3) |> x -> permutedims(x, (1, 3, 2)) 
    end
end

layers = (10, 2, 5)
model = RNNClassifierCompact(layers...)
ps, st = Lux.setup(Random.default_rng(), model) 

x = rand(10, 50, 128)
train_state = Training.TrainState(model, ps, st, Adam(0.01f0))
st_ = Lux.testmode(train_state.states)
ŷ, st_ = model(x, train_state.parameters, st_)

size(ŷ) # (5, 50, 128)

```

However, the way I handle the output of the RNN cell seems just wrong. In Flux I was wrapping all the layers with `Recurrence` and that was sufficient. Any hint on how do it in Lux in the cleanest and most performant way?

update:  
@avikpal maybe you have a quick answer on this? I saw you answered many of RNN related questions. Is `Chain` the correct solution here? How do you put a `Chain` in the `@compact` ?

Thanks!
