mahiyama/splade-ja-310m を CPU 推論向けに ONNX 化 + 動的 INT8 量子化したバリアントです。
オンライン検索のクエリエンコーダ用途を想定し、1 リクエストの低レイテンシを最優先に最適化しています。
CPU PyTorch (FP32) から CPU ONNX INT8 で P50 を 4.33 倍高速化し、モデルサイズは 4 分の 1 に圧縮しています。
GPU PyTorch (RTX 3080) の P50 が 31.4 ms であり、本モデル (44.2 ms) は CPU 推論で GPU と同等オーダーのレイテンシを達成しています。
1# pip install "optimum[onnxruntime]" "transformers>=4.48,<5" "onnxruntime<1.20" torch
2import torch
3import onnxruntime as ort
4from optimum.onnxruntime import ORTModelForMaskedLM
5from transformers import AutoTokenizer
6
7MODEL_ID = "mahiyama/splade-ja-310m-onnx-int8"
8
9# --- 1. モデル & トークナイザのロード (起動時に 1 回) -------------------------
10opts = ort.SessionOptions()
11opts.intra_op_num_threads = 4 # 物理コア数 or 1〜4 で実測して最速を選ぶ
12opts.inter_op_num_threads = 1
13opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
14
15model = ORTModelForMaskedLM.from_pretrained(
16 MODEL_ID,
17 file_name="model_quantized.onnx",
18 provider="CPUExecutionProvider",
19 session_options=opts,
20)
21tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
22
23
24# --- 2. クエリ 1 件をスパースベクトルに変換 ----------------------------------
25@torch.no_grad()
26def encode_query(text: str, top_k: int = 128, max_length: int = 512):
27 """SPLADE スパースベクトル (indices, values) を返す。
28
29 top_k: 量子化ノイズで非ゼロ次元が増えるため、本番では上位 K 次元だけ
30 残すと後段のスパース検索が速くなる。
31 """
32 inputs = tokenizer(
33 text, return_tensors="pt", truncation=True, max_length=max_length,
34 )
35 logits = model(**inputs).logits # (1, L, V=102400)
36
37 # SPLADE pooling: max_i log(1 + ReLU(logits))
38 activated = torch.log1p(torch.relu(logits)) # (1, L, V)
39 activated *= inputs["attention_mask"].unsqueeze(-1) # mask padding
40 pooled = activated.max(dim=1).values.squeeze(0) # (V,)
41
42 # Top-K プルーニング (推奨)
43 values, idx = pooled.topk(top_k)
44 mask = values > 0
45 return idx[mask].numpy(), values[mask].numpy()
46
47
48# --- 3. 動かしてみる ----------------------------------------------------------
49indices, values = encode_query("日本の首都はどこですか?")
50print(f"non-zero dims : {len(indices)}")
51print(f"top tokens : "
52 f"{[tokenizer.decode([i]) for i in indices[(-values).argsort()[:10]]]}")
1import numpy as np
2import onnxruntime as ort
3from huggingface_hub import snapshot_download
4from transformers import AutoTokenizer
5
6MODEL_ID = "mahiyama/splade-ja-310m-onnx-int8"
7local_dir = snapshot_download(MODEL_ID)
8
9opts = ort.SessionOptions()
10opts.intra_op_num_threads = 4
11opts.inter_op_num_threads = 1
12opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
13
14session = ort.InferenceSession(
15 f"{local_dir}/model_quantized.onnx",
16 sess_options=opts,
17 providers=["CPUExecutionProvider"],
18)
19tokenizer = AutoTokenizer.from_pretrained(local_dir)
20input_names = {i.name for i in session.get_inputs()}
21
22
23def encode_query(text, top_k=128, max_length=512):
24 inputs = tokenizer(
25 text, return_tensors="np", truncation=True, max_length=max_length,
26 )
27 feed = {k: v for k, v in inputs.items() if k in input_names}
28 logits = session.run(None, feed)[0] # (1, L, V)
29 activated = np.log1p(np.maximum(logits, 0.0))
30 activated *= inputs["attention_mask"][:, :, None]
31 pooled = activated.max(axis=1).squeeze(0) # (V,)
32 idx = np.argpartition(-pooled, top_k)[:top_k]
33 mask = pooled[idx] > 0
34 return idx[mask], pooled[idx][mask]
ModernBERT (vocab = 102400) では per_channel=True が極端に遅く (15 分以上で未終了)、本モデルは per_channel=False を採用しています。精度面ではどちらも recall@10 へのインパクトは僅少なため、運用上問題はありません。
このモデルは avx512_vnni プロファイルでビルドしてあります。
古い世代の CPU でうまく動かない場合は、onnxruntime 側で provider オプションを変えるか、自前で再量子化してください。
オリジナルの mahiyama/splade-ja-310m は transformers 5.x でアップロードされており、tokenizer_config.json の tokenizer_class が TokenizersBackend (5.x 専用クラス) になっています。本リポでは optimum 2.1 が transformers 5.x 非対応のため、tokenizer_class を PreTrainedTokenizerFast に書き換えた transformers 4.x 互換版を同梱しています。挙動は完全に同じです。