# Building simple sequence-to-one RNN with Flux

**URL:** <https://discourse.julialang.org/t/building-simple-sequence-to-one-rnn-with-flux/56361>\
**Category:** New to Julia\
**Tags:** flux\
**Created:** [March 2, 2021, 8:20pm UTC](https://discourse.julialang.org/t/building-simple-sequence-to-one-rnn-with-flux/56361 "2021-03-02T20:20:27Z")\
**Posts on this page:** 1\
**Showing post:** 8

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [March 3, 2021, 10:31pm UTC](https://discourse.julialang.org/t/building-simple-sequence-to-one-rnn-with-flux/56361/8 "2021-03-03T22:31:29Z")

</div>

> [@asyrov](#):
>
> I would prefer to have reset in one place, it should belong to model to me, not the loss, because model can be called separately, and it would be strange to ask for additional requirement to call reset! before (or after) calling for `model(new data)` .

I mean, this is already the case if you use `model = Chain(...)`. There is no functionality in Flux for auto-calling `reset!`, so you will have to do it yourself at some point. That said, doing so is pretty straightforward:

```julia
rnn = LSTM(10, 15)
fc = Dense(15, 5)

function model(seq)
  reset!(rnn)
  x = rnn.(seq)[end] # or use map, or just a loop
  return fc(x)
end

function loss(seq, y, ...)
  y_hat = model(seq)
  return loss_func(y, y_hat)
end

```

Now you can use model without needing to reset manually or putting `reset!` into the loss function.

> [@asyrov](#):
>
> But there is a bigger problem that I got. With similar to your above suggestion, I created the code where each of each `x` is array of data for time step `i` . This is pretty much what you have above, but is not working. And here is why:
> 
> If you look at the code of RNNCell, you will notice that `h` field is a vector, but it should be matrix in my case, as there must be separate hidden state for each input in minibatch.
> 
> Specifically `h::V` is initialized as `zeros(out)` . But think what happens if I pass x which is _i-th time step of all samples in minibatch_ , single hidden state `h` will be broadcasted to all samples, which is not what I want.
> 
> Does this make sense? Or how would you do minibatching then?

Have you actually tried calling the RNN with a minibatched input like you describe? It’s a little confusing and I think we could make it less so, but everything works as you’d expect:

```juliarepl
julia> rnn = RNN(10, 3)
Recur(RNNCell(10, 3, tanh))

julia> rnn.state
3-element Vector{Float32}:
 0.0
 0.0
 0.0

julia> rnn.cell.h
3-element Vector{Float32}:
 0.0
 0.0
 0.0

julia> x = rand(Float32, 10, 8);

julia> rnn(x)
3×8 Matrix{Float32}:
 0.121863 0.0712726 0.468342 0.0159795 -0.50595 0.217166 0.321759 0.0969098
 0.78138 0.0184485 0.309471 -0.131435 -0.0146722 0.552875 0.227291 0.191328
 0.938252 0.981406 0.826487 0.98748 0.974808 0.960942 0.963614 0.964724

julia> rnn.state
3×8 Matrix{Float32}:
 0.121863 0.0712726 0.468342 0.0159795 -0.50595 0.217166 0.321759 0.0969098
 0.78138 0.0184485 0.309471 -0.131435 -0.0146722 0.552875 0.227291 0.191328
 0.938252 0.981406 0.826487 0.98748 0.974808 0.960942 0.963614 0.964724

julia> rnn.cell.h
3-element Vector{Float32}:
 0.0
 0.0
 0.0

julia> Flux.reset!(rnn)
3-element Vector{Float32}:
 0.0
 0.0
 0.0

julia> rnn.state
3-element Vector{Float32}:
 0.0
 0.0
 0.0

```

As you can see, the hidden state is actually stored in `Recur` and not the RNN cell. That hidden state does start off as a vector, but will be overwritten as a matrix with the right number of samples if you pass it a minibatched input.

---

_[View the full topic](https://discourse.julialang.org/t/building-simple-sequence-to-one-rnn-with-flux/56361)._
