# Simple Flux LSTM for Time Series

**URL:** https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494
**Category:** Machine Learning
**Tags:** question, flux, time-series, machine-learning
**Created:** [March 4, 2020, 1:10am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494 "2020-03-04T01:10:12Z")
**Posts on this page:** 20
**Page:** 2

<div class="post-metadata">

### Author: ![Volker](https://avatars.discourse-cdn.com/v4/letter/v/77aa72/32.png) [@Volker](https://discourse.julialang.org/u/Volker)
#### Post date: [October 14, 2020, 8:36am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/21 "2020-10-14T08:36:16Z")

</div>

If I have understood it correctly, the loss function should look something like this:

```julia
loss(inputs, output) = sum(abs2.(output.-vcat(model.(inputs)...)[:, end]))

```

if you just would like to compare the model output and the end of the batches. Will the inner state be resetted after each batch? Like in `stateful = false`. Otherwise, would it not make more sense to pass the data as one big batch? Because after a couple of time steps the state would be close to the next time step target value and to repeat the model calculation wouldn´t be necessary.

If the states would be resetted, is the above loss function efficient or is there a better way to calculate the mse?

---

<div class="post-metadata">

### Author: ![reemmasoud123](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/reemmasoud123/32/26181_2.png) [@reemmasoud123](https://discourse.julialang.org/u/reemmasoud123)
#### Post date: [June 20, 2021, 2:01pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/22 "2021-06-20T14:01:33Z")

</div>

can you please demonstrate how training can be done after this step using a batch size of 30 for example? I tried the below but it didn’t work for the 3D data:

train\_loader = DataLoader((trainX, trainY), batchsize=30 shuffle=true)

---

<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: [June 20, 2021, 3:12pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/23 "2021-06-20T15:12:18Z")

</div>

It’s impossible to answer this without seeing what `trainX` and `trainY` are. That is, the full type, dimensionality, etc.

---

<div class="post-metadata">

### Author: ![reemmasoud123](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/reemmasoud123/32/26181_2.png) [@reemmasoud123](https://discourse.julialang.org/u/reemmasoud123)
#### Post date: [June 20, 2021, 3:33pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/24 "2021-06-20T15:33:19Z")

</div>

> [@LSTM training for a sequence of multiple features using a batch size 30](https://discourse.julialang.org/t/lstm-training-for-a-sequence-of-multiple-features-using-a-batch-size-30/63238):
>
> I am trying to do batch training using LSTM for a time series data with multiple features. Assuming I have 5000 samples and 5 features for each sample. The input uses 14 days into the past and the output is a single value on the 15th day. (My time step is 14). The size of my data is the following: xtrain: (5000,14,5) ytrain: (5000,1,1) My model is below. How do I train my data by using a batch size of 30? I tried using DataLoader and Flux.train but they are both not working with this input s…

I posted a question in the above link with the details.

---

<div class="post-metadata">

### Author: ![sherlock\_holmes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sherlock_holmes/32/32280_2.png) [@sherlock\_holmes](https://discourse.julialang.org/u/sherlock_holmes)
#### Post date: [January 3, 2022, 2:43pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/25 "2022-01-03T14:43:41Z")

</div>

Hi, sorry for bumping the topic.  
I have tried to run code you’ve written which is below:

```julia
using Flux

m = Chain(LSTM(3,2), Dense(2,1))

inputs = rand(3,4)

for t in 1:4
    output = m(inputs[:,t])
    @show output
end

```

I don’t know why it doesn’t work? It’s a basic code and there is no reason to not being able to run it. I get the error below when I run the code:

```julia

ERROR: LoadError: MethodError: no method matching (::Flux.LSTMCell{Matrix{Float32}, Vector{Float32}, Tuple{Matrix{Float32}, Matrix{Float32}}})(::Tuple{Matrix{Float32}, Matrix{Float32}}, ::Vector{Float64})
Closest candidates are:
  (::Flux.LSTMCell{A, V, <:Tuple{AbstractMatrix{T}, AbstractMatrix{T}}})(::Any, ::Union{AbstractVector{T}, AbstractMatrix{T}, Flux.OneHotArray}) where {A, V, T} at ~/.julia/packages/Flux/BPPNj/src/layers/recurrent.jl:157
Stacktrace:
 [1] (::Flux.Recur{Flux.LSTMCell{Matrix{Float32}, Vector{Float32}, Tuple{Matrix{Float32}, Matrix{Float32}}}, Tuple{Matrix{Float32}, Matrix{Float32}}})(x::Vector{Float64})
   @ Flux ~/.julia/packages/Flux/BPPNj/src/layers/recurrent.jl:47
 [2] applychain(fs::Tuple{Flux.Recur{Flux.LSTMCell{Matrix{Float32}, Vector{Float32}, Tuple{Matrix{Float32}, Matrix{Float32}}}, Tuple{Matrix{Float32}, Matrix{Float32}}}, Dense{typeof(identity), Matrix{Float32}, Vector{Float32}}}, x::Vector{Float64})
   @ Flux ~/.julia/packages/Flux/BPPNj/src/layers/basic.jl:47
 [3] (::Chain{Tuple{Flux.Recur{Flux.LSTMCell{Matrix{Float32}, Vector{Float32}, Tuple{Matrix{Float32}, Matrix{Float32}}}, Tuple{Matrix{Float32}, Matrix{Float32}}}, Dense{typeof(identity), Matrix{Float32}, Vector{Float32}}}})(x::Vector{Float64})
   @ Flux ~/.julia/packages/Flux/BPPNj/src/layers/basic.jl:49
 [4] top-level scope
   @ ~/Desktop/b/new.jl:8
in expression starting at /home/user/Desktop/b/new.jl:7

```

What can be the reason? Thanks in advance.

---

<div class="post-metadata">

### Author: ![lazarusA](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lazarusa/32/6571_2.png) [@lazarusA](https://discourse.julialang.org/u/lazarusA)
#### Post date: [January 3, 2022, 3:30pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/26 "2022-01-03T15:30:11Z")

</div>

just do `inputs = rand(Float32, 3,4)`, things nowadays need to be Float32 from the start.

---

<div class="post-metadata">

### Author: ![sherlock\_holmes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sherlock_holmes/32/32280_2.png) [@sherlock\_holmes](https://discourse.julialang.org/u/sherlock_holmes)
#### Post date: [January 3, 2022, 3:35pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/27 "2022-01-03T15:35:20Z")

</div>

It works, thank you so much!

---

<div class="post-metadata">

### Author: ![compleat](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/compleat/32/8958_2.png) [@compleat](https://discourse.julialang.org/u/compleat)
#### Post date: [April 6, 2022, 9:32am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/28 "2022-04-06T09:32:15Z")

</div>

> [@sherlock\_holmes](#):
>
> ```julia
> using Flux
> 
> m = Chain(LSTM(3,2), Dense(2,1))
> 
> inputs = rand(Float32,3,4)
> 
> for t in 1:4
> output = m(inputs[:,t])
> @show output
> end
> 
> ```

I tried the above and it does indeed work, but when I follow by

```julia
Flux.reset!(m)
m(inputs)

```

I also get results (a 1 by 4 row vector) where the first element matches the result from the loop but the other numbers don’t. Please pardon my ignorance, but could you please explain why this is?

Another (not entirely unrelated?) question: does Flux reset during training with each new pattern?

Thanks in advance

---

<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: [April 6, 2022, 9:51am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/29 "2022-04-06T09:51:51Z")

</div>

Maybe this explains it? I would guess the internal state is updated each time the model is run, not matter if it is on a single datapoint or a batch. So in the first for loop the first data has the reset state of the LSTM, but later encounters a state that is based on previous data. In the batched case or the loop with the reset, all datapoints are calculated based on the reset state of the LSTM.

```julia
julia> for t in 1:4
           output = m(inputs[:,t])
           @show output
       end
output = Float32[0.052006032]
output = Float32[0.12330223]
output = Float32[0.21572198]
output = Float32[0.20931965]

julia> Flux.reset!(m)

julia> m(inputs)
1×4 Matrix{Float32}:
 0.052006 0.0927443 0.137698 0.0787673

julia> for t in 1:4
           Flux.reset!(m)
           output = m(inputs[:,t])
           @show output
       end
output = Float32[0.052006032]
output = Float32[0.09274433]
output = Float32[0.13769753]
output = Float32[0.07876728]

```

---

<div class="post-metadata">

### Author: ![compleat](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/compleat/32/8958_2.png) [@compleat](https://discourse.julialang.org/u/compleat)
#### Post date: [April 6, 2022, 10:08am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/30 "2022-04-06T10:08:27Z")

</div>

> [@albheim](#):
>
> Maybe this explains it? I would guess the internal state is updated each time the model is run, not matter if it is on a single datapoint or a batch. So in the first for loop the first data has the reset state of the LSTM, but later encounters a state that is based on previous data. In the batched case or the loop with the reset, all datapoints are calculated based on the reset state of the LSTM.
> 
> ```julia
> 
> ```

Thanks so much for that. I see. But that leads to my question about training… during training surely the forecast for a given sequence should not depend on which sequence (or batch) came before it(?) I had understood that the sequential structure was assumed only within a given vector of input matrices - that there was assumed no spatial or time-relation between these vectors. Have I got it wrong?  
Thanks again

---

<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: [April 6, 2022, 11:01am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/31 "2022-04-06T11:01:08Z")

</div>

I’ve not really used this myself, but a quick glance at the [docs](https://fluxml.ai/Flux.jl/stable/models/recurrence/) seem to suggest you have to do that yourself.

> In many situations, such as when dealing with a language model, the sentences in each batch are independent (i.e. the last item of the first sentence of the first batch is independent from the first item of the first sentence of the second batch), so we cannot handle the model as if each batch was the direct continuation of the previous one. To handle such situations, we need to reset the state of the model between each batch, which can be conveniently performed within the loss function:
> 
> ```julia
> function loss(x, y)
> Flux.reset!(m)
> sum(mse(m(xi), yi) for (xi, yi) in zip(x, y))
> end
> 
> ```

---

<div class="post-metadata">

### Author: ![compleat](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/compleat/32/8958_2.png) [@compleat](https://discourse.julialang.org/u/compleat)
#### Post date: [April 6, 2022, 11:20am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/32 "2022-04-06T11:20:16Z")

</div>

> [@albheim](#):
>
> I’ve not really used this myself, but a quick glance at the [docs](https://fluxml.ai/Flux.jl/stable/models/recurrence/) seem to suggest you have to do that yourself.

Thank you so much, that is exactly what I was looking for. I didn’t see that last bit of the docs which you quoted at the end there!

I understand a bit more, but I’m still looking for a simple notebook (zoo?) example where someone has applied this to a basic time series model.

Thanks again for all your help.

---

<div class="post-metadata">

### Author: ![mkschleg](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mkschleg/32/8035_2.png) [@mkschleg](https://discourse.julialang.org/u/mkschleg)
#### Post date: [April 6, 2022, 4:18pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/33 "2022-04-06T16:18:26Z")

</div>

Currently data is assumed to be in the following shapes for the recurrent layers:

```julia

X = [rand(Float32, in) for _ in 1:T] # not batched vector over time steps
Flux.reset!(m) # reset the hidden state to m.cell.state0
res = [m(x) for x in X] # each element is out \by 1

X = [rand(Float32, in, batch) for _in 1:T] # batched vector over time steps
Flux.reset!(m) # reset the hidden state to m.cell.state0
res = [m(x) for x in X] # used same as above. Each element is out \by batch

X = rand(Float32, in, batch, T)
Flux.reset!(m)
res = m(X) # should produce a matrix of out \by batch \by T

```

So in your example, the input `inputs = rand(Float32,3,4)` is assumed to be a single batch.

---

<div class="post-metadata">

### Author: ![mkschleg](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mkschleg/32/8035_2.png) [@mkschleg](https://discourse.julialang.org/u/mkschleg)
#### Post date: [April 6, 2022, 4:21pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/34 "2022-04-06T16:21:49Z")

</div>

I think the [flux model zoo](https://github.com/FluxML/model-zoo) should have what you are looking for. I also have a [jupyter notebook](https://github.com/mkschleg/FluxBooks.jl/blob/main/RNNs.ipynb) as an example, but it is getting pretty out of date. It should still work but the `map(rnn, x)[end]` should be replaced with `[rnn(_x) for _x in x]` due to guaranteed ordering.

---

<div class="post-metadata">

### Author: ![compleat](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/compleat/32/8958_2.png) [@compleat](https://discourse.julialang.org/u/compleat)
#### Post date: [April 6, 2022, 4:45pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/35 "2022-04-06T16:45:02Z")

</div>

Thanks for your suggestions. I am trying to apply LSTM to a time series (scalar) and I don’t see anything even remotely relevant to that in the zoo, and your MNIST notebook is completely different [I don’t know why recurrence would be useful in this case, actually, but that is just my ignorance, I’m sure]

I just want input sequences of 30 (days) each to forecast the next day, and want to assess my model on the performance on the quality of the forecast on the 30th day(only). This should be a prototypical time series forecasting problem, but I see nothing like this in any examples anywhere.  
Anyway, thanks again!

---

<div class="post-metadata">

### Author: ![mcreel](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcreel/32/30088_2.png) [@mcreel](https://discourse.julialang.org/u/mcreel)
#### Post date: [April 6, 2022, 5:23pm UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/36 "2022-04-06T17:23:41Z")

</div>

A working example of time series regression has never been in the model zoo, I believe. I finally got hold of an nvidia gpu, so I will work on making a simple example when I finish teaching this term, if one doesn’t appear before.

---

<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: [April 7, 2022, 1:43am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/37 "2022-04-07T01:43:40Z")

</div>

There are at least a few discussions and examples/MWEs around this kind of univariate time series forecasting with Flux RNNs floating around community forums. [How to train Flux to learn a sequence conditional to some initial "seeds"?](https://discourse.julialang.org/t/how-to-train-flux-to-learn-a-sequence-conditional-to-some-initial-seeds/78330) is a recent example I just remembered. The reason such a thing does not exist in the model zoo is probably two-fold:

1. Model zoo entries don’t write themselves 😛
2. A LSTM is a big hammer to model a 30-sample univariate timeseries forecasting problem with. Generally we try to strike a balance between clear, brief files and sufficiently “common” or “interesting” datasets and tasks in the model zoo to differentiate it from tutorials. In this case, perhaps something like forecasting with a UCI benchmark dataset would be appropriate.

---

<div class="post-metadata">

### Author: ![mcreel](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcreel/32/30088_2.png) [@mcreel](https://discourse.julialang.org/u/mcreel)
#### Post date: [April 7, 2022, 8:05am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/38 "2022-04-07T08:05:55Z")

</div>

Here’s a simple LSTM model that forecasts AR1 or MA1 data pretty well. I’m going to put this in a github archive for further work, but here’s an initial version. BTW, this needs a train/test split. At the moment, it’s probably over-fitting.

```julia
using Flux, Plots, Statistics
using Base.Iterators

# DGPs: AR1 is more forecastable than MA1
function MA1(n, σ)
    e = randn(n+1) .* σ 
    y = e[2:n+1] + 0.9*e[1:n]
end

function AR1(n, σ)
    y = zeros(n)
    for t = 2:n
        y[t] = 0.9*y[t-1] + σ*randn()
    end
    y
end    

# generate the data
n = 10000 # sample size
σ = 1.0 # true std. dev. of the shocks
data = Float32.(AR1(n, σ))

# set up the training
batchsize = 2 # remember, MA model is only predictable one step out
epochs = 100 # number of training loops through data

# the model: this is just an initial guess
# need to experiment with this
m = Chain(LSTM(batchsize, 10), Dense(10,2, tanh), Dense(2,batchsize))

function loss(x,y)
    Flux.reset!(m)
    Flux.mse(m(x),y)
end

# the first element of the batched data is one lag
# of the second element, in chunks of batchsize. So,
# we are doing one-step-ahead forecasting, conditioning
# on batchsize lags
batches = [(data[ind .- 1], data[ind]) for ind in partition(2:size(data,1), batchsize)]
batches = batches[1:end-1] # drop the last, which may not have full size
Flux.@epochs epochs Flux.train!(loss,Flux.params(m), batches, ADAM())

function predict(data, batchsize)
    n = size(data,1)
    yhat = zeros(n)
    for t = batchsize+1:n
        x = data[t-batchsize:t-1]
        Flux.reset!(m)
        yhat[t] = m(x)[end]
    end
    yhat
end

pred = predict(data,batchsize)
error = data - pred
plot(1:n, [data error])
println("true std. error of noise: ", σ)
println("std. error of forecast: ", std(error))
println("std. error of data: ", std(data))

```

---

<div class="post-metadata">

### Author: ![compleat](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/compleat/32/8958_2.png) [@compleat](https://discourse.julialang.org/u/compleat)
#### Post date: [April 7, 2022, 8:35am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/39 "2022-04-07T08:35:44Z")

</div>

> [@mcreel](#):
>
> Here’s a simple LSTM model that forecasts AR1 or MA1 data pretty well.

Thank you so much! That’s exactly the kind of example code I was hoping to see and play around with!

---

<div class="post-metadata">

### Author: ![CarloLucibello](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carlolucibello/32/3278_2.png) [@CarloLucibello](https://discourse.julialang.org/u/CarloLucibello)
#### Post date: [April 7, 2022, 8:56am UTC](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494/40 "2022-04-07T08:56:32Z")

</div>

@mcreel that would be a welcome contribution to the [model-zoo](https://github.com/FluxML/model-zoo)

[Previous page](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494.md?page=1)

[Next page](https://discourse.julialang.org/t/simple-flux-lstm-for-time-series/35494.md?page=3)
