pile·
AI / ML·스캐터랩스캐터랩 (이루다)·

꼼꼼하고 이해하기 쉬운 XLNet 논문 리뷰

XLNet이 BERT의 masked language modeling과 GPT 계열 autoregressive 모델의 한계를 어떻게 결합해 풀었는지 수식과 architecture 수준에서 해설한다. 핵심은 token의 실제 위치는 유지하면서 예측 순서만 섞는 permutation language modeling과 target 위치를 알되 정답 content는 보지 않는 two-stream attention이다. Transformer-XL의 상대 위치 encoding과 segment recurrence를 더해 긴 문맥 task에서 큰 성능 향상을 얻었다.

핵심 포인트
  • 단방향 AR 모델은 target 앞의 문맥만 보며, BERT는 양방향을 보지만 여러 `[MASK]` token을 서로 독립적으로 예측하고 pretrain·fine-tune 입력이 다르다.
  • XLNet은 sequence index의 permutation마다 autoregressive factorization을 적용해 mask 없이 양방향 context와 target 간 dependency를 모두 학습한다.
  • 같은 context에서 서로 다른 target을 구분하려면 target position-aware representation이 필요해 query stream과 content stream을 분리했다.
  • Transformer-XL의 relative positional encoding과 cached segment state를 사용해 이전 segment를 다시 계산하지 않고 긴 문맥을 이어간다.
  • 32.89B token을 512 TPU v3에서 2.5일간 학습했으며 RACE, SQuAD, text classification, GLUE, document ranking에서 높은 성능을 냈다.
  • Ablation 결과 permutation objective 자체와 memory caching, span prediction, bidirectional pipeline이 중요했고 next sentence prediction은 대부분의 task에서 오히려 해가 됐다.
상세 정리
  • AR 한계: GPT식 language model은 이전 token으로 다음 token을 예측해 target 사이 dependency를 자연스럽게 학습한다. 그러나 forward나 backward 한 방향만 정해야 하므로 QA처럼 양쪽 문맥이 필요한 표현에 약하다.
  • BERT 한계: `[MASK]`를 복원할 때 양방향 self-attention을 쓸 수 있지만 여러 masked token의 확률을 독립이라고 가정한다. 실제 fine-tuning에는 `[MASK]`가 없어 pretraining 입력과 불일치한다.
  • Permutation objective: 길이 T인 sequence의 가능한 index 순서를 고려하고 각 순서에서 AR likelihood를 최대화한다. token과 positional encoding의 실제 배열은 바꾸지 않고 factorization order만 바꾼다.
  • 양방향 AR: 특정 token을 예측할 때 다른 token의 여러 부분집합을 context로 만난다. 모든 순서를 sampling하면 왼쪽과 오른쪽 정보를 모두 쓰면서도 앞서 예측한 target이 다음 target의 조건이 된다.
  • Target 위치 문제: context가 `x2, x3`로 같아도 다음 target이 `x1`인지 `x4`인지에 따라 분포가 달라야 한다. context만 encoding한 standard Transformer hidden state로는 둘을 구분할 수 없다.
  • Query stream: 현재 target content는 제외하고 이전 permutation step의 content와 target position만 attention한다. 최종 query representation으로 현재 token을 예측해 정답 누출을 막는다.
  • Content stream: 현재 token을 포함한 이전 content를 standard self-attention처럼 encoding한다. 이후 target을 예측할 때 현재 token 정보를 context로 제공한다.
  • Relative position: 절대 위치를 segment마다 0부터 반복하면 과거와 현재의 같은 offset을 구분하지 못한다. attention score에 token 간 상대 거리 encoding을 넣어 content·position bias를 나눠 계산했다.
  • Segment recurrence: 이전 segment의 layer별 content representation을 cache하고 현재 segment의 key·value에 결합한다. 과거 factorization order와 독립적으로 memory를 재사용해 긴 sequence를 처리한다.
  • 학습 규모: BooksCorpus, Wikipedia, Giga5, ClueWeb2012-B, Common Crawl을 정제해 총 32.89B token을 만들었다. XLNet-Large는 batch 2,048, 약 500K step, Adam과 linear decay로 학습됐다.
  • 실험 결과: 긴 passage의 RACE와 SQuAD에서 GPT·BERT보다 큰 폭으로 개선됐다. 여러 text classification의 error rate를 낮추고 GLUE 9개 중 7개, ClueWeb document reranking에서도 당시 최고 수준을 기록했다.
  • Ablation: memory를 빼면 긴 RACE에서 특히 크게 떨어졌다. span-based prediction과 bidirectional input도 기여했으며 next sentence prediction은 RACE를 제외하면 성능을 낮춰 XLNet-Large 학습에서 제외했다.
왜 읽나BERT와 autoregressive language model의 objective 차이가 attention 구조와 긴 문맥 성능에 어떻게 이어지는지 깊게 이해하려는 NLP 엔지니어에게 유용하다.
스캐터랩
스캐터랩 (이루다) 블로그
원문은 여기서 이어서 읽을 수 있어요
원문 읽기
읽음 (0)

이 글과 비슷한

  1. AI / ML·LY CorporationLY Corporation·

    Grafana에서 자연어로 장애 원인을 분석하기: LLM 에이전트 기반 SRELens 개발기

    LY Corporation Home SRE 팀이 장애 분석 시 메트릭·로그·트레이스가 각각 다른 화면에 흩어져 있는 문제를 해결하기 위해 Grafana 플러그인 SRELens를 개발했다. SRELens는 LLM 에이전트가 자연어 질의를 받아 실제 관측성 데이터를 조회하고, 근거와 함께 장애 원인 후보를 정리해 주는 도구다. LGTM-P 스택(Loki·Grafana·Tempo·Mimir·Pyroscope)과 FlavaMCP 게이트웨이를 통합해 단일 채팅 인터페이스에서 멀티시그널 분석이 가능하다.

    #llm-app#mcp#observability+2