Instructions to use mahiyama/Sazanami-ColBERT-310m with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use mahiyama/Sazanami-ColBERT-310m with sentence-transformers:
from sentence_transformers import MultiVectorEncoder model = MultiVectorEncoder("mahiyama/Sazanami-ColBERT-310m") queries = ["Which planet is known as the Red Planet?"] documents = [ "Venus is often called Earth's twin because of its similar size and proximity.", "Mars, known for its reddish appearance, is often referred to as the Red Planet.", "Jupiter, the largest planet in our solar system, has a prominent red spot.", ] query_embeddings = model.encode_query(queries) document_embeddings = model.encode_document(documents) similarities = model.similarity(query_embeddings, document_embeddings) print(similarities) - Notebooks
- Google Colab
- Kaggle
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 つごとに、文書のトークンのうち最も似ているものとの類似度を取り、それを質問トークン全体で合計します。文書の一部だけが質問に関係する場合でも、その部分の一致がスコアに現れます。
教師モデルの採点を使って学習させる考え方は 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
Model tree for mahiyama/Sazanami-ColBERT-310m
Base model
sbintuitions/modernbert-ja-310m