# Variable sequence length RNN in Flux

**URL:** https://discourse.julialang.org/t/variable-sequence-length-rnn-in-flux/70715
**Category:** Machine Learning
**Tags:** flux
**Created:** [November 1, 2021, 7:05am UTC](https://discourse.julialang.org/t/variable-sequence-length-rnn-in-flux/70715 "2021-11-01T07:05:14Z")
**Posts on this page:** 3
**Page:** 1

<div class="post-metadata">

### Author: ![paalon](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/paalon/32/5784_2.png) [@paalon](https://discourse.julialang.org/u/paalon)
#### Post date: [November 1, 2021, 7:05am UTC](https://discourse.julialang.org/t/variable-sequence-length-rnn-in-flux/70715/1 "2021-11-01T07:05:14Z")

</div>

How to implement RNN for variable sequence length data with minibatching in Flux?  
According to Flux doc,

> In Flux, those 3 dimensions are provided through a vector of seq length containing a matrix `(features, samples)` .  
> [Recurrence · Flux](https://fluxml.ai/Flux.jl/stable/models/recurrence/)

but it’s impossible to create such a vector of matrices for variable sequence length data.

---

<div class="post-metadata">

### Author: ![paalon](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/paalon/32/5784_2.png) [@paalon](https://discourse.julialang.org/u/paalon)
#### Post date: [November 3, 2021, 1:42am UTC](https://discourse.julialang.org/t/variable-sequence-length-rnn-in-flux/70715/2 "2021-11-03T01:42:59Z")

</div>

For example, I’m assuming the following situation:

```julia
using Flux

dim_feature = 4
dim_sample = 100
min_seq_length = 3
max_seq_length = 6

dim_output = 2

x_seq_length = rand(min_seq_length:max_seq_length, dim_sample)

# input dataset
x = []
for i = 1:dim_sample
	 push!(x, rand(dim_feature, x_seq_length[i]))
end

network = LSTM(dim_feature, dim_output)

# want to apply minibatch of input dataset `x` to `network` in minibatch SGD

```

---

<div class="post-metadata">

### Author: ![albheim](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/albheim/32/34660_2.png) [@albheim](https://discourse.julialang.org/u/albheim)
#### Post date: [November 3, 2021, 9:57am UTC](https://discourse.julialang.org/t/variable-sequence-length-rnn-in-flux/70715/3 "2021-11-03T09:57:18Z")

</div>

I haven’t used the recursive cells in Flux before, so not sure this is the best way of doing things, but this seems to learn at least when plotting the loss. And it can handle batches containing sequences of different lengths.

I’m a bit unsure about the `reset!` in the loss function, since there was recently a post about something similar where I remember it was suggested that this might not be good practice, but I can’t see why it is a problem.

```julia
using Flux, Plots

function loss(m, xs, ys)
    loss = 0f0
    for (x, y) in zip(xs, ys)
        Flux.reset!(m) # Reset the state of the recursive cell for each new sequence
        loss += sum(exp2.(m(x)[:, end] - y)) 
    end
    loss 
end

dim_feature = 4
dim_sample = 100
min_seq_length = 3
max_seq_length = 6
dim_output = 2

x_seq_length = rand(min_seq_length:max_seq_length, dim_sample)
x = rand.(Float32, dim_feature, x_seq_length)
y = [rand(dim_output) for _ in 1:dim_sample]

m = LSTM(dim_feature, dim_output)
	
data = Flux.DataLoader((x, y), batchsize=4)
opt = ADAM()
losses = [loss(m, x, y)]
Flux.@epochs 40 begin
	Flux.train!((x, y) -> loss(m, x, y), Flux.params(m), data, opt)
	push!(losses, loss(m, x, y))
end
plot(losses)

```
