Sazanami-ColBERT-310m

sbintuitions/modernbert-ja-310m を初期値に学習した、日本語の Late Interaction 検索モデルです。質問と文書をそれぞれトークンごとのベクトル列に変換し、MaxSim で類似度を計算します。文書を 1 本のベクトルにまとめないため、長い文書の一部だけが質問に関係する場合でも、その一致を残したまま順位を付けられます。

同じベースモデル・同じ学習データで学習した 4 つの検索方式のうち、文書を圧縮しない構成にあたります。文書表現を小さくした 2 つのモデルと単一ベクトルのモデルは、いずれも同じ条件で別に学習したものです。

1. モデルの概要

項目 内容
ベースモデル sbintuitions/modernbert-ja-310m
アーキテクチャ Transformer → トークン単位の線形射影 (768 → 128、バイアスなし) → パディング位置の除外 → トークン単位の L2 正規化
パラメータ数 約 315M (うち射影層 98,304)
ベクトルの次元 128
文書あたりのベクトル本数 有効トークン数と同数。文書の長さで変わり、JaGovFaqs-22k では平均 100.5 本、Jaqket では平均 741.4 本
推奨入力長 質問 256 トークン / 文書 1,024 トークン (モデルに保存された既定値。上限は 8,192)
スコア関数 MaxSim (質問トークンごとの最大類似度の合計)
接頭辞 質問に 検索クエリ: 、文書に 検索文書:
トークナイザ SentencePiece (ModernBERT-Ja の語彙をそのまま使用)
主要フレームワーク sentence-transformers 6.0.0 以降 (MultiVectorEncoder)
学習時の混合精度 bf16

関連モデル

4 モデルはすべて sbintuitions/modernbert-ja-310m を初期値に、同じ学習データで学習しています。違うのは、文書 1 件を何本のベクトルで表すかです。Sazanami-Embed-310m は単一ベクトルのモデルで、比較のために同じ条件で学習しました。

モデル 文書あたりのベクトル 選ぶ目安
Sazanami-ColBERT-310m (本モデル) 本文の有効トークン数と同数 長い文書を多数の候補から探す用途で最も精度が高い
Sazanami-AGC-310m 最大 32 本 データ量を抑えたいとき
Sazanami-MetaEmbed-310m 最大 64 本 (使う本数を減らせる) 類似度計算を速くしたいとき
Sazanami-Embed-310m 1 本 (768 次元) 単一ベクトル。データ量と計算量を最小にしたいとき

2. 元となる研究

Late Interaction は、質問と文書を別々にエンコードしておき、類似度計算のときだけトークン単位で突き合わせる検索方式です。ColBERT が提案しました (Khattab and Zaharia, 2020)。文書側のエンコーディングを事前に済ませられるため、質問と文書を一緒に読み込む Cross-Encoder より高速で、文書を 1 本のベクトルにまとめる Bi-Encoder より細かい一致を残せます。

スコアは MaxSim です。質問のトークン 1 つごとに、文書のトークンのうち最も似ているものとの類似度を取り、それを質問トークン全体で合計します。文書の一部だけが質問に関係する場合でも、その部分の一致がスコアに現れます。

MaxSim(q,d)=∑imax⁡j  qi⋅dj \text{MaxSim}(q, d) = \sum_{i} \max_{j} \; q_i \cdot d_j

教師モデルの採点を使って学習させる考え方は ColBERTv2 に近いものです (Santhanam et al., 2022)。一方で、本モデルは ColBERT の原論文の実装と次の点が違います。

  • 質問を [MASK] で埋める拡張を使いません。原論文は質問を固定長まで [MASK] で埋めて類似度計算の対象を増やしますが、本モデルは有効トークンだけを使います。
  • 句読点を除外する処理を入れていません。原論文の実装は記号トークンを類似度計算から外しますが、本モデルは除外しません。

ベースモデルは ModernBERT を日本語で事前学習した ModernBERT-Ja です。位置埋め込みの上限が 8,192 トークンあるため、長い文書をそのまま読み込めます。

3. 使い方

sentence-transformers 6.0.0 以降の MultiVectorEncoder で読み込みます。

from sentence_transformers import MultiVectorEncoder

model = MultiVectorEncoder("mahiyama/Sazanami-ColBERT-310m")

queries = [
    "住民票はどこで取れますか",
    "国民健康保険の加入手続きを知りたい",
]
documents = [
    "住民票の写しは、市区町村の窓口またはコンビニ交付サービスで取得できます。",
    "国民健康保険の加入手続きは、転入や退職から 14 日以内に市区町村の窓口で行ってください。",
    "粗大ごみの収集は、事前の申し込みが必要です。",
]

# 接頭辞はモデルに保存されており、本文に付ける必要はない
query_vectors = model.encode_query(queries)
document_vectors = model.encode_document(documents)

# 保存された設定にしたがって MaxSim で採点される
scores = model.similarity(query_vectors, document_vectors)
print(scores)

4. トレーニング方法

ベースモデル

sbintuitions/modernbert-ja-310m を初期値にしています。穴埋めの事前学習だけを済ませたモデルで、検索向けの学習は入っていません。読み込んだ時点でプーリング層が外れ、トークン単位の線形射影 (768 → 128)、パディング位置の除外、トークン単位の L2 正規化が付きます。射影層は新規に初期化されます。

学習データ

質問 1 件に正例 1 件と負例 5 件を並べた n-tuples と、正例だけを並べた pairs の 2 種類を使いました。

出典 形式 件数
mahiyama/auto-wiki-qa n-tuples 100,000
mahiyama/mqa-ja n-tuples 100,000
mahiyama/mmarco-ja n-tuples 50,000
mahiyama/civicqa-ja (非公開) n-tuples 43,382
mahiyama/quiz-no-mori n-tuples 13,422
mahiyama/quiz-works n-tuples 12,502
mahiyama/amagasaki-qna n-tuples 11,069
mahiyama/miracl-retrieval n-tuples 4,603
mahiyama/mrtydi n-tuples 3,602
mahiyama/anlp-meeting-retrieval (title-abs) n-tuples 1,926
mahiyama/anlp-meeting-retrieval (title-intro) n-tuples 1,952
mahiyama/civicqa-ja (非公開) pairs 28,653
mahiyama/amagasaki-qna pairs 7,323
mahiyama/kosodate-faq-pairs-ja pairs 595
合計 379,029

mahiyama/civicqa-ja は、日本の国の機関と自治体が公開している FAQ から収集した質問と回答です。現在は非公開です。

評価に使うデータは学習から外しています。JaGovFaqs-22k は学習に一切使っていません。NFKC 正規化して空白を除いた質問が JaGovFaqs-22k の評価分割の質問と一致する行、または同じ処理をした正例がその評価分割の文書と一致する行を、各段階で除外しています。

訓練方法

n-tuples には、対照損失と蒸留損失を足したものを使います。両方を 1 回のスコア計算から求めます。組み込みの損失を 2 つ並べると同じ文書を 2 回エンコードすることになるため、1 つの損失にまとめています。

  • 対照損失 — バッチ内のすべての文書 (自分の候補 6 件と、同じバッチの他の行の文書) を選択肢とする交差エントロピー。
  • 蒸留損失 — 自分の候補 6 件に対する MaxSim の確率分布と、教師スコアの確率分布の KL ダイバージェンス。教師の温度は 2.0、生徒の温度は 1.0 です。

pairs には、バッチ内の他の行の文書を負例とする対照損失を使います。

項目 値
learning_rate (Transformer) 2e-5
learning_rate (射影層) 1e-4
num_train_epochs 1
per_device_train_batch_size 64
lr_scheduler_type linear
warmup_steps 全体の 5%
weight_decay 0.01
学習時の入力長 質問 128 トークン / 文書 512 トークン
混合精度 bf16
batch_sampler NO_DUPLICATES
勾配チェックポイント 有効
教師の温度 / 生徒の温度 2.0 / 1.0
乱数シード 42

訓練時間は NVIDIA A100 (40 GB) 1 枚で約 6 時間 18 分 (5,923 ステップ)、GPU メモリの最大使用量は 21.6 GB でした。

5. Sentence Transformers での学習コード

MultiVectorEncoder で初期値を読み込み、MaxSim を使う損失を定義し、MultiVectorEncoderTrainer に渡します。train_dataset の作り方は含めていません。

import torch
import torch.nn.functional as F
from torch import nn
from sentence_transformers import (
    MultiVectorEncoder,
    MultiVectorEncoderTrainer,
    MultiVectorEncoderTrainingArguments,
)
from sentence_transformers.base.losses.merged_forward import embed_columns_padded
from sentence_transformers.base.sampler import BatchSamplers
from sentence_transformers.multi_vector_encoder.scoring import colbert_scores
from sentence_transformers.util import stack_padded_token_embeddings

QUERY_PREFIX = "検索クエリ: "
DOCUMENT_PREFIX = "検索文書: "


class ContrastiveDistillLoss(nn.Module):
    """対照損失と教師スコアの蒸留を、1 回の MaxSim 計算から求める損失。

    入力: 質問と候補文書の特徴、教師スコア (バッチ数, 候補数)
    出力: 対照損失と蒸留損失の和 (スカラー)
    """

    def __init__(self, model, teacher_temperature=2.0, student_temperature=1.0):
        super().__init__()
        self.model = model
        self.teacher_temperature = teacher_temperature
        self.student_temperature = student_temperature
        self.kl = nn.KLDivLoss(reduction="batchmean", log_target=True)

    def forward(self, sentence_features, labels):
        sentence_features = list(sentence_features)

        # 質問をエンコードする。パディング位置は MaxSim の合計から除く。
        query_features = sentence_features[0]
        query_outputs = self.model(query_features, task="query")
        queries = query_outputs["token_embeddings"]
        query_mask = query_outputs["attention_mask"].bool()

        # 文書の列をまとめてエンコードし、(バッチ数, 候補数, トークン数, 次元) に積む。
        doc_embeddings, doc_masks = embed_columns_padded(
            self.model, sentence_features[1:], None, task_default="document"
        )
        docs, docs_mask = stack_padded_token_embeddings(doc_embeddings, doc_masks)

        # バッチ内の全文書に対する MaxSim のスコア行列を作り、対照損失を取る。
        scores = colbert_scores(queries, docs, queries_mask=query_mask, documents_mask=docs_mask)
        batch_size = queries.size(0)
        n_ways = len(sentence_features) - 1
        targets = torch.arange(batch_size, device=scores.device) * n_ways
        contrastive = F.cross_entropy(scores, targets)

        # 自分の候補だけを取り出し、教師の分布との KL を取る。
        rows = torch.arange(batch_size)
        own = scores.view(batch_size, batch_size, n_ways)[rows, rows]
        student = F.log_softmax(own / self.student_temperature, dim=-1)
        teacher = F.log_softmax(labels.detach().float() / self.teacher_temperature, dim=-1)
        distill = self.kl(student, teacher) * (self.student_temperature**2)

        return contrastive + distill


# 初期値を読み込む。プーリング層が外れ、射影層と正規化層が新規に付く。
model = MultiVectorEncoder(
    "sbintuitions/modernbert-ja-310m",
    prompts={"query": QUERY_PREFIX, "document": DOCUMENT_PREFIX},
    model_kwargs={"torch_dtype": torch.float32},
    similarity_fn_name="maxsim",
)

# 保存後の推論で使う入力長を、この時点で指定する。
model[0].query_length = 256
model[0].document_length = 1024

# 勾配チェックポイントは、学習引数ではなく Transformer に直接有効化する。
model[0].auto_model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})

# 学習時の接頭辞は、列名ごとに指定する必要がある。
prompts = {"query": QUERY_PREFIX, "positive": DOCUMENT_PREFIX}
prompts.update({f"negative_{i}": DOCUMENT_PREFIX for i in range(1, 6)})

training_args = MultiVectorEncoderTrainingArguments(
    output_dir="runs/sazanami-colbert",
    num_train_epochs=1,
    per_device_train_batch_size=64,
    learning_rate=2e-5,
    learning_rate_mapping={r"\.1\.linear\.": 1e-4},   # 射影層だけ学習率を上げる
    warmup_steps=0.05,                               # 比率を float で渡す
    lr_scheduler_type="linear",
    weight_decay=0.01,
    bf16=True,
    batch_sampler=BatchSamplers.NO_DUPLICATES,
    prompts=prompts,
    max_length={"query": 128, "document": 512},
    seed=42,
)

trainer = MultiVectorEncoderTrainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,                     # query, positive, negative_1..5, label を持つ Dataset
    loss=ContrastiveDistillLoss(model),
)
trainer.train()
model.save_pretrained("runs/sazanami-colbert/final")

次の 3 点は、sentence-transformers 6.1.0 と transformers 5.16.1 で確認した挙動です。

  • 学習時の接頭辞は自動では付きません。モデルに保存した prompts は推論でしか使われないため、MultiVectorEncoderTrainingArguments(prompts={列名: 接頭辞}) で列ごとに指定しない場合、接頭辞なしで学習されます。
  • gradient_checkpointing=True を学習引数に渡すと失敗します。引数の受け渡しが噛み合わないため、model[0].auto_model.gradient_checkpointing_enable(...) を直接呼びます。これを入れない場合、バッチ 64 × 文書 6 件 × 512 トークンでメモリが足りません。
  • warmup_ratio は使えません。warmup_steps に比率を float で渡します。

6. 評価結果

MTEB 2.21.8 の JMTEB(v2) に含まれる Retrieval 11 タスクのうち、文書の長さと候補数が異なる 3 件を測りました。指標は NDCG@10 です。文書 1,024 トークン、質問 256 トークンの設定です。

ベンチマークの規模

3 タスクの文書数、質問数、文書の長さです。

タスク 文書数 クエリ数 文書の平均トークン
Mintaka 1,592 2,312 10.0
JaGovFaqs 22,794 2,048 100.5
Jaqket 114,229 997 741.4

文書の平均トークンは、全トークンを保持する Sazanami-ColBERT-310m が実際に作ったベクトル本数から求めた値です (1 本 = 1 トークン)。文書 1,024 トークンの設定で測っています。

検索精度

モデル Mintaka JaGovFaqs Jaqket
Sazanami-ColBERT-310m (本モデル) 0.18070 0.77986 0.77868
Sazanami-AGC-310m 0.17302 0.78513 0.72986
Sazanami-MetaEmbed-310m 0.22697 0.79316 0.72855
Sazanami-Embed-310m 0.21771 0.78883 0.69682

データ量と類似度計算の時間

文書 114,229 件の Jaqket で測った実測値です。文書 1,024 トークンの設定です。エンコーディングは文書の本文をベクトルに変える処理で、全文書を処理した時間を 114,229 で割った値です (索引を作るときの一括処理なので、GPU のバッチ処理で均した値になります)。類似度計算の時間は、評価質問 997 件すべてを処理した時間を 997 で割った値です。

モデル 文書あたりのベクトル本数 次元 データ量の比 文書 1 件のエンコーディング 質問 1 件の類似度計算
Sazanami-ColBERT-310m (本モデル) 741.4 128 1 11.57 ms 43.03 ms
Sazanami-AGC-310m 32.0 128 約 1/23 12.22 ms 4.01 ms
Sazanami-MetaEmbed-310m 64.0 128 約 1/12 12.56 ms 0.70 ms
Sazanami-Embed-310m 1.0 768 約 1/124 12.83 ms 0.90 ms

7. 評価結果の考察

  • 文書が長く候補が多いほど、本モデルの優位が出ます。文書が平均 741 トークン・114,229 件の Jaqket では、Sazanami-AGC-310m より 0.049、Sazanami-MetaEmbed-310m より 0.050、Sazanami-Embed-310m より 0.082 高い NDCG@10 です。質問に関係する箇所が文書の一部にある場合、トークン単位の一致を残す MaxSim が有利に働くと解釈できます。
  • 文書が短い課題では、他の方式が上回ります。文書が平均 100 トークンの JaGovFaqs では Sazanami-MetaEmbed-310m が 0.013、Sazanami-Embed-310m が 0.009 上回ります。文書が平均 10 トークンの Mintaka では Sazanami-MetaEmbed-310m が 0.046、Sazanami-Embed-310m が 0.037 上回ります。多ベクトルを使えば常に精度が高くなる、という結果ではありません。

8. ライセンスと学習データの制約

本モデルの重みは MIT License で提供します。ベースモデルの sbintuitions/modernbert-ja-310m も MIT License です。

学習データに MS MARCO の日本語訳 (mahiyama/mmarco-ja) を含みます。MS MARCO は非商用の研究目的での利用を前提に公開されており、商用利用のライセンスは明確ではありません。本モデルを商用で利用する場合は、利用者ご自身の責任で MS MARCO の利用規約を確認してください。

学習データの一部に、日本の国の機関と自治体が公開している FAQ から収集した質問と回答を含みます。各出典の利用条件に従うデータで、本モデルのリポジトリでは再配布していません。本モデルを利用する場合は、利用者ご自身の責任で用途を判断してください。

Downloads last month
33
Safetensors
Model size
0.3B params
Tensor type
F32
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for mahiyama/Sazanami-ColBERT-310m

Finetuned
(21)
this model

Datasets used to train mahiyama/Sazanami-ColBERT-310m

Collection including mahiyama/Sazanami-ColBERT-310m

Papers for mahiyama/Sazanami-ColBERT-310m