# Best way to implement shortcut connections for feed-forward neural networks in Flux.jl

**URL:** <https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633>\
**Category:** Machine Learning\
**Tags:** question, package\
**Created:** [May 1, 2018, 9:07am UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633 "2018-05-01T09:07:50Z")\
**Posts on this page:** 11\
**Page:** 1

<div class="post-metadata">

**Author:** ![Azamat](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/azamat/32/6892_2.png) [@Azamat](https://discourse.julialang.org/u/Azamat)\
**Post date:** [May 1, 2018, 9:07am UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/1 "2018-05-01T09:07:50Z")

</div>

Using Flux.jl, I would like to build Feed-Forward Neural Network, but with shortcut connections like those found in ResNet. What would be the best way to do this?  
I know that something like `Chain(Dense(10, 5, relu), Dense(5, 4, relu), Dense(4, 2, relu))` would give me feed-forward NN, but how can I incorporate shortcut connections into that?

---

<div class="post-metadata">

**Author:** ![vvjn](https://avatars.discourse-cdn.com/v4/letter/v/5f9b8f/32.png) [@vvjn](https://discourse.julialang.org/u/vvjn)\
**Post date:** [May 1, 2018, 5:42pm UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/2 "2018-05-01T17:42:35Z")

</div>

I’m more familiar with Knet but I’ll give this a shot. If you look at the docs ([http://fluxml.ai/Flux.jl/stable/models/basics.html](http://fluxml.ai/Flux.jl/stable/models/basics.html)), you will see that you can define your own prediction and loss functions like so

```julia
W = param(rand(2, 5))                                                                                         
b = param(rand(2))                                                                                            
predict(x) = W*x .+ b                                                                                         
loss(x, y) = sum((predict(x) .- y).^2)                                                                        

```

You can use `vcat` in Knet and it looks like you can use it in Flux too. Note that I have not verified whether the gradients are correct when using `vcat` in Flux.

```julia
W1 = param(rand(2, 5))                                                                                        
b1 = param(rand(2))                                                                                           
W2 = param(rand(2, 7))                                                                                        
b2 = param(rand(2))                                                                                           
predict(x) = W2 * vcat(W1*x .+ b1, x) .+ b2                                                                   
loss(x, y) = sum((predict(x) .- y).^2)                                                                        
                                                                                                              
x, y = rand(5), rand(2)                                                                                       
l = loss(x,y)                                                                                                 
Flux.back!(l)                                                                                                 
W1.grad                                                                                                       
b1.grad                                                                                                       
W2.grad                                                                                                       
b2.grad                                                                                                       

```

You can think of a layer like `layer1 = Dense(5,2,σ)` and `layer2 = Dense(7,2,σ)` as functions and define predict similarly using `vcat`:

```julia
predict(x) = layer2(vcat(layer1(x), x))                                                                       

```

---

<div class="post-metadata">

**Author:** ![Evizero](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/evizero/32/10118_2.png) [@Evizero](https://discourse.julialang.org/u/Evizero)\
**Post date:** [May 1, 2018, 7:32pm UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/3 "2018-05-01T19:32:30Z")

</div>

I feel like there are two partial questions here. One is about hooking into this `Chain`, `Dense` syntax, and the other is about resnet skip connections

I am going to naively sketch two types of skip connections, which I will call `IdentitySkip` and `CatSkip`. The first is based on the later resnet variation sometimes referred to as the pre-activation version (see [[1603.05027] Identity Mappings in Deep Residual Networks](https://arxiv.org/abs/1603.05027)). The later is the based on concatenation, similar how Dense Conv Nets do it ([[1608.06993] Densely Connected Convolutional Networks](https://arxiv.org/abs/1608.06993)) and @vvjn sketched. My examples use simple feature matrices instead of higher dimensional arrays, but I hope you get the idea.

(Note that I don’t use Flux much, so take this with a grain of salt. It seems to work though)

```julia
julia> using Flux

julia> struct IdentitySkip
           inner
       end

julia> struct CatSkip
           inner
       end

julia> (m::IdentitySkip)(x) = m.inner(x) .+ x

julia> (m::CatSkip)(x) = vcat(m.inner(x), x)

julia> m = Chain(Dense(2,3), IdentitySkip(Dense(3, 3)), Dense(3,4))
Chain(Dense(2, 3), IdentitySkip(Dense(3, 3)), Dense(3, 4))

julia> m(rand(2,5))
Tracked 4×5 Array{Float64,2}:
  0.806883 -0.0375264 0.139005 0.441874 0.0202739
 -0.447715 0.549833 0.349582 -0.0181219 0.0610884
 -0.474843 0.299503 0.141969 -0.140966 0.0260037
 -0.0398341 0.400103 0.314375 0.149121 0.0534565

```

---

<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:** [May 2, 2018, 11:50am UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/4 "2018-05-02T11:50:48Z")

</div>

One step to add here is to tell Flux where to find the parameters inside these types:

```julia
julia> Flux.params(m) |> length
4

julia> Flux.treelike(IdentitySkip)

julia> Flux.params(m) |> length
6

```

---

<div class="post-metadata">

**Author:** ![osofr](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/osofr/32/3857_2.png) [@osofr](https://discourse.julialang.org/u/osofr)\
**Post date:** [May 4, 2018, 5:53pm UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/5 "2018-05-04T17:53:55Z")

</div>

@Evizero, the `IdentitySkip ` will work for skipping the most recent transformation, k=1. Do you think it might be possible to modify this struct to allow skipping for some k layers past, where k\>1?

---

<div class="post-metadata">

**Author:** ![Evizero](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/evizero/32/10118_2.png) [@Evizero](https://discourse.julialang.org/u/Evizero)\
**Post date:** [May 5, 2018, 6:35pm UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/6 "2018-05-05T18:35:16Z")

</div>

There shouldn’t be anything special about just having one inner operation. Without testing it myself, I am guessing it should probably just work fine if you write `IdentitySkip(Chain(Dense(3,3),Dense(3,3)))`. Alternatively, add more member variables or replace it with a vector/tuple of inner operations.

---

<div class="post-metadata">

**Author:** ![Azamat](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/azamat/32/6892_2.png) [@Azamat](https://discourse.julialang.org/u/Azamat)\
**Post date:** [May 18, 2018, 12:50pm UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/7 "2018-05-18T12:50:35Z")

</div>

Amazing! Thanks a lot @Evizero!

---

<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:** [October 13, 2018, 11:28pm UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/8 "2018-10-13T23:28:28Z")

</div>

Would it better to allow for the specification of an activation function?

```julia
struct IdentitySkip
   inner
   activation
end

(m::IdentitySkip)(x) = m.activation.(m.inner(x) .+ x)

```

This is how I understand it from Andrew Ng’s lectures

```julia
policy = Flux.Chain(
  Dense(16*14, 128, relu),
  IdentitySkip(Dense(128, 128), relu),
  Dense(128, 32, relu),
  IdentitySkip(Dense(32, 32), identity),
  Dense(32, 4),
  softmax
  )

```

---

<div class="post-metadata">

**Author:** ![Evizero](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/evizero/32/10118_2.png) [@Evizero](https://discourse.julialang.org/u/Evizero)\
**Post date:** [October 14, 2018, 9:53am UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/9 "2018-10-14T09:53:11Z")

</div>

> [@xiaodai](#):
>
> Would it better to allow for the specification of an activation function?

You can combine it however you desire. I didn’t include any activiation functions simply because the preactivation formulation used for the identity skip i reference uses a different ordering of things (which would just make everything look more complicated than it needs to).

Note though that the version you propose (and indeed the one Prof. Ng discusses) is from the original Resnet paper, and not the identity skip version from the later revision that i reference. see the paper I linked in my earlier post.

---

<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:** [January 23, 2019, 11:16am UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/10 "2019-01-23T11:16:56Z")

</div>

> [@Evizero](#):
>
> (Note that I don’t use Flux much, so take this with a grain of salt. It seems to work though)
> 
> ```julia
> julia> using Flux
> 
> julia> struct IdentitySkip
> inner
> end
> 
> ```

Causing a crashing if I use the GPU. See

```julia
using Flux, Flux.Data.MNIST, Statistics
using Flux: onehotbatch, onecold, crossentropy, throttle
using Base.Iterators: repeated, partition
using CuArrays
# Classify MNIST digits with a convolutional network

imgs = MNIST.images()

labels = onehotbatch(MNIST.labels(), 0:9)

# Partition into batches of size 32
train = [(cat(float.(imgs[i])..., dims = 4), labels[:,i])
         for i in partition(1:60_000, 32)]

train = gpu.(train)

# Prepare test set (first 1,000 images)
#tX = cat(float.(MNIST.images(:test)[1:1000])..., dims = 4) |> gpu
tX = reshape(reduce(hcat, vec.(float.(MNIST.images(:test)))),28,28,1,10_000) |> gpu
tY = onehotbatch(MNIST.labels(:test), 0:9) |> gpu

trainX = reshape(reduce(hcat, vec.(float.(MNIST.images()))),28,28,1,60_000) |> gpu
trainY = onehotbatch(MNIST.labels(), 0:9) |> gpu

struct IdentitySkip
   inner
end

(m::IdentitySkip)(x) = m.inner(x) .+ x

m = Chain(
    Conv((2, 2), 1=>32, relu),
    x -> maxpool(x, (2,2)),
    Conv((2, 2), 32=>32, relu),
    x -> maxpool(x, (2,2)),
    Conv((2, 2), 32=>32, relu),
    x -> reshape(x, :, size(x, 4)),
    Dense(800, 100, relu),
    IdentitySkip(Dense(100, 100, relu)),
    Dense(100, 10),
    softmax) |> gpu

```

---

<div class="post-metadata">

**Author:** ![bdeonovic](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bdeonovic/32/3928_2.png) [@bdeonovic](https://discourse.julialang.org/u/bdeonovic)\
**Post date:** [February 21, 2020, 3:19pm UTC](https://discourse.julialang.org/t/best-way-to-implement-shortcut-connections-for-feed-forward-neural-networks-in-flux-jl/10633/11 "2020-02-21T15:19:32Z")

</div>

It looks like this was implemented directly in Flux ([Model Reference · Flux](https://fluxml.ai/Flux.jl/stable/models/layers/#Flux.SkipConnection)) can you give that a try?
