Lightning-AI/torchmetrics
Machine learning metrics for distributed, scalable PyTorch applications.
What it solves
TorchMetrics 提供了一種標準化的方式來為 PyTorch 應用程式計算機器學習指標。它消除了在多個批次和分散式裝置(例如多個 GPU 或節點)之間累積與同步指標通常所需的樣板程式碼,確保結果具有可重現性和可擴展性。
How it works
該函式庫提供兩種主要的指標計算方式:
- 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: 整合了繪圖支援,以視覺化指標隨時間的進展。