Sentence Transformers를 사용하여 400배 더 빠른 정적 임베딩 모델 학습하기

TL;DR

Hugging Face는 표준 트랜스포머 기반 임베딩 품질의 최소 85%를 유지하면서 CPU에서 100배~400배 더 빠르게 실행되는 정적 임베딩 모델 학습 레시피를 공개했습니다. 또한 두 개의 구체적인 모델(영어 검색용 static-retrieval-mrl-en-v1 및 다국어 유사도용 static-similarity-mrl-multilingual-v1)을 학습 스크립트, 평가 결과 및 Weights & Biases 로그와 함께 발표했습니다.

방법

이 접근 방식은 대조 학습(contrastive learning)을 MultipleNegativesRankingLoss 및 선택 사항인 Matryoshka Representation Learning과 결합하여, 어텐션 기반 인코딩 대신 단순한 토큰 임베딩 룩업을 수행하는 정적 인코더를 학습시킵니다. 이를 통해 품질 저하를 최소화하면서 수십 배의 속도 향상을 달성합니다.

학습 상세 정보

요구 사항

학습에는 dataset, loss function, training arguments, evaluator, trainer 구성 요소를 포함하는 Sentence Transformers 라이브러리가 사용됩니다.

모델 영감

두 가지 모델을 목표로 했습니다: 영어 전용 검색 모델과 다국어 일반 유사도 모델이며, 두 모델 모두 BERT 기반 토크나이저(bert-base-uncased 또는 bert-base-multilingual-uncased)와 1024의 임베딩 차원을 래핑하는 StaticEmbedding 모듈로 초기화되었습니다.

데이터셋 선택

영어 검색을 위해 gooaq, msmarco (triplet), squad, s2orc (title-abstract-pair), allnli (triplet), paq, trivia_qa, msmarco_10m, swim_ir (en), pubmedqa (triplet-20), miracl (en-triplet-all), mldr (en-triplet-all), mr_tydi (en-triplet-all)를 포함한 30개의 데이터셋이 선택되었습니다. 다국어 유사도의 경우, 병렬 문장 데이터셋(wikititles, tatoeba, talks, europarl, global_voices, jw300, muse, wikimatrix, opensubtitles)과 양성 쌍(positive-pair) 데이터셋(stackexchange-duplicates, quora-duplicates, wikianswers-duplicates, allnli, simple_wiki, altlex, flickr30k_captions, coco_captions, nli_for_simcse, negation)을 결합하여 총 30개의 학습 데이터셋을 구성했습니다.

손실 함수

MultipleNegativesRankingLoss가 선택된 이유는 2048의 배치 크기가 RTX 3090에 적합하여 CachedMultipleNegativesRankingLoss의 오버헤드와 GISTEmbedLoss의 가이드 모델로 인한 속도 저하를 피할 수 있기 때문입니다. 그 위에 MatryoshkaLoss를 적용했으며, 차원은 [32, 64, 128, 256, 512, 1024]입니다.

학습 인자 (Training Arguments)

두 모델 모두 다음을 사용했습니다: num_train_epochs=1, per_device_train_batch_size=2048, per_device_eval_batch_size=2048, learning_rate=2e-1, warmup_ratio=0.1, bf16=True, batch_sampler=BatchSamplers.NO_DUPLICATES, multi_dataset_batch_sampler=MultiDatasetBatchSamplers.PROPORTIONAL, eval_strategy=steps, eval_steps=250 (영어) 또는 1000 (다국어), save_strategy=steps, eval_steps와 일치하는 save_steps, save_total_limit=2, eval_steps와 일치하는 logging_steps, logging_first_step=True, 그리고 실행별 출력 디렉토리를 사용했습니다.

평가기 (Evaluator)

영어 검색 모델은 제로샷 검색 평가를 위해 NanoBEIREvaluator를 사용했으며, 다국어 모델은 평가를 위해 MTEB 태스크(STS, Classification, Pair Classification)에 의존했습니다.

하드웨어

학습은 RTX 3090 GPU, i7-13700K CPU 및 32GB RAM에서 수행되었습니다.

전체 학습 스크립트

제공된 스크립트는 데이터셋을 로드하고, StaticEmbedding 모델을 인스턴스화하며, MatryoshkaLoss로 손실 함수를 구성하고, 학습 인자를 설정하며, 선택적으로 평가기를 실행하고, SentenceTransformerTrainer로 학습한 후 최종 모델을 저장합니다. 영어 검색 스크립트는 17.8시간이 소요되었으며 2.6kWh를 소비하고 1kg의 CO₂를 배출했습니다. 다국어 스크립트는 3.1시간이 소요되었으며 0.5kWh를 소비하고 0.2kg의 CO₂를 배출했습니다.

사용법

두 모델 모두 모델 이름과 device="cpu"를 사용하여 SentenceTransformer를 통해 로드됩니다. 추론 방식은 표준 Sentence Transformers와 동일합니다: model.encode는 임베딩을 반환하고, model.similarity는 코사인 유사도를 계산합니다. truncate_dim 인자를 통해 Matryoshka 스타일의 차원 축소를 활성화할 수 있습니다 (예: truncate_dim=256). 이 모델들은 LangChain, LlamaIndex, Haystack, txtai와 함께 즉시 사용할 수 있습니다.

성능

영어 검색

NanoBEIR에서 static-retrieval-mrl-en-v1은 NDCG@10 0.5032를 달성했으며, 이는 all-mpnet-base-v2 점수(0.5757)의 87.4%입니다. CPU에서는 초당 107,419.51개의 문장을 처리하여 all-mpnet-base-v2(270.40 sentences/s)보다 397배 빠릅니다. GPU에서는 초당 97,171.47개의 문장을 처리하여 all-mpnet-base-v2(4043.13 sentences/s)보다 24배 빠릅니다. Matryoshka 평가 결과, 차원을 512로 절반 줄여도 NDCG@10은 1.47%만 감소합니다 (0.5032 → 0.4957).

다국어 유사도

multilingual-e5-small과 비교했을 때, static-similarity-mrl-multilingual-v1은 STS에서 92.3%, Pair Classification에서 95.52%, Classification에서 86.52%의 점수를 기록했습니다. 이는 multilingual-e5-small보다 CPU에서 약 125배, GPU에서 10배 더 빠릅니다. Matryoshka 평가에 따르면 차원을 256(4배 축소)으로 줄여도 영어 STS 성능 저하는 0.56%에 불과합니다.

결론

제시된 레시피로 학습된 정적 임베딩 모델은 일반적인 트랜스포머 기반 임베딩 품질의 최소 85%를 유지하면서 CPU에서 100배400배, GPU에서 10배25배의 속도 향상을 제공합니다. 공개된 모델은 최소한의 정확도 손실로 효율적인 온디바이스(on-device), 인브라우저(in-browser) 및 엣지 컴퓨팅 유스케이스를 가능하게 합니다.

다음 단계

사용자는 기존 Sentence Transformer 모델을 static-retrieval-mrl-en-v1 또는 static-similarity-mrl-multilingual-v1으로 교체하거나, 특정 작업 데이터로 자신만의 정적 임베딩을 학습할 수 있습니다. 잠재적인 개선 사항으로는 hard-negative mining, model souping, curriculum learning, guided false-in-batch negatives filtering, seed-optimized random initialization, tokenizer retraining, CachedMultipleNegativesRankingLoss를 통한 gradient caching, 그리고 더 큰 인코더로부터의 모델 증류(distillation) 등이 있습니다.

Sources