# How to efficiently build AD-compatible matrices line by line

**URL:** <https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632>\
**Category:** Machine Learning\
**Created:** [January 14, 2022, 4:21pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632 "2022-01-14T16:21:14Z")\
**Posts on this page:** 18\
**Page:** 1

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 14, 2022, 4:21pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/1 "2022-01-14T16:21:14Z")

</div>

I find myself frequently needing a construct like the following:

```
a = Array{Float32}(undef, 100, 5)
for i=1:100
    a[i, :] = <a vector>
end

```

which works fine in plain Julia, and is highly efficient, but gives Flux fits when used inside a layer or a loss function because it screws up the AD.

What I am currently doing instead is the following:

```
a = <first vector>
for i=1:100
    a = hcat(a, <next vector>)
end

```

which is terribly inefficient when writing vanilla Julia because of all the allocations that have to be performed.

Is there any way to get the efficiency of memory preallocation using Flux?

---

<div class="post-metadata">

**Author:** ![Oscar\_Smith](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oscar_smith/32/25343_2.png) [@Oscar\_Smith](https://discourse.julialang.org/u/Oscar_Smith)\
**Post date:** [January 14, 2022, 4:23pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/2 "2022-01-14T16:23:11Z")

</div>

Use `a = Array{T}(undef, 100, 5)` where `T` is a type parameter based on the input type of the function. This way, when ad runs on the code, `T` will be a `Dual{Float32}` and everything will work.

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [January 14, 2022, 4:38pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/3 "2022-01-14T16:38:35Z")

</div>

For reverse mode, I think you will want to use `reduce(hcat, all_the_vectors)`, which should be efficient.

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 14, 2022, 5:38pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/4 "2022-01-14T17:38:17Z")

</div>

In trying this, I get an error I’m not able to diagnose:

```
Cannot `convert` an object of type Float32 to an object of type CuArray{Float32, 2, CUDA.Mem.DeviceBuffer}

```

when I try to write to a column of `a`.

Here’s my code:

```
function (m::MyLayer)(x::T) where T
    xi = Array{T}(undef, <sizex>, <sizey>)
    for i=1:size(x,1)
        xi[i, :] = m.W * x[i, :] .+ m.b
    end
    return xi
end

```

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 14, 2022, 6:07pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/5 "2022-01-14T18:07:16Z")

</div>

Ah, got it. `eltype(T)` gives the element type of the array.

---

<div class="post-metadata">

**Author:** ![baggepinnen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/baggepinnen/32/693_2.png) [@baggepinnen](https://discourse.julialang.org/u/baggepinnen)\
**Post date:** [January 14, 2022, 6:10pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/6 "2022-01-14T18:10:04Z")

</div>

Keep in mind that Julia uses column major memory layout, so it might be more efficient to store your data so you can slice it like  
`x[:, i]` instead of `x[i, :]`.  
Also have a look at `@views`

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [January 14, 2022, 6:18pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/7 "2022-01-14T18:18:19Z")

</div>

> [@GlenHenshaw](#):
>
> ```julia
> for i=1:size(x,1)
> xi[i, :] = m.W * x[i, :] .+ m.b
> end
> 
> ```

Also, note that this loop is `ξ = x * m.W' .+ b'`, which will be quicker. But perhaps it’s a toy example.

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 14, 2022, 10:09pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/8 "2022-01-14T22:09:23Z")

</div>

OK, so… here’s the code as it currently stands. It doesn’t work.

```
function (m::MyLayer)(x::T) where T
    xi = Array{eltype{T}}(undef, <sizex>, <sizey>)
    for i=1:size(x,1)
        xi[i, :] = m.W * x[i, :] .+ m.b
    end
    return xi
end

```

The reason it doesn’t work is (I think) that `x` is a `CuArray{Float32}`, but `xi` is just a `Matrix{Float32}`. So the layer gets executed, but everything gets pulled back onto the CPU, which screws up the nest layer down which is still on the GPU.

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 14, 2022, 10:33pm UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/9 "2022-01-14T22:33:03Z")

</div>

Sigh. OK, I got the routine to run… and I’m back to the same AD error. Here’s the code:

```
function (m::MyLayer)(x::T) where T
    xi = similar(T, (<sizex>, <sizey>))
    for i=1:size(x,1)
        xi[i, :] = m.W * x[i, :] .+ m.b
    end
    return xi
end

```

and the error I get is the same one I was getting originally:

```
ERROR: LoadError: Mutating arrays is not supported -- called set_index!(::CuArray{Float32, 2, CUDA.Mem.DeviceBuffer}, _...)

```

Note that `typeof{T}` returns `CuArray{Float32, 2, CUDA.Mem.DeviceBuffer}`. There is not a `Dual{Float32}` in sight.

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 15, 2022, 12:33am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/10 "2022-01-15T00:33:41Z")

</div>

For future reference, the correct answer appears to be using the Zygote.Buffer data structure, which lets you create something that acts like an array but is mutable:

```
using Zygote
function (m::MyLayer)(x::AbstractArray)
    xi = Zygote.Buffer(x, (<sizex>, <sizey>))
    for i=1:size(x,1)
        xi[i, :] = m.W * x[i, :] .+ m.b
    end
    return copy(xi)
end

```

It is not yet clear to me that this is actually more efficient than the `vcat` solution I was originally using, but it does in fact run. It’s kind of hackish that it requires explicit coding for AD…

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [January 15, 2022, 12:36am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/11 "2022-01-15T00:36:10Z")

</div>

The correct solution is `*`, as above.

Unless this is a warm-up problem for something harder. In which case the correct solution is probably something like `reduce(hcat, map(f, xs))`.

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 15, 2022, 12:58am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/12 "2022-01-15T00:58:27Z")

</div>

> [@mcabbott](#):
>
> The correct solution is `*` , as above.

I’m not sure what you mean by this?

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [January 15, 2022, 1:54am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/13 "2022-01-15T01:54:12Z")

</div>

Sorry, perhaps too compressed. But this function does matrix multiplication, by hand. This will be slow, 20x on my computer at the size below, and even worse under AD.

```julia
function loop(W, x::T, b) where T
    xi = Array{eltype(T)}(undef, size(x,1), size(W,1))
    for i=1:size(x,1)
        xi[i, :] = W * x[i, :] .+ b
    end
    return xi
end
blas(W, x, b) = x * W' .+ b'
blas2(W, x, b) = muladd(x, W', b')

x = rand(50,40); W = rand(30,40); b = rand(30);
loop(W,x,b) ≈ blas(W, x, b) ≈ blas2(W, x, b) # true

```

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 15, 2022, 1:56am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/14 "2022-01-15T01:56:27Z")

</div>

> [@mcabbott](#):
>
> Unless this is a warm-up problem for something harder. In which case the correct solution is probably something like `reduce(hcat, map(f, xs))` .

I actually have two different use cases. One does indeed fit into this framework, and I have adopted it.

The other I’m not so sure. What I need to do is the following:

```
x = xi # an nx1 vector
for i=1:l
    x = vcat(x, K*x[end-n:end, :])
end

```

e.g. I’m propagating a linear difference equation for `l` steps and recording the trajectory in x. I can’t figure out how to fit a recursive relationship like this into a `reduce()` framework.

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [January 15, 2022, 2:04am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/15 "2022-01-15T02:04:05Z")

</div>

Well maybe this is messier. Can you make a function that runs, and sample data? `end-n:end` is out of bounds if `x` is an n-vector, what’s the intended behaviour?

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 15, 2022, 2:09am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/16 "2022-01-15T02:09:05Z")

</div>

Yeah, sorry, I should have been more precise. Here’s a running example:

```
function f(K, xi, d)
    x = xi
    for i = 2:d
        x = hcat(x, K*x[:, i-1])
    end
    return x
end

K = rand(3,3)
xi = rand(3,1)
f(K, xi, 50)

```

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [January 15, 2022, 2:46am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/17 "2022-01-15T02:46:46Z")

</div>

OK, this is much messier, I’m not sure there is a really nice answer. It can be written as `accumulate`, which is a bit faster, BUT it forgets the gradient of the `init` keyword, which might be bad.

You could write an efficient gradient for this by hand, in which you just allocate the right size output once, and accumulate in the reverse pass. All these `cat`s and indexing use a lot of memory.

```julia
function f2(K, xi, d::Int)
    xs = accumulate(1:d-1; init=xi) do x, i
        K * x
    end
    hcat(xi, reduce(hcat, xs))
end

function f3(K, xi, d::Int)
    xs = accumulate(vcat([xi], 1:d-1)) do x, i # avoiding init, type-unstable
        K * x
    end
    reduce(hcat, xs)
end

function f4(K, xi, d)
    xs = [xi]
    for i = 2:d
        xs = vcat(xs, [K*xs[i-1]])
    end
    reduce(hcat, xs)
end

f4(K, xi, 50) ≈ f3(K, xi, 50) ≈ f2(K, xi, 50) ≈ f(K, xi, 50)

using BenchmarkTools, Zygote
@btime f($K, $xi, 50);
@btime f2($K, $xi, 50); # twice as quick
@btime f3($K, $xi, 50); # a bit slower
@btime f4($K, $xi, 50);

```

```julia
julia> gradient(sum∘f, K, xi, 10)
([63.45016309970954 50.40609159573776 101.36588271461751; 23.572874731387856 18.379315265377535 35.224999619160954; 31.033286457367566 24.176359057416636 46.03455941092244], [48.455853839178204; 14.765466919614408; 18.845362109436827;;], nothing)

julia> gradient(sum∘f2, K, xi, 10) # NB the gradient for init=xi is missing!
([63.45016309970953 50.40609159573775 101.3658827146175; 23.572874731387852 18.379315265377535 35.22499961916095; 31.033286457367552 24.17635905741663 46.03455941092242], Fill(1.0, 3, 1), nothing)

julia> gradient(sum∘f3, K, xi, 10)
([63.45016309970953 50.40609159573775 101.3658827146175; 23.572874731387852 18.379315265377535 35.22499961916095; 31.033286457367552 24.17635905741663 46.03455941092242], [48.455853839178204; 14.765466919614404; 18.845362109436827;;], nothing)

julia> gradient(sum∘f4, K, xi, 10)
([63.45016309970953 50.40609159573775 101.3658827146175; 23.572874731387852 18.379315265377535 35.22499961916095; 31.033286457367552 24.17635905741663 46.03455941092242], [48.455853839178204; 14.765466919614404; 18.845362109436827;;], nothing)

julia> @btime gradient(sum∘f, $K, $xi, $10);
  min 30.291 μs, mean 34.062 μs (369 allocations, 24.08 KiB)

julia> @btime gradient(sum∘f2, $K, $xi, $10);
  min 21.583 μs, mean 28.247 μs (275 allocations, 36.14 KiB)

julia> @btime gradient(sum∘f3, $K, $xi, $10);
  min 49.917 μs, mean 58.936 μs (544 allocations, 48.08 KiB)

julia> @btime gradient(sum∘f4, $K, $xi, $10);
  min 76.375 μs, mean 89.648 μs (894 allocations, 63.39 KiB)

```

---

<div class="post-metadata">

**Author:** ![GlenHenshaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/glenhenshaw/32/5269_2.png) [@GlenHenshaw](https://discourse.julialang.org/u/GlenHenshaw)\
**Post date:** [January 15, 2022, 3:46am UTC](https://discourse.julialang.org/t/how-to-efficiently-build-ad-compatible-matrices-line-by-line/74632/18 "2022-01-15T03:46:01Z")

</div>

This is amazing, thanks very much for the help!
