이커머스 검색을 위한 희소 임베딩 파인튜닝 | 2부: Modal에서 SPLADE 훈련하기

이커머스 검색을 위한 희소 임베딩 파인튜닝 | 2부: Modal에서 SPLADE 훈련하기 (Fine-Tuning Sparse Embeddings for E-Commerce Search | Part 2: Training SPLADE on Modal)

1부에서 희소 임베딩이 왜 이커머스에 강력한지 봤다면, 2부에서는 그걸 코드로 옮겨요. Amazon ESCI 데이터셋을 로딩하고, SPLADE 모델을 만들어 Modal의 서버리스 GPU에서 파인튜닝하는 전체 훈련 파이프라인을 살펴볼게요. 코드와 하이퍼파라미터, CLI 명령은 원문 그대로 유지했어요.

출처: 공식문서

이커머스 검색을 위한 희소 임베딩 파인튜닝에 관한 5부작 시리즈의 2부입니다. 1부에서 희소 임베딩이 왜 이커머스에서 BM25를 이기는지 다뤘어요. 이제 훈련 파이프라인을 구축합니다.

시리즈:


지난 글에서 이커머스 검색에 희소 임베딩이 필요한 이유를 설명했어요. 이제 코드를 씁니다. 모든 소스 코드는 GitHub 저장소에 있고, 파인튜닝된 모델을 HuggingFace에서 써 볼 수 있어요. 바로 자신의 데이터로 파인튜닝하고 싶다면 sparse-finetune CLI를 보세요. 이 글을 끝내면 여러분은 Amazon의 ESCI 데이터셋으로 훈련된 SPLADE 모델을 얻고, 그것을 Modal의 서버리스 GPU에서 실행하며, 체크포인트를 영구 저장소에 저장하게 될 거예요.

데이터셋: Amazon ESCI

KDD Cup 2022를 위해 공개된 Amazon의 ESCI 데이터셋(Shopping Queries Dataset)을 사용해요. 이것은 가장 현실적인 이커머스 검색 벤치마크 중 하나예요.

  • 사람이 주석을 단 관련성 레이블이 있는 120만 개 이상의 질의-제품 쌍
  • 네 가지 관련성 등급: Exact(E), Substitute(S), Complement(C), Irrelevant(I)
  • 풍부한 제품 메타데이터: 제목, 설명, 불릿 포인트, 브랜드

등급화된 관련성(graded relevance)이 ESCI를 흥미롭게 만드는 부분이에요.

훈련을 위해 우리는 Exact와 Substitute 쌍을 긍정(positive)으로 사용해요. 이는 정확한 제품과 합리적인 대안이 모두 관련 있다는 걸 모델에게 가르치는데, 실제 쇼핑객이 생각하는 방식과 정확히 맞아떨어져요.

데이터 로딩

from datasets import load_dataset
from src.data.text_builder import build_product_text

def load_esci_training_data(max_samples=None):
    """Load ESCI dataset as anchor-positive pairs for contrastive training."""
    dataset = load_dataset("tasksource/esci", split="train")

    pairs = []
    for row in dataset:
        if row["relevance_label"] not in ("E", "S"):
            continue

        query = row["query"]
        product_text = build_product_text(
            title=row["product_title"],
            brand=row.get("product_brand", ""),
            description=row.get("product_description", ""),
            bullets=row.get("product_bullet_point", []),
        )
        pairs.append({"anchor": query, "positive": product_text})

        if max_samples and len(pairs) >= max_samples:
            break

    return pairs

제품 텍스트 포맷팅

제품 텍스트를 어떻게 포맷하는지는 희소 임베딩에 중요해요. 광범위한 의미를 포착하는 밀집 모델과 달리, SPLADE는 어휘적으로 근거를 둡니다. 텍스트의 특정 토큰이 어떤 어휘 차원이 활성화될지를 결정하거든요.

def build_product_text(title, brand="", description="", bullets=None, max_length=512):
    """Consistent product text formatting for SPLADE."""
    parts = []

    # Brand in brackets makes it a distinct signal
    if brand:
        parts.append(f"[{brand}]")

    parts.append(title)

    # Pipe separators help the model distinguish sections
    if description:
        parts.append(f"| {description[:200]}")

    if bullets:
        parts.append(f"| {' | '.join(bullets[:3])}")

    text = " ".join(parts)
    return text[:max_length]

# Example output:
# "[Sony] WH-1000XM5 Wireless Headphones | Industry-leading noise
#  cancellation | 30hr battery | Hi-Res Audio"

브랜드의 대괄호 표기, 섹션 사이의 파이프 구분자, 문자 수 제한은 모두 의도적이에요. SPLADE가 학습할 수 있는 어휘 시그널을 보존합니다. 브랜드명, 제품 속성, 핵심 기능이 텍스트의 벽으로 흐려지지 않고 별개의 토큰으로 유지되도록요.

Modal은 서버리스 GPU를 제공해요. 프로비저닝도, 유휴 하드웨어도 없고, 초 단위 과금이에요. 앱 구성은 다음과 같아요.

import modal

app = modal.App("esci-sparse-encoder")

# Persistent storage for checkpoints and datasets
checkpoint_volume = modal.Volume.from_name(
    "esci-sparse-checkpoints", create_if_missing=True
)
dataset_volume = modal.Volume.from_name(
    "esci-datasets", create_if_missing=True
)

# Docker image with dependencies
image = (
    modal.Image.debian_slim(python_version="3.11")
    .pip_install(
        "sentence-transformers>=5.0.0",
        "torch>=2.2.0",
        "transformers>=4.45.0",
        "datasets>=2.20.0",
        "qdrant-client>=1.12.0",
        "accelerate>=0.30.0",
    )
)

여기서 중요한 게 두 가지 있어요.

영구 볼륨(Persistent volumes). 훈련 실행은 몇 시간이 걸릴 수 있어요. SSH 연결이 끊기거나 컨테이너가 재시작되면 체크포인트를 잃고 싶지 않을 거예요. Modal 볼륨은 실행 간에 데이터를 영구 유지해요. 볼륨을 경로에 마운트하고 로컬 파일시스템처럼 쓰면 됩니다.

분리 실행(Detached runs). 긴 훈련 작업은 --detach로 시작하고 그 자리를 떠나면 돼요.

# Start training and disconnect
uv run modal run --detach modal_app.py --mode train

# Come back later, check your checkpoints
uv run modal volume ls esci-sparse-checkpoints /checkpoints/

S3 업로드도, 체크포인트 관리 코드도, 잃어버린 훈련 실행도 없어요.

SPLADE 모델 만들기

Sentence Transformers v5는 SparseEncoder를 도입해 SPLADE 훈련을 간단하게 만들었어요. 모델은 두 구성 요소를 가져요.

  1. MLMTransformer: 전체 어휘에 걸쳐 로짓(logits)을 출력하는 마스크 언어 모델 헤드를 가진 트랜스포머
  2. SpladePooling: 토큰 수준 로짓에 ReLU + 로그 포화를 적용하고 위치 전체에 걸쳐 맥스 풀링
from sentence_transformers import SparseEncoder
from sentence_transformers.sparse_encoder.models import (
    MLMTransformer,
    SpladePooling,
)

def create_sparse_encoder(base_model="distilbert/distilbert-base-uncased"):
    """Create a SPLADE model from a base transformer."""

    # MLM transformer outputs logits over vocabulary
    mlm = MLMTransformer(base_model)

    # SPLADE pooling: max over tokens, ReLU activation
    pooling = SpladePooling(pooling_strategy="max")

    return SparseEncoder(modules=[mlm, pooling])

우리는 사전 훈련된 SPLADE 체크포인트(예: naver/splade-v3)보다 DistilBERT에서 시작해요. 이것은 의도적인 선택이에요. 웹 검색 데이터로 이미 훈련된 모델이 아니라 일반 언어 모델에서 시작할 때, 도메인 특화 파인튜닝이 얼마나 도움이 되는지 측정하고 싶었어요.

훈련 함수

Modal 함수로 데코레이션된 핵심 훈련 로직은 다음과 같아요.

@app.function(
    image=image,
    gpu="A100",
    volumes={
        "/checkpoints": checkpoint_volume,
        "/datasets": dataset_volume,
    },
    timeout=3600 * 6,
)
def train_sparse_encoder(config: dict):
    from sentence_transformers import SparseEncoder
    from sentence_transformers.sparse_encoder import SparseEncoderTrainer
    from sentence_transformers.training_args import SparseEncoderTrainingArguments
    from sentence_transformers.losses import SpladeLoss, SparseMultipleNegativesRankingLoss

    # Create model
    model = create_sparse_encoder(config["base_model"])

    # Load ESCI dataset (anchor-positive pairs)
    train_dataset = load_esci_training_data(
        max_samples=config.get("max_samples")
    )

    # SPLADE loss combines contrastive learning with sparsity regularization
    loss = SpladeLoss(
        model=model,
        loss=SparseMultipleNegativesRankingLoss(model=model),
        query_regularizer_weight=float(config.get("query_regularizer_weight", 5e-5)),
        document_regularizer_weight=float(config.get("document_regularizer_weight", 3e-5)),
    )

    # Training arguments
    args = SparseEncoderTrainingArguments(
        output_dir=f"/checkpoints/{config['run_name']}",
        num_train_epochs=config.get("num_epochs", 1),
        per_device_train_batch_size=config.get("batch_size", 32),
        learning_rate=float(config.get("learning_rate", 2e-5)),
        warmup_ratio=0.1,
        fp16=True,
        save_steps=1000,
        logging_steps=100,
    )

    # Train
    trainer = SparseEncoderTrainer(
        model=model,
        args=args,
        train_dataset=train_dataset,
        loss=loss,
    )
    trainer.train()

    # Save final model
    model.save_pretrained(f"/checkpoints/{config['run_name']}/final")

    return f"/checkpoints/{config['run_name']}/final"

SpladeLoss 이해하기

SpladeLoss는 두 가지 목표를 감쌉니다.

대조 손실(Contrastive loss) (SparseMultipleNegativesRankingLoss): (질의, 제품) 쌍의 배치가 주어지면, 배치 안의 다른 제품을 네거티브로 취급해요. 관련 있는 질의-제품 쌍은 당기고, 관련 없는 쌍은 밀어내죠. 이는 밀집 임베딩 훈련에 쓰이는 것과 같은 인배치 네거티브(in-batch negative) 방식이며, 대부분의 무작위 제품이 주어진 질의에 관련이 없기 때문에 잘 동작해요.

희소성 정규화(Sparsity regularization): 효율성을 유지하기 위해 밀집한 출력에 패널티를 줘요. 없으면 모델이 모든 입력에 대해 30,000개 어휘 차원을 전부 활성화할 거예요. 매칭에는 기술적으로 최적이지만 검색 속도와 저장에는 쓸모없죠.

정규화 가중치가 이 트레이드오프를 제어해요.

Parameter Value Effect
query_regularizer_weight 5e-5 Higher = sparser queries
document_regularizer_weight 3e-5 Higher = sparser documents

최적 지점은 벡터당 활성 용어 100~300개예요. 정규화가 너무 높으면 거의 빈 벡터가 생겨요(빠르지만 재현율이 낮음). 너무 낮으면 수천 개의 용어가 생기죠(느리고 인덱스가 거대해짐).

문서 정규화가 질의 정규화보다 낮은 이유는, 제품 설명이 모든 관련 속성을 포착하려면 더 많은 용어가 필요하기 때문이에요. 헤드폰 제품 목록은 "audio", "wireless", "bluetooth", "noise", "canceling" 같은 용어를 활성화해야 합니다 — 일반 질의의 3~4개 단어보다 더 많죠.

YAML로 구성하기

실험을 쉽게 하기 위해 하이퍼파라미터를 YAML 파일에 보관해요.

# configs/splade_standard.yaml
run_name: splade_standard
base_model: distilbert/distilbert-base-uncased
architecture: splade
batch_size: 32
learning_rate: 2e-5
num_epochs: 1
query_regularizer_weight: 5e-5
document_regularizer_weight: 3e-5
max_samples: 100000

10만 개 샘플은 A100에서 약 6분 안에 훈련되고 Modal에서 1달러 미만이 들어요. 여러 에폭의 전체 120만 데이터셋은 몇 시간이 걸리지만, 예약 GPU 인스턴스에 비하면 여전히 저렴해요.

병렬 하이퍼파라미터 스윕

Modal의 강점 중 하나는 부끄러울 정도로 병렬적인(embarrassingly parallel) 워크로드예요. 하이퍼파라미터 스윕은 자연스러운 적용 대상이죠. spawn()은 구성 하나당 하나의 GPU를 실행합니다.

@app.function(gpu="A100")
def train_single_experiment(config: dict):
    """Train one configuration."""
    model = create_sparse_encoder(config["base_model"])
    # ... training code ...
    return {"config": config, "ndcg": evaluate(model)}

@app.local_entrypoint()
def run_hyperparameter_sweep():
    """Launch all experiments in parallel."""
    configs = [
        {"learning_rate": 1e-5, "regularizer_weight": 3e-5},
        {"learning_rate": 2e-5, "regularizer_weight": 3e-5},
        {"learning_rate": 2e-5, "regularizer_weight": 5e-5},
        {"learning_rate": 5e-5, "regularizer_weight": 5e-5},
        # ... more configurations ...
    ]

    # Launch all experiments simultaneously
    handles = [train_single_experiment.spawn(c) for c in configs]

    # Collect results as they complete
    results = [h.get() for h in handles]
    best = max(results, key=lambda r: r["ndcg"])
    print(f"Best config: {best}")

24개 실험의 스윕은 단일 훈련 실행 하나가 걸리는 시간 안에 끝나요. 각 실험은 자신의 A100을 받아요. 대기열에서 기다리는 유휴 GPU 비용이 아니라, 실제로 사용한 컴퓨팅 시간만 지불합니다.

하지 말아야 할 것: 추론-프리 SPLADE 함정

지연 시간을 줄이기 위해 질의 쪽 트랜스포머를 정적 임베딩 룩업으로 바꿔 보았어요. 아이디어는 매력적이었죠. 질의는 짧으니까, 왜 전체 트랜스포머를 돌리겠어요?

# DON'T DO THIS (for e-commerce)
router = Router.for_query_document(
    query_modules=[
        SparseStaticEmbedding(tokenizer=mlm.tokenizer)  # Fast but weak
    ],
    document_modules=[
        mlm,
        SpladePooling(pooling_strategy="max"),
    ],
)

결과는 처참했어요.

Architecture nDCG@10
Standard SPLADE (contextual) 0.389
Inference-Free (static) 0.065

참고: 이 지표들은 모든 관련 문서가 포함된 10만 개 제품과 1만 개 질의의 하위 표본에서 측정된 것이에요. 공식 Amazon ESCI 벤치마크와 직접 비교할 수 없으며, 비교적 신호로만 취급해야 해요.

문맥 인코딩이 없으면 6배나 나빠져요.

정적 임베딩이 완전히 실패한 이유는 이커머스 질의가 매우 문맥 의존적이기 때문이에요. "Apple"은 "apple iphone"과 "apple fruit"에서 서로 다른 뜻이에요. 정적 임베딩은 이걸 구분하지 못해요. "apple"을 룩업해서 문맥과 무관하게 같은 벡터를 반환하죠.

트랜스포머가 질의당 약 15ms로 병목이지만, 검색에 15ms는 완벽히 수용 가능해요. 모델을 동작하게 만드는 구성 요소를 성급하게 최적화해 없애지 마세요.

훈련 실행하기

모든 게 준비되면 훈련을 시작합니다.

# Quick test run (100K samples)
uv run modal run modal_app.py \
    --config-path configs/splade_standard.yaml \
    --mode train

# Full dataset, detached
uv run modal run --detach modal_app.py \
    --config-path configs/splade_standard.yaml \
    --mode train

모델 체크포인트는 /checkpoints/splade_standard/final의 영구 볼륨에 저장돼요. 또한 훈련된 모델을 splade-ecommerce-esci로 HuggingFace에 공개해서, 훈련을 건너뛰고 바로 사용할 수 있게 했어요. 다음 글에서 이 모델을 로드하고 Qdrant에 제품을 인덱싱한 뒤 검색 벤치마크를 실행해 BM25보다 얼마나 개선됐는지 정확히 볼 거예요.

핵심 요점

  • ESCI의 등급화된 관련성(Exact, Substitute, Complement, Irrelevant)은 모델에게 이분법적인 관련/비관련이 아니라 미묘한 매칭을 가르쳐요.
  • 제품 텍스트 포맷팅은 희소 모델에 중요하다. 구조화된 포맷팅으로 어휘 시그널을 뚜렷하게 유지하세요.
  • SpladeLoss는 두 목표의 균형을 잡는다: 관련성을 위한 대조 학습과 희소성을 위한 정규화. 정규화 가중치가 바로 튜닝할 주요 손잡이예요.
  • Modal의 영구 볼륨이 체크포인트 관리 문제를 해결해요. 분리 실행은 SSH 끊김에도 안전합니다.
  • 질의 트랜스포머를 건너뛰지 마세요. 15ms의 지연 시간이 정적 임베딩 대비 6배의 품질 향상을 사줍니다.

다음: Part 3 - 평가, 하드 네거티브, 결과

더 알아보기 (Learn more)