[nlp][논문 리뷰] BART: Denoising Sequence-to-Sequence Pre-training for Natural Language Generation, Translation, and Comprehension

이영락·2024년 10월 23일

CV & NLP 논문 리뷰

목록 보기
9/14

status: Finished!
스터디 주제: 주제: BART 응용 및 실습
논문:

• BART 모델 구조 및 특징 이해
• denoising 기법의 적용방법 및 효과
• time step 별 이해 (croos attention > encoer > decoder)
• 어떠한 finetuning 태스크에 적합한지 조사 및 학습

실습: BART 를 활용한 텍스트 요약 실습
주차: 4주차

참고자료


BART 논문 리뷰 / BART: Denoising Sequence-to-Sequence Pre-training for Natural Language Generation, Translation, and Comprehension

1 | 논문정리


BART: Denoising Sequence-to-Sequence Pre-training for Natural Language Generation, Translation, and Comprehension

1. Introduction

Self-supervised learning의 배경

최근 NLP 분야에서 자기 지도 학습(self-supervised learning) 방법이 큰 성공을 거둠.

  • Mikolov et al. (2013)

  • Peters et al. (2018)

  • Devlin et al. (2019)

  • Joshi et al. (2019)

  • Yang et al. (2019)

  • Liu et al. (2019)

가장 성공적인 접근 방식은 마스크된 언어 모델(masked language models).

🤔

What is Masked Lanuage model ??

: 노이즈가 있는 데이터를 복원하는 방식의 오토인코더

  • 입력 텍스트에서 일부 단어를 무작위로 마스크 처리하고 이를 복원하도록 훈련.

기존 연구의 한계와 개선 사항

기존 연구

  1. 마스크 토큰의 분포 개선(Joshi et al., 2019)
  2. 예측 순서 개선(Yang et al., 2019)
  3. 마스크 토큰을 대체하는 데 사용할 문맥 개선(Dong et al., 2019)

→ 이러한 방법들은 특정 작업(예: 스팬 예측, 텍스트 생성 등)에 집중하는 경향이 있어 그 적용 범위가 제한적이라는 한계가 있습니다.

BART의 도입

🤔

BART??

: Bidirectional(양방향)Auto-Regressive(자기 회귀) 트랜스포머를 결합한 모델

: 시퀀스-투-시퀀스(sequence-to-sequence) 모델을 사용한 디노이징 오토인코더

STEP1. 텍스트에 임의의 노이즈가 추가.

STEP2. 시퀀스-투-시퀀스 모델을 통해 원래 텍스트를 복원하는 작업이 진행

BART는 BERT(양방향 인코더 사용)와 GPT(좌측에서 우측으로 디코더 사용)를 일반화한 모델로, 더 최신의 여러 사전 훈련 방식들 역시 포괄합니다.

BART의 주요 특징: 노이징 유연성

👀

주요 장점 : 원본 텍스트에 대해 임의의 변환을 적용할 수 있는 노이징 유연성

  • 텍스트의 길이를 변경하는 것과 같은 임의의 변환도 가능
  • 텍스트의 문장 순서를 무작위로 섞거나, 길이에 상관없이 텍스트 스팬을 마스크 토큰으로 대체하는 새로운 인필링(in-filling) 방식을 사용

→ BERT의 단어 마스킹과 다음 문장 예측 목적을 일반화하여, 모델이 문장 길이를 더 깊이 있게 추론하고 입력에 대해 더 넓은 범위의 변환을 수행

BART의 성능

BART는 특히 텍스트 생성 작업에서 강력한 성능을 발휘하지만, 이해(comprehension) 작업에서도 우수한 성능을 보입니다.

GLUE(Wang et al., 2018) 및 SQuAD(Rajpurkar et al., 2016)와 같은 데이터셋에서 BART는 RoBERTa(Liu et al., 2019)와 비슷한 성능을 보였으며, 추상화 대화, 질문 응답, 요약 작업에서 새로운 최첨단 성능을 달성했습니다. 예를 들어, XSum(Narayan et al., 2018)에서 이전 연구보다 6 ROUGE 점수가 개선되었습니다.

새로운 파인튜닝 방식

BART 모델을 몇 개의 추가적인 트랜스포머 레이어 위에 쌓아 기계 번역 작업에서 사용했습니다.

이 레이어들은 외국어를 노이즈가 포함된 영어로 번역하도록 훈련되었으며, 이를 통해 BART는 사전 훈련된 타깃 언어 모델로 활용됩니다. 이 방법은 WMT 루마니아어-영어 벤치마크에서 강력한 역번역(Back-Translation) MT 기준보다 1.1 BLEU 점수가 향상되었습니다.

Ablation 분석

이 논문에서는 다양한 훈련 목적을 재현하는 ablation 분석을 수행하여 BART의 성능에 영향을 미치는 요인들을 평가. 이 연구는 데이터와 최적화 파라미터를 통제하여, 특정 훈련 목적 선택뿐만 아니라 이러한 요소들이 전체 성능에 중요한 영향을 미친다는 것을 보여줍니다.

BART는 이 논문에서 고려한 모든 작업에서 일관되게 강력한 성능을 보였습니다.

https://dladustn95.github.io/assets/images/bart_figure1.png


🚨

각 시퀀스에 대해 라벨 yy 를 예측하는 확률 수식

P(yx1,,xm)=softmax(hmlWy)P(y | x_1, \dots, x_m) = \text{softmax}(h_m^l W_y)

  • Softmax는 모델의 출력값을 확률 값으로 변환하는 역할을 하며, 이를 통해 모델이 특정 작업의 라벨 예측을 수행할 수 있습니다.

2.Model

2.1 Architecture

  • seq2seq 트랜스포머 구조를 사용 → 손상된 텍스트에 대해 양방향 인코더를 사용하고, 좌에서 우로 진행하는 자기 회귀 디코더를 사용함
  • ReLU activation function을 GeLUs로 변경
  • 기본 모델(base model): 인코더와 디코더에 각각 6개의 레이어를 사용합니다.
  • 대형 모델(large model): 인코더와 디코더에 각각 12개의 레이어를 사용합니다.
  • 문서의 부정 로그 가능도(negative log likelihood)를 최적화

BART의 아키텍처는 BERT와 매우 유사하지만, 다음과 같은 차이가 있습니다:

  1. 디코더의 각 레이어가 인코더의 최종 히든 레이어에 대해 교차 주의(cross-attention)를 수행합니다. 이는 트랜스포머 시퀀스-투-시퀀스 모델에서 주로 사용되는 방식입니다.
  2. BERT는 단어 예측 전에 추가적인 피드포워드 네트워크를 사용하는 반면, BART는 이를 사용하지 않습니다.

결과적으로, BART는 동일한 크기의 BERT 모델보다 약 10% 더 많은 파라미터를 포함합니다.

🤔

Cross Attention??

: 노이즈가 있는 데이터를 복원하는 방식의 오토인코더

  • 입력 텍스트에서 일부 단어를 무작위로 마스크 처리하고 이를 복원하도록 훈련.

2.2 Pretraining BART

BART는 손상된 text로 학습하는데 디코더의 출력과 원본 text의 loss를 줄이도록 한다. 다른 Auto-Encoder 모델과 다르게 모든 종류의 noise를 적용할 수 있다.

즉, BART는 문서를 손상시키고, 이를 복원하는 손실 함수(디코더의 출력과 원래 문서 간의 교차 엔트로피)를 최적화하여 훈련됩니다. 기존의 디노이징 오토인코더는 특정 노이즈 스키마에 맞춰져 있지만, BART는 문서 손상에 다양한 방식으로 대응할 수 있습니다. 가장 극단적인 경우에는, 원본에 대한 모든 정보가 손실되면 BART는 일반적인 언어 모델과 동일하게 동작합니다

https://dladustn95.github.io/assets/images/bart_figure2.png

이 논문에서는 그림과 같이 5가지의 noise 기법을 사용했다.

👀

Noise 기법 5가지

  1. Token Masking : BERT처럼 랜덤 토큰을 masking하고 이를 복구하는 방식이다.(무작위로 선택된 토큰들을 [MASK] 요소로 대체.)
  2. Token Deletion: 랜덤 토큰을 삭제하고 이를 복구하는 방식이다. 토큰 마스킹과 달리, 모델은 어떤 위치의 입력이 빠졌는지 추정해야 한다.
  3. Text Infilling: 포아송 분포((\lambda = 3))로부터 샘플링된 길이의 텍스트 스팬을 선택하고, 각 스팬을 단일 [MASK] 토큰으로 대체.
    1. 길이가 0인 스팬은 [MASK] 토큰이 삽입된 것을 의미.
    2. 텍스트 인필링은 SpanBERT(Joshi et al., 2019)에서 영감을 받았지만, SpanBERT는 클램프된 기하 분포에서 샘플링하여 동일한 길이의 [MASK] 토큰으로 각 스팬을 대체.
    3. 텍스트 인필링은 모델이 스팬에서 얼마나 많은 토큰이 누락되었는지를 예측하도록 학습
  4. Sentence Permutaion: Document를 문장 단위로 나눠서 섞는 방법이다.
  5. Document Rotation: 무작위로 선택된 토큰을 문서의 시작으로 하여 문서가 회전. 이 작업은 모델이 문서의 시작점을 식별하는 능력을 훈련.

3. Fine-tuning BART

3. BART의 사전 학습 (Pre-training BART)

  • 사전 학습 방법:
    • BART는 문서 손상(damage)복원(reconstruction)의 두 가지 주요 단계로 구성.
      • 먼저 입력 문서는 특정 규칙에 따라 손상됩니다. 예를 들어, 문서의 일부 토큰이 마스킹되거나 문장이 무작위로 재배열됩니다.
      • 그 후 모델은 손상된 입력을 복원하도록 학습됩니다.
    • 음의 로그 우도(Negative Log-Likelihood)를 최적화하는 방식으로 학습이 이루어지며, 이는 모델이 원본 문서를 최대한 정확하게 재구성할 수 있도록 합니다.

3.1 시퀀스 분류 작업 (Sequence Classification Tasks)

https://dladustn95.github.io/assets/images/bart_figure3.png

시퀀스 분류 작업에서는 동일한 입력을 인코더와 디코더에 넣고, 마지막 디코더 토큰의 최종 히든 상태를 새로운 다중 클래스 선형 분류기에 전달합니다. 이 방식은 BERT에서 사용된 CLS 토큰과 유사하지만, 우리는 추가 토큰을 끝에 추가하여 디코더에서 해당 토큰의 표현이 전체 입력에서 온 디코더 상태에 attend할 수 있도록 합니다 (Figure 3a 참조).

3.2 토큰 분류 작업 (Token Classification Tasks)

토큰 분류 작업(예: SQuAD에서의 답변 끝점 분류 작업)에서는 전체 문서를 인코더와 디코더에 입력하고, 디코더의 최상위 히든 상태를 각 단어에 대한 표현으로 사용합니다. 이 표현은 토큰을 분류하는 데 사용됩니다.

3.3 시퀀스 생성 작업 (Sequence Generation Tasks)

BART는 자기 회귀 디코더를 사용하기 때문에, 추상적 질문 응답 및 요약과 같은 시퀀스 생성 작업에 직접적으로 파인 튜닝할 수 있습니다. 이 두 작업에서는 입력에서 정보를 복사하지만, 입력을 조작(manipulate)하는데, 이는 BART의 디노이징 사전 훈련 목표와 밀접한 관련이 있습니다. 여기서 인코더 입력은 입력 시퀀스이고, 디코더는 자기 회귀 방식으로 출력을 생성합니다.

3.4 기계 번역 (Machine Translation)

https://dladustn95.github.io/assets/images/bart_figure4.png

BART는 영어로 번역하는 기계 번역 디코더를 개선하는 데에도 사용. 전체 BART 모델(인코더와 디코더 모두)을 하나의 사전 훈련된 디코더로 사용하여 기계 번역 성능을 향상시킬 수 있음을 보여줍니다. 이를 위해 새로운 인코더 파라미터를 추가하고, 이를 병렬 텍스트로부터 학습합니다 (Figure 3b 참조).

구체적으로는, BART의 인코더 임베딩 레이어를 새로운 임의 초기화된 인코더로 대체합니다. 모델은 처음부터 끝까지(end-to-end) 학습되며, 새로운 인코더는 외국어를 영어로 디노이즈할 수 있도록 학습됩니다. 이 새로운 인코더는 BART 모델과 별개의 어휘(vocabulary)를 사용할 수 있습니다.

인코더는 두 단계로 학습됩니다:

  1. 첫 번째 단계에서는 BART의 대부분의 파라미터를 고정하고, 임의로 초기화된 소스 인코더, BART의 위치 임베딩, 그리고 BART 인코더의 첫 번째 레이어에서 사용되는 자기 주의 입력 투영 행렬만 업데이트합니다.
  2. 두 번째 단계에서는 모델의 모든 파라미터를 소수의 반복(iteration) 동안 학습합니다.

이 과정을 통해 BART는 기계 번역 작업에서도 효율적으로 활용될 수 있습니다.

4. ****Comparing Pre-training Objectives

BART는 이전 연구에 비해 사전 훈련 중 사용할 수 있는 노이징 스키마의 범위가 매우 넓습니다. 다양한 사전 훈련 방식에 대해 BART의 성능을 비교하기 위해, 우리는 6개의 인코더 레이어와 6개의 디코더 레이어, 히든 크기 768로 구성된 Base 모델을 사용하여 실험을 진행했습니다. 이 실험은 다양한 작업에 대해 BART의 성능을 비교하며, 특히 사전 훈련 목적에 따른 성능 차이를 분석합니다.

4.1 비교 목표 (Comparison Objectives)

사전 훈련에 대해 제안된 많은 목적들이 있지만, 공정한 비교는 여러 요인들로 인해 어려웠습니다. 이러한 요인에는 훈련 데이터의 차이, 훈련 자원의 차이, 모델의 아키텍처 차이 및 파인 튜닝 절차의 차이가 포함됩니다. 우리는 가능한 한 이러한 차이를 통제하려 했으며, 성능을 개선하기 위해 각 목표에 맞게 학습률레이어 정규화를 조정했습니다.

아래에서 우리는 BERT와 같은 100만 스텝 동안 훈련된 모델들과의 비교를 통해 몇 가지 사전 훈련 목표를 평가했습니다.

  1. Language Model (언어 모델):

    GPT와 유사하게 왼쪽에서 오른쪽으로 문맥을 예측하는 트랜스포머 기반 언어 모델을 훈련합니다. 이 모델은 BART의 디코더와 유사하지만 크로스 어텐션이 없습니다.

  2. Permuted Language Model (순서 뒤섞인 언어 모델):

    XLNet에 기반하여 6분의 1의 토큰을 샘플링하고, 이를 무작위 순서로 자기 회귀적으로 생성합니다.

  3. Masked Language Model (마스크된 언어 모델):

    BERT와 유사하게 15%의 토큰을 [MASK] 심볼로 대체하고, 모델이 원래 토큰을 독립적으로 예측하도록 훈련합니다.

  4. Multitask Masked Language Model (멀티태스크 마스크된 언어 모델):

    UniLM의 접근 방식을 사용해 마스크된 언어 모델을 훈련하되, 추가적인 자기 주의 마스크를 사용합니다. 마스크는 왼쪽에서 오른쪽, 오른쪽에서 왼쪽, 비마스크 상태, 50% 마스크 상태 등의 비율로 무작위로 선택됩니다.

  5. Masked Seq-to-Seq (마스크된 시퀀스-투-시퀀스 모델):

    MASS 모델에서 영감을 받아 전체 토큰의 50%가 포함된 스팬을 마스킹하고, 이 마스킹된 토큰들을 예측하도록 시퀀스-투-시퀀스 모델을 훈련합니다.

비교 방법: 우리는 각 모델의 파인 튜닝 목표를 효과적으로 모델링하는 능력을 비교하기 위해 퍼플렉시티(perplexity)를 보고합니다. 이를 통해 각 모델이 문장을 어떻게 이해하고 생성하는지 평가합니다.

4.2 작업(Task)

  • SQuAD: 위키피디아 문단을 기반으로 한 추출적 질문 응답 작업입니다. 모델은 문서에서 답변이 포함된 텍스트 스팬을 예측해야 합니다. BERT와 유사하게 질문과 문맥을 인코더에 입력하고, 답변 시작과 종료 지점을 예측하는 분류기를 사용합니다.
  • MNLI: 이중 텍스트 분류 작업으로, 한 문장이 다른 문장을 포함하거나 반박하는지를 예측합니다.
  • ELI5: 장문 형태의 추상적 질문 응답 데이터셋으로, 질문과 지원 문서를 기반으로 답변을 생성하는 작업입니다.
  • XSum: 뉴스 요약 작업으로, 매우 추상적인 요약이 필요한 데이터셋입니다.
  • ConvAI2: 대화 응답 생성 작업으로, 대화 문맥과 개인화된 정보를 기반으로 응답을 생성해야 합니다.
  • CNN/DM: 뉴스 요약 작업으로, 요약은 주로 원본 문장과 유사한 형태로 생성됩니다.

4.3 결과 (Results)

스크린샷 2024-09-25 오후 4.27.45.png

  1. 작업에 따른 사전 훈련 방법의 성능 차이:

    사전 훈련 방법의 효과는 작업에 따라 크게 달라짐.

    예를 들어, 단순한 언어 모델은 ELI5 작업에서 가장 높은 성능을 보였지만, SQuAD에서는 최악의 성능을 보였습니다.

  2. 토큰 마스킹의 중요성:

    문서 회전 또는 문장 순서 뒤섞기를 기반으로 한 사전 훈련 방법은 독립적으로는 성능이 낮았습니다. 성공적인 방법들은 토큰 삭제 또는 마스킹, 또는 자기 주의 마스크를 사용했습니다. 토큰 삭제가 생성 작업에서 마스킹보다 더 나은 성능을 보여주었습니다.

  3. 좌에서 우로의 사전 훈련이 생성 작업을 개선:

    마스크된 언어 모델순서 뒤섞인 언어 모델은 생성 작업에서 상대적으로 낮은 성능을 보였습니다. 이 모델들은 사전 훈련 중 좌에서 우로의 자기 회귀적 언어 모델링을 포함하지 않았습니다.

  4. 양방향 인코더가 SQuAD에 중요:

    이전 연구에서 밝혀진 것처럼, 좌에서 우로의 디코더만으로는 SQuAD에서 성능이 떨어집니다. 이는 미래 문맥이 분류 작업에서 중요한 역할을 하기 때문입니다. 그러나 BART는 양방향 레이어 수가 절반임에도 유사한 성능을 달성했습니다.

  5. 사전 훈련 목표 외에도 아키텍처가 중요:

    Permuted Language Model은 XLNet보다 성능이 낮았는데, 이는 XLNet이 상대적 위치 임베딩 또는 세그먼트 수준 반복성과 같은 추가 아키텍처 개선 사항을 포함했기 때문입니다.

  6. ELI5에서는 순수한 언어 모델이 가장 성능이 좋음:

    ELI5 데이터셋은 다른 작업들에 비해 훨씬 높은 퍼플렉시티를 기록했으며, BART가 아닌 순수한 언어 모델이 최고의 성능을 기록했습니다. 이는 BART가 입력 문장과 출력 간의 관계가 느슨한 작업에서는 덜 효과적임을 시사합니다.

  7. BART는 가장 일관되게 강력한 성능을 발휘:

    텍스트 인필링을 사용한 BART 모델은 거의 모든 작업에서 좋은 성능을 보였으며, 특히 생성 작업에서 두각을 나타냈습니다.

5. Large-scale Pre-training Experiments

5.1 실험 설정 (Experimental Setup)

우리는 인코더와 디코더에 각각 12개의 레이어, 1024개의 히든 크기를 가진 대형 모델을 사전 훈련했습니다. RoBERTa와 동일하게 배치 크기 8000을 사용하여 50만 단계 동안 모델을 훈련했습니다. 문서는 GPT-2에서 사용된 것과 동일한 바이트 쌍 인코딩(byte-pair encoding)으로 토크나이즈되었습니다(Radford et al., 2019). 섹션 4의 결과를 바탕으로 텍스트 인필링과 문장 순서 뒤섞기를 결합하여 사용했습니다. 각 문서에서 30%의 토큰을 마스크 처리하고, 모든 문장의 순서를 섞었습니다.

우리는 훈련의 마지막 10% 동안 드롭아웃을 비활성화하여 모델이 데이터에 더 잘 맞도록 했습니다. 사전 훈련 데이터는 Liu et al.(2019)이 사용한 것과 동일한 160GB의 뉴스, 책, 이야기, 웹 텍스트로 구성되었습니다.

5.2 판별 작업 (Discriminative Tasks)

BART는 SQuAD 및 GLUE와 같은 널리 연구된 판별 작업에서 여러 최신 모델들과 비교되었습니다. BART의 가장 직접적인 비교 대상은 RoBERTa입니다. RoBERTa는 동일한 자원을 사용하여 훈련되었지만, 다른 목적을 가지고 훈련되었습니다. 전반적으로 BART는 대부분의 작업에서 유사한 성능을 보였으며, 이는 BART의 생성 작업에서의 성능 향상이 분류 작업 성능에 손실을 주지 않았음을 시사합니다. (Table 2 참조)

5.3 생성 작업 (Generation Tasks)

BART는 여러 텍스트 생성 작업에서도 실험되었습니다. BART는 시퀀스-투-시퀀스 모델로 입력 텍스트에서 출력 텍스트로 파인 튜닝되었습니다. 파인 튜닝 시 레이블 스무딩 교차 엔트로피 손실(label smoothed cross entropy loss)을 사용했으며, 스무딩 매개변수는 0.1로 설정되었습니다. 생성 시 빔 사이즈는 5로 설정했고, 빔 검색 중 중복된 트라이그램을 제거하며, 검증 세트에서 min-len, max-len, length penalty 등을 조정했습니다.

요약 작업 (Summarization)

CNN/DailyMail과 XSum이라는 두 개의 요약 데이터셋에서 실험 결과를 제시했습니다. CNN/DailyMail은 소스 문장과 유사한 요약을 생성하며, 추출적 모델이 이 작업에서 잘 수행됩니다. 반면, XSum은 더 추상적이기 때문에 추출적 모델은 이 작업에서 성능이 저조합니다. BART는 두 데이터셋 모두에서 이전 연구들보다 우수한 성능을 보였으며, 특히 XSum에서는 ROUGE 메트릭에서 약 6점의 성능 향상을 기록했습니다. (Table 3 참조)

대화 생성 작업 (Dialogue Generation)

BART는 CONVAI2에서 대화 응답 생성 성능도 평가되었습니다. 여기서 BART는 두 가지 자동화된 메트릭에서 이전 연구들보다 우수한 성능을 보였습니다. (Table 4 참조)

추상적 질문 응답 (Abstractive Question Answering)

ELI5 데이터셋을 사용하여 모델의 자유 형식(free-form)의 긴 답변 생성 능력을 테스트했습니다. BART는 ROUGE-L에서 이전 연구보다 1.2점 더 높은 성능을 기록했습니다. (Table 5 참조)

5.4 번역 작업 (Translation)

BART는 WMT16 루마니아어-영어 번역 작업에서도 평가되었습니다. 백번역 데이터를 사용한 실험에서 BART는 기존의 강력한 백번역 모델보다 성능이 향상되었습니다. 이 실험에서 우리는 6개의 레이어로 구성된 트랜스포머 소스 인코더를 사용하여 루마니아어를 BART가 영어로 디노이즈할 수 있는 표현으로 변환했습니다. (Table 6 참조)

https://dladustn95.github.io/assets/images/bart_figure8.png

6. 질적 분석 (Qualitative Analysis)

BART는 요약 작업에서 이전 최신 모델보다 최대 6점 더 높은 성능을 기록하며 큰 성능 향상을 보여주었습니다. 자동화된 메트릭 외에도, BART의 성능을 보다 깊이 이해하기 위해 생성된 요약을 질적으로 분석했습니다.

Table 7은 BART가 생성한 요약의 예시를 보여줍니다. 예시는 위키뉴스(WikiNews) 기사에서 가져왔으며, 이 기사는 사전 훈련 코퍼스 생성 이후에 출판된 기사들입니다. 이는 해당 기사에 묘사된 사건이 모델의 훈련 데이터에 포함되지 않았을 가능성을 제거하기 위한 조치입니다. Narayan et al.(2018)을 참고하여, 요약 전에 기사 첫 번째 문장을 제거하였는데, 이는 문서에 대한 추출적 요약이 쉽지 않도록 하기 위함입니다.

BART의 출력은 예상대로 유창하고 문법적으로 정확한 영어입니다. 그러나 출력은 매우 추상적(abstractive)이며, 입력에서 문구를 거의 복사하지 않았습니다. 또한, 출력은 대체로 사실에 부합하며, 입력 문서 전반에서 지원 증거를 통합하고 배경 지식을 사용하여 정확한 내용을 생성합니다. 예를 들어, 첫 번째 예시에서는 PG&E가 캘리포니아에서 운영된다는 점을 추론하거나, 물고기가 지구 온난화로부터 산호초를 보호한다는 점을 추론하는 등의 복잡한 추론을 필요로 합니다. 그러나 '해당 연구가 Science 저널에 게재되었다'는 주장은 원본 문서에서 지원되지 않습니다.

이러한 예시는 BART의 사전 훈련이 자연어 이해와 생성 능력의 강력한 조합을 학습했음을 보여줍니다.

7. 각 작업별 결과 (Results on Tasks)

초기 사전 훈련 방법은 주로 언어 모델을 기반으로 했습니다. GPT(Radford et al., 2018)는 왼쪽 문맥만 모델링하여 특정 작업에 문제를 일으킬 수 있었습니다. ELMo(Peters et al., 2018)는 왼쪽과 오른쪽 문맥을 각각 결합하지만, 이들 간의 상호작용을 사전 훈련하지 않았습니다. Radford et al. (2019)은 매우 큰 언어 모델이 비지도 다중 작업 모델로 작동할 수 있음을 보여주었습니다.

BERT(Devlin et al., 2019)는 마스크된 언어 모델링(masked language modeling)을 도입하여 왼쪽과 오른쪽 문맥 간 상호작용을 학습할 수 있게 했습니다. 이후 연구에서는 더 긴 훈련(Liu et al., 2019), 레이어 간 파라미터 공유(Lan et al., 2019), 단어 대신 스팬 마스킹(Joshi et al., 2019) 등의 방법으로 매우 강력한 성능을 얻을 수 있음을 보여주었습니다. 그러나 BERT는 예측을 자기 회귀 방식으로 수행하지 않기 때문에 생성 작업에 적합하지 않은 단점이 있습니다.

UniLM(Dong et al., 2019)은 마스크 집합을 사용하여 BERT를 미세 조정하는 방법을 제시했으며, 일부 마스크는 왼쪽 문맥만 허용합니다. UniLM은 BART와 마찬가지로 생성 작업과 판별 작업 모두에 사용될 수 있습니다. 그러나 UniLM의 예측은 조건부 독립적이지만, BART는 자기 회귀 방식으로 예측을 수행합니다. BART는 디코더가 항상 손상되지 않은 문맥에서 훈련되기 때문에, 사전 훈련과 생성 작업 간의 불일치를 줄일 수 있습니다.

MASS(Song et al., 2019)는 BART와 가장 유사한 모델입니다. 연속된 스팬이 마스킹된 입력 시퀀스를 누락된 토큰 시퀀스로 매핑합니다. 그러나 MASS는 판별 작업에 덜 효과적입니다. 이는 분리된 토큰 집합이 인코더와 디코더에 입력되기 때문입니다.

XLNet(Yang et al., 2019)은 마스크된 토큰을 자기 회귀 방식으로 예측하는 방식으로 BERT를 확장했습니다. 이 목표는 예측이 왼쪽과 오른쪽 문맥 모두에 의존할 수 있게 해줍니다. BART는 생성 시 디코더가 왼쪽에서 오른쪽으로 작동하는 방식으로 사전 훈련됩니다.

여러 연구는 사전 훈련된 표현을 사용하여 기계 번역 성능을 향상시키는 방법을 탐구했습니다. 가장 큰 성능 향상은 소스 및 타겟 언어 모두에서 사전 훈련을 수행했을 때 나타났습니다(Song et al., 2019; Lample & Conneau, 2019). 그러나 이는 모든 대상 언어에 대해 사전 훈련이 필요합니다. 다른 연구에서는 사전 훈련된 표현을 사용하여 인코더를 개선할 수 있음을 보여주었지만(Edunov et al., 2019), 디코더에서의 성능 향상은 제한적이었습니다. 우리는 BART가 기계 번역 디코더를 개선하는 방법을 제시합니다.

10. Conclusion (결론)

BART의 기여

  • BART는 손상된 문서를 원래 상태로 복원하는 사전 학습 방법을 도입하여, 텍스트 생성과 이해 작업에서 강력한 성능을 발휘했습니다.
  • BART는 RoBERTa와 유사한 성능을 이해 작업에서 달성하면서도, 여러 텍스트 생성 작업에서는 새로운 최첨단(SOTA) 성능을 기록했습니다.
  • 앞으로의 연구는 특정 작업에 맞춰 문서를 손상시키는 새로운 방법을 탐색하여, 사전 학습 방법을 더욱 발전시키

02 | 논문 탐구


🚥 주제 1 : BART 모델 구조 및 특징 이해

BART(Bidirectional and Auto-Regressive Transformer)는 텍스트 생성 및 변환 작업에서 뛰어난 성능을 발휘하는 사전 훈련된 모델. BART는 손상된 문서를 원래 문서로 복원하는 방식으로 훈련된 디노이징 오토인코더(denoising autoencoder)이다. BERTGPT의 장점을 결합한 모델로, BERT는 양방향 인코더를 사용하고 GPT는 좌측에서 우측으로 진행되는 자기 회귀 디코더를 사용.

BART의 주요 특징:

  • 양방향 인코딩: BART의 인코더는 문서의 양방향 문맥을 이해할 수 있도록 설계되었습니다. 이는 BERT와 유사하며, 입력 텍스트의 모든 부분을 전체적으로 이해하는 데 도움이 됩니다.
  • 자기 회귀 디코더: BART의 디코더는 GPT처럼 왼쪽에서 오른쪽으로 토큰을 생성합니다. 이로 인해 시퀀스 생성 작업에서 특히 강력한 성능을 발휘합니다.
  • 시퀀스-투-시퀀스 모델: BART는 인코더-디코더 구조를 사용하여 시퀀스를 입력받고 새로운 시퀀스를 생성할 수 있습니다. 이 구조는 텍스트 요약, 번역, 생성 작업에서 특히 유용합니다.

BART의 구성:

  • 인코더: 입력 문서를 받아들여 양방향으로 문맥을 이해하고, 손상된 텍스트의 정보를 인코딩합니다.
  • 디코더: 인코더의 최종 히든 상태를 받아들이고, 손상된 문서를 복원하거나 새로운 텍스트를 생성합니다.
  • 크로스 어텐션(Cross-Attention): 디코더가 인코더의 출력에 집중하여 각 시간 단계에서 이전 단계의 정보와 입력 문서의 모든 정보를 함께 고려하여 더 나은 예측을 합니다.

BART는 기존 BERT와 달리 마스크된 토큰만 복원하는 것이 아니라, 손상된 문서 전체를 복원하는 더 복잡한 작업을 수행하기 때문에 문서의 더 넓은 맥락을 이해하고 예측하는 데 효과적입니다.

🚥 주제 2 : denoising 기법의 적용방법 및 효과

BART의 주요 사전 훈련 목표는 손상된 입력을 복원하는 것입니다. 이를 위해 여러 디노이징 기법이 적용됩니다. BART는 임의의 문서 손상 방식(noising scheme)을 사용할 수 있으며, 이는 사전 훈련 작업에서 매우 유연하게 적용됩니다.

BART에서 사용되는 주요 손상 방법:

  1. 토큰 마스킹(Token Masking): BERT 방식처럼 입력 문서에서 무작위로 토큰을 선택하고, 이를 [MASK] 토큰으로 대체합니다.
  2. 토큰 삭제(Token Deletion): 입력 문서에서 무작위로 선택된 토큰을 삭제하여, 모델이 어떤 토큰이 빠졌는지 추론하도록 학습합니다.
  3. 텍스트 인필링(Text Infilling): 포아송 분포로 길이를 샘플링하여 임의의 스팬(span)을 선택한 후, 해당 스팬을 단일 [MASK] 토큰으로 대체합니다. 이 방법은 모델이 스팬의 길이와 내용에 대해 더 넓은 추론을 하도록 돕습니다.
  4. 문장 순서 뒤섞기(Sentence Permutation): 문서 내 문장들의 순서를 무작위로 섞어 모델이 순서가 없는 문장에서 의미를 파악하도록 훈련시킵니다.
  5. 문서 회전(Document Rotation): 문서 내 임의의 위치에서 시작하도록 문서를 회전시켜, 모델이 문서의 시작점과 문맥을 추론하도록 학습합니다.

디노이징 기법의 효과:

  • 이러한 기법들은 BART가 더 복잡한 변환 작업을 학습하도록 돕습니다. 단순한 단어 마스킹과 달리, 문장의 순서나 스팬을 바꾸는 방식은 문서의 전반적인 구조를 이해하고 문맥을 보다 정확하게 예측할 수 있게 만듭니다.
  • 텍스트 인필링은 모델이 문서의 길이와 내용을 추론하는 데 도움을 주며, 문장 순서 뒤섞기 및 문서 회전은 문서의 구조적 이해를 높여줍니다.
  • BART가 다양한 작업에서 문서의 추론과 생성 능력을 크게 향상시키는 데 기여합니다.
🚥 주제 3 : time step 별 이해 (croos attention > encoer > decoder)

BART의 인코딩과 디코딩 과정은 크로스 어텐션, 인코더, 디코더로 구성됩니다.

  1. 인코더(Encoder)
  • BART의 인코더는 입력된 문서 전체를 받아들이고, 양방향으로 문맥을 파악하여 텍스트의 표현을 인코딩합니다. 이 과정에서 입력 문서의 각 토큰이 서로 상호작용하며, 문서 전체의 구조적, 의미적 맥락을 이해하게 됩니다.
  • 인코더의 각 레이어는 셀프 어텐션(self-attention) 메커니즘을 사용하여, 각 토큰이 문서의 다른 부분과 상호작용하도록 합니다.
  1. 크로스 어텐션(Cross-attention)
  • BART의 크로스 어텐션은 디코더와 인코더 사이에서 발생합니다. 디코더가 새로운 토큰을 생성할 때, 이전 단계에서 생성된 토큰뿐만 아니라 인코더가 출력한 전체 문서의 히든 상태를 참조합니다. 이를 통해 디코더는 더 정교하게 문서의 모든 정보를 통합하여 토큰을 생성할 수 있습니다.
  • 크로스 어텐션은 인코더의 출력을 디코더의 입력으로 연결하여 디코더가 더 풍부한 문맥을 기반으로 추론하고 예측하게 만듭니다.
  1. 디코더(Decoder)
  • 디코더는 인코더에서 나온 정보를 바탕으로 새로운 텍스트를 생성하는 단계입니다. 이때, 디코더는 자기 회귀적 방식으로 왼쪽에서 오른쪽으로 순차적으로 토큰을 예측합니다.
  • 디코더의 첫 번째 입력은 [CLS] 토큰과 같으며, 이후 단계에서는 이전에 생성된 토큰들이 디코더의 입력으로 사용됩니다.
  • 디코더는 주어진 문맥 내에서 가장 적절한 다음 단어를 예측하며, 이 과정이 반복되어 문서 전체를 생성하게 됩니다.
🚥 주제 4 : 어떠한 finetuning 태스크에 적합한지 조사 및 학습
  1. 시퀀스 분류(Sequence Classification)
  • 시퀀스 분류 작업에서 BART는 문서 전체를 인코더와 디코더에 입력한 후, 디코더의 마지막 토큰의 히든 상태를 새로운 다중 클래스 분류기에 전달합니다. 이는 BERT의 CLS 토큰과 유사한 방식입니다.
  • 이러한 방식은 감정 분석, 주제 분류 등의 작업에 적합합니다.
  1. 토큰 분류(Token Classification)
  • 토큰 분류 작업에서는 문서 전체를 인코더와 디코더에 입력하고, 디코더의 최상위 히든 상태를 각 토큰의 표현으로 사용하여 토큰을 분류합니다.
  • SQuAD와 같은 질문 응답 태스크에서 사용됩니다.
  1. 시퀀스 생성(Sequence Generation)
  • BART는 자기 회귀적 디코더를 사용하기 때문에 텍스트 요약, 추상적 질문 응답, 대화 생성과 같은 시퀀스 생성 작업에 매우 적합합니다.
  • 생성 작업에서 BART는 입력 문서를 인코더에 입력하고, 디코더가 이를 바탕으로 새로운 시퀀스를 생성합니다. 요약 작업에서는 입력 문서를 압축하여 요약본을 생성하며, 질문 응답에서는 주어진 질문에 맞는 답변을 생성합니다.
  1. 기계 번역(Machine Translation)
  • BART는 영어와 같은 타겟 언어로 번역하는 기계 번역 디코더로도 사용할 수 있습니다. 사전 훈련된 인코더와 디코더를 사용하여 번역 품질을 향상시킬 수 있으며, 백번역(back-translation) 데이터로도 학습이 가능합니다.
  1. 대화 응답 생성(Dialogue Response Generation)
  • BART는 대화 생성 작업에서 매우 좋은 성능을 발휘합니다. 특히, CONVAI2와 같은 대화 응답 생성 작업에서 BART는 자동화된 메트릭 기준에서 이전 모델들보다 우수한 성능을 기록했습니다. 대화 생성 작업에서 BART는 주어진 대화 문맥과 대화 참여자의 성격(persona)을 기반으로 자연스럽고 적절한 응답을 생성할 수 있습니다.

BART가 적합한 Fine-tuning 태스크:

  1. 텍스트 요약 (Text Summarization)
  • BART는 특히 텍스트 요약 작업에 강력한 성능을 보여줍니다. CNN/DailyMail 및 XSum 같은 데이터셋에서 매우 추상적인 요약을 생성할 수 있으며, 이전 연구들과 비교해 모든 ROUGE 지표에서 뛰어난 성과를 기록했습니다.
  • CNN/DailyMail에서는 소스 문장을 압축한 추출적 요약 방식에서 좋은 성과를 냈고, XSum에서는 완전히 새로운 문장으로 요약하는 추상적 요약 방식에서 뛰어난 성과를 보여줍니다.
  • BART는 추상적 요약을 잘 수행하며, 기존 문장 구조에 구애받지 않고 문서의 중요한 내용을 요약해낼 수 있습니다.
  1. 추상적 질문 응답 (Abstractive Question Answering)
  • ELI5 데이터셋에서 BART는 긴 형식의 답변을 생성하는 능력을 테스트했습니다. 이 데이터셋은 자유로운 형식의 긴 답변을 요구하는데, BART는 이 작업에서 이전의 연구보다 1.2점 더 높은 ROUGE-L 성과를 보였습니다.
  • 이와 같이, BART는 추상적 질문 응답에 매우 적합하며, 긴 형식의 답변을 생성하는 데 뛰어난 능력을 보입니다. 답변을 생성할 때 문맥에 기반한 추론을 잘 수행하며, 복잡한 질문에도 일관된 답변을 제공합니다.
  1. 기계 번역 (Machine Translation)
  • BART는 기계 번역 작업에서도 사용될 수 있습니다. WMT16 Romanian-English 번역 작업에서 백번역 데이터를 사용한 실험을 통해, BART는 강력한 백번역 모델을 능가하는 성과를 기록했습니다.
  • 특히 BART는 소스 인코더를 학습하여 외국어 문장을 영어로 디노이즈하는 데 사용하며, 트랜스포머 인코더-디코더 구조로 성능을 높일 수 있습니다.
  1. 대화 응답 생성 (Dialogue Response Generation)
  • BART는 대화 응답 생성 작업에서 CONVAI2 데이터셋을 기반으로 실험되었습니다. 이 작업에서 BART는 자동화된 F1 및 PPL 메트릭에서 이전 연구들보다 더 높은 성과를 기록했습니다.
  • BART는 대화 문맥과 참여자의 배경 정보(persona)를 기반으로 자연스럽고 적절한 대화 응답을 생성할 수 있습니다. 이로 인해 대화 시스템, 챗봇 응용 프로그램에서 활용할 수 있는 가능성이 큽니다.
  1. 텍스트 생성 및 변환 작업 (Text Generation and Transformation Tasks)
  • 텍스트 생성변환 작업은 BART의 가장 강력한 분야입니다. 텍스트 생성 작업에서 BART는 입력 문서의 내용을 바탕으로 자연스럽고 문법적으로 정확한 텍스트를 생성할 수 있습니다.
  • BART는 자기 회귀 방식으로 디코더를 훈련하므로, 이전에 생성된 토큰을 기반으로 다음 토큰을 생성하며, 이는 텍스트 변환 작업에서도 강력한 성능을 발휘하게 합니다.

결론

BART는 디노이징 오토인코더 방식으로 문서를 복원하는 작업에서 매우 강력한 사전 훈련 모델입니다. 이 모델은 BERT와 GPT의 장점을 결합하여 문서의 맥락을 깊이 있게 이해하고 생성할 수 있으며, 다양한 자연어 처리 작업에 적합한 시퀀스-투-시퀀스 구조를 사용합니다. 특히 텍스트 생성, 요약, 질문 응답, 기계 번역, 대화 생성과 같은 작업에서 뛰어난 성능을 발휘하며, 크로스 어텐션자기 회귀적 디코딩 메커니즘이 이 성능의 핵심 역할을 합니다.

향후 연구는 BART의 사전 훈련을 위한 새로운 디노이징 방법을 탐구하거나, 특정 작업에 최적화된 훈련 방식을 개발하는 방향으로 확장될 수 있을 것입니다.

03 | 실습 : gpt 코드 분


Huggingface Pipeline과 BART를 이용한 텍스트 요약 - 인하대학교 인트아이

https://github.com/ljm565/chatbot-BART/blob/main/README.md

Translation


bart.py


import torch
import torch.nn as nn
from transformers import BartForConditionalGeneration

# BART
class BART(nn.Module):
    def __init__(self, config, tokenizer):
        super(BART, self).__init__()
        self.pretrained = config.pretrained
        self.bart = BartForConditionalGeneration.from_pretrained('gogamza/kobart-base-v2')
        if not self.pretrained:
            bart_config = self.bart.config
            self.bart = BartForConditionalGeneration(bart_config)
        self.tokenizer = tokenizer

    def make_mask(self, src):
        mask = torch.where(src==self.tokenizer.pad_token_id, 0, 1)
        return mask
        

    def forward(self, src, trg):
        enc_mask, dec_mask = self.make_mask(src), self.make_mask(trg)
        output = self.bart(input_ids=src, attention_mask=enc_mask, decoder_input_ids=trg, decoder_attention_mask=dec_mask).logits
        return output

2. UTILS


tokenizer.py

from transformers import PreTrainedTokenizerFast

class BARTTokenizer:
    def __init__(self):
        self.tokenizer = PreTrainedTokenizerFast.from_pretrained('gogamza/kobart-base-v2')

        self.pad_token, self.pad_token_id = self.tokenizer.pad_token, self.tokenizer.pad_token_id
        self.cls_token, self.cls_token_id = self.tokenizer.convert_ids_to_tokens(0), self.tokenizer.convert_tokens_to_ids('<s>')
        self.sep_token, self.sep_token_id = self.tokenizer.convert_ids_to_tokens(1), self.tokenizer.convert_tokens_to_ids('</s>')
        self.unk_token, self.unk_token_id = self.tokenizer.unk_token, self.tokenizer.unk_token_id

        self.vocab_size = len(self.tokenizer)

    def tokenize(self, s):
        return self.tokenizer.tokenize(s)

    def encode(self, s):
        return self.tokenizer.encode(s)

    def decode(self, tok):
        try:
            tok = tok[:tok.index(self.sep_token_id)]
        except ValueError:
            try:
                tok = tok[:tok.index(self.pad_token_id)]
            except:
                pass
        return self.tokenizer.decode(tok)

train.py

import torch
import torch.nn as nn
import torch.optim as optim
from torch.optim.lr_scheduler import OneCycleLR
from torch.utils.data import DataLoader, random_split

import re
import time
import pickle
import random
from tqdm import tqdm
from transformers import top_k_top_p_filtering

from models.bart import BART
from utils.utils_func import *
from utils.config import Config
from tokenizer import BARTTokenizer
from utils.utils_data import DLoader

class Trainer:
    def __init__(self, config:Config, device:torch.device, mode:str, continuous:int):
        self.config = config
        self.device = device
        self.mode = mode
        self.continuous = continuous
        self.dataloaders = {}

        # if continuous, load previous training info
        if self.continuous:
            with open(self.config.loss_data_path, 'rb') as f:
                self.loss_data = pickle.load(f)

        # path, data params
        self.base_path = self.config.base_path
        self.data_path = self.config.data_path
        self.model_path = self.config.model_path
        if self.mode != 'train':
            self.model_path = self.model_path[:-3] + '_' + self.config.model_type + '.pt'
 
        # train params
        self.batch_size = self.config.batch_size
        self.epochs = self.config.epochs
        self.lr = self.config.lr
        self.q_max_len = self.config.q_max_len
        self.a_max_len = self.config.a_max_len
        self.result_num = self.config.result_num

        # define tokenizer
        self.tokenizer = BARTTokenizer()
        self.config.vocab_size = self.tokenizer.vocab_size

        # dataloader
        if self.mode != 'chatting':
            torch.manual_seed(999)  # for reproducibility
            self.dataset = DLoader(load_dataset(self.data_path), self.config, self.tokenizer)
            data_size = len(self.dataset)
            train_size = int(data_size * 0.95)
            val_size = int(data_size * 0.03)
            test_size = data_size - train_size - val_size

            self.trainset, self.valset, self.testset = random_split(self.dataset, [train_size, val_size, test_size])
            if self.mode == 'train':
                self.dataset = {'train': self.trainset, 'val': self.valset, 'test': self.testset}
                self.dataloaders = {
                    s: DataLoader(d, self.batch_size, shuffle=True) if s == 'train' else DataLoader(d, self.batch_size, shuffle=False)
                    for s, d in self.dataset.items()}
            else:
                self.dataset = {'test': self.testset}
                self.dataloaders = {s: DataLoader(d, self.batch_size, shuffle=False) for s, d in self.dataset.items() if s == 'test'}

        # model, optimizer, loss
        self.model = BART(self.config, self.tokenizer).to(self.device)
        self.criterion = nn.CrossEntropyLoss(ignore_index=self.tokenizer.pad_token_id)
    
        if self.mode == 'train':
            total_steps = len(self.dataloaders['train']) * 100#self.epochs
            pct_start = 100 / total_steps
            final_div_factor = self.lr / 25 / 1e-7    # OneCycleLR default value is 25
            self.optimizer = optim.Adam(self.model.parameters(), lr=self.lr)
            self.scheduler = OneCycleLR(self.optimizer, max_lr=self.lr, total_steps=total_steps, pct_start=pct_start, final_div_factor=final_div_factor)
            if self.continuous:
                self.check_point = torch.load(self.model_path, map_location=self.device)
                self.model.load_state_dict(self.check_point['model'])
                self.optimizer.load_state_dict(self.check_point['optimizer'])
                self.scheduler.load_state_dict(self.check_point['scheduler'])
                del self.check_point
                torch.cuda.empty_cache()
        else:
            self.check_point = torch.load(self.model_path, map_location=self.device)
            self.model.load_state_dict(self.check_point['model'])    
            self.model.eval()
            del self.check_point
            torch.cuda.empty_cache()

        
    def training(self):
        early_stop = 0
        best_val_loss = float('inf')
        best_val_bleu = 0 if not self.continuous else self.loss_data['best_val_bleu']
        train_loss_history = [] if not self.continuous else self.loss_data['train_loss_history']
        val_loss_history = [] if not self.continuous else self.loss_data['val_loss_history']
        val_score_history = {'bleu2': [], 'bleu4': [], 'nist2': [], 'nist4': []} if not self.continuous else self.loss_data['val_score_history']
        best_epoch_info = 0 if not self.continuous else self.loss_data['best_epoch']

        for epoch in range(self.epochs):
            start = time.time()
            print(epoch+1, '/', self.epochs)
            print('-'*10)
            for phase in ['train', 'val']:
                print('Phase: {}'.format(phase))
                if phase == 'train':
                    epoch_loss = self.train(phase, epoch)
                    train_loss_history.append(epoch_loss)
                else:
                    epoch_loss = self.test(phase)
                    bleu2, bleu4, nist2, nist4 = self.inference(phase, self.result_num)
                    if phase == 'val':
                        val_loss_history.append(epoch_loss)
                        val_score_history['bleu2'].append(bleu2)
                        val_score_history['bleu4'].append(bleu4)
                        val_score_history['nist2'].append(nist2)
                        val_score_history['nist4'].append(nist4)

                        # save best model for bleu4
                        if  val_score_history['bleu4'][-1] > best_val_bleu:
                            early_stop = 0
                            best_val_bleu = val_score_history['bleu4'][-1]
                            save_checkpoint(self.model_path, self.model, self.optimizer, self.scheduler, 'bleu4')

                        # save best model for loss
                        early_stop += 1
                        if  epoch_loss < best_val_loss:
                            early_stop = 0
                            best_val_loss = epoch_loss
                            best_epoch = best_epoch_info + epoch + 1
                            save_checkpoint(self.model_path, self.model, self.optimizer, self.scheduler, 'loss')
                            
                            self.loss_data = {'best_epoch': best_epoch, 'best_val_bleu': best_val_bleu, 'train_loss_history': train_loss_history, 'val_loss_history': val_loss_history, 'val_score_history': val_score_history}
                            print('Saving the loss related data...')
                            with open(self.config.loss_data_path, 'wb') as f:
                                pickle.dump(self.loss_data, f)

            print("time: {} s\n".format(time.time() - start))
            print('\n'*2)

            # early stopping
            if early_stop == self.config.early_stop_criterion:
                break

        print('best val bleu: {:4f}, best epoch: {:d}\n'.format(best_val_bleu, best_epoch))
        self.loss_data = {'best_epoch': best_epoch, 'best_val_bleu': best_val_bleu, 'train_loss_history': train_loss_history, 'val_score_history': val_score_history}
        return self.loss_data

    def train(self, phase, epoch):
        self.model.train()
        epoch_loss = 0

        for i, (src, trg) in enumerate(self.dataloaders[phase]):
            self.optimizer.zero_grad()
            batch_size = src.size(0)
            src, trg = src.to(self.device), trg.to(self.device)

            with torch.set_grad_enabled(phase=='train'):
                output = self.model(src, trg)
                loss = self.criterion(output[:, :-1, :].reshape(-1, output.size(-1)), trg[:, 1:].reshape(-1))
                loss.backward()
                self.optimizer.step()
                self.scheduler.step()

            epoch_loss += loss.item() * batch_size
           
            if i % 200 == 0:
                print('Epoch {}: {}/{} step loss: {}'.format(epoch+1, i, len(self.dataloaders[phase]), loss.item()))

        epoch_loss = epoch_loss / len(self.dataloaders[phase].dataset)

        print('{} loss: {}\n'.format(phase, epoch_loss))

        return epoch_loss

    def test(self, phase):
        self.model.eval()
        epoch_loss = 0

        with torch.no_grad():
            for src, trg in tqdm(self.dataloaders[phase], desc=phase + ' testing..'):
                batch_size = src.size(0)
                src, trg = src.to(self.device), trg.to(self.device)

                output = self.model(src, trg)
                loss = self.criterion(output[:, :-1, :].reshape(-1, output.size(-1)), trg[:, 1:].reshape(-1))
                epoch_loss += loss.item() * batch_size

        # calculate epoch loss
        epoch_loss = epoch_loss / len(self.dataloaders[phase].dataset)
        print('{} loss: {:4f}\n'.format(phase, epoch_loss))

        return epoch_loss

    def inference(self, phase, result_num=3):
        self.model.eval()
        all_trg, all_output = [], []
        
        with torch.no_grad():
            for src, trg in tqdm(self.dataloaders[phase], desc=phase+' inferencing..'):
                src, trg = src.to(self.device), trg.to(self.device)
                all_trg.append(trg.detach().cpu())
            
                decoder_all_output = []
                for j in range(self.a_max_len):
                    if j == 0:
                        trg = trg[:, j].unsqueeze(1)
                        output = self.model(src, trg)
                        trg = torch.cat((trg, torch.argmax(output[:, -1], dim=-1).unsqueeze(1)), dim=1)
                    else:
                        output = self.model(src, trg)
                        trg = torch.cat((trg, torch.argmax(output[:, -1], dim=-1).unsqueeze(1)), dim=1)
                    decoder_all_output.append(output[:, -1].unsqueeze(1).detach().cpu())
                        
                all_output.append(torch.argmax(torch.cat(decoder_all_output, dim=1), dim=-1))

        # calculate scores
        all_ref, all_pred = tensor2list(all_trg, all_output, self.tokenizer)
        bleu2 = cal_scores(all_ref, all_pred, 'bleu', 2)
        bleu4 = cal_scores(all_ref, all_pred, 'bleu', 4)
        nist2 = cal_scores(all_ref, all_pred, 'nist', 2)
        nist4 = cal_scores(all_ref, all_pred, 'nist', 4)
        print('\nInference Score')
        print('bleu2: {}, bleu4: {}, nist2: {}, nist4: {}'.format(bleu2, bleu4, nist2, nist4))

        # print samples
        ids = random.sample(list(range(len(all_pred))), result_num)
        print_samples(all_ref, all_pred, ids, self.tokenizer)

        return bleu2, bleu4, nist2, nist4

    def chatting(self, query, num_result=1):
        phone = re.compile('[0-9]{2,3}[- ]?[0-9]{3,4}[- ]?[0-9]{4}')
        replace_num = '###-####-####'

        with torch.no_grad():
            query = [self.tokenizer.cls_token_id] + self.tokenizer.encode(query)[:self.q_max_len-2] + [self.tokenizer.sep_token_id]
            query = query + [self.tokenizer.pad_token_id] * (self.q_max_len - len(query))
            
            query = torch.LongTensor(query).expand(num_result, -1).to(self.device)
            trg = torch.LongTensor([self.tokenizer.cls_token_id]).expand(num_result, -1).to(self.device)

            for _ in range(self.a_max_len):
                output = self.model(query, trg)
                if self.config.greedy:
                    output = torch.argmax(output[:, -1], dim=-1).unsqueeze(1)
                else:
                    output = output[:, -1] / self.config.temperature
                    output = top_k_top_p_filtering(output, top_k=self.config.topk, top_p=self.config.topp)
                    output = torch.multinomial(torch.softmax(output, dim=-1), num_samples=1)
                
                trg = torch.cat((trg, output), dim=1)

                if num_result == 1 and output[0, 0].item() == self.tokenizer.sep_token_id:
                    break

        trg = [self.tokenizer.decode(s[1:].tolist()) for s in trg.detach().cpu()]

        # phone number filtering
        numbers = [phone.findall(s) for s in trg]
        for i in range(len(trg)):
            for number in numbers[i]:
                trg[i] = trg[i].replace(number, replace_num)
            
        return trg

layer_norm

sublayer.py

3. Transformer

4. Three Embeddings

position.py

segment.py

token.py

bert.py

: 세가지 임베딩을 최종적으로 모델에 입력할 수 있는 형태로 전달하는 class

5. Bert model

Bert model

MLM(masked language model)

NSP(next sentence prediction)

languagae_model.py (MLM + NSP 학습)

6. Training

질문


profile
AI Engineer / 의료인공지능

0개의 댓글