본문 바로가기
C.W.K.
Stream
Lesson 07 of 07 · published

실전 — 사용자 정의 Attention 레이어

~10 min · layers

Level 0Keras 도제
0 XP0/97 lessons0/20 achievements
0/120 XP to next level120 XP to go0% complete

제품 코드가 아니라 원리를 이해하려고 직접 만들어

Keras에는 이미 MultiHeadAttention이 있으므로 직접 구현할 이유는 이해에 있어. Attention은 복잡해 보이지만 핵심은 몇 번의 행렬곱이야. keras.ops만으로 만들어보면 Transformer 논문의 수식이 구체적인 코드로 읽히기 시작해.

사용자 정의 레이어의 세 부분

__init__에는 projection 폭인 units 같은 설정을 저장해. build(input_shape)에서는 입력 모양을 안 뒤 query, key, value projection 행렬을 add_weight로 만들고, call(inputs)에서 순전파를 계산해. 생성과 build를 나눈 덕분에 들어오는 입력에 맞춰 가중치 모양을 늦게 정할 수 있어.

Scaled dot-product attention을 순서대로 계산해

입력을 Q, K, V로 투영한 뒤 matmul(q, transpose(k))로 모든 query와 key의 점수를 구해. 이 값을 sqrt(units)로 나누어 차원이 커져도 softmax가 포화하지 않게 하고, softmax 결과를 Attention 가중치로 사용해 value의 가중합을 계산해. 모든 연산을 keras.ops로 작성하면 같은 레이어가 TensorFlow, PyTorch, JAX에서 실행되고 model.fit()으로 projection 가중치도 학습돼.

Code

SimpleAttention — scaled dot-product attention 직접 만들기·python
class SimpleAttention(keras.layers.Layer):
    def __init__(self, units, **kwargs):
        super().__init__(**kwargs)
        self.units = units

    def build(self, input_shape):
        self.W_q = self.add_weight(
            shape=(input_shape[-1], self.units), name="query_weight"
        )
        self.W_k = self.add_weight(
            shape=(input_shape[-1], self.units), name="key_weight"
        )
        self.W_v = self.add_weight(
            shape=(input_shape[-1], self.units), name="value_weight"
        )

    def call(self, inputs):
        q = keras.ops.matmul(inputs, self.W_q)
        k = keras.ops.matmul(inputs, self.W_k)
        v = keras.ops.matmul(inputs, self.W_v)

        # Scaled dot-product attention
        scale = keras.ops.sqrt(
            keras.ops.cast(self.units, dtype="float32")
        )
        scores = keras.ops.matmul(q, keras.ops.transpose(k)) / scale
        weights = keras.ops.nn.softmax(scores)
        return keras.ops.matmul(weights, v)

External links

Exercise

SimpleAttention을 keras.ops만 사용한 multi-head self-attention 레이어로 확장해. Q/K/V projection → num_heads로 reshape → 헤드별 scaled dot-product attention → 이어 붙이기 → 최종 projection 순서로 만들고, 같은 가중치에서 keras.layers.MultiHeadAttention과 출력이 허용 오차 안에서 일치하는지 확인해.
Hint
(batch, seq, units)를 (batch, num_heads, seq, head_dim)로 바꾸고 다시 되돌리는 축 순서를 단계마다 확인해.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.