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

딥러닝 모델 서비스 A-Z 2편: Knowledge Distillation

서비스하기 무거운 BERT-large 기반 대화 응답 검색 모델을 Knowledge Distillation로 경량화한 사례다. Teacher의 예측 확률뿐 아니라 문맥·답변 embedding 자체를 Student에 전달하고, Teacher layer 일부로 Student를 초기화했다. 학습 데이터 1/40만 사용하면서 추론을 약 3배 빠르게 하고 Top 10 정확도의 99.9%를 유지했다.

핵심 포인트
  • Knowledge Distillation은 정답뿐 아니라 Teacher가 만든 soft target의 클래스 간 확률 관계를 Student가 배우게 한다.
  • Teacher는 대화 문맥과 답변을 각각 BERT로 encoding하고 projection 뒤 cosine similarity로 적합도를 계산한다.
  • Prediction Logit Distillation은 temperature를 적용한 Teacher·Student logit 분포의 cross entropy를 최소화한다.
  • Embedding Distillation은 두 encoder 표현을 MSE로 맞추며 차원이 다르면 학습 가능한 projection matrix를 둔다.
  • BERT-large 24개 layer에서 3개마다 하나를 가져와 8-layer Student를 초기화해 안정적인 속도와 성능을 확보했다.
  • Embedding distillation, logit distillation, 일반 fine-tuning 순으로 학습해 Top 1 96.2%, Top 5 99.3%, Top 10 99.9%를 보존했다.
상세 정리
  • 문제: 내부 대형 언어 모델은 메모리와 연산량이 커 실제 대화 서비스에 바로 투입하기 어려웠다. 대표적 model compression 기법인 distillation로 작은 모델을 만들었다.
  • Soft target 의미: Teacher가 숫자 2에 낮은 확률로 3과 7을 부여하는 비율은 클래스 사이의 유사 구조를 담는다. Student는 hard label만 학습할 때 얻지 못하는 관계를 전달받는다.
  • 대상 task: 일상 대화에서 문맥에 맞는 답변을 고르는 검색 모델을 경량화했다. 빠른 ANN 검색을 위해 context와 response encoder의 BERT hidden state를 고정 차원으로 projection했다.
  • Teacher 구조: context BERT와 response BERT가 별도로 embedding을 만든다. 두 결과의 cosine similarity를 기준으로 답변 적합도를 판단해 Faiss 같은 ANN index에 연결할 수 있다.
  • Logit distillation: Student와 Teacher의 classification logit에 temperature를 적용해 softmax distribution의 cross entropy를 계산했다. 여러 값 중 이 실험에서는 temperature 1이 가장 좋았다.
  • Embedding distillation: 같은 문장에 대한 Student와 Teacher hidden representation을 MSE로 맞췄다. hidden dimension이 다르면 Student embedding에 학습 가능한 matrix를 곱해 Teacher 차원으로 투영했다.
  • 초기화 선택지: Teacher layer를 일정 간격으로 가져오면 transformer block 크기는 그대로지만 layer 수에 비례한 예측 가능한 가속을 얻는다. 24개 중 8개를 쓰면 약 3배 속도를 기대할 수 있다.
  • 더 작은 Student: 모바일처럼 자원이 더 제한되면 compact model을 별도로 pretrain하는 방식을 고려할 수 있다. block 자체도 줄일 수 있지만 사전학습 비용이 추가된다.
  • 학습 순서: BERT-large Teacher에서 세 layer마다 하나를 선택해 8-layer Student를 만들었다. 먼저 embedding을 맞추고 다음으로 prediction logit을 distill한 뒤 Teacher와 같은 방식으로 fine-tuning했다.
  • 데이터 효율: 세 단계 모두 Teacher 학습 데이터의 1/40만 사용했다. 적은 데이터로도 Teacher가 이미 학습한 표현과 의사결정 경계를 전달받았다.
  • 결과: 추론 속도는 약 3배 개선됐다. 정확도는 Teacher 대비 Top 1 96.2%, Top 5 99.3%, Top 10 99.9%를 유지해 실제 후보 검색 품질 손실을 제한했다.
  • 한계와 다음 단계: 모델 architecture와 size를 더 공격적으로 줄이지 못했다. 긴 context용 Sparse Transformer, TensorFlow Serving quantization, transformer block 폭 축소를 후속 후보로 제시했다.
왜 읽나대형 dual-encoder를 실제 검색 서비스에 넣기 위해 속도와 정확도 손실을 수치로 비교하며 distillation pipeline을 설계하려는 ML 엔지니어에게 적합하다.
스캐터랩
스캐터랩 (이루다) 블로그
원문은 여기서 이어서 읽을 수 있어요
원문 읽기
읽음 (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