# How to create tracked \`cumsum\` using Flux.jl?

**URL:** <https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772>\
**Category:** Machine Learning\
**Tags:** first-steps, flux\
**Created:** [November 20, 2018, 1:57pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772 "2018-11-20T13:57:45Z")\
**Posts on this page:** 13\
**Page:** 1

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 20, 2018, 1:57pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/1 "2018-11-20T13:57:45Z")

</div>

I wanted to create a `cumsum` of tracked values. How do I do that?

I tried the below which failed

```julia
W = param(rand(12,1))
cumsum(Tracker.data(W), dims = 1) # works
cumsum(W, dims = 1) # this fails

```

I got it to work using matrix multiplication but that’s not ideal as it’s more cumbersome

```julia
mw = Matrix{Float64}(undef, 12, 12)
mw .= 1
for i = 1:12
    for j = i+1:12
        mw[i,j]=0
    end
end
mw
mw*W # this works

```

---

<div class="post-metadata">

**Author:** ![improbable22](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/improbable22/32/5464_2.png) [@improbable22](https://discourse.julialang.org/u/improbable22)\
**Post date:** [November 20, 2018, 4:41pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/2 "2018-11-20T16:41:44Z")

</div>

For the case of a vector, I think this is right:

```julia
using Flux
using Flux.Tracker: TrackedVector, @grad, track

Base.cumsum(x::TrackedVector) = track(cumsum, x)

@grad function cumsum(x::TrackedVector)
    cumsum(x.data), Δ -> ( reverse(cumsum(reverse(Δ))) ,)
end
    
Tracker.gradcheck(x -> sum(sin, cumsum(x)), randn(3))

```

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 20, 2018, 9:09pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/3 "2018-11-20T21:09:04Z")

</div>

I think I am starting to get it.

But I can’t seem to make it work for the case of an array. See cod ebelow

```julia
Base.cumsum(x::TrackedArray; dims=1) = track(x->cumsum(x, dims = dims), x)

@grad function cumsum(x::TrackedArray; dims=1)
    cumsum(x.data, dims=dims), Δ -> ( reverse(cumsum(reverse(Δ), dims=dims)) ,)
end
    
Tracker.gradcheck(x -> sum(cumsum(x, dims=1)), randn(3,1))

```

---

<div class="post-metadata">

**Author:** ![improbable22](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/improbable22/32/5464_2.png) [@improbable22](https://discourse.julialang.org/u/improbable22)\
**Post date:** [November 21, 2018, 8:53am UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/4 "2018-11-21T08:53:00Z")

</div>

There’s a trick to get around keyword arguments, and you need to tell `reverse` what dimension too. I think this is correct now… if you agree, perhaps worth making a Flux PR?

```julia
using Flux
using Flux.Tracker: TrackedArray, @grad, track

Base.cumsum(x::TrackedArray; dims=1) = track(cumsum, x, dims)

@grad cumsum(x::TrackedArray, dims) = 
    cumsum(x.data, dims=dims), Δ -> ( reverse(cumsum(reverse(Δ, dims=dims), dims=dims), dims=dims) , nothing)
    
Tracker.gradcheck(x -> sum(sin, cumsum(x)), randn(3))
Tracker.gradcheck(x -> sum(sin, cumsum(x, dims=1)), randn(3))

Tracker.gradcheck(x -> sum(sin, cumsum(x, dims=1)), randn(3,4))
Tracker.gradcheck(x -> sum(sin, cumsum(x, dims=2)), randn(3,4))

```

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 21, 2018, 8:59am UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/5 "2018-11-21T08:59:27Z")

</div>

> [@improbable22](#):
>
> There’s a trick to get around keyword arguments, and you need to tell `reverse` what dimension too.

How did you figure out this? Did you look at the source code?

---

<div class="post-metadata">

**Author:** ![improbable22](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/improbable22/32/5464_2.png) [@improbable22](https://discourse.julialang.org/u/improbable22)\
**Post date:** [November 21, 2018, 9:04am UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/6 "2018-11-21T09:04:24Z")

</div>

I don’t remember, maybe? Or perhaps from one of @MikeInnes’s posts on here?

Perhaps the docs could use a PR too, to explain this.

---

<div class="post-metadata">

**Author:** ![MikeInnes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mikeinnes/32/3656_2.png) [@MikeInnes](https://discourse.julialang.org/u/MikeInnes)\
**Post date:** [November 21, 2018, 12:59pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/7 "2018-11-21T12:59:55Z")

</div>

Tracking now supports keyword arguments just fine, [e.g.](https://github.com/FluxML/Flux.jl/blob/4cba46c2936e1cc35386c7c9341ffedeab36c28e/src/tracker/lib/array.jl#L275). Let me know if you have any issues getting this together (there’s actually also a [PR](https://github.com/FluxML/Flux.jl/pull/388) on it that I need to get to).

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [November 21, 2018, 10:14pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/8 "2018-11-21T22:14:23Z")

</div>

I wad doing it on Juliabox so not using latest version. Mayeb thats why?

---

<div class="post-metadata">

**Author:** ![MikeInnes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mikeinnes/32/3656_2.png) [@MikeInnes](https://discourse.julialang.org/u/MikeInnes)\
**Post date:** [November 26, 2018, 11:46am UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/9 "2018-11-26T11:46:02Z")

</div>

Very possibly. You should be able to `add Flux#master` on JuliaBox if you want the latest.

---

<div class="post-metadata">

**Author:** ![floswald](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/floswald/32/195_2.png) [@floswald](https://discourse.julialang.org/u/floswald)\
**Post date:** [November 30, 2018, 8:30pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/10 "2018-11-30T20:30:58Z")

</div>

hey @MikeInnes stupid question: what is the `:` in

```julia
Base.sum(xs::TrackedArray; dims = :) 

```

?  
where to read up on that? thanks.

---

<div class="post-metadata">

**Author:** ![zsunberg](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/zsunberg/32/1883_2.png) [@zsunberg](https://discourse.julialang.org/u/zsunberg)\
**Post date:** [November 30, 2018, 9:00pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/11 "2018-11-30T21:00:00Z")

</div>

@floswald, I think it basically means “all dimensions” here.

[https://docs.julialang.org/en/v1/base/punctuation/index.html](https://docs.julialang.org/en/v1/base/punctuation/index.html)

As a general tip, when googling for this kind of stuff, I find it useful to include the word “julialang”

[https://www.google.com/search?q=julialang+colon+symbol](https://www.google.com/search?q=julialang+colon+symbol)

---

<div class="post-metadata">

**Author:** ![floswald](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/floswald/32/195_2.png) [@floswald](https://discourse.julialang.org/u/floswald)\
**Post date:** [November 30, 2018, 9:21pm UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/12 "2018-11-30T21:21:38Z")

</div>

oh yeah, of course. i mean I know of course what `colon()` means, but i never saw it as an argument of a function. but then again, why not? 🙂

---

<div class="post-metadata">

**Author:** ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)\
**Post date:** [December 1, 2018, 7:35am UTC](https://discourse.julialang.org/t/how-to-create-tracked-cumsum-using-flux-jl/17772/13 "2018-12-01T07:35:27Z")

</div>

Also, in general you can find the source with

```julia
using Flux; methods(sum, (Flux.TrackedArray, ))

```

or, even better,

```julia
edit(first(methods(sum, (Flux.TrackedArray, ))))

```

will open it in your editor directly.
