splade-ja-310m-onnx-int8

mahiyama/splade-ja-310m を CPU 推論向けに ONNX 化 + 動的 INT8 量子化したバリアントです。 オンライン検索のクエリエンコーダ用途を想定し、1 リクエストの低レイテンシを最優先に最適化しています。

ハイライト

指標 元の PyTorch (CPU) 本モデル (ONNX INT8 / CPU) 改善
P50 latency 191.3 ms 44.2 ms 4.33× 高速
P95 latency 250.7 ms 68.5 ms 3.66× 高速
P99 latency 297.4 ms 85.4 ms 3.48× 高速
mean latency 193.5 ms 46.6 ms 4.15× 高速
Model size 1576 MB 395 MB 4× 圧縮

CPU PyTorch (FP32) から CPU ONNX INT8 で P50 を 4.33 倍高速化し、モデルサイズは 4 分の 1 に圧縮しています。 GPU PyTorch (RTX 3080) の P50 が 31.4 ms であり、本モデル (44.2 ms) は CPU 推論で GPU と同等オーダーのレイテンシを達成しています。

使用例

# pip install "optimum[onnxruntime]" "transformers>=4.48,<5" "onnxruntime<1.20" torch
import torch
import onnxruntime as ort
from optimum.onnxruntime import ORTModelForMaskedLM
from transformers import AutoTokenizer

MODEL_ID = "mahiyama/splade-ja-310m-onnx-int8"

# --- 1. モデル & トークナイザのロード (起動時に 1 回) -------------------------
opts = ort.SessionOptions()
opts.intra_op_num_threads = 4   # 物理コア数 or 1〜4 で実測して最速を選ぶ
opts.inter_op_num_threads = 1
opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL

model = ORTModelForMaskedLM.from_pretrained(
    MODEL_ID,
    file_name="model_quantized.onnx",
    provider="CPUExecutionProvider",
    session_options=opts,
)
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)


# --- 2. クエリ 1 件をスパースベクトルに変換 ----------------------------------
@torch.no_grad()
def encode_query(text: str, top_k: int = 128, max_length: int = 512):
    """SPLADE スパースベクトル (indices, values) を返す。

    top_k: 量子化ノイズで非ゼロ次元が増えるため、本番では上位 K 次元だけ
           残すと後段のスパース検索が速くなる。
    """
    inputs = tokenizer(
        text, return_tensors="pt", truncation=True, max_length=max_length,
    )
    logits = model(**inputs).logits                          # (1, L, V=102400)

    # SPLADE pooling: max_i log(1 + ReLU(logits))
    activated = torch.log1p(torch.relu(logits))              # (1, L, V)
    activated *= inputs["attention_mask"].unsqueeze(-1)      # mask padding
    pooled = activated.max(dim=1).values.squeeze(0)          # (V,)

    # Top-K プルーニング (推奨)
    values, idx = pooled.topk(top_k)
    mask = values > 0
    return idx[mask].numpy(), values[mask].numpy()


# --- 3. 動かしてみる ----------------------------------------------------------
indices, values = encode_query("日本の首都はどこですか?")
print(f"non-zero dims : {len(indices)}")
print(f"top tokens    : "
      f"{[tokenizer.decode([i]) for i in indices[(-values).argsort()[:10]]]}")

onnxruntime 直叩き (optimum 不要の最小例)

import numpy as np
import onnxruntime as ort
from huggingface_hub import snapshot_download
from transformers import AutoTokenizer

MODEL_ID = "mahiyama/splade-ja-310m-onnx-int8"
local_dir = snapshot_download(MODEL_ID)

opts = ort.SessionOptions()
opts.intra_op_num_threads = 4
opts.inter_op_num_threads = 1
opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL

session = ort.InferenceSession(
    f"{local_dir}/model_quantized.onnx",
    sess_options=opts,
    providers=["CPUExecutionProvider"],
)
tokenizer = AutoTokenizer.from_pretrained(local_dir)
input_names = {i.name for i in session.get_inputs()}


def encode_query(text, top_k=128, max_length=512):
    inputs = tokenizer(
        text, return_tensors="np", truncation=True, max_length=max_length,
    )
    feed = {k: v for k, v in inputs.items() if k in input_names}
    logits = session.run(None, feed)[0]                          # (1, L, V)
    activated = np.log1p(np.maximum(logits, 0.0))
    activated *= inputs["attention_mask"][:, :, None]
    pooled = activated.max(axis=1).squeeze(0)                    # (V,)
    idx = np.argpartition(-pooled, top_k)[:top_k]
    mask = pooled[idx] > 0
    return idx[mask], pooled[idx][mask]

量子化の詳細

項目
手法 ONNX Runtime 動的 INT8 量子化 (post-training)
プロファイル avx512_vnni (Ice Lake / Sapphire Rapids 以降の CPU 想定)
per_channel False (per-tensor)
キャリブレーション 不要 (動的量子化)
量子化対象 MatMul / Add (LayerNorm は除外)
量子化所要時間 約 17 秒 (本モデルで実測)

ModernBERT (vocab = 102400) では per_channel=True が極端に遅く (15 分以上で未終了)、本モデルは per_channel=False を採用しています。精度面ではどちらも recall@10 へのインパクトは僅少なため、運用上問題はありません。

CPU 世代別の推奨

このモデルは avx512_vnni プロファイルでビルドしてあります。 古い世代の CPU でうまく動かない場合は、onnxruntime 側で provider オプションを変えるか、自前で再量子化してください。

CPU 世代 プロファイル
Ice Lake / Sapphire Rapids 以降 avx512_vnni (本モデル)
Skylake / Cascade Lake avx512
Haswell / Broadwell avx2
Apple Silicon / Graviton arm64

トークナイザの注意

オリジナルの 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 互換版を同梱しています。挙動は完全に同じです。

ベンチマーク詳細

レイテンシ計測 (本モデルで実測した値)

構成 P50 (ms) P95 (ms) P99 (ms) mean (ms) min (ms) max (ms)
GPU PyTorch FP32 (RTX 3080) 31.4 39.3 42.3 33.0 27.7 44.1
CPU PyTorch FP32 191.3 250.7 297.4 193.5 117.1 299.4
CPU ONNX INT8 (本モデル) 44.2 68.5 85.4 46.6 26.9 108.9

計測条件

  • バッチサイズ: 1 (オンラインクエリ想定)
  • max_length: 512
  • threads: 4
  • 100 trials + 10 warmup
  • 評価データセット: JaCWIR-mini (queries 1,000) からランダムサンプリング
  • 環境: Windows 11, 24 vCore CPU (AMD Ryzen 9 3900XT 12-Core Processor) + RTX 3080, Python 3.12, onnxruntime 1.19, optimum 2.1, transformers 4.57

制約事項

  • 量子化ノイズで非ゼロ次元 (query_active_dims) が増えます。検索に渡す前に Top-K プルーニング (k=128 程度) を強く推奨します。
  • 本モデルは CPU 専用にチューニングされており、GPU では速度メリットがありません。GPU 用途は元の PyTorch モデル (mahiyama/splade-ja-310m) を使ってください。
  • 量子化は CPU 推論コストの削減目的なので、オフラインで一度だけ計算するコーパス側は元モデル (GPU FP32) で処理する構成が最も品質が出ます。

ライセンス

MIT

Downloads last month
5
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for mahiyama/splade-ja-310m-onnx-int8

Quantized
(1)
this model

Dataset used to train mahiyama/splade-ja-310m-onnx-int8