한 줄 요약
Transformer의 행렬 곱셈 규칙을 대수적 구조로 대체하여 연산량을 줄이며 생성 속도를 6.2–7.8% 향상시킨다.
핵심 기여도
- 기존 행렬 곱셈 대신 **associative unital algebra** 기반의 새로운 곱셈 규칙(예: `𝒬₂`)을 도입.
- `𝒬₂`는 8개 대신 **6개의 block GEMM**만 사용하며, **Alder–Strassen bound**에 따라 최적.
- 110M 파라미터 모델에서 **6.2–7.8%의 end-to-end generation throughput 증가**를 달성.
- 파라미터 수는 동일하면서도, **causal masking 및 KV-cached decoding**과 호환되는 구조를 제시.
핵심 아이디어
기존 Transformer는 행렬 곱셈 `Y = XW`를 기본으로 하며, 이는 행과 열의 블록 간 곱셈 규칙(`row–column rule`)에 의존한다. 본 연구는 이 규칙 자체를 **associative algebra**로 대체함으로써, 동일한 가중치 블록을 사용하면서도 **더 적은 연산량**으로 동작하는 새로운 곱셈 구조를 제안한다.
예를 들어, `q = 2`인 경우, 기존 8개의 block GEMM 대신 **6개의 block GEMM**만 사용하는 `𝒬₂`라는 곱셈 규칙을 정의한다. 이는 특정 블록 간 곱셈(`UW_v`, `W_uA`)을 제거함으로써 달성되며, **off-diagonal cycle**을 닫지 않도록 설계된다.
이 새로운 규칙은 **associative unital algebra**의 성질을 만족하며, 이는 **compositionality**를 보장한다. 즉, `Φ_W(X) := X ⋆₂ W` 형태의 선형 변환은 여러 레이어를 통해 합성되더라도 동일한 형태를 유지한다.
기술적 접근법
- **Associative algebra layer**: 기존 행렬 곱셈 대신, `𝒬₂`와 같은 **associative multiplication table**을 사용.
- **Block GEMM reduction**: 8개 대신 **6개의 block GEMM**만 사용.
- **Alder–Strassen bound**: `R ≥ 2m − v` (m: 슬롯 수, v: 정점 수)를 만족하며, `𝒬₂`는 `R ≥ 6`을 만족하므로 **최적**.
- **Row-typed rectangular projections**: 각 토큰 위치가 특정 row type을 가지며, `type-1`과 `type-2` 토큰이 각각 `(A, U)`와 `(V, D)`를 사용.
- **Causal masking 및 KV-cached decoding**과 호환.
- **Feed-forward layer**에서만 `𝒬₂`를 사용한 실험 수행.
주요 결과
- **110M 파라미터** 모델에서, `𝒬₂` 기반 모델은 **6.2–7.8%의 end-to-end generation throughput 증가**를 기록.
- **3개의 downstream task**에서 점수는 **모두 낮음**.
- **12.3B 토큰**의 학습 데이터를 사용한 동일한 학습 레시피로 모델 학습.
- **파라미터 수**는 동일(`약 110M`)로 유지.
의의 및 한계
- **행렬 곱셈 규칙 자체를 아키텍처 선택 요소로 도입**함으로써, 연산 효율성과 모델 유연성을 동시에 추구.
- `𝒬₂`는 **associativity**를 유지하면서도 연산량을 줄이는 구조로, **Transformer의 연산 최적화**에 기여.
- 그러나 **downstream 성능 저하**가 관찰되어, **정확도-속도 트레이드오프**가 명확하지 않음.
- **한 모델 크기, 한 법칙, 한 학습 런**만 실험하여, 일반화 가능성 제한.
- **더 큰 algebra, attention 메커니즘 적용, 반복 실험**이 필요.
실용적 활용
- **대규모 언어 모델의 decoding 속도 향상**에 활용 가능.
- **GPU 연산 최적화**가 필요한 실시간 추론 시스템에 적합.
- **파라미터 수를 유지하면서 연산량을 줄이는 모델 압축** 기법으로 활용 가능.