Sublayer Connection

Transformer Sublayer Connection 완벽 이해

Transformer 구조를 공부하다 보면 SublayerConnection이라는 클래스를 자주 만나게 됩니다.

대표적으로 다음과 같은 코드입니다.

import torch
import torch.nn as nn

class SublayerConnection(nn.Module):
    def __init__(self, size, dropout=0.1):
        super().__init__()

        self.norm = LayerNorm(size)
        self.dropout = nn.Dropout(p=dropout)

    def forward(self, x, sublayer):

        return self.norm(
            x + self.dropout(sublayer(x))
        )

class LayerNorm(nn.Module):
    def __init__(self, features, eps=1e-5):
        super().__init__()

        # γ (gamma): Scale
        self.gamma = nn.Parameter(torch.ones(features))

        # β (beta): Shift
        self.beta = nn.Parameter(torch.zeros(features))

        self.eps = eps

    def forward(self, x):
        mean = x.mean(-1, keepdim=True)
        var = x.var(-1, keepdim=True, unbiased=False)

        return self.gamma * (x - mean) / torch.sqrt(var + self.eps) + self.beta

Layer Norm:https://deepbored.com/transformer-layer-norm/

FFN:https://deepbored.com/transformer-ffn/

MHA:https://deepbored.com/transformer-multi-head-attention-mha/

Input(Embedding+Position Embedding):https://deepbored.com/pytorch-transformer-input/

Full Code: https://deepbored.com/transformer-sublayer-connection/#full-code

처음 보면 단순히 LayerNorm, Dropout, Residual Connection을 하나로 묶은 코드처럼 보입니다.

하지만 자세히 보면 중요한 질문이 하나 생깁니다.

LayerNorm을 정확히 어디에 적용해야 할까?

LayerNorm의 위치에 따라 대표적으로 Post-Norm, Pre-Norm, 그리고 Post-Norm in Residual과 같은 구조로 나눠볼 수 있습니다.

이번 글에서는 세 구조의 차이를 PyTorch 코드와 함께 간단하게 정리해보겠습니다.


1. Sublayer Connection이란?

Transformer의 Encoder와 Decoder는 여러 개의 Sublayer로 구성됩니다.

예를 들어 대표적인 Sublayer에는 다음과 같은 것들이 있습니다.

Multi-Head Attention
Feed Forward Network
Masked Multi-Head Attention

여기서 중요한 점은 Transformer가 이러한 Sublayer의 출력값을 그대로 다음 층에 전달하지 않는다는 것입니다.

일반적으로 다음 요소들이 함께 사용됩니다.

Sublayer
+
Residual Connection
+
Layer Normalization
+
Dropout

즉, SublayerConnection은 이름 그대로 Sublayer 주변의 연결 구조를 담당하는 클래스라고 이해하면 됩니다.


2. Post-Norm

먼저 Post-Norm 구조입니다.

구조를 단순화하면 다음과 같습니다.

핵심은 Residual Connection으로 더한 다음 LayerNorm을 수행한다는 것입니다.

수식으로 표현하면 다음과 같습니다.

y=LayerNorm(x+Dropout(Sublayer(x)))y = \text{LayerNorm} (x + \text{Dropout}(\text{Sublayer}(x)))

PyTorch 코드로 구현하면 다음과 같습니다.

class PostNormSublayerConnection(nn.Module):
    def __init__(self, size, dropout=0.1):
        super().__init__()

        self.norm = LayerNorm(size)
        self.dropout = nn.Dropout(p=dropout)

    def forward(self, x, sublayer):

        return self.norm(
            x + self.dropout(sublayer(x))
        )

동작 과정을 순서대로 보면,

① x → Sublayer
② Sublayer → Dropout
③ 원래 x와 더하기
④ LayerNorm

입니다.

즉,

Sublayer → Add → Norm

이라고 기억하면 쉽습니다.

특히 2017년 발표된 Transformer 원 논문 Attention Is All You Need에서 설명한 구조가 바로 이 Post-Norm 방식입니다.


3. Pre-Norm

두 번째는 Pre-Norm입니다.

Post-Norm과 비교하면 LayerNorm의 위치가 달라졌습니다.

이번에는 Sublayer에 들어가기 전에 먼저 LayerNorm을 수행합니다.

수식은 다음과 같습니다.

y=x+Dropout(Sublayer(LayerNorm(x)))y = x + \text{Dropout} (\text{Sublayer}(\text{LayerNorm}(x)))

코드는 다음과 같습니다.

class PreNormSublayerConnection(nn.Module):
    def __init__(self, size, dropout=0.1):
        super().__init__()

        self.norm = LayerNorm(size)
        self.dropout = nn.Dropout(p=dropout)

    def forward(self, x, sublayer):

        return x + self.dropout(
            sublayer(self.norm(x))
        )

순서대로 보면,

① x → LayerNorm
② LayerNorm 결과 → Sublayer
③ Sublayer → Dropout
④ 원래 x와 더하기

입니다.

따라서 Pre-Norm은

Norm → Sublayer → Add

라고 기억하면 됩니다.

Transformer 구현을 공부할 때 자주 보게 되는 다음 코드,

return x + self.dropout(
    sublayer(self.norm(x))
)

가 바로 Pre-Norm 구조입니다.

원래 Transformer 논문의 Post-Norm과 코드가 비슷해 보이기 때문에 처음 공부할 때 특히 헷갈리기 쉬운 부분입니다.


4. Post-Norm in Residual

세 번째는 Post-Norm in Residual 구조입니다.

이번에는 Sublayer의 출력에 LayerNorm을 적용한 후 원래 입력 x와 더합니다.

수식으로 표현하면,

y=x+Dropout(LayerNorm(Sublayer(x)))y = x + \text{Dropout} (\text{LayerNorm}(\text{Sublayer}(x)))

입니다.

코드는 다음과 같이 구현할 수 있습니다.

class PostNormInResidualSublayerConnection(nn.Module):
    def __init__(self, size, dropout=0.1):
        super().__init__()

        self.norm = LayerNorm(size)
        self.dropout = nn.Dropout(p=dropout)

    def forward(self, x, sublayer):

        return x + self.dropout(
            self.norm(sublayer(x))
        )

동작 순서는,

① x → Sublayer
② Sublayer 결과 → LayerNorm
③ LayerNorm 결과 → Dropout
④ 원래 x와 더하기

입니다.

따라서,

Sublayer → Norm → Add

라고 기억하면 됩니다.


5. 세 가지 구조 비교

결국 세 구조의 가장 중요한 차이는 LayerNorm이 어디에 위치하는가입니다.

구조연산 순서핵심 코드
Post-NormSublayer → Add → Normnorm(x + sublayer(x))
Pre-NormNorm → Sublayer → Addx + sublayer(norm(x))
Post-Norm in ResidualSublayer → Norm → Addx + norm(sublayer(x))

Dropout까지 포함한 실제 코드를 한 번에 비교하면 더욱 명확합니다.

# ① Post-Norm
return self.norm(
    x + self.dropout(sublayer(x))
)


# ② Pre-Norm
return x + self.dropout(
    sublayer(self.norm(x))
)


# ③ Post-Norm in Residual
return x + self.dropout(
    self.norm(sublayer(x))
)

코드 자체는 거의 비슷합니다.

하지만 self.norm()의 위치가 각각 다릅니다.


6. 그림처럼 생각하면 더 쉽다

세 구조를 가장 간단하게 표현하면 다음과 같습니다.

[ Post-Norm ]

Sublayer
   ↓
   +
   ↓
 Norm
[ Pre-Norm ]

 Norm
   ↓
Sublayer
   ↓
   +
[ Post-Norm in Residual ]

Sublayer
   ↓
 Norm
   ↓
   +

즉, 세 구조를 처음부터 각각 외울 필요는 없습니다.

기준을

Sublayer
   ↓
   +

라고 생각한 뒤,

Norm을 어디에 넣었는가?

만 확인하면 됩니다.


7. Transformer 원 논문은 어떤 구조일까?

여기서 특히 주의해야 할 부분이 있습니다.

2017년 Transformer 원 논문 Attention Is All You Need에서는 각 Sublayer에 Residual Connection을 적용한 뒤 Layer Normalization을 수행한다고 설명합니다.

따라서 원 논문의 구조는

Sublayer
   ↓
Dropout
   ↓
Residual Add
   ↓
LayerNorm

즉,

Post-Norm

입니다.

코드로 다시 표현하면,

return self.norm(
    x + self.dropout(sublayer(x))
)

입니다.

반면 Transformer를 공부하면서 자주 접하게 되는,

return x + self.dropout(
    sublayer(self.norm(x))
)

코드는 Pre-Norm입니다.

따라서 어떤 Transformer 구현 코드를 볼 때 단순히 SublayerConnection이라는 클래스 이름만 보고 원 논문과 동일하다고 생각하기보다는, LayerNorm이 실제로 어디에 위치하는지 확인하는 것이 중요합니다.


8. 정리

2017년 Transformer 원 논문 Attention Is All You Need에서는 Post-Norm 방식을 사용했습니다.

Transformer의 Sublayer Connection에서는 다음 세 가지 요소의 순서를 이해하는 것이 핵심입니다.

Sublayer
Residual Connection (+)
LayerNorm

그리고 세 구조는 다음처럼 기억할 수 있습니다.

Post-Norm
Sublayer → Add → Norm

Pre-Norm
Norm → Sublayer → Add

Post-Norm in Residual
Sublayer → Norm → Add

특히 가장 중요한 차이는 다음과 같습니다.

Post-Norm

norm(x + sublayer(x))

Pre-Norm

x + sublayer(norm(x))

Post-Norm in Residual

x + norm(sublayer(x))

결국 이름을 복잡하게 외우는 것보다 Residual Connection의 +와 Sublayer를 기준으로 LayerNorm이 어디에 있는지 보는 것이 훨씬 이해하기 쉽습니다.

전체 코드:
import torch
import torch.nn as nn
import torch.nn.functional as F

import math
import copy

from torch.autograd import Variable


# ==========================================
# 1. Embedding
# ==========================================

class Embeddings(nn.Module):
    def __init__(self, d_model, vocab):
        super().__init__()

        # vocab개의 Token을 d_model 차원의 Vector로 변환
        self.lut = nn.Embedding(
            vocab,
            d_model
        )

        self.d_model = d_model

    def forward(self, x):

        # Transformer에서는 Embedding 결과에 sqrt(d_model)을 곱함
        return self.lut(x) * math.sqrt(
            self.d_model
        )


# ==========================================
# 2. Positional Encoding
# ==========================================

class PositionalEncoding(nn.Module):
    def __init__(
        self,
        d_model,
        dropout,
        max_len=5000
    ):
        super().__init__()

        # Dropout Layer
        self.dropout = nn.Dropout(
            p=dropout
        )

        # 위치 인코딩 저장용 Tensor
        # [max_len, d_model]
        pe = torch.zeros(
            max_len,
            d_model
        )

        # Token의 절대 위치
        # [max_len, 1]
        position = torch.arange(
            0,
            max_len,
            dtype=torch.float
        ).unsqueeze(1)

        # 차원마다 서로 다른 주기를 만들기 위한 값
        div_term = torch.exp(
            torch.arange(
                0,
                d_model,
                2
            )
            * (
                -math.log(10000.0)
                / d_model
            )
        )

        # 짝수 차원
        pe[:, 0::2] = torch.sin(
            position * div_term
        )

        # 홀수 차원
        pe[:, 1::2] = torch.cos(
            position * div_term
        )

        # Batch Dimension 추가
        #
        # [max_len, d_model]
        # →
        # [1, max_len, d_model]
        pe = pe.unsqueeze(0)

        # 학습 Parameter가 아닌 Buffer로 등록
        self.register_buffer(
            'pe',
            pe
        )

    def forward(self, x):

        # 입력 Sequence Length만큼 위치 인코딩을 잘라 더함
        x = x + Variable(
            self.pe[:, :x.size(1)],
            requires_grad=False
        )

        return self.dropout(x)


# ==========================================
# 3. Scaled Dot-Product Attention
# ==========================================

def attention(
    query,
    key,
    value,
    mask=None,
    dropout=None
):

    # 각 Head의 Dimension
    d_k = query.size(-1)

    # Attention Score 계산
    scores = torch.matmul(
        query,
        key.transpose(-2, -1)
    ) / math.sqrt(d_k)

    # Mask 적용
    if mask is not None:

        # Mask가 0인 위치를 매우 작은 값으로 변경
        scores = scores.masked_fill(
            mask == 0,
            -1e9
        )

    # Attention Weight
    p_attn = F.softmax(
        scores,
        dim=-1
    )

    # Attention Weight에 Dropout
    if dropout is not None:
        p_attn = dropout(p_attn)

    # Attention Output과 Attention Weight 반환
    return (
        torch.matmul(
            p_attn,
            value
        ),
        p_attn
    )


# ==========================================
# 4. Layer Clone 함수
# ==========================================

def clones(module, N):

    # 같은 구조지만 서로 독립적인 Parameter를 가지는 Layer를 N개 생성
    return nn.ModuleList([
        copy.deepcopy(module)
        for _ in range(N)
    ])


# ==========================================
# 5. Multi-Head Attention
# ==========================================

class MultiHeadedAttention(nn.Module):
    def __init__(
        self,
        head,
        embedding_dim,
        dropout=0.1
    ):
        super().__init__()

        # d_model이 Head 개수로 나누어 떨어져야 함
        assert embedding_dim % head == 0

        # 하나의 Head가 사용할 Dimension
        self.d_k = (
            embedding_dim // head
        )

        # Head 개수
        self.head = head

        # Q, K, V, Output용 Linear Layer 총 4개
        self.linears = clones(
            nn.Linear(
                embedding_dim,
                embedding_dim
            ),
            4
        )

        # Attention Weight 저장
        self.attn = None

        # Dropout
        self.dropout = nn.Dropout(
            p=dropout
        )

    def forward(
        self,
        query,
        key,
        value,
        mask=None
    ):

        # Head Dimension을 위한 Dimension 추가
        #
        # [batch, seq, seq]
        # →
        # [batch, 1, seq, seq]
        if mask is not None:
            mask = mask.unsqueeze(1)

        # Batch Size
        batch_size = query.size(0)

        # Q, K, V를 각각 Linear Layer에 전달한 후
        # Head별로 분리
        query, key, value = [
            model(x)
            .view(
                batch_size,
                -1,
                self.head,
                self.d_k
            )
            .transpose(1, 2)

            for model, x in zip(
                self.linears,
                (
                    query,
                    key,
                    value
                )
            )
        ]

        # 각 Head에서 Attention 계산
        x, self.attn = attention(
            query,
            key,
            value,
            mask=mask,
            dropout=self.dropout
        )

        # Head Dimension을 다시 합침
        x = (
            x.transpose(1, 2)
            .contiguous()
            .view(
                batch_size,
                -1,
                self.head * self.d_k
            )
        )

        # 최종 Output Linear Layer
        return self.linears[-1](x)


# ==========================================
# 6. Layer Normalization
# ==========================================

class LayerNorm(nn.Module):
    def __init__(
        self,
        features,
        eps=1e-5
    ):
        super().__init__()

        # gamma:
        # 정규화 결과의 크기를 조정
        self.gamma = nn.Parameter(
            torch.ones(features)
        )

        # beta:
        # 정규화 결과를 이동
        self.beta = nn.Parameter(
            torch.zeros(features)
        )

        self.eps = eps

    def forward(self, x):

        # 마지막 Dimension을 기준으로 평균 계산
        mean = x.mean(
            -1,
            keepdim=True
        )

        # 마지막 Dimension을 기준으로 분산 계산
        var = x.var(
            -1,
            keepdim=True,
            unbiased=False
        )

        return (
            self.gamma
            * (x - mean)
            / torch.sqrt(
                var + self.eps
            )
            + self.beta
        )


# ==========================================
# 7. Sublayer Connection
# ==========================================

class SublayerConnection(nn.Module):
    def __init__(
        self,
        size,
        dropout=0.1
    ):
        super().__init__()

        # Layer Normalization
        self.norm = LayerNorm(size)

        # Dropout
        self.dropout = nn.Dropout(
            p=dropout
        )

    def forward(
        self,
        x,
        sublayer
    ):

        # Post-Norm
        #
        # LayerNorm(
        #     x + Dropout(Sublayer(x))
        # )
        return self.norm(
            x
            + self.dropout(
                sublayer(x)
            )
        )


# ==========================================
# 8. 입력 데이터 생성
# ==========================================

d_model = 512
dropout = 0.1
max_len = 60

vocab = 1000

x = Variable(
    torch.LongTensor([
        [200, 20, 412, 608],
        [572, 968, 10, 322]
    ])
)

print("원본 Token ID:")
print(x)

print("원본 shape:")
print(x.shape)


# ==========================================
# 9. Embedding
# ==========================================

emb = Embeddings(
    d_model,
    vocab
)

embr = emb(x)

print("\nEmbedding 결과:")
print(embr)

print("Embedding shape:")
print(embr.shape)


# ==========================================
# 10. Positional Encoding
# ==========================================

pe = PositionalEncoding(
    d_model,
    dropout,
    max_len
)

pe_result = pe(embr)

print("\nPositional Encoding 결과:")
print(pe_result)

print("Positional Encoding shape:")
print(pe_result.shape)


# ==========================================
# 11. Multi-Head Self-Attention
# ==========================================

size = 512
dropout = 0.2

head = 8
d_model = 512

x = pe_result


# 현재 Batch Size는 2
# Sequence Length는 4
#
# 1 = Attention 허용
# 0 = Attention 차단
#
# 현재는 Padding이 없기 때문에 모든 위치를 1로 설정
mask = torch.ones(
    2,
    4,
    4
)


# Multi-Head Attention 생성
self_attn = MultiHeadedAttention(
    head,
    d_model
)


# Encoder Self-Attention이므로
# Q, K, V가 모두 같은 x
sublayer = lambda x: self_attn(
    x,
    x,
    x,
    mask=mask
)


# SublayerConnection 생성
sc = SublayerConnection(
    size,
    dropout
)


# Attention
# +
# Residual Connection
# +
# LayerNorm
sc_result = sc(
    x,
    sublayer
)


print("\nSublayerConnection 결과:")
print(sc_result)

print("SublayerConnection shape:")
print(sc_result.shape)

실행 결과:

댓글 남기기

이메일 주소는 공개되지 않습니다. 필수 필드는 *로 표시됩니다