Transformers v5에서 제거된 JAX 환경의 문장 임베딩 계산 방법
요약
Hugging Face Transformers v5가 JAX 코드를 제거함에 따라, 본 포스트는 오픈 소스 라이브러리 eqx-zoo를 사용하여 JAX 환경에서 문장 임베딩을 계산하는 방법을 안내합니다. 이 방법은 Hugging Face 체크포인트를 Equinox 모듈로 로드하며, sentence-transformers의 결과와 일치함을 검증했습니다.
핵심 포인트
- Transformers v5가 PyTorch에 집중하여 JAX 코드를 제거함.
- eqx-zoo 라이브러리를 사용하여 JAX 환경에서 임베딩 계산이 가능함.
- Hugging Face 체크포인트를 Equinox 모듈로 로드하는 방법을 제시함.
- jax.vmap과 mask를 활용하여 배치 처리 및 패딩을 효과적으로 관리할 수 있음.
Hugging Face Transformers v5가 PyTorch에 집중하기 위해 TensorFlow와 JAX 코드를 제거했습니다. 만약 FlaxBertModel 또는 FlaxAutoModel을 사용하여 JAX에서 문장 임베딩을 계산했다면, 해당 클래스들은 사라졌습니다.
본 포스트에서는 오픈 소스 라이브러리인 eqx-zoo를 사용하여 JAX로 문장 임베딩을 계산하는 방법을 보여줍니다. 이 라이브러리는 Hugging Face 체크포인트를 일반 Equinox 모듈로 로드합니다. 생성된 임베딩은 sentence-transformers의 결과와 float32 반올림 오차 범위 내에서 일치하며, 포스트는 이 검증 과정을 마지막에 다룹니다.
우리는 영어 및 다국어 모델, 간단한 시맨틱 검색, 그리고 언어 모델을 기반으로 구축된 최신 임베딩 모델인 Qwen3-Embedding에 대해 다룰 것입니다.
설치
pip install eqx-zoo tokenizers
eqx-zoo는 JAX와 Equinox를 가져오며, tokenizers는 텍스트를 토큰 ID로 변환하는 데 사용할 Hugging Face의 빠른 토크나이저 라이브러리입니다. PyTorch가 필요하지 않습니다: eqx-zoo가 체크포인트의 safetensors 파일을 직접 읽기 때문입니다.
첫 번째 임베딩 계산
all-MiniLM-L6-v2는 시맨틱 검색의 일반적인 기본 모델인 작고 빠른 영어 모델입니다:
import jax
import jax.numpy as jnp
from tokenizers import Tokenizer
...
이 임베딩들은 단위 길이(unit length)를 가지므로, 이들의 내적(dot product)은 코사인 유사도(cosine similarities)가 됩니다:
[[0.9999997 0.55840427 0.05399836]
[0.55840427 1.0000001 0.06003189]
[0.05399836 0.06003189 0.99999946]]
두 문장으로 구성된(cat) 문장들은 서로 0.56의 점수를 받고, 주식 시장 관련 문장과는 공유하는 단어가
두 가지 주목할 점이 있습니다. model.embed는 하나의 문장에 대해 작동하고, jax.vmap은 이를 배치(batch) 전체에 매핑합니다. 그리고 mask는 어떤 토큰이 실제인지 표시해 줍니다. 즉, 더 짧은 문장들은 가장 긴 문장에 맞춰 패딩(padding) 처리되고, 이 패딩 부분은 임베딩 계산에서 제외됩니다.
각 체크포인트마다 고유한 방식(recipe)을 가지고 있습니다
임베딩 모델은 단순히 트랜스포머(transformer) 그 이상입니다. sentence-transformers의 체크포인트는 또한 토큰별 출력(per-token outputs)을 단일 벡터로 변환하는 방법까지 명시합니다. 이는 modules.json 파일과 풀링 설정(pooling config)에 담겨 있습니다. eqx-zoo의 from_pretrained 함수는 이 정보들을 읽어오기 때문에, 임베딩 계산은 각 체크포인트가 가진 고유한 방식을 따릅니다:
| 체크포인트 | 풀링 방식 (Pooling) | 정규화 여부 (Normalised) |
|---|---|---|
| all-MiniLM-L6-v2 | 실제 토큰 평균 (Mean over real tokens) | 예 (Yes) |
| ... |
이것이 중요한 이유는, 잘못된 풀링 방식을 사용하면 그럴듯해 보이지만 모델이 학습하도록 훈련되지 않은 임베딩을 생성하기 때문입니다. 만약 어떤 체크포인트가 eqx-zoo가 아직 구현하지 못한 방식, 예를 들어 다른 풀링 모드나 추가적인 투영 레이어(projection layer)를 요구한다면, 조용히 다른 것을 계산하는 대신 명확한 NotImplementedError와 함께 로딩에 실패합니다.
무엇이 읽혀왔는지(model.pooling, model.normalize)를 확인하고, 대신 model(ids, mask)을 사용하여 토큰별 은닉 상태(per-token hidden states)를 얻을 수 있습니다.
작은 의미론적 검색 (A tiny semantic search)
모델들은 일반적인 JAX pytree이므로, 평소와 같은 변환(transformation)들이 적용됩니다. 여기서는 몇 개의 문서를 대상으로 검색을 수행하며, 임베딩 단계는 JIT-컴파일되었습니다:
import equinox as eqx
@eqx.filter_jit
...
0.533 섬으로 가는 페리 시간표가 아홉 시에 출발합니다.
0.207 등대지기는 매일 일지에 기록했습니다.
0.065 레몬 케이크를 위한 새로운 레시피입니다.
...
페리 시간표 관련 내용이 가장 상위에 나타났지만, 쿼리와 공유하는 단어는
eqx.filter_jit은 입력 형태(input shape)마다 함수를 한 번 컴파일합니다. enable_padding()을 사용하면 각 배치(batch)가 자체적으로 가장 긴 문장 길이로 패딩되므로, 새로운 길이는 새로운 컴파일을 의미합니다. 안정적인 처리량(throughput)을 위해서는 대신 고정된 길이로 패딩하는 것이 좋습니다. 예를 들어 tokenizer.enable_padding(length=128)와 같이 설정할 수 있습니다.
다국어 임베딩 (Multilingual embeddings)
multilingual-e5 모델은 약 100개 언어를 지원하며 동일한 Encoder API를 사용합니다. 레포지토리(repository)만 교체하면 되며, 코드는 다음과 같이 변경됩니다:
repo = "intfloat/multilingual-e5-base"
tokenizer = Tokenizer.from_pretrained(repo)
tokenizer.enable_padding()
...
이전과 마찬가지로 토큰화하고 임베딩합니다. 영어 쿼리(query)에 대해 프랑스 등대 지문은 0.741점을 받고, 금리 관련 독일 지문은 약 0.670점을 받습니다.
두 가지 세부 사항은 모델 카드에서 가져온 것이며, 둘 다 중요합니다:
- 모든 입력은 반드시
query:또는passage:로 시작해야 합니다. 모델이 훈련된 방식이며, 접두사(prefix)를 생략하면 결과가 저하됩니다. 동일한 종류의 텍스트 간 유사도를 측정할 때는 카드에서 양쪽에 모두query:를 권장합니다. - 점수(Scores)는 주로 0.7에서 1.0 사이로 높게 형성되는데, 이는 모델이 훈련된 방식 때문입니다. 중요한 것은 절대적인 값이 아니라 점수의 순서입니다.
기본 모델은 XLM-RoBERTa이며, multilingual-e5-small은 다국어 어휘(vocabulary)를 가진 BERT입니다. eqx-zoo는 두 아키텍처를 모두 지원하므로, 어떤 것이 무엇인지 알 필요가 없습니다.
Qwen3-Embedding: 임베더로서의 언어 모델
Qwen3-Embedding은 BERT 스타일 모델과는 다르게 작동합니다. 이는 채팅 모델처럼 디코더(decoder)이며, 인과적 어텐션(causal attention)을 사용합니다: 각 토큰은 자신 앞의 토큰들만 볼 수 있습니다. 따라서 임베딩은 전체 입력을 본 유일한 토큰인 마지막 토큰의 은닉 상태(hidden state)가 됩니다. eqx-zoo는 DecoderEmbedder를 사용하여 이를 로드하며, 이 클래스는 동일한 embed 메서드를 가지고 있습니다:
from eqx_zoo import DecoderEmbedder
repo = "Qwen/Qwen3-Embedding-0.6B"
...
e5와 마찬가지로 따라야 할 규칙이 있습니다: 쿼리(queries)에는 지침 프롬프트(instruction prompt)가 붙고, 문서(documents)에는 붙지 않습니다. 이 프롬프트는 체크포인트의 sentence-transformers 설정에서 가져오며, 검색에 맞게 작업 설명(task description)을 변경할 수 있습니다.
토크나이저(tokenizer)는 또한 모든 입력에 끝 문자 토큰(end-of-text token)을 추가하며, 이 토큰의 은닉 상태가 임베딩이 됩니다. 독립적인 tokenizers 라이브러리가 sentence-transformers가 하는 것처럼 자동으로 이를 추가합니다.
올바른 방법을 아는 방법
로드하여 실행할 수 있는 포트(port)가 반드시 정확한 것은 아닙니다: 누락된 바이어스(bias)나 잘못된 풀링(pooling)도 여전히 합리적으로 보이는 벡터를 생성할 수 있습니다. 따라서 지원되는 모든 체크포인트는 모든 풀 리퀘스트(pull request)에서 참조 구현체와 비교하여 확인됩니다:
- 레이어별(Layer by layer): 각 레이어의 출력은 Hugging Face Transformers에서 float32로 일치합니다.
- 종단 간(End to end): 임베딩은 길이가 다른 패딩된 문장 배치에 대해 sentence-transformers 자체와 일치합니다. 본 게시물의 예시에서는, all-MiniLM-L6-v2의 경우 최대 차이가 1.7e-7이며, query 프롬프트를 사용한 Qwen3-Embedding의 경우 float32 반올림으로 인해 4.7e-7입니다.
- bfloat16에서: float32 대비 오차는 참조 라이브러리 자체의 bf16 오차 범위 내 2배 이내여야 합니다. 측정된 비율은 인코더의 경우 0.84에서 1.51 사이이며, Qwen3-Embedding의 경우 최대 1.38입니다.
- 블록별(Block by block): LayerNorm과 같이 정밀도에 민감한 부분은 bf16에서 직접 테스트됩니다. 왜냐하면 이러한 곳의 미묘한 버그는 전체 모델의 출력을 거의 움직이지 않게 하지만, 블록 수준에서는 수천 번의 반올림 단계를 벗어나기 때문입니다.
- 매우 작은 랜덤 모델(Tiny random models): 각 아키텍처의 작고 무작위로 초기화된 버전은 단일 체크포인트로는 테스트할 수 없는 코드 경로를 커버하며, 몇 초 만에 실행됩니다.
지금까지 검증된 체크포인트는 all-MiniLM-L6-v2, bge-small-en-v1.5, multilingual-e5-small, multilingual-e5-base 및 Qwen3-Embedding-0.6B입니다. 이들은 eqx-zoo에서 Verified 상태로 수집되었습니다. 동일한 아키텍처를 가진 다른 체크포인트(BERT, RoBERTa, XLM-RoBERTa, Qwen3)도 동일한 방식으로 로드됩니다.
제한 사항 및 사용 시 주의점
eqx-zoo는 비교적 새로운 프로젝트이므로 몇 가지 솔직한 주의사항을 알려드립니다:
- 테스트 스위트는 CPU에서 실행됩니다. JAX는 GPU와 TPU 모두에서 동일한 코드를 실행하지만, 수치 검사는 아직 체계적으로 수행되지 않았습니다. GPU 및 TPU 보고서는 매우 환영하며, 저장소에는 이를 위한 이슈 템플릿이 있습니다.
- 임베딩 모델은 위의 아키텍처로 제한됩니다. all-mpnet-base-v2와 같은 MPNet 기반 모델이나 사용자 정의 코드가 필요한 체크포인트는 아직 지원되지 않습니다.
- 토큰화는 사용자가 결정합니다. eqx-zoo는 토큰 ID를 받기 때문에, 어떻게 토큰화(tokenize), 자르기(truncate), 패딩(pad)할지 사용자가 선택해야 합니다. 여기서 사용된 tokenizers 라이브러리는 sentence-transformers가 수행하는 방식과 일치합니다.
사용 방법:
pip install eqx-zoo tokenizers
코드는 GitHub의 xquantize/eqx-zoo에 있으며, 여기에 Llama, Qwen 및 [Qwen3-MoE](https://huggingface.co/docs/transformers/en/model_doc/qwen3_moe] 언어 모델도 동일한 방식으로 검증되었습니다. 버그 보고부터 새로운 체크포인트까지, 이슈와 풀 리퀘스트를 환영합니다.
JAX에서 어떤 임베딩 모델을 가장 사용하고 싶으신가요? 댓글로 알려주세요.
AI 자동 생성 콘텐츠
본 콘텐츠는 Dev.to AI tag의 원문을 AI가 자동으로 요약·번역·분석한 것입니다. 원 저작권은 원저작자에게 있으며, 정확한 내용은 반드시 원문을 확인해 주세요.
원문 바로가기