Lightning-AI/torchmetrics

Machine learning metrics for distributed, scalable PyTorch applications.

What it solves

TorchMetricsは、PyTorchアプリケーション向けに機械学習メトリクスを計算するための標準化された方法を提供します。複数のバッチや分散デバイス(複数のGPUやノードなど)にわたるメトリクスの累積と同期に通常必要となるボイラープレートコードを排除し、結果の再現性とスケーラビリティを確保します。

How it works

このライブラリは、主に2つのメトリクス計算方法を提供します:

  • Module-based metrics: これらはPyTorchモジュールのように動作し、内部状態を保持することでバッチを越えてデータを自動的に追跡・累積します。複数のデバイス間での同期を自動的に処理するため、CPU、単一GPU、またはマルチGPU構成に対応しています。
  • Functional metrics: これらは単純なPython関数であり、テンソルを入力として受け取り、即座にメトリクス値を返します。状態は保持しません。

ユーザーは、torchmetrics.Metric をサブクラス化して、メトリクスがどのように状態を更新し、最終結果を計算するかを定義することで、カスタムメトリクスを作成することも可能です。

Who it’s for

PyTorch開発者や機械学習エンジニア向けに設計されており、音声、テキスト、画像などの多様なドメインにおけるモデルの性能を追跡する必要があるユーザーや、大規模な分散トレーニングを行っているユーザーに適しています。

Highlights

  • Extensive Library: 分類、回帰、セグメンテーション、音声、テキスト、マルチモーダルデータを含む100以上の組み込みメトリクスが含まれています。
  • Distributed Support: マルチデバイストレーニングのための組み込みの自動同期と累積機能を提供します。
  • Customizable: 基底クラスをサブクラス化することで、カスタムメトリクスを実装するための簡単なAPIを提供します。
  • Visualization: メトリクスの進捗を可視化するための統合されたプロットサポートを提供します。