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

RetNet은 무엇일까?

~13 min · retnet, retention, microsoft-research

Level 0관찰자
0 XP0/50 lessons0/14 achievements
0/100 XP to next level100 XP to go0% complete

Retention은 지수 감쇠를 품은 순환식이야

RetNet(Sun et al., 2023.7, arXiv:2307.08621, Microsoft Research Asia)은 retention을 도입했어. 순환식으로도, 제약된 어텐션으로도 볼 수 있는 하나의 기본 연산이야. 갱신식은 단순해. s_n = γ · s_{n-1} + K_n^T · V_n이고 출력은 o_n = Q_n · s_n이야. 상태 s는 SSM의 hidden state처럼 크기가 고정되어 있고 과거 토큰 전체를 요약해.

핵심 설계 선택은 감쇠 γ야. RetNet에서는 어텐션 head마다 γ가 고정돼 있고 head마다 값이 달라. γ가 1에 가까우면 천천히 잊어 장기 기억을 유지하고, 0에 가까우면 빨리 잊어 단기 기억을 맡아. 여러 head를 나란히 두면 여러 시간 규모의 기억 계층이 자연스럽게 생겨. 어떤 head는 가까운 과거를 보고 다른 head는 더 오래된 과거를 붙잡는 식이야.

위치 부호화에는 xPos를 써

RetNet은 위치 부호화로 xPos를 써. RoPE와 가까운 복소 지수 회전이지만 retention의 순환 형태에 맞춰 설계됐어. 절대 위치보다 상대적 간격에 집중하도록 이동 불변성을 주면서 순환 계산에도 잘 맞아. 전체 개념을 바꾸는 요소는 아니지만 구현에서는 아주 중요한 세부 사항이야.

순환식과 어텐션이 한 식에서 만났어

retention은 토큰마다 상태를 갱신하면 순환식이 되고, 식을 펼치면 지수 감쇠 마스크를 가진 어텐션 행렬이 돼. 둘이 본질적으로 다른 연산이 아니라 하나의 연속선 위에 놓인 서로 다른 표현임을 구체적으로 보여 준 셈이지. 이 통합이 RetNet의 오래 남은 지적 기여야. 뒤에 나온 linear attention과 SSM 아키텍처는 거의 모두 이 통찰의 한 가지 형태를 받아들였어.

Code

RetNet retention update·python
# Per timestep n, given input x_n:
# Q_n, K_n, V_n produced from x_n by linear projections
# gamma is FIXED per-head (no input dependence — the key constraint)

# Recurrent form (constant memory at inference)
s_n = gamma * s_prev + K_n.transpose(-1, -2) @ V_n
o_n = Q_n @ s_n

# Parallel form (training): equivalent to attention with exponential-decay mask
# attn_mask[i, j] = gamma ** (i - j) for j <= i, else 0

External links

Exercise

retention을 PyTorch로 병렬 형태와 순환 형태 두 가지로 구현해. 작은 (1, 64, 8, 16) Q/K/V/gamma 입력에서 출력이 같은지 확인해 봐. 같은 매개변수가 실행 방식이 달라도 같은 결과를 낸다는 점 덕분에 RetNet을 제품에서 유연하게 활용할 수 있어. 이런 이중 구조가 없는 아키텍처에는 저절로 생기지 않는 장점이야.

Progress

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

댓글 0

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

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