# Implementation of self-attention in Transformers.jl?

**URL:** <https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732>\
**Category:** Machine Learning\
**Created:** [February 16, 2023, 5:02pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732 "2023-02-16T17:02:50Z")\
**Posts on this page:** 10\
**Page:** 1

<div class="post-metadata">

**Author:** ![rkube](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rkube/32/211198_2.png) [@rkube](https://discourse.julialang.org/u/rkube)\
**Post date:** [February 16, 2023, 5:02pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/1 "2023-02-16T17:02:50Z")

</div>

Hi,  
I’m following [this tutorial](https://www.youtube.com/watch?v=kCc8FmEb1nY) on implementing transformer architectures and have a question on how to best implement a self-attention block.

My current code for a single attention head looks like this:

```julia
# Self-attention head
key = Dense(n_embed, head_size, bias=false)
query = Dense(n_embed, head_size, bias=false)
value = Dense(n_embed, head_size, bias=false)
tril_mask = tril(ones(block_size, block_size)) .== 0

function head(x)
    C, T, B = size(x)
    k = key(x) # (head_size, Token, Batch)
    q = query(x) # (head_size, T, B)
    v = value(x) # (head_size, T, B)
    wts3 = Transformers.batchedmul(q, k, transA=true) ./ sqrt(1f0 * C)
    wts3[tril_mask[1:T, 1:T], :] .= -1f10
    wts3 = softmax(wts3; dims=2) # size (T, T, B)
    out = permutedims(Transformers.batchedmul(wts3, v, transB=true), (2, 1, 3)) #
end

```

The lower-triangular mask `tril` restricts information flow between tokens to preceding tokens. That is, for a 4 token list (T=4) the output looks like

```julia
julia> wts3 = Transformers.batchedmul(q, k, transA=true) ./ sqrt(1f0 * C);

julia> wts3
4×4×1 Array{Float32, 3}:
[:, :, 1] =
  1.03351 2.32426 3.46095 2.27605
 -0.873403 4.64232 1.02749 -3.3062
 -0.329188 2.06334 1.62721 0.66499
  2.08628 -2.60349 -0.965947 1.51771

julia> wts3[tril_mask[1:T, 1:T], :] .= -1f0;

julia> wts3
4×4×1 Array{Float32, 3}:
[:, :, 1] =
  1.03351 -1.0f10 -1.0f10 -1.0f10
 -0.873403 4.64232 -1.0f10 -1.0f10
 -0.329188 2.06334 1.62721 -1.0f10
  2.08628 -2.60349 -0.965947 1.51771

```

The forward pass works, but because of the array mutation, Zygote needs a Buffer to take the gradient.  
Ideally, I’d like to avoid this and looked through [Transformers.jl](https://github.com/chengchingwen/Transformers.jl) to understand how this library implements it.

Also, comparable pytorch implementations [1](https://github.com/karpathy/ng-video-lecture/blob/52201428ed7b46804849dea0b3ccf0de9df1a5c3/gpt.py#L84) [2](http://nlp.seas.harvard.edu/annotated-transformer/#encoder-and-decoder-stacks) use the `masked_fill` operation.

Looking through Transformers.jl, self -attention seems is be implemented [here](https://github.com/chengchingwen/Transformers.jl/blob/29bb6b0407a691264990e43363fe9b1e98ce872c/src/layers/layer.jl#L278-L283)

For causal attention, using the mask above, it looks like `causal=true` is the right argument. So `atten_op_constr = CausalMultiheadQKVAttenOp` in [this line](https://github.com/chengchingwen/Transformers.jl/blob/29bb6b0407a691264990e43363fe9b1e98ce872c/src/layers/layer.jl#L293).

In a forward call, self-attention block applies attention to the Q,K,V projections [here](https://github.com/chengchingwen/Transformers.jl/blob/29bb6b0407a691264990e43363fe9b1e98ce872c/src/layers/layer.jl#L137)

This resolves [here](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/68f7c058a5c90625a41170560279d7fe96d86914/src/functional/attention.jl#L40)  
where the mask in in `args` and `mxiginf = weighted_sum_mixing`.

Followed by [this](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/68f7c058a5c90625a41170560279d7fe96d86914/src/functional/attention.jl#L26)

Attention calculation is deferred to [mixing](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/2aca35e464c93b42d6b1aefda90a9294d65df8ef/src/functional/mixing.jl#L3)  
where `f=mixingf`.

Finally, `weighted_sum_mixing` goes back back `scaled_matmul` [here](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/2aca35e464c93b42d6b1aefda90a9294d65df8ef/src/functional/mixing.jl#L3)

I’m digging down in the code, but can’t find the place where the causal mask is actually applied.  
Does anybody maybe have a pointer?

---

<div class="post-metadata">

**Author:** ![jling](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jling/32/212909_2.png) [@jling](https://discourse.julialang.org/u/jling)\
**Post date:** [February 16, 2023, 5:07pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/2 "2023-02-16T17:07:37Z")

</div>

@chengchingwen

---

<div class="post-metadata">

**Author:** ![reachtarunhere](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/reachtarunhere/32/38358_2.png) [@reachtarunhere](https://discourse.julialang.org/u/reachtarunhere)\
**Post date:** [February 16, 2023, 5:18pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/3 "2023-02-16T17:18:43Z")

</div>

Here is how you can mask without mutation:

```julia
make_decoder_mask(block_size) = tril(fill(Float32(-1f8), block_size, block_size), -1)
mask = make_decoder_mask(block_size)

```

Now after calculating the attention matrix:

`A .+ mask`

If your mask is constant like for training a decoder you would probably save it as some field of your Self Attention layer and make sure that only relevant params are trainable using something like this:

`Flux.trainable(m::MHSelfAttention) = (m.MH_QKV, m.MH_O)`

I have a messy implementation here which supports custom masks ones you might want to use for encoder etc. but it should give you an idea [MakeMore.jl/model.jl at main · reachtarunhere/MakeMore.jl · GitHub](https://github.com/reachtarunhere/MakeMore.jl/blob/main/src/model.jl)

---

<div class="post-metadata">

**Author:** ![rkube](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rkube/32/211198_2.png) [@rkube](https://discourse.julialang.org/u/rkube)\
**Post date:** [February 16, 2023, 5:19pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/4 "2023-02-16T17:19:19Z")

</div>

Instead of a Zygote buffer to mutate the array, one can also just add an upper triangular matrix with large negative values. That has the same effect:

```julia
wts3 = wts3 .+ triu(ones(eltype(wts3), T, T), 1) .* -1f10

```

---

<div class="post-metadata">

**Author:** ![rkube](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rkube/32/211198_2.png) [@rkube](https://discourse.julialang.org/u/rkube)\
**Post date:** [February 16, 2023, 5:23pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/5 "2023-02-16T17:23:00Z")

</div>

Thanks @reachtarunhere . That solves the implementation part ( I had the same idea after spelling it out for this post).

I’m still interested to understand the implementation in NeuralAttention.jl though…

---

<div class="post-metadata">

**Author:** ![reachtarunhere](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/reachtarunhere/32/38358_2.png) [@reachtarunhere](https://discourse.julialang.org/u/reachtarunhere)\
**Post date:** [February 16, 2023, 5:24pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/6 "2023-02-16T17:24:47Z")

</div>

Do checkout [NNlib.jl/attention.jl at master · FluxML/NNlib.jl · GitHub](https://github.com/FluxML/NNlib.jl/blob/master/src/attention.jl) too as the attention op is now part of NNlib and has very readable implementation which supports masking, bias, dropout etc.

---

<div class="post-metadata">

**Author:** ![chengchingwen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chengchingwen/32/8390_2.png) [@chengchingwen](https://discourse.julialang.org/u/chengchingwen)\
**Post date:** [February 16, 2023, 10:15pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/7 "2023-02-16T22:15:11Z")

</div>

The `CausalMultiheadQKVAttenOp` is calling `NeuralAttentionlib.multihead_qkv_attention` with [`NeuralAttentionlib.CausalMask`](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/master/src/mask/dataless.jl#L7), which create a non-allocating broadcastable object indicating the position that would involve in the computation. There’re lots of different kinds of [mask in NeuralAttentionlib.jl](https://github.com/chengchingwen/NeuralAttentionlib.jl/tree/master/src/mask) and the applying part is nothing but [some broadcast](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/master/src/mask/mask.jl#L89). The part you are missing is just the [`masked_score`](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/master/src/functional/score.jl#L22).

---

<div class="post-metadata">

**Author:** ![rkube](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rkube/32/211198_2.png) [@rkube](https://discourse.julialang.org/u/rkube)\
**Post date:** [February 19, 2023, 2:16pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/8 "2023-02-19T14:16:22Z")

</div>

Thanks @chengchingwen

---

<div class="post-metadata">

**Author:** ![rkube](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rkube/32/211198_2.png) [@rkube](https://discourse.julialang.org/u/rkube)\
**Post date:** [February 20, 2023, 8:30pm UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/9 "2023-02-20T20:30:55Z")

</div>

What is the syntax for using masks? @chengchingwen  
This script uses the syntax in the [documentation](https://chengchingwen.github.io/NeuralAttentionlib.jl/stable/api/#NeuralAttentionlib.naive_qkv_attention)  
but throws an error:

```julia
using NeuralAttentionlib

q = ones(Float32, 4)
k = ones(Float32, 4)
v = ones(Float32, 4)

NeuralAttentionlib.naive_qkv_attention(q, k, v; mask=NeuralAttentionlib.:CausalMask)

ERROR: LoadError: MethodError: no method matching naive_qkv_attention(::Vector{Float32}, ::Vector{Float32}, ::Vector{Float32}; mask=NeuralAttentionlib.CausalMask)
Closest candidates are:
  naive_qkv_attention(::AbstractArray, ::AbstractArray, ::AbstractArray, ::Any...) at ~/.julia/packages/NeuralAttentionlib/F0XsF/src/functional/attention.jl:34 got unsupported keyword argument "mask"
  naive_qkv_attention(::typeof(NeuralAttentionlib.score_returning), ::AbstractArray, ::AbstractArray, ::AbstractArray, ::Any...) at ~/.julia/packages/NeuralAttentionlib/F0XsF/src/functional/attention.jl:37 got unsupported keyword argument "mask"
Stacktrace:
 [1] top-level scope
   @ ~/source/repos/tinygpt/src/test_neuralattn.jl:14
in expression starting at /Users/ralph/source/repos/tinygpt/src/test_neuralattn.jl:14

```

I’ve also followed the link to the [source](https://github.com/chengchingwen/NeuralAttentionlib.jl/blob/9894abd75c5279fc305505833210eadbb66a083d/src/module/functional.jl#L39-L66) but that didn’t help.

---

<div class="post-metadata">

**Author:** ![chengchingwen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chengchingwen/32/8390_2.png) [@chengchingwen](https://discourse.julialang.org/u/chengchingwen)\
**Post date:** [February 21, 2023, 4:51am UTC](https://discourse.julialang.org/t/implementation-of-self-attention-in-transformers-jl/94732/10 "2023-02-21T04:51:03Z")

</div>

```julia
using NeuralAttentionlib

q = ones(Float32, 10, 7)
k = ones(Float32, 10, 4)
v = ones(Float32, 10, 4)

NeuralAttentionlib.naive_qkv_attention(q, k, v, NeuralAttentionlib.CausalMask())
NeuralAttentionlib.naive_qkv_attention(NeuralAttentionlib.score_returning, q, k, v, NeuralAttentionlib.CausalMask()).attention_score

```

1. mask should be an object, not type.
2. mask is passed as position argument, not keyword.
3. q/k/v would have the shape of `(hidden size, length 1 size, ..., length n size, batch size)` so `Vector` input is treated as single element (which can’t really see any effect of the attention).
