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: 集成了绘图支持,以可视化指标随时间的变化趋势。