Improve: Disable enable_cpu_mem_arena for CPU inference only to prevent excessive memory consumption during the inference session of the BERT model

This commit is contained in:
tsukumi
2024-12-22 09:17:31 +09:00
parent 3c218c7087
commit 810ca43615
3 changed files with 61 additions and 25 deletions

View File

@@ -43,6 +43,7 @@ def load_model(
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = [("CPUExecutionProvider", {"arena_extend_strategy": "kSameAsRequested"})],
cache_dir: Optional[str] = None,
revision: str = "main",
enable_cpu_mem_arena: bool | None = None,
) -> onnxruntime.InferenceSession: # fmt: skip
"""
指定された言語の ONNX 版 BERT モデルをロードし、ロード済みの ONNX 版 BERT モデルを返す。
@@ -61,6 +62,7 @@ def load_model(
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
enable_cpu_mem_arena (bool | None): CPU 推論時にもメモリアリーナを有効化するかどうか。デフォルトでは GPU 推論時のみ有効化される (デフォルト: None)
Returns:
onnxruntime.InferenceSession: ロード済みの BERT モデル
@@ -101,16 +103,40 @@ def load_model(
else:
model_path = Path(pretrained_model_name_or_path).resolve() / "model_fp16.onnx"
start_time = time.time()
# 推論時に一番優先される ExecutionProvider の名前を取得
assert len(onnx_providers) > 0
first_provider_name = (
onnx_providers[0]
if type(onnx_providers[0]) is str
else onnx_providers[0][0]
)
# 推論セッションの設定
sess_options = onnxruntime.SessionOptions()
# ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする
## ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip
# エラー以外のログを出力しない
# 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している
## エラー以外のログを出力しない
## 本来は log_severity_level = 3 だけで効くはずだが、なぜか CUDA 系のログが抑制できないので set_default_logger_severity() も呼び出している
sess_options.log_severity_level = 3
onnxruntime.set_default_logger_severity(3)
# CPU 推論時のみ enable_cpu_mem_arena を無効化し、BERT モデルの推論セッションより富豪的なメモリ消費を防止する
## 既に RunOptions の memory.enable_memory_arena_shrinkage や、ProviderOptions の "arena_extend_strategy": "kSameAsRequested" を指定して
## InferenceSession が構築するメモリアリーナを推論後に縮小するよう構成し、メモリアリーナによるメモリ消費量が漸進的に増加することを防いでいる
## しかし、CPU 推論時の BERT モデルに関しては入力長や入力内容次第では依然大量のメモリが確保される傾向にあるため、CPU 推論時のみメモリアリーナ自体を無効化する
## BERT 特徴量の抽出処理が 0.数秒遅くなるトレードオフがあるが、元々 CPU 推論は CUDA 推論よりかなり遅いこと、
## BERT 特徴量の抽出処理自体は音声合成処理よりも遥かに軽量なこと、低メモリ環境での OOM エラー回避の観点から有益だと判断した
## メモリアリーナを無効化することで、若干の速度低下と引き換えに、多量の推論処理を行ってもメモリリークのような挙動が発生しなくなる
## なお CUDA 推論時は独自に VRAM 管理が行われているようで、CPU 推論時のように過剰に VRAM が消費されることはない
## 明示的に enable_cpu_mem_arena が指定されている場合は、指定された値を利用する
if enable_cpu_mem_arena is not None:
sess_options.enable_cpu_mem_arena = enable_cpu_mem_arena
## 明示的に enable_cpu_mem_arena が指定されていない場合は、推論セッションが CPUExecutionProvider の場合のみメモリアリーナを無効化する
elif first_provider_name == "CPUExecutionProvider":
sess_options.enable_cpu_mem_arena = False
# BERT モデルをロードし、辞書に格納して返す
start_time = time.time()
__loaded_models[language] = onnxruntime.InferenceSession(
model_path,
sess_options=sess_options,