diff --git a/.gitattributes b/.gitattributes index fc95bd6fe798ba04ceaecb0fbd49194ad0383303..f6436e783239fae0c5f7c3c860d0df80b6b7a860 100644 --- a/.gitattributes +++ b/.gitattributes @@ -5333,3 +5333,4 @@ rtme/lib/python3.10/site-packages/babel/locale-data/so.dat filter=lfs diff=lfs m rtme/lib/python3.10/site-packages/huggingface_hub/inference/_generated/__pycache__/_async_client.cpython-310.pyc filter=lfs diff=lfs merge=lfs -text rtme/lib/python3.10/site-packages/babel/locale-data/vi.dat filter=lfs diff=lfs merge=lfs -text rtme/lib/python3.10/site-packages/babel/locale-data/to.dat filter=lfs diff=lfs merge=lfs -text +rtme/lib/python3.10/site-packages/babel/locale-data/sv.dat filter=lfs diff=lfs merge=lfs -text diff --git a/rtme/lib/python3.10/site-packages/babel/locale-data/sv.dat b/rtme/lib/python3.10/site-packages/babel/locale-data/sv.dat new file mode 100644 index 0000000000000000000000000000000000000000..40e146e7d18edae02d109a83c819a469bee6631a --- /dev/null +++ b/rtme/lib/python3.10/site-packages/babel/locale-data/sv.dat @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:453bd0da58bd70a1c811b565edf23c10d56b99c09ad601119374bd3b66c72a60 +size 198967 diff --git a/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/LICENSE b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..bee14a47d2d83f2d4e3aca6d5a41e587ebf3c516 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/LICENSE @@ -0,0 +1,27 @@ +pycparser -- A C parser in Python + +Copyright (c) 2008-2022, Eli Bendersky +All rights reserved. + +Redistribution and use in source and binary forms, with or without modification, +are permitted provided that the following conditions are met: + +* Redistributions of source code must retain the above copyright notice, this + list of conditions and the following disclaimer. +* Redistributions in binary form must reproduce the above copyright notice, + this list of conditions and the following disclaimer in the documentation + and/or other materials provided with the distribution. +* Neither the name of the copyright holder nor the names of its contributors may + be used to endorse or promote products derived from this software without + specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND +ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED +WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE +GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) +HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT +LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT +OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. diff --git a/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/METADATA b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/METADATA new file mode 100644 index 0000000000000000000000000000000000000000..2c8038a3c2362879d22d50c54b1ff0cc479990d3 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/METADATA @@ -0,0 +1,28 @@ +Metadata-Version: 2.1 +Name: pycparser +Version: 2.22 +Summary: C parser in Python +Home-page: https://github.com/eliben/pycparser +Author: Eli Bendersky +Author-email: eliben@gmail.com +Maintainer: Eli Bendersky +License: BSD-3-Clause +Platform: Cross Platform +Classifier: Development Status :: 5 - Production/Stable +Classifier: License :: OSI Approved :: BSD License +Classifier: Programming Language :: Python :: 3 +Classifier: Programming Language :: Python :: 3.8 +Classifier: Programming Language :: Python :: 3.9 +Classifier: Programming Language :: Python :: 3.10 +Classifier: Programming Language :: Python :: 3.11 +Classifier: Programming Language :: Python :: 3.12 +Requires-Python: >=3.8 +License-File: LICENSE + + + pycparser is a complete parser of the C language, written in + pure Python using the PLY parsing library. + It parses C code into an AST and can serve as a front-end for + C compilers or analysis tools. + + diff --git a/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/RECORD b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/RECORD new file mode 100644 index 0000000000000000000000000000000000000000..358c64d8b9638d86ed73cb443b7cdbb2eda4fe6d --- /dev/null +++ b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/RECORD @@ -0,0 +1,25 @@ +pycparser-2.22.dist-info/INSTALLER,sha256=5hhM4Q4mYTT9z6QB6PGpUAW81PGNFrYrdXMj4oM_6ak,2 +pycparser-2.22.dist-info/LICENSE,sha256=DIRjmTaep23de1xE_m0WSXQV_PAV9cu1CMJL-YuBxbE,1543 +pycparser-2.22.dist-info/METADATA,sha256=3XOB8nggH4ijl17DCjUhk7g6qioMJLprUlEkwYgZvW8,943 +pycparser-2.22.dist-info/RECORD,, +pycparser-2.22.dist-info/REQUESTED,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0 +pycparser-2.22.dist-info/WHEEL,sha256=G16H4A3IeoQmnOrYV4ueZGKSjhipXx8zc8nu9FGlvMA,92 +pycparser-2.22.dist-info/top_level.txt,sha256=c-lPcS74L_8KoH7IE6PQF5ofyirRQNV4VhkbSFIPeWM,10 +pycparser/__init__.py,sha256=hrf-AyuVYNHQGTD0Nv2bywxoTN3N1ZCs03m-9-QDS14,2918 +pycparser/_ast_gen.py,sha256=0JRVnDW-Jw-3IjVlo8je9rbAcp6Ko7toHAnB5zi7h0Q,10555 +pycparser/_build_tables.py,sha256=4d_UkIxJ4YfHTVn6xBzBA52wDo7qxg1B6aZAJYJas9Q,1087 +pycparser/_c_ast.cfg,sha256=ld5ezE9yzIJFIVAUfw7ezJSlMi4nXKNCzfmqjOyQTNo,4255 +pycparser/ast_transforms.py,sha256=GTMYlUgWmXd5wJVyovXY1qzzAqjxzCpVVg0664dKGBs,5691 +pycparser/c_ast.py,sha256=HWeOrfYdCY0u5XaYhE1i60uVyE3yMWdcxzECUX-DqJw,31445 +pycparser/c_generator.py,sha256=yi6Mcqxv88J5ue8k5-mVGxh3iJ37iD4QyF-sWcGjC-8,17772 +pycparser/c_lexer.py,sha256=RSUjq0SRH8dkvwrQslBIZY2AXOrpQpe-oO1udJXotZk,17186 +pycparser/c_parser.py,sha256=WUnIHNydl32QBuRUqrqk-F2lyB6WRP4BUYFELqVETyw,74282 +pycparser/lextab.py,sha256=Nc3I0_D8Xlf-BOpfOKkEvFw-rPuFPPwAjkcLubwTCU4,8554 +pycparser/ply/__init__.py,sha256=q4s86QwRsYRa20L9ueSxfh-hPihpftBjDOvYa2_SS2Y,102 +pycparser/ply/cpp.py,sha256=UtC3ylTWp5_1MKA-PLCuwKQR8zSOnlGuGGIdzj8xS98,33282 +pycparser/ply/ctokens.py,sha256=MKksnN40TehPhgVfxCJhjj_BjL943apreABKYz-bl0Y,3177 +pycparser/ply/lex.py,sha256=rCMi0yjlZmjH5SNXj_Yds1VxSDkaG2thS7351YvfN-I,42926 +pycparser/ply/yacc.py,sha256=eatSDkRLgRr6X3-hoDk_SQQv065R0BdL2K7fQ54CgVM,137323 +pycparser/ply/ygen.py,sha256=2JYNeYtrPz1JzLSLO3d4GsS8zJU8jY_I_CR1VI9gWrA,2251 +pycparser/plyparser.py,sha256=8tLOoEytcapvWrr1JfCf7Dog-wulBtS1YrDs8S7JfMo,4875 +pycparser/yacctab.py,sha256=B6ck8QEPnRi04VSxKEL6xHaP8sEEsTbWtwsjfKHABgM,209738 diff --git a/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/REQUESTED b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/REQUESTED new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/WHEEL b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/WHEEL new file mode 100644 index 0000000000000000000000000000000000000000..becc9a66ea739ba941d48a749e248761cc6e658a --- /dev/null +++ b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/WHEEL @@ -0,0 +1,5 @@ +Wheel-Version: 1.0 +Generator: bdist_wheel (0.37.1) +Root-Is-Purelib: true +Tag: py3-none-any + diff --git a/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/top_level.txt b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/top_level.txt new file mode 100644 index 0000000000000000000000000000000000000000..dc1c9e101ad9ccd943b359338ef42c342ebc84a1 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/pycparser-2.22.dist-info/top_level.txt @@ -0,0 +1 @@ +pycparser diff --git a/rtme/lib/python3.10/site-packages/send2trash/__init__.py b/rtme/lib/python3.10/site-packages/send2trash/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..d09faa3adbd34a6bf792607dd247714c1ab6291b --- /dev/null +++ b/rtme/lib/python3.10/site-packages/send2trash/__init__.py @@ -0,0 +1,21 @@ +# Copyright 2013 Hardcoded Software (http://www.hardcoded.net) + +# This software is licensed under the "BSD" License as described in the "LICENSE" file, +# which should be included with this package. The terms are also available at +# http://www.hardcoded.net/licenses/bsd_license + +import sys + +from send2trash.exceptions import TrashPermissionError # noqa: F401 + +if sys.platform == "darwin": + from send2trash.mac import send2trash +elif sys.platform == "win32": + from send2trash.win import send2trash +else: + try: + # If we can use gio, let's use it + from send2trash.plat_gio import send2trash + except ImportError: + # Oh well, let's fallback to our own Freedesktop trash implementation + from send2trash.plat_other import send2trash # noqa: F401 diff --git a/rtme/lib/python3.10/site-packages/send2trash/__main__.py b/rtme/lib/python3.10/site-packages/send2trash/__main__.py new file mode 100644 index 0000000000000000000000000000000000000000..a733e82bed7936207140f98f54d12f0e64d2bc6e --- /dev/null +++ b/rtme/lib/python3.10/site-packages/send2trash/__main__.py @@ -0,0 +1,33 @@ +# encoding: utf-8 +# Copyright 2017 Virgil Dupras + +# This software is licensed under the "BSD" License as described in the "LICENSE" file, +# which should be included with this package. The terms are also available at +# http://www.hardcoded.net/licenses/bsd_license + +from __future__ import print_function + +import sys + +from argparse import ArgumentParser +from send2trash import send2trash + + +def main(args=None): + parser = ArgumentParser(description="Tool to send files to trash") + parser.add_argument("files", nargs="+") + parser.add_argument("-v", "--verbose", action="store_true", help="Print deleted files") + args = parser.parse_args(args) + + for filename in args.files: + try: + send2trash(filename) + if args.verbose: + print("Trashed «" + filename + "»") + except OSError as e: + print(str(e), file=sys.stderr) + sys.exit(1) + + +if __name__ == "__main__": + main() diff --git a/rtme/lib/python3.10/site-packages/send2trash/compat.py b/rtme/lib/python3.10/site-packages/send2trash/compat.py new file mode 100644 index 0000000000000000000000000000000000000000..a3043a4eb3d197588ba5b8121cd694c4fe41a2f5 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/send2trash/compat.py @@ -0,0 +1,25 @@ +# Copyright 2017 Virgil Dupras + +# This software is licensed under the "BSD" License as described in the "LICENSE" file, +# which should be included with this package. The terms are also available at +# http://www.hardcoded.net/licenses/bsd_license + +import sys +import os + +PY3 = sys.version_info[0] >= 3 +if PY3: + text_type = str + binary_type = bytes + if os.supports_bytes_environ: + # environb will be unset under Windows, but then again we're not supposed to use it. + environb = os.environb +else: + text_type = unicode # noqa: F821 + binary_type = str + environb = os.environ + +try: + from collections.abc import Iterable as iterable_type +except ImportError: + from collections import Iterable as iterable_type # noqa: F401 diff --git a/rtme/lib/python3.10/site-packages/send2trash/plat_gio.py b/rtme/lib/python3.10/site-packages/send2trash/plat_gio.py new file mode 100644 index 0000000000000000000000000000000000000000..258e4ef98050b9e7287d88366ce06e2ec262283f --- /dev/null +++ b/rtme/lib/python3.10/site-packages/send2trash/plat_gio.py @@ -0,0 +1,23 @@ +# Copyright 2017 Virgil Dupras + +# This software is licensed under the "BSD" License as described in the "LICENSE" file, +# which should be included with this package. The terms are also available at +# http://www.hardcoded.net/licenses/bsd_license + +from gi.repository import GObject, Gio +from send2trash.exceptions import TrashPermissionError +from send2trash.util import preprocess_paths + + +def send2trash(paths): + paths = preprocess_paths(paths) + for path in paths: + try: + f = Gio.File.new_for_path(path) + f.trash(cancellable=None) + except GObject.GError as e: + if e.code == Gio.IOErrorEnum.NOT_SUPPORTED: + # We get here if we can't create a trash directory on the same + # device. I don't know if other errors can result in NOT_SUPPORTED. + raise TrashPermissionError("") + raise OSError(e.message) diff --git a/rtme/lib/python3.10/site-packages/send2trash/plat_other.py b/rtme/lib/python3.10/site-packages/send2trash/plat_other.py new file mode 100644 index 0000000000000000000000000000000000000000..ace7b1307bfe814476149f330367798c5b02d8f9 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/send2trash/plat_other.py @@ -0,0 +1,218 @@ +# Copyright 2017 Virgil Dupras + +# This software is licensed under the "BSD" License as described in the "LICENSE" file, +# which should be included with this package. The terms are also available at +# http://www.hardcoded.net/licenses/bsd_license + +# This is a reimplementation of plat_other.py with reference to the +# freedesktop.org trash specification: +# [1] http://www.freedesktop.org/wiki/Specifications/trash-spec +# [2] http://www.ramendik.ru/docs/trashspec.html +# See also: +# [3] http://standards.freedesktop.org/basedir-spec/basedir-spec-latest.html +# +# For external volumes this implementation will raise an exception if it can't +# find or create the user's trash directory. + +from __future__ import unicode_literals + +import errno +import sys +import os +import shutil +import os.path as op +from datetime import datetime +import stat + +try: + from urllib.parse import quote +except ImportError: + # Python 2 + from urllib import quote + +from send2trash.compat import text_type, environb +from send2trash.util import preprocess_paths +from send2trash.exceptions import TrashPermissionError + +try: + fsencode = os.fsencode # Python 3 + fsdecode = os.fsdecode +except AttributeError: + + def fsencode(u): # Python 2 + return u.encode(sys.getfilesystemencoding()) + + def fsdecode(b): + return b.decode(sys.getfilesystemencoding()) + + # The Python 3 versions are a bit smarter, handling surrogate escapes, + # but these should work in most cases. + +FILES_DIR = b"files" +INFO_DIR = b"info" +INFO_SUFFIX = b".trashinfo" + +# Default of ~/.local/share [3] +XDG_DATA_HOME = op.expanduser(environb.get(b"XDG_DATA_HOME", b"~/.local/share")) +HOMETRASH_B = op.join(XDG_DATA_HOME, b"Trash") +HOMETRASH = fsdecode(HOMETRASH_B) + +uid = os.getuid() +TOPDIR_TRASH = b".Trash" +TOPDIR_FALLBACK = b".Trash-" + text_type(uid).encode("ascii") + + +def is_parent(parent, path): + path = op.realpath(path) # In case it's a symlink + if isinstance(path, text_type): + path = fsencode(path) + parent = op.realpath(parent) + if isinstance(parent, text_type): + parent = fsencode(parent) + return path.startswith(parent) + + +def format_date(date): + return date.strftime("%Y-%m-%dT%H:%M:%S") + + +def info_for(src, topdir): + # ...it MUST not include a ".." directory, and for files not "under" that + # directory, absolute pathnames must be used. [2] + if topdir is None or not is_parent(topdir, src): + src = op.abspath(src) + else: + src = op.relpath(src, topdir) + + info = "[Trash Info]\n" + info += "Path=" + quote(src) + "\n" + info += "DeletionDate=" + format_date(datetime.now()) + "\n" + return info + + +def check_create(dir): + # use 0700 for paths [3] + if not op.exists(dir): + os.makedirs(dir, 0o700) + + +def trash_move(src, dst, topdir=None, cross_dev=False): + filename = op.basename(src) + filespath = op.join(dst, FILES_DIR) + infopath = op.join(dst, INFO_DIR) + base_name, ext = op.splitext(filename) + + counter = 0 + destname = filename + while op.exists(op.join(filespath, destname)) or op.exists(op.join(infopath, destname + INFO_SUFFIX)): + counter += 1 + destname = base_name + b" " + text_type(counter).encode("ascii") + ext + + check_create(filespath) + check_create(infopath) + + with open(op.join(infopath, destname + INFO_SUFFIX), "w") as f: + f.write(info_for(src, topdir)) + destpath = op.join(filespath, destname) + if cross_dev: + shutil.move(fsdecode(src), fsdecode(destpath)) + else: + os.rename(src, destpath) + + +def find_mount_point(path): + # Even if something's wrong, "/" is a mount point, so the loop will exit. + # Use realpath in case it's a symlink + path = op.realpath(path) # Required to avoid infinite loop + while not op.ismount(path): # Note ismount() does not always detect mounts + path = op.split(path)[0] + return path + + +def find_ext_volume_global_trash(volume_root): + # from [2] Trash directories (1) check for a .Trash dir with the right + # permissions set. + trash_dir = op.join(volume_root, TOPDIR_TRASH) + if not op.exists(trash_dir): + return None + + mode = os.lstat(trash_dir).st_mode + # vol/.Trash must be a directory, cannot be a symlink, and must have the + # sticky bit set. + if not op.isdir(trash_dir) or op.islink(trash_dir) or not (mode & stat.S_ISVTX): + return None + + trash_dir = op.join(trash_dir, text_type(uid).encode("ascii")) + try: + check_create(trash_dir) + except OSError: + return None + return trash_dir + + +def find_ext_volume_fallback_trash(volume_root): + # from [2] Trash directories (1) create a .Trash-$uid dir. + trash_dir = op.join(volume_root, TOPDIR_FALLBACK) + # Try to make the directory, if we lack permission, raise TrashPermissionError + try: + check_create(trash_dir) + except OSError as e: + if e.errno == errno.EACCES: + raise TrashPermissionError(e.filename) + raise + return trash_dir + + +def find_ext_volume_trash(volume_root): + trash_dir = find_ext_volume_global_trash(volume_root) + if trash_dir is None: + trash_dir = find_ext_volume_fallback_trash(volume_root) + return trash_dir + + +# Pull this out so it's easy to stub (to avoid stubbing lstat itself) +def get_dev(path): + return os.lstat(path).st_dev + + +def send2trash(paths): + paths = preprocess_paths(paths) + for path in paths: + if isinstance(path, text_type): + path_b = fsencode(path) + elif isinstance(path, bytes): + path_b = path + else: + raise TypeError("str, bytes or PathLike expected, not %r" % type(path)) + + if not op.exists(path_b): + raise OSError(errno.ENOENT, "File not found: %s" % path) + # ...should check whether the user has the necessary permissions to delete + # it, before starting the trashing operation itself. [2] + if not os.access(path_b, os.W_OK): + raise OSError(errno.EACCES, "Permission denied: %s" % path) + + path_dev = get_dev(path_b) + # If XDG_DATA_HOME or HOMETRASH do not yet exist we need to stat the + # home directory, and these paths will be created further on if needed. + trash_dev = get_dev(op.expanduser(b"~")) + + # if the file to be trashed is on the same device as HOMETRASH we + # want to move it there. + if path_dev == trash_dev: + topdir = XDG_DATA_HOME + dest_trash = HOMETRASH_B + else: + topdir = find_mount_point(path_b) + trash_dev = get_dev(topdir) + if trash_dev != path_dev: + raise OSError("Couldn't find mount point for %s" % path) + dest_trash = find_ext_volume_trash(topdir) + try: + trash_move(path_b, dest_trash, topdir) + except OSError as error: + # Cross link errors default back to HOMETRASH + if error.errno == errno.EXDEV: + trash_move(path_b, HOMETRASH_B, XDG_DATA_HOME, cross_dev=True) + else: + raise diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/f_beta.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/f_beta.py new file mode 100644 index 0000000000000000000000000000000000000000..dc853c5c40b90e060da30beecbcbb658603eca4e --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/f_beta.py @@ -0,0 +1,1221 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.stat_scores import BinaryStatScores, MulticlassStatScores, MultilabelStatScores +from torchmetrics.functional.classification.f_beta import ( + _binary_fbeta_score_arg_validation, + _fbeta_reduce, + _multiclass_fbeta_score_arg_validation, + _multilabel_fbeta_score_arg_validation, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryFBetaScore.plot", + "MulticlassFBetaScore.plot", + "MultilabelFBetaScore.plot", + "BinaryF1Score.plot", + "MulticlassF1Score.plot", + "MultilabelF1Score.plot", + ] + + +class BinaryFBetaScore(BinaryStatScores): + r"""Compute `F-score`_ metric for binary tasks. + + .. math:: + F_{\beta} = (1 + \beta^2) * \frac{\text{precision} * \text{recall}} + {(\beta^2 * \text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered a score of `zero_division` + (0 or 1, default is 0) is returned. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor or float tensor of shape ``(N, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bfbs`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``multidim_average`` argument: + + - If ``multidim_average`` is set to ``global`` the output will be a scalar tensor + - If ``multidim_average`` is set to ``samplewise`` the output will be a tensor of shape ``(N,)`` consisting of + a scalar value per sample. + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + beta: Weighting between precision and recall in calculation. Setting to 1 corresponds to equal weight + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when + :math:`\text{TP} + \text{FP} = 0 \wedge \text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryFBetaScore + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinaryFBetaScore(beta=2.0) + >>> metric(preds, target) + tensor(0.6667) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryFBetaScore + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinaryFBetaScore(beta=2.0) + >>> metric(preds, target) + tensor(0.6667) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinaryFBetaScore + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinaryFBetaScore(beta=2.0, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.5882, 0.0000]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + beta: float, + threshold: float = 0.5, + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + threshold=threshold, + multidim_average=multidim_average, + ignore_index=ignore_index, + validate_args=False, + **kwargs, + ) + if validate_args: + _binary_fbeta_score_arg_validation(beta, threshold, multidim_average, ignore_index, zero_division) + self.validate_args = validate_args + self.zero_division = zero_division + self.beta = beta + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _fbeta_reduce( + tp, + fp, + tn, + fn, + self.beta, + average="binary", + multidim_average=self.multidim_average, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryFBetaScore + >>> metric = BinaryFBetaScore(beta=2.0) + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryFBetaScore + >>> metric = BinaryFBetaScore(beta=2.0) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassFBetaScore(MulticlassStatScores): + r"""Compute `F-score`_ metric for multiclass tasks. + + .. math:: + F_{\beta} = (1 + \beta^2) * \frac{\text{precision} * \text{recall}} + {(\beta^2 * \text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered for any class, the metric for that class + will be set to `zero_division` (0 or 1, default is 0) and the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcfbs`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``average`` and + ``multidim_average`` arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + beta: Weighting between precision and recall in calculation. Setting to 1 corresponds to equal weight + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + top_k: + + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when + :math:`\text{TP} + \text{FP} = 0 \wedge \text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassFBetaScore + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassFBetaScore(beta=2.0, num_classes=3) + >>> metric(preds, target) + tensor(0.7963) + >>> mcfbs = MulticlassFBetaScore(beta=2.0, num_classes=3, average=None) + >>> mcfbs(preds, target) + tensor([0.5556, 0.8333, 1.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassFBetaScore + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassFBetaScore(beta=2.0, num_classes=3) + >>> metric(preds, target) + tensor(0.7963) + >>> mcfbs = MulticlassFBetaScore(beta=2.0, num_classes=3, average=None) + >>> mcfbs(preds, target) + tensor([0.5556, 0.8333, 1.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassFBetaScore + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassFBetaScore(beta=2.0, num_classes=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.4697, 0.2706]) + >>> mcfbs = MulticlassFBetaScore(beta=2.0, num_classes=3, multidim_average='samplewise', average=None) + >>> mcfbs(preds, target) + tensor([[0.9091, 0.0000, 0.5000], + [0.0000, 0.3571, 0.4545]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + beta: float, + num_classes: int, + top_k: int = 1, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + num_classes=num_classes, + top_k=top_k, + average=average, + multidim_average=multidim_average, + ignore_index=ignore_index, + validate_args=False, + **kwargs, + ) + if validate_args: + _multiclass_fbeta_score_arg_validation( + beta, num_classes, top_k, average, multidim_average, ignore_index, zero_division + ) + self.validate_args = validate_args + self.zero_division = zero_division + self.beta = beta + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _fbeta_reduce( + tp, + fp, + tn, + fn, + self.beta, + average=self.average, + multidim_average=self.multidim_average, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassFBetaScore + >>> metric = MulticlassFBetaScore(num_classes=3, beta=2.0, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassFBetaScore + >>> metric = MulticlassFBetaScore(num_classes=3, beta=2.0, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelFBetaScore(MultilabelStatScores): + r"""Compute `F-score`_ metric for multilabel tasks. + + .. math:: + F_{\beta} = (1 + \beta^2) * \frac{\text{precision} * \text{recall}} + {(\beta^2 * \text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered for any label, the metric for that label + will be set to `zero_division` (0 or 1, default is 0) and the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, C, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlfbs`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``average`` and + ``multidim_average`` arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + beta: Weighting between precision and recall in calculation. Setting to 1 corresponds to equal weight + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when + :math:`\text{TP} + \text{FP} = 0 \wedge \text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelFBetaScore + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelFBetaScore(beta=2.0, num_labels=3) + >>> metric(preds, target) + tensor(0.6111) + >>> mlfbs = MultilabelFBetaScore(beta=2.0, num_labels=3, average=None) + >>> mlfbs(preds, target) + tensor([1.0000, 0.0000, 0.8333]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelFBetaScore + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelFBetaScore(beta=2.0, num_labels=3) + >>> metric(preds, target) + tensor(0.6111) + >>> mlfbs = MultilabelFBetaScore(beta=2.0, num_labels=3, average=None) + >>> mlfbs(preds, target) + tensor([1.0000, 0.0000, 0.8333]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelFBetaScore + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelFBetaScore(num_labels=3, beta=2.0, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.5556, 0.0000]) + >>> mlfbs = MultilabelFBetaScore(num_labels=3, beta=2.0, multidim_average='samplewise', average=None) + >>> mlfbs(preds, target) + tensor([[0.8333, 0.8333, 0.0000], + [0.0000, 0.0000, 0.0000]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + beta: float, + num_labels: int, + threshold: float = 0.5, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + num_labels=num_labels, + threshold=threshold, + average=average, + multidim_average=multidim_average, + ignore_index=ignore_index, + validate_args=False, + **kwargs, + ) + if validate_args: + _multilabel_fbeta_score_arg_validation( + beta, num_labels, threshold, average, multidim_average, ignore_index, zero_division + ) + self.validate_args = validate_args + self.zero_division = zero_division + self.beta = beta + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _fbeta_reduce( + tp, + fp, + tn, + fn, + self.beta, + average=self.average, + multidim_average=self.multidim_average, + multilabel=True, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelFBetaScore + >>> metric = MultilabelFBetaScore(num_labels=3, beta=2.0) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelFBetaScore + >>> metric = MultilabelFBetaScore(num_labels=3, beta=2.0) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class BinaryF1Score(BinaryFBetaScore): + r"""Compute F-1 score for binary tasks. + + .. math:: + F_{1} = 2\frac{\text{precision} * \text{recall}}{(\text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered a score of `zero_division` + (0 or 1, default is 0) is returned. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, ...)``. If preds is a floating point + tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid per + element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bf1s`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``multidim_average`` argument: + + - If ``multidim_average`` is set to ``global``, the metric returns a scalar value. + - If ``multidim_average`` is set to ``samplewise``, the metric returns ``(N,)`` vector consisting of a scalar + value per sample. + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when + :math:`\text{TP} + \text{FP} = 0 \wedge \text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryF1Score + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinaryF1Score() + >>> metric(preds, target) + tensor(0.6667) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryF1Score + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinaryF1Score() + >>> metric(preds, target) + tensor(0.6667) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinaryF1Score + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinaryF1Score(multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.5000, 0.0000]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + threshold: float = 0.5, + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + beta=1.0, + threshold=threshold, + multidim_average=multidim_average, + ignore_index=ignore_index, + validate_args=validate_args, + zero_division=zero_division, + **kwargs, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryF1Score + >>> metric = BinaryF1Score() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryF1Score + >>> metric = BinaryF1Score() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassF1Score(MulticlassFBetaScore): + r"""Compute F-1 score for multiclass tasks. + + .. math:: + F_{1} = 2\frac{\text{precision} * \text{recall}}{(\text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered for any class, the metric for that class + will be set to `zero_division` (0 or 1, default is 0) and the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcf1s`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``average`` and + ``multidim_average`` arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + preds: Tensor with predictions + target: Tensor with true labels + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + top_k: + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when + :math:`\text{TP} + \text{FP} = 0 \wedge \text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassF1Score + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassF1Score(num_classes=3) + >>> metric(preds, target) + tensor(0.7778) + >>> mcf1s = MulticlassF1Score(num_classes=3, average=None) + >>> mcf1s(preds, target) + tensor([0.6667, 0.6667, 1.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassF1Score + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassF1Score(num_classes=3) + >>> metric(preds, target) + tensor(0.7778) + >>> mcf1s = MulticlassF1Score(num_classes=3, average=None) + >>> mcf1s(preds, target) + tensor([0.6667, 0.6667, 1.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassF1Score + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassF1Score(num_classes=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.4333, 0.2667]) + >>> mcf1s = MulticlassF1Score(num_classes=3, multidim_average='samplewise', average=None) + >>> mcf1s(preds, target) + tensor([[0.8000, 0.0000, 0.5000], + [0.0000, 0.4000, 0.4000]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + top_k: int = 1, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + beta=1.0, + num_classes=num_classes, + top_k=top_k, + average=average, + multidim_average=multidim_average, + ignore_index=ignore_index, + validate_args=validate_args, + zero_division=zero_division, + **kwargs, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassF1Score + >>> metric = MulticlassF1Score(num_classes=3, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassF1Score + >>> metric = MulticlassF1Score(num_classes=3, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelF1Score(MultilabelFBetaScore): + r"""Compute F-1 score for multilabel tasks. + + .. math:: + F_{1} = 2\frac{\text{precision} * \text{recall}}{(\text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered for any label, the metric for that label + will be set to `zero_division` (0 or 1, default is 0) and the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, C, ...)``. + If preds is a floating point tensor with values outside [0,1] range we consider the input to be logits and + will auto apply sigmoid per element. Additionally, we convert to int tensor with thresholding using the value + in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlf1s`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``average`` and + ``multidim_average`` arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)``` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when + :math:`\text{TP} + \text{FP} = 0 \wedge \text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelF1Score + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelF1Score(num_labels=3) + >>> metric(preds, target) + tensor(0.5556) + >>> mlf1s = MultilabelF1Score(num_labels=3, average=None) + >>> mlf1s(preds, target) + tensor([1.0000, 0.0000, 0.6667]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelF1Score + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelF1Score(num_labels=3) + >>> metric(preds, target) + tensor(0.5556) + >>> mlf1s = MultilabelF1Score(num_labels=3, average=None) + >>> mlf1s(preds, target) + tensor([1.0000, 0.0000, 0.6667]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelF1Score + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelF1Score(num_labels=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.4444, 0.0000]) + >>> mlf1s = MultilabelF1Score(num_labels=3, multidim_average='samplewise', average=None) + >>> mlf1s(preds, target) + tensor([[0.6667, 0.6667, 0.0000], + [0.0000, 0.0000, 0.0000]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + threshold: float = 0.5, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + beta=1.0, + num_labels=num_labels, + threshold=threshold, + average=average, + multidim_average=multidim_average, + ignore_index=ignore_index, + validate_args=validate_args, + zero_division=zero_division, + **kwargs, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelF1Score + >>> metric = MultilabelF1Score(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelF1Score + >>> metric = MultilabelF1Score(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class FBetaScore(_ClassificationTaskWrapper): + r"""Compute `F-score`_ metric. + + .. math:: + F_{\beta} = (1 + \beta^2) * \frac{\text{precision} * \text{recall}} + {(\beta^2 * \text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered for any class/label, the metric for that + class/label will be set to `zero_division` (0 or 1, default is 0) and the overall metric may therefore be + affected in turn. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryFBetaScore`, + :class:`~torchmetrics.classification.MulticlassFBetaScore` and + :class:`~torchmetrics.classification.MultilabelFBetaScore` for the specific details of each argument influence + and examples. + + Legcy Example: + >>> from torch import tensor + >>> target = tensor([0, 1, 2, 0, 1, 2]) + >>> preds = tensor([0, 2, 1, 0, 0, 1]) + >>> f_beta = FBetaScore(task="multiclass", num_classes=3, beta=0.5) + >>> f_beta(preds, target) + tensor(0.3333) + + """ + + def __new__( # type: ignore[misc] + cls: type["FBetaScore"], + task: Literal["binary", "multiclass", "multilabel"], + beta: float = 1.0, + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + "zero_division": zero_division, + }) + if task == ClassificationTask.BINARY: + return BinaryFBetaScore(beta, threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassFBetaScore(beta, num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelFBetaScore(beta, num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") + + +class F1Score(_ClassificationTaskWrapper): + r"""Compute F-1 score. + + .. math:: + F_{1} = 2\frac{\text{precision} * \text{recall}}{(\text{precision}) + \text{recall}} + + The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0 \wedge \text{TP} + \text{FN} \neq 0` + where :math:`\text{TP}`, :math:`\text{FP}` and :math:`\text{FN}` represent the number of true positives, false + positives and false negatives respectively. If this case is encountered for any class/label, the metric for that + class/label will be set to `zero_division` (0 or 1, default is 0) and the overall metric may therefore be + affected in turn. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryF1Score`, :class:`~torchmetrics.classification.MulticlassF1Score` and + :class:`~torchmetrics.classification.MultilabelF1Score` for the specific details of each argument influence and + examples. + + Legacy Example: + >>> from torch import tensor + >>> target = tensor([0, 1, 2, 0, 1, 2]) + >>> preds = tensor([0, 2, 1, 0, 0, 1]) + >>> f1 = F1Score(task="multiclass", num_classes=3) + >>> f1(preds, target) + tensor(0.3333) + + """ + + def __new__( # type: ignore[misc] + cls: type["F1Score"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + "zero_division": zero_division, + }) + if task == ClassificationTask.BINARY: + return BinaryF1Score(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassF1Score(num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelF1Score(num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/group_fairness.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/group_fairness.py new file mode 100644 index 0000000000000000000000000000000000000000..b43d18c6b133698633705a2ce9be81118adc53c2 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/group_fairness.py @@ -0,0 +1,326 @@ +# Copyright The PyTorch Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.classification.group_fairness import ( + _binary_groups_stat_scores, + _compute_binary_demographic_parity, + _compute_binary_equal_opportunity, +) +from torchmetrics.functional.classification.stat_scores import _binary_stat_scores_arg_validation +from torchmetrics.metric import Metric +from torchmetrics.utilities import rank_zero_warn +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["BinaryFairness.plot"] + + +class _AbstractGroupStatScores(Metric): + """Create and update states for computing group stats tp, fp, tn and fn.""" + + tp: Tensor + fp: Tensor + tn: Tensor + fn: Tensor + + def _create_states(self, num_groups: int) -> None: + default = lambda: torch.zeros(num_groups, dtype=torch.long) + self.add_state("tp", default(), dist_reduce_fx="sum") + self.add_state("fp", default(), dist_reduce_fx="sum") + self.add_state("tn", default(), dist_reduce_fx="sum") + self.add_state("fn", default(), dist_reduce_fx="sum") + + def _update_states(self, group_stats: list[tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]]) -> None: + for group, stats in enumerate(group_stats): + tp, fp, tn, fn = stats + self.tp[group] += tp + self.fp[group] += fp + self.tn[group] += tn + self.fn[group] += fn + + +class BinaryGroupStatRates(_AbstractGroupStatScores): + r"""Computes the true/false positives and true/false negatives rates for binary classification by group. + + Related to `Type I and Type II errors`_. + + Accepts the following input tensors: + + - ``preds`` (int or float tensor): ``(N, ...)``. If preds is a floating point tensor with values outside + [0,1] range we consider the input to be logits and will auto apply sigmoid per element. Additionally, + we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (int tensor): ``(N, ...)``. + - ``groups`` (int tensor): ``(N, ...)``. The group identifiers should be ``0, 1, ..., (num_groups - 1)``. + + The additional dimensions are flatted along the batch dimension. + + Args: + num_groups: The number of groups. + threshold: Threshold for transforming probability to binary {0,1} predictions. + ignore_index: Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + The metric returns a dict with a group identifier as key and a tensor with the tp, fp, tn and fn rates as value. + + Example (preds is int tensor): + >>> from torchmetrics.classification import BinaryGroupStatRates + >>> target = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> preds = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> groups = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> metric = BinaryGroupStatRates(num_groups=2) + >>> metric(preds, target, groups) + {'group_0': tensor([0., 0., 1., 0.]), 'group_1': tensor([1., 0., 0., 0.])} + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryGroupStatRates + >>> target = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> preds = torch.tensor([0.11, 0.84, 0.22, 0.73, 0.33, 0.92]) + >>> groups = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> metric = BinaryGroupStatRates(num_groups=2) + >>> metric(preds, target, groups) + {'group_0': tensor([0., 0., 1., 0.]), 'group_1': tensor([1., 0., 0., 0.])} + + """ + + is_differentiable: bool = False + higher_is_better: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + num_groups: int, + threshold: float = 0.5, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__() + + if validate_args: + _binary_stat_scores_arg_validation(threshold, "global", ignore_index) + + if not isinstance(num_groups, int) and num_groups < 2: + raise ValueError(f"Expected argument `num_groups` to be an int larger than 1, but got {num_groups}") + self.num_groups = num_groups + self.threshold = threshold + self.ignore_index = ignore_index + self.validate_args = validate_args + + self._create_states(self.num_groups) + + def update(self, preds: Tensor, target: Tensor, groups: Tensor) -> None: + """Update state with predictions, target and group identifiers. + + Args: + preds: Tensor with predictions. + target: Tensor with true labels. + groups: Tensor with group identifiers. The group identifiers should be ``0, 1, ..., (num_groups - 1)``. + + """ + group_stats = _binary_groups_stat_scores( + preds, target, groups, self.num_groups, self.threshold, self.ignore_index, self.validate_args + ) + + self._update_states(group_stats) + + def compute( + self, + ) -> dict[str, Tensor]: + """Compute tp, fp, tn and fn rates based on inputs passed in to ``update`` previously.""" + results = torch.stack((self.tp, self.fp, self.tn, self.fn), dim=1) + + return {f"group_{i}": group / group.sum() for i, group in enumerate(results)} + + +class BinaryFairness(_AbstractGroupStatScores): + r"""Computes `Demographic parity`_ and `Equal opportunity`_ ratio for binary classification problems. + + Accepts the following input tensors: + + - ``preds`` (int or float tensor): ``(N, ...)``. If preds is a floating point tensor with values outside + [0,1] range we consider the input to be logits and will auto apply sigmoid per element. Additionally, + we convert to int tensor with thresholding using the value in ``threshold``. + - ``groups`` (int tensor): ``(N, ...)``. The group identifiers should be ``0, 1, ..., (num_groups - 1)``. + - ``target`` (int tensor): ``(N, ...)``. + + The additional dimensions are flatted along the batch dimension. + + This class computes the ratio between positivity rates and true positives rates for different groups. + If more than two groups are present, the disparity between the lowest and highest group is reported. + A disparity between positivity rates indicates a potential violation of demographic parity, and between + true positive rates indicates a potential violation of equal opportunity. + + The lowest rate is divided by the highest, so a lower value means more discrimination against the numerator. + In the results this is also indicated as the key of dict is {metric}_{identifier_low_group}_{identifier_high_group}. + + Args: + num_groups: The number of groups. + task: The task to compute. Can be either ``demographic_parity`` or ``equal_opportunity`` or ``all``. + threshold: Threshold for transforming probability to binary {0,1} predictions. + ignore_index: Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + The metric returns a dict where the key identifies the metric and groups with the lowest and highest true + positives rates as follows: {metric}__{identifier_low_group}_{identifier_high_group}. + The value is a tensor with the disparity rate. + + Example (preds is int tensor): + >>> from torchmetrics.classification import BinaryFairness + >>> target = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> preds = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> groups = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> metric = BinaryFairness(2) + >>> metric(preds, target, groups) + {'DP_0_1': tensor(0.), 'EO_0_1': tensor(0.)} + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryFairness + >>> target = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> preds = torch.tensor([0.11, 0.84, 0.22, 0.73, 0.33, 0.92]) + >>> groups = torch.tensor([0, 1, 0, 1, 0, 1]) + >>> metric = BinaryFairness(2) + >>> metric(preds, target, groups) + {'DP_0_1': tensor(0.), 'EO_0_1': tensor(0.)} + + """ + + is_differentiable: bool = False + higher_is_better: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + num_groups: int, + task: Literal["demographic_parity", "equal_opportunity", "all"] = "all", + threshold: float = 0.5, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__() + + if task not in ["demographic_parity", "equal_opportunity", "all"]: + raise ValueError( + f"Expected argument `task` to either be ``demographic_parity``," + f"``equal_opportunity`` or ``all`` but got {task}." + ) + + if validate_args: + _binary_stat_scores_arg_validation(threshold, "global", ignore_index) + + if not isinstance(num_groups, int) and num_groups < 2: + raise ValueError(f"Expected argument `num_groups` to be an int larger than 1, but got {num_groups}") + self.num_groups = num_groups + self.task = task + self.threshold = threshold + self.ignore_index = ignore_index + self.validate_args = validate_args + + self._create_states(self.num_groups) + + def update(self, preds: Tensor, target: Tensor, groups: Tensor) -> None: + """Update state with predictions, groups, and target. + + Args: + preds: Tensor with predictions. + target: Tensor with true labels. + groups: Tensor with group identifiers. The group identifiers should be ``0, 1, ..., (num_groups - 1)``. + + """ + if self.task == "demographic_parity": + if target is not None: + rank_zero_warn("The task demographic_parity does not require a target.", UserWarning) + target = torch.zeros(preds.shape) + + group_stats = _binary_groups_stat_scores( + preds, target, groups, self.num_groups, self.threshold, self.ignore_index, self.validate_args + ) + + self._update_states(group_stats) + + def compute( + self, + ) -> dict[str, torch.Tensor]: + """Compute fairness criteria based on inputs passed in to ``update`` previously.""" + if self.task == "demographic_parity": + return _compute_binary_demographic_parity(self.tp, self.fp, self.tn, self.fn) + + if self.task == "equal_opportunity": + return _compute_binary_equal_opportunity(self.tp, self.fp, self.tn, self.fn) + + if self.task == "all": + return { + **_compute_binary_demographic_parity(self.tp, self.fp, self.tn, self.fn), + **_compute_binary_equal_opportunity(self.tp, self.fp, self.tn, self.fn), + } + return None + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import ones, rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryFairness + >>> metric = BinaryFairness(2) + >>> metric.update(rand(50), randint(2, (50,)), ones(50).long()) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import ones, rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryFairness + >>> metric = BinaryFairness(2) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(50), randint(2, (50,) ), ones(50).long())) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/hamming.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/hamming.py new file mode 100644 index 0000000000000000000000000000000000000000..20a1d9d2c718d818078ac60f1afc333d41fb81f8 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/hamming.py @@ -0,0 +1,529 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.stat_scores import BinaryStatScores, MulticlassStatScores, MultilabelStatScores +from torchmetrics.functional.classification.hamming import _hamming_distance_reduce +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryHammingDistance.plot", + "MulticlassHammingDistance.plot", + "MultilabelHammingDistance.plot", + ] + + +class BinaryHammingDistance(BinaryStatScores): + r"""Compute the average `Hamming distance`_ (also known as Hamming loss) for binary tasks. + + .. math:: + \text{Hamming distance} = \frac{1}{N \cdot L} \sum_i^N \sum_l^L 1(y_{il} \neq \hat{y}_{il}) + + Where :math:`y` is a tensor of target values, :math:`\hat{y}` is a tensor of predictions, + and :math:`\bullet_{il}` refers to the :math:`l`-th label of the :math:`i`-th sample of that + tensor. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, ...)``. If preds is a floating point + tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid per + element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bhd`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``multidim_average`` arguments: + + - If ``multidim_average`` is set to ``global``, the metric returns a scalar value. + - If ``multidim_average`` is set to ``samplewise``, the metric returns ``(N,)`` vector consisting of a + scalar value per sample. + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryHammingDistance + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinaryHammingDistance() + >>> metric(preds, target) + tensor(0.3333) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryHammingDistance + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinaryHammingDistance() + >>> metric(preds, target) + tensor(0.3333) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinaryHammingDistance + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinaryHammingDistance(multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.6667, 0.8333]) + + """ + + is_differentiable: bool = False + higher_is_better: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _hamming_distance_reduce(tp, fp, tn, fn, average="binary", multidim_average=self.multidim_average) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryHammingDistance + >>> metric = BinaryHammingDistance() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryHammingDistance + >>> metric = BinaryHammingDistance() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassHammingDistance(MulticlassStatScores): + r"""Compute the average `Hamming distance`_ (also known as Hamming loss) for multiclass tasks. + + .. math:: + \text{Hamming distance} = \frac{1}{N \cdot L} \sum_i^N \sum_l^L 1(y_{il} \neq \hat{y}_{il}) + + Where :math:`y` is a tensor of target values, :math:`\hat{y}` is a tensor of predictions, + and :math:`\bullet_{il}` refers to the :math:`l`-th label of the :math:`i`-th sample of that + tensor. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mchd`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``average`` and + ``multidim_average`` arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + top_k: + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassHammingDistance + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassHammingDistance(num_classes=3) + >>> metric(preds, target) + tensor(0.1667) + >>> mchd = MulticlassHammingDistance(num_classes=3, average=None) + >>> mchd(preds, target) + tensor([0.5000, 0.0000, 0.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassHammingDistance + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassHammingDistance(num_classes=3) + >>> metric(preds, target) + tensor(0.1667) + >>> mchd = MulticlassHammingDistance(num_classes=3, average=None) + >>> mchd(preds, target) + tensor([0.5000, 0.0000, 0.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassHammingDistance + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassHammingDistance(num_classes=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.5000, 0.7222]) + >>> mchd = MulticlassHammingDistance(num_classes=3, multidim_average='samplewise', average=None) + >>> mchd(preds, target) + tensor([[0.0000, 1.0000, 0.5000], + [1.0000, 0.6667, 0.5000]]) + + """ + + is_differentiable: bool = False + higher_is_better: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _hamming_distance_reduce(tp, fp, tn, fn, average=self.average, multidim_average=self.multidim_average) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value per class + >>> from torch import randint + >>> from torchmetrics.classification import MulticlassHammingDistance + >>> metric = MulticlassHammingDistance(num_classes=3, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting a multiple values per class + >>> from torch import randint + >>> from torchmetrics.classification import MulticlassHammingDistance + >>> metric = MulticlassHammingDistance(num_classes=3, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelHammingDistance(MultilabelStatScores): + r"""Compute the average `Hamming distance`_ (also known as Hamming loss) for multilabel tasks. + + .. math:: + \text{Hamming distance} = \frac{1}{N \cdot L} \sum_i^N \sum_l^L 1(y_{il} \neq \hat{y}_{il}) + + Where :math:`y` is a tensor of target values, :math:`\hat{y}` is a tensor of predictions, + and :math:`\bullet_{il}` refers to the :math:`l`-th label of the :math:`i`-th sample of that + tensor. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor or float tensor of shape ``(N, C, ...)``. If preds is a + floating point tensor with values outside [0,1] range we consider the input to be logits and will auto + apply sigmoid per element. Additionally, we convert to int tensor with thresholding using the value in + ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlhd`` (:class:`~torch.Tensor`): A tensor whose returned shape depends on the ``average`` and + ``multidim_average`` arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelHammingDistance + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelHammingDistance(num_labels=3) + >>> metric(preds, target) + tensor(0.3333) + >>> mlhd = MultilabelHammingDistance(num_labels=3, average=None) + >>> mlhd(preds, target) + tensor([0.0000, 0.5000, 0.5000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelHammingDistance + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelHammingDistance(num_labels=3) + >>> metric(preds, target) + tensor(0.3333) + >>> mlhd = MultilabelHammingDistance(num_labels=3, average=None) + >>> mlhd(preds, target) + tensor([0.0000, 0.5000, 0.5000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelHammingDistance + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelHammingDistance(num_labels=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.6667, 0.8333]) + >>> mlhd = MultilabelHammingDistance(num_labels=3, multidim_average='samplewise', average=None) + >>> mlhd(preds, target) + tensor([[0.5000, 0.5000, 1.0000], + [1.0000, 1.0000, 0.5000]]) + + """ + + is_differentiable: bool = False + higher_is_better: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _hamming_distance_reduce( + tp, fp, tn, fn, average=self.average, multidim_average=self.multidim_average, multilabel=True + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelHammingDistance + >>> metric = MultilabelHammingDistance(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelHammingDistance + >>> metric = MultilabelHammingDistance(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class HammingDistance(_ClassificationTaskWrapper): + r"""Compute the average `Hamming distance`_ (also known as Hamming loss). + + .. math:: + \text{Hamming distance} = \frac{1}{N \cdot L} \sum_i^N \sum_l^L 1(y_{il} \neq \hat{y}_{il}) + + Where :math:`y` is a tensor of target values, :math:`\hat{y}` is a tensor of predictions, + and :math:`\bullet_{il}` refers to the :math:`l`-th label of the :math:`i`-th sample of that + tensor. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryHammingDistance`, + :class:`~torchmetrics.classification.MulticlassHammingDistance` and + :class:`~torchmetrics.classification.MultilabelHammingDistance` for the specific details of each argument influence + and examples. + + Legacy Example: + >>> from torch import tensor + >>> target = tensor([[0, 1], [1, 1]]) + >>> preds = tensor([[0, 1], [0, 1]]) + >>> hamming_distance = HammingDistance(task="multilabel", num_labels=2) + >>> hamming_distance(preds, target) + tensor(0.2500) + + """ + + def __new__( # type: ignore[misc] + cls: type["HammingDistance"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + }) + if task == ClassificationTask.BINARY: + return BinaryHammingDistance(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassHammingDistance(num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelHammingDistance(num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/hinge.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/hinge.py new file mode 100644 index 0000000000000000000000000000000000000000..878ea2710494f9876b5bb4cbb67aaf6e85b6959d --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/hinge.py @@ -0,0 +1,380 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.functional.classification.hinge import ( + _binary_confusion_matrix_format, + _binary_hinge_loss_arg_validation, + _binary_hinge_loss_tensor_validation, + _binary_hinge_loss_update, + _hinge_loss_compute, + _multiclass_confusion_matrix_format, + _multiclass_hinge_loss_arg_validation, + _multiclass_hinge_loss_tensor_validation, + _multiclass_hinge_loss_update, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTaskNoMultilabel +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["BinaryHingeLoss.plot", "MulticlassHingeLoss.plot"] + + +class BinaryHingeLoss(Metric): + r"""Compute the mean `Hinge loss`_ typically used for Support Vector Machines (SVMs) for binary tasks. + + .. math:: + \text{Hinge loss} = \max(0, 1 - y \times \hat{y}) + + Where :math:`y \in {-1, 1}` is the target, and :math:`\hat{y} \in \mathbb{R}` is the prediction. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, ...)``. Preds should be a tensor containing + probabilities or logits for each observation. If preds has values outside [0,1] range we consider the input + to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). The value + 1 always encodes the positive class. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bhl`` (:class:`~torch.Tensor`): A tensor containing the hinge loss. + + Args: + squared: + If True, this will compute the squared hinge loss. Otherwise, computes the regular hinge loss. + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torchmetrics.classification import BinaryHingeLoss + >>> preds = torch.tensor([0.25, 0.25, 0.55, 0.75, 0.75]) + >>> target = torch.tensor([0, 0, 1, 1, 1]) + >>> bhl = BinaryHingeLoss() + >>> bhl(preds, target) + tensor(0.6900) + >>> bhl = BinaryHingeLoss(squared=True) + >>> bhl(preds, target) + tensor(0.6905) + + """ + + is_differentiable: bool = True + higher_is_better: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + measures: Tensor + total: Tensor + + def __init__( + self, + squared: bool = False, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _binary_hinge_loss_arg_validation(squared, ignore_index) + self.validate_args = validate_args + self.squared = squared + self.ignore_index = ignore_index + + self.add_state("measures", default=torch.tensor(0.0), dist_reduce_fx="sum") + self.add_state("total", default=torch.tensor(0), dist_reduce_fx="sum") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric state.""" + if self.validate_args: + _binary_hinge_loss_tensor_validation(preds, target, self.ignore_index) + preds, target = _binary_confusion_matrix_format( + preds, target, threshold=0.0, ignore_index=self.ignore_index, convert_to_labels=False + ) + measures, total = _binary_hinge_loss_update(preds, target, self.squared) + self.measures += measures + self.total += total + + def compute(self) -> Tensor: + """Compute metric.""" + return _hinge_loss_compute(self.measures, self.total) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryHingeLoss + >>> metric = BinaryHingeLoss() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryHingeLoss + >>> metric = BinaryHingeLoss() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassHingeLoss(Metric): + r"""Compute the mean `Hinge loss`_ typically used for Support Vector Machines (SVMs) for multiclass tasks. + + The metric can be computed in two ways. Either, the definition by Crammer and Singer is used: + + .. math:: + \text{Hinge loss} = \max\left(0, 1 - \hat{y}_y + \max_{i \ne y} (\hat{y}_i)\right) + + Where :math:`y \in {0, ..., \mathrm{C}}` is the target class (where :math:`\mathrm{C}` is the number of classes), + and :math:`\hat{y} \in \mathbb{R}^\mathrm{C}` is the predicted output per class. Alternatively, the metric can + also be computed in one-vs-all approach, where each class is valued against all other classes in a binary fashion. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply softmax per sample. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain values in the [0, n_classes-1] range (except if `ignore_index` + is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mchl`` (:class:`~torch.Tensor`): A tensor containing the multi-class hinge loss. + + Args: + num_classes: Integer specifying the number of classes + squared: + If True, this will compute the squared hinge loss. Otherwise, computes the regular hinge loss. + multiclass_mode: + Determines how to compute the metric + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torchmetrics.classification import MulticlassHingeLoss + >>> preds = torch.tensor([[0.25, 0.20, 0.55], + ... [0.55, 0.05, 0.40], + ... [0.10, 0.30, 0.60], + ... [0.90, 0.05, 0.05]]) + >>> target = torch.tensor([0, 1, 2, 0]) + >>> mchl = MulticlassHingeLoss(num_classes=3) + >>> mchl(preds, target) + tensor(0.9125) + >>> mchl = MulticlassHingeLoss(num_classes=3, squared=True) + >>> mchl(preds, target) + tensor(1.1131) + >>> mchl = MulticlassHingeLoss(num_classes=3, multiclass_mode='one-vs-all') + >>> mchl(preds, target) + tensor([0.8750, 1.1250, 1.1000]) + + """ + + is_differentiable: bool = True + higher_is_better: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + measures: Tensor + total: Tensor + + def __init__( + self, + num_classes: int, + squared: bool = False, + multiclass_mode: Literal["crammer-singer", "one-vs-all"] = "crammer-singer", + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _multiclass_hinge_loss_arg_validation(num_classes, squared, multiclass_mode, ignore_index) + self.validate_args = validate_args + self.num_classes = num_classes + self.squared = squared + self.multiclass_mode = multiclass_mode + self.ignore_index = ignore_index + + self.add_state( + "measures", + default=torch.tensor(0.0) + if self.multiclass_mode == "crammer-singer" + else torch.zeros( + num_classes, + ), + dist_reduce_fx="sum", + ) + self.add_state("total", default=torch.tensor(0), dist_reduce_fx="sum") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric state.""" + if self.validate_args: + _multiclass_hinge_loss_tensor_validation(preds, target, self.num_classes, self.ignore_index) + preds, target = _multiclass_confusion_matrix_format(preds, target, self.ignore_index, convert_to_labels=False) + measures, total = _multiclass_hinge_loss_update(preds, target, self.squared, self.multiclass_mode) + self.measures += measures + self.total += total + + def compute(self) -> Tensor: + """Compute metric.""" + return _hinge_loss_compute(self.measures, self.total) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value per class + >>> from torch import randint, randn + >>> from torchmetrics.classification import MulticlassHingeLoss + >>> metric = MulticlassHingeLoss(num_classes=3) + >>> metric.update(randn(20, 3), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting a multiple values per class + >>> from torch import randint, randn + >>> from torchmetrics.classification import MulticlassHingeLoss + >>> metric = MulticlassHingeLoss(num_classes=3) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randn(20, 3), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class HingeLoss(_ClassificationTaskWrapper): + r"""Compute the mean `Hinge loss`_ typically used for Support Vector Machines (SVMs). + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'`` or ``'multiclass'``. See the documentation of + :class:`~torchmetrics.classification.BinaryHingeLoss` and :class:`~torchmetrics.classification.MulticlassHingeLoss` + for the specific details of each argument influence and examples. + + Legacy Example: + >>> from torch import tensor + >>> target = tensor([0, 1, 1]) + >>> preds = tensor([0.5, 0.7, 0.1]) + >>> hinge = HingeLoss(task="binary") + >>> hinge(preds, target) + tensor(0.9000) + + >>> target = tensor([0, 1, 2]) + >>> preds = tensor([[-1.0, 0.9, 0.2], [0.5, -1.1, 0.8], [2.2, -0.5, 0.3]]) + >>> hinge = HingeLoss(task="multiclass", num_classes=3) + >>> hinge(preds, target) + tensor(1.5551) + + >>> target = tensor([0, 1, 2]) + >>> preds = tensor([[-1.0, 0.9, 0.2], [0.5, -1.1, 0.8], [2.2, -0.5, 0.3]]) + >>> hinge = HingeLoss(task="multiclass", num_classes=3, multiclass_mode="one-vs-all") + >>> hinge(preds, target) + tensor([1.3743, 1.1945, 1.2359]) + + """ + + def __new__( # type: ignore[misc] + cls: type["HingeLoss"], + task: Literal["binary", "multiclass"], + num_classes: Optional[int] = None, + squared: bool = False, + multiclass_mode: Optional[Literal["crammer-singer", "one-vs-all"]] = "crammer-singer", + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTaskNoMultilabel.from_str(task) + kwargs.update({"ignore_index": ignore_index, "validate_args": validate_args}) + if task == ClassificationTaskNoMultilabel.BINARY: + return BinaryHingeLoss(squared, **kwargs) + if task == ClassificationTaskNoMultilabel.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if multiclass_mode not in ("crammer-singer", "one-vs-all"): + raise ValueError( + f"`multiclass_mode` is expected to be one of 'crammer-singer' or 'one-vs-all' but " + f"`{multiclass_mode}` was passed." + ) + return MulticlassHingeLoss(num_classes, squared, multiclass_mode, **kwargs) + raise ValueError(f"Unsupported task `{task}`") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/jaccard.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/jaccard.py new file mode 100644 index 0000000000000000000000000000000000000000..50ad2af3bc31f4ee873c197895c131d1545b0618 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/jaccard.py @@ -0,0 +1,485 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.confusion_matrix import ( + BinaryConfusionMatrix, + MulticlassConfusionMatrix, + MultilabelConfusionMatrix, +) +from torchmetrics.functional.classification.jaccard import ( + _jaccard_index_reduce, + _multiclass_jaccard_index_arg_validation, + _multilabel_jaccard_index_arg_validation, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["BinaryJaccardIndex.plot", "MulticlassJaccardIndex.plot", "MultilabelJaccardIndex.plot"] + + +class BinaryJaccardIndex(BinaryConfusionMatrix): + r"""Calculate the Jaccard index for binary tasks. + + The `Jaccard index`_ (also known as the intersection over union or jaccard similarity coefficient) is an statistic + that can be used to determine the similarity and diversity of a sample set. It is defined as the size of the + intersection divided by the union of the sample sets: + + .. math:: J(A,B) = \frac{|A\cap B|}{|A\cup B|} + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A int or float tensor of shape ``(N, ...)``. If preds is a floating point + tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid per element. + Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bji`` (:class:`~torch.Tensor`): A tensor containing the Binary Jaccard Index. + + Args: + threshold: Threshold for transforming probability to binary (0,1) predictions + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: + Value to replace when there is a division by zero. Should be `0` or `1`. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryJaccardIndex + >>> target = tensor([1, 1, 0, 0]) + >>> preds = tensor([0, 1, 0, 0]) + >>> metric = BinaryJaccardIndex() + >>> metric(preds, target) + tensor(0.5000) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryJaccardIndex + >>> target = tensor([1, 1, 0, 0]) + >>> preds = tensor([0.35, 0.85, 0.48, 0.01]) + >>> metric = BinaryJaccardIndex() + >>> metric(preds, target) + tensor(0.5000) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + threshold: float = 0.5, + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + threshold=threshold, ignore_index=ignore_index, normalize=None, validate_args=validate_args, **kwargs + ) + self.zero_division = zero_division + + def compute(self) -> Tensor: + """Compute metric.""" + return _jaccard_index_reduce(self.confmat, average="binary", zero_division=self.zero_division) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryJaccardIndex + >>> metric = BinaryJaccardIndex() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryJaccardIndex + >>> metric = BinaryJaccardIndex() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassJaccardIndex(MulticlassConfusionMatrix): + r"""Calculate the Jaccard index for multiclass tasks. + + The `Jaccard index`_ (also known as the intersection over union or jaccard similarity coefficient) is an statistic + that can be used to determine the similarity and diversity of a sample set. It is defined as the size of the + intersection divided by the union of the sample sets: + + .. math:: J(A,B) = \frac{|A\cap B|}{|A\cup B|} + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcji`` (:class:`~torch.Tensor`): A tensor containing the Multi-class Jaccard Index. + + Args: + num_classes: Integer specifying the number of classes + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: + Value to replace when there is a division by zero. Should be `0` or `1`. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (pred is integer tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassJaccardIndex + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassJaccardIndex(num_classes=3) + >>> metric(preds, target) + tensor(0.6667) + + Example (pred is float tensor): + >>> from torchmetrics.classification import MulticlassJaccardIndex + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassJaccardIndex(num_classes=3) + >>> metric(preds, target) + tensor(0.6667) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + num_classes=num_classes, ignore_index=ignore_index, normalize=None, validate_args=False, **kwargs + ) + if validate_args: + _multiclass_jaccard_index_arg_validation(num_classes, ignore_index, average) + self.validate_args = validate_args + self.average = average + self.zero_division = zero_division + + def compute(self) -> Tensor: + """Compute metric.""" + return _jaccard_index_reduce( + self.confmat, average=self.average, ignore_index=self.ignore_index, zero_division=self.zero_division + ) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value per class + >>> from torch import randint + >>> from torchmetrics.classification import MulticlassJaccardIndex + >>> metric = MulticlassJaccardIndex(num_classes=3, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting a multiple values per class + >>> from torch import randint + >>> from torchmetrics.classification import MulticlassJaccardIndex + >>> metric = MulticlassJaccardIndex(num_classes=3, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelJaccardIndex(MultilabelConfusionMatrix): + r"""Calculate the Jaccard index for multilabel tasks. + + The `Jaccard index`_ (also known as the intersection over union or jaccard similarity coefficient) is an statistic + that can be used to determine the similarity and diversity of a sample set. It is defined as the size of the + intersection divided by the union of the sample sets: + + .. math:: J(A,B) = \frac{|A\cap B|}{|A\cup B|} + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A int tensor or float tensor of shape ``(N, C, ...)``. If preds is a + floating point tensor with values outside [0,1] range we consider the input to be logits and will auto apply + sigmoid per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)`` + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlji`` (:class:`~torch.Tensor`): A tensor containing the Multi-label Jaccard Index loss. + + Args: + num_classes: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: + Value to replace when there is a division by zero. Should be `0` or `1`. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelJaccardIndex + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelJaccardIndex(num_labels=3) + >>> metric(preds, target) + tensor(0.5000) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelJaccardIndex + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelJaccardIndex(num_labels=3) + >>> metric(preds, target) + tensor(0.5000) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + threshold: float = 0.5, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + ignore_index: Optional[int] = None, + validate_args: bool = True, + zero_division: float = 0, + **kwargs: Any, + ) -> None: + super().__init__( + num_labels=num_labels, + threshold=threshold, + ignore_index=ignore_index, + normalize=None, + validate_args=False, + **kwargs, + ) + if validate_args: + _multilabel_jaccard_index_arg_validation(num_labels, threshold, ignore_index, average) + self.validate_args = validate_args + self.average = average + self.zero_division = zero_division + + def compute(self) -> Tensor: + """Compute metric.""" + return _jaccard_index_reduce(self.confmat, average=self.average, zero_division=self.zero_division) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelJaccardIndex + >>> metric = MultilabelJaccardIndex(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelJaccardIndex + >>> metric = MultilabelJaccardIndex(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class JaccardIndex(_ClassificationTaskWrapper): + r"""Calculate the Jaccard index for multilabel tasks. + + The `Jaccard index`_ (also known as the intersection over union or jaccard similarity coefficient) is an statistic + that can be used to determine the similarity and diversity of a sample set. It is defined as the size of the + intersection divided by the union of the sample sets: + + .. math:: J(A,B) = \frac{|A\cap B|}{|A\cup B|} + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryJaccardIndex`, + :class:`~torchmetrics.classification.MulticlassJaccardIndex` and + :class:`~torchmetrics.classification.MultilabelJaccardIndex` for the specific details of each argument influence + and examples. + + Legacy Example: + >>> from torch import randint, tensor + >>> target = randint(0, 2, (10, 25, 25)) + >>> pred = tensor(target) + >>> pred[2:5, 7:13, 9:15] = 1 - pred[2:5, 7:13, 9:15] + >>> jaccard = JaccardIndex(task="multiclass", num_classes=2) + >>> jaccard(pred, target) + tensor(0.9660) + + """ + + def __new__( # type: ignore[misc] + cls: type["JaccardIndex"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + kwargs.update({"ignore_index": ignore_index, "validate_args": validate_args}) + if task == ClassificationTask.BINARY: + return BinaryJaccardIndex(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassJaccardIndex(num_classes, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelJaccardIndex(num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/logauc.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/logauc.py new file mode 100644 index 0000000000000000000000000000000000000000..6ccab43bb48908a4cc7fa045ed65bd36f1777a28 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/logauc.py @@ -0,0 +1,507 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Any, List, Optional, Sequence, Tuple, Type, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.roc import BinaryROC, MulticlassROC, MultilabelROC +from torchmetrics.functional.classification.logauc import ( + _binary_logauc_compute, + _reduce_logauc, + _validate_fpr_range, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["BinaryLogAUC.plot", "MulticlassLogAUC.plot", "MultilabelLogAUC.plot"] + + +class BinaryLogAUC(BinaryROC): + r"""Compute the `Log AUC`_ score for binary classification tasks. + + The score is computed by first computing the ROC curve, which then is interpolated to the specified range of false + positive rates (FPR) and then the log is taken of the FPR before the area under the curve (AUC) is computed. The + score is commonly used in applications where the positive and negative are imbalanced and a low false positive rate + is of high importance. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, ...)`` containing probabilities or logits for + each observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply + sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` containing ground truth labels, and + therefore only contain {0,1} values (except if `ignore_index` is specified). The value 1 always encodes the + positive class. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``logauc`` (:class:`~torch.Tensor`): A single scalar with the logauc score. + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + Args: + fpr_range: 2-element tuple with the lower and upper bound of the false positive rate range to compute the log + AUC score. + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryLogAUC + >>> preds = tensor([0.75, 0.05, 0.05, 0.05, 0.05]) + >>> target = tensor([1, 0, 0, 0, 0]) + >>> metric = BinaryLogAUC() + >>> metric(preds, target) + tensor(1.) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + fpr_range: Tuple[float, float] = (0.001, 0.1), + thresholds: Optional[Union[int, List[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = False, + **kwargs: Any, + ) -> None: + super().__init__(thresholds=thresholds, ignore_index=ignore_index, validate_args=validate_args, **kwargs) + if validate_args: + _validate_fpr_range(fpr_range) + self.fpr_range = fpr_range + + def compute(self) -> Tensor: # type: ignore[override] + """Computes the log AUC score.""" + fpr, tpr, _ = super().compute() + return _binary_logauc_compute(fpr, tpr, fpr_range=self.fpr_range) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single + >>> import torch + >>> from torchmetrics.classification import BinaryLogAUC + >>> metric = BinaryLogAUC() + >>> metric.update(torch.rand(20,), torch.randint(2, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.classification import BinaryLogAUC + >>> metric = BinaryLogAUC() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.rand(20,), torch.randint(2, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassLogAUC(MulticlassROC): + r"""Compute the `Log AUC`_ score for multiclass classification tasks. + + The score is computed by first computing the ROC curve, which then is interpolated to the specified range of false + positive rates (FPR) and then the log is taken of the FPR before the area under the curve (AUC) is computed. The + score is commonly used in applications where the positive and negative are imbalanced and a low false positive rate + is of high importance. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)`` containing probabilities or logits + for each observation. If preds has values outside [0,1] range we consider the input to be logits and will auto + apply softmax per sample. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` containing ground truth labels, and + therefore only contain values in the [0, n_classes-1] range (except if `ignore_index` is specified). + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``logauc`` (:class:`~torch.Tensor`): If `average=None|"none"` then a 1d tensor of shape (n_classes, ) will + be returned with logauc score per class. If `average="macro"` then a single scalar is returned. + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + Args: + num_classes: Integer specifying the number of classes + fpr_range: 2-element tuple with the lower and upper bound of the false positive rate range to compute the log + AUC score. + average: + Defines the reduction that is applied over classes. Should be one of the following: + + - ``"macro"``: Calculate score for each class and average them + - ``"weighted"``: calculates score for each class and computes weighted average using their support + - ``"none"`` or ``None``: calculates score for each class and applies no reduction + + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassLogAUC + >>> preds = tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = tensor([0, 1, 3, 2]) + >>> metric = MulticlassLogAUC(num_classes=5, average="macro", thresholds=None) + >>> metric(preds, target) + tensor(0.4000) + >>> metric = MulticlassLogAUC(num_classes=5, average=None, thresholds=None) + >>> metric(preds, target) + tensor([1., 1., 0., 0., 0.]) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + fpr_range: Tuple[float, float] = (0.001, 0.1), + average: Optional[Literal["macro", "none"]] = None, + thresholds: Optional[Union[int, List[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_classes=num_classes, + thresholds=thresholds, + average=None, + ignore_index=ignore_index, + validate_args=validate_args, + **kwargs, + ) + if validate_args: + _validate_fpr_range(fpr_range) + self.fpr_range = fpr_range + self.average2 = average # self.average is already used by parent class + + def compute(self) -> Tensor: # type: ignore[override] + """Computes the log AUC score.""" + fpr, tpr, _ = super().compute() + return _reduce_logauc(fpr, tpr, fpr_range=self.fpr_range, average=self.average2) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single + >>> import torch + >>> from torchmetrics.classification import MulticlassLogAUC + >>> metric = MulticlassLogAUC(num_classes=3) + >>> metric.update(torch.randn(20, 3), torch.randint(3,(20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.classification import MulticlassLogAUC + >>> metric = MulticlassLogAUC(num_classes=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randn(20, 3), torch.randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelLogAUC(MultilabelROC): + r"""Compute the `Log AUC`_ score for multiclass classification tasks. + + The score is computed by first computing the ROC curve, which then is interpolated to the specified range of false + positive rates (FPR) and then the log is taken of the FPR before the area under the curve (AUC) is computed. The + score is commonly used in applications where the positive and negative are imbalanced and a low false positive rate + is of high importance. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)`` containing probabilities or logits + for each observation. If preds has values outside [0,1] range we consider the input to be logits and will auto + apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)`` containing ground truth labels, and + therefore only contain {0,1} values (except if `ignore_index` is specified). + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``logauc`` (:class:`~torch.Tensor`): If `average=None|"none"` then a 1d tensor of shape (num_labels, ) will + be returned with logauc score per class. If `average="macro"` then a single scalar is returned. + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{labels})` (constant memory). + + Args: + num_labels: Integer specifying the number of labels + fpr_range: 2-element tuple with the lower and upper bound of the false positive rate range to compute the log + AUC score. + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``"macro"``: Calculate the score for each label and average them + - ``"none"`` or ``None``: calculates score for each label and applies no reduction + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelLogAUC + >>> preds = tensor([[0.75, 0.05, 0.35], + ... [0.45, 0.75, 0.05], + ... [0.05, 0.55, 0.75], + ... [0.05, 0.65, 0.05]]) + >>> target = tensor([[1, 0, 1], + ... [0, 0, 0], + ... [0, 1, 1], + ... [1, 1, 1]]) + >>> metric = MultilabelLogAUC(num_labels=3, average="macro", thresholds=None) + >>> metric(preds, target) + tensor(0.3945) + >>> metric = MultilabelLogAUC(num_labels=3, average=None, thresholds=None) + >>> metric(preds, target) + tensor([0.5000, 0.0000, 0.6835]) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + fpr_range: Tuple[float, float] = (0.001, 0.1), + average: Optional[Literal["macro", "none"]] = None, + thresholds: Optional[Union[int, List[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + if validate_args: + _validate_fpr_range(fpr_range) + self.fpr_range = fpr_range + self.average2 = average # self.average is already used by parent class + super().__init__( + num_labels=num_labels, + thresholds=thresholds, + ignore_index=ignore_index, + validate_args=validate_args, + **kwargs, + ) + + def compute(self) -> Tensor: # type: ignore[override] + """Computes the log AUC score.""" + fpr, tpr, _ = super().compute() + return _reduce_logauc(fpr, tpr, fpr_range=self.fpr_range, average=self.average2) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single + >>> import torch + >>> from torchmetrics.classification import MultilabelLogAUC + >>> metric = MultilabelLogAUC(num_labels=3) + >>> metric.update(torch.rand(20,3), torch.randint(2, (20,3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.classification import MultilabelLogAUC + >>> metric = MultilabelLogAUC(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.rand(20,3), torch.randint(2, (20,3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class LogAUC(_ClassificationTaskWrapper): + r"""Compute the `Log AUC`_ score for multiclass classification tasks. + + The score is computed by first computing the ROC curve, which then is interpolated to the specified range of false + positive rates (FPR) and then the log is taken of the FPR before the area under the curve (AUC) is computed. The + score is commonly used in applications where the positive and negative are imbalanced and a low false positive rate + is of high importance. + + This module is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryLogAUC`, :class:`~torchmetrics.classification.MulticlassLogAUC` and + :class:`~torchmetrics.classification.MultilabelLogAUC` for the specific details of each argument influence and + examples. + + """ + + def __new__( # type: ignore[misc] + cls: Type["LogAUC"], + task: Literal["binary", "multiclass", "multilabel"], + thresholds: Optional[Union[int, List[float], Tensor]] = None, + fpr_range: Optional[Tuple[float, float]] = (0.001, 0.1), + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + kwargs.update({ + "thresholds": thresholds, + "fpr_range": fpr_range, + "ignore_index": ignore_index, + "validate_args": validate_args, + }) + if task == ClassificationTask.BINARY: + return BinaryLogAUC(**kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassLogAUC(num_classes, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelLogAUC(num_labels, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/matthews_corrcoef.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/matthews_corrcoef.py new file mode 100644 index 0000000000000000000000000000000000000000..2f26b452e25732c1ee6fe89297f83e89a6a8c2d9 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/matthews_corrcoef.py @@ -0,0 +1,416 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.confusion_matrix import ( + BinaryConfusionMatrix, + MulticlassConfusionMatrix, + MultilabelConfusionMatrix, +) +from torchmetrics.functional.classification.matthews_corrcoef import _matthews_corrcoef_reduce +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryMatthewsCorrCoef.plot", + "MulticlassMatthewsCorrCoef.plot", + "MultilabelMatthewsCorrCoef.plot", + ] + + +class BinaryMatthewsCorrCoef(BinaryConfusionMatrix): + r"""Calculate `Matthews correlation coefficient`_ for binary tasks. + + This metric measures the general correlation or quality of a classification. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A int tensor or float tensor of shape ``(N, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bmcc`` (:class:`~torch.Tensor`): A tensor containing the Binary Matthews Correlation Coefficient. + + Args: + threshold: Threshold for transforming probability to binary (0,1) predictions + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryMatthewsCorrCoef + >>> target = tensor([1, 1, 0, 0]) + >>> preds = tensor([0, 1, 0, 0]) + >>> metric = BinaryMatthewsCorrCoef() + >>> metric(preds, target) + tensor(0.5774) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryMatthewsCorrCoef + >>> target = tensor([1, 1, 0, 0]) + >>> preds = tensor([0.35, 0.85, 0.48, 0.01]) + >>> metric = BinaryMatthewsCorrCoef() + >>> metric(preds, target) + tensor(0.5774) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + threshold: float = 0.5, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(threshold, ignore_index, normalize=None, validate_args=validate_args, **kwargs) + + def compute(self) -> Tensor: + """Compute metric.""" + return _matthews_corrcoef_reduce(self.confmat) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryMatthewsCorrCoef + >>> metric = BinaryMatthewsCorrCoef() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryMatthewsCorrCoef + >>> metric = BinaryMatthewsCorrCoef() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassMatthewsCorrCoef(MulticlassConfusionMatrix): + r"""Calculate `Matthews correlation coefficient`_ for multiclass tasks. + + This metric measures the general correlation or quality of a classification. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcmcc`` (:class:`~torch.Tensor`): A tensor containing the Multi-class Matthews Correlation Coefficient. + + Args: + num_classes: Integer specifying the number of classes + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (pred is integer tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassMatthewsCorrCoef + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassMatthewsCorrCoef(num_classes=3) + >>> metric(preds, target) + tensor(0.7000) + + Example (pred is float tensor): + >>> from torchmetrics.classification import MulticlassMatthewsCorrCoef + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassMatthewsCorrCoef(num_classes=3) + >>> metric(preds, target) + tensor(0.7000) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(num_classes, ignore_index, normalize=None, validate_args=validate_args, **kwargs) + + def compute(self) -> Tensor: + """Compute metric.""" + return _matthews_corrcoef_reduce(self.confmat) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassMatthewsCorrCoef + >>> metric = MulticlassMatthewsCorrCoef(num_classes=3) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassMatthewsCorrCoef + >>> metric = MulticlassMatthewsCorrCoef(num_classes=3) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelMatthewsCorrCoef(MultilabelConfusionMatrix): + r"""Calculate `Matthews correlation coefficient`_ for multilabel tasks. + + This metric measures the general correlation or quality of a classification. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, C, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)`` + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlmcc`` (:class:`~torch.Tensor`): A tensor containing the Multi-label Matthews Correlation Coefficient. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelMatthewsCorrCoef + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelMatthewsCorrCoef(num_labels=3) + >>> metric(preds, target) + tensor(0.3333) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelMatthewsCorrCoef + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelMatthewsCorrCoef(num_labels=3) + >>> metric(preds, target) + tensor(0.3333) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + threshold: float = 0.5, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(num_labels, threshold, ignore_index, normalize=None, validate_args=validate_args, **kwargs) + + def compute(self) -> Tensor: + """Compute metric.""" + return _matthews_corrcoef_reduce(self.confmat) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelMatthewsCorrCoef + >>> metric = MultilabelMatthewsCorrCoef(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelMatthewsCorrCoef + >>> metric = MultilabelMatthewsCorrCoef(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MatthewsCorrCoef(_ClassificationTaskWrapper): + r"""Calculate `Matthews correlation coefficient`_ . + + This metric measures the general correlation or quality of a classification. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryMatthewsCorrCoef`, + :class:`~torchmetrics.classification.MulticlassMatthewsCorrCoef` and + :class:`~torchmetrics.classification.MultilabelMatthewsCorrCoef` for the specific details of each argument influence + and examples. + + Legacy Example: + >>> from torch import tensor + >>> target = tensor([1, 1, 0, 0]) + >>> preds = tensor([0, 1, 0, 0]) + >>> matthews_corrcoef = MatthewsCorrCoef(task='binary') + >>> matthews_corrcoef(preds, target) + tensor(0.5774) + + """ + + def __new__( # type: ignore[misc] + cls: type["MatthewsCorrCoef"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + kwargs.update({"ignore_index": ignore_index, "validate_args": validate_args}) + if task == ClassificationTask.BINARY: + return BinaryMatthewsCorrCoef(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassMatthewsCorrCoef(num_classes, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelMatthewsCorrCoef(num_labels, threshold, **kwargs) + raise ValueError(f"Not handled value: {task}") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/negative_predictive_value.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/negative_predictive_value.py new file mode 100644 index 0000000000000000000000000000000000000000..6d4471bd257beb3927889b7539c297ae2ea7bf2b --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/negative_predictive_value.py @@ -0,0 +1,522 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.stat_scores import BinaryStatScores, MulticlassStatScores, MultilabelStatScores +from torchmetrics.functional.classification.negative_predictive_value import _negative_predictive_value_reduce +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryNegativePredictiveValue.plot", + "MulticlassNegativePredictiveValue.plot", + "MultilabelNegativePredictiveValue.plot", + ] + + +class BinaryNegativePredictiveValue(BinaryStatScores): + r"""Compute `Negative Predictive Value`_ for binary tasks. + + .. math:: \text{Negative Predictive Value} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TN} + \text{FP} \neq 0`. If this case is + encountered a score of 0 is returned. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, ...)``. If preds is a floating point + tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid per + element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``npv`` (:class:`~torch.Tensor`): If ``multidim_average`` is set to ``global``, the metric returns a scalar value. + If ``multidim_average`` is set to ``samplewise``, the metric returns ``(N,)`` vector consisting of a scalar value + per sample. + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryNegativePredictiveValue + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinaryNegativePredictiveValue() + >>> metric(preds, target) + tensor(0.6667) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryNegativePredictiveValue + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinaryNegativePredictiveValue() + >>> metric(preds, target) + tensor(0.6667) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinaryNegativePredictiveValue + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinaryNegativePredictiveValue(multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.0000, 0.2500]) + + """ + + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _negative_predictive_value_reduce( + tp, fp, tn, fn, average="binary", multidim_average=self.multidim_average + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryNegativePredictiveValue + >>> metric = BinaryNegativePredictiveValue() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryNegativePredictiveValue + >>> metric = BinaryNegativePredictiveValue() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassNegativePredictiveValue(MulticlassStatScores): + r"""Compute `Negative Predictive Value`_ for multiclass tasks. + + .. math:: \text{Negative Predictive Value} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TN} + \text{FP} \neq 0`. If this case is + encountered for any class, the metric for that class will be set to 0 and the overall metric may therefore be + affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``npv`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + top_k: + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassNegativePredictiveValue + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassNegativePredictiveValue(num_classes=3) + >>> metric(preds, target) + tensor(0.8889) + >>> metric = MulticlassNegativePredictiveValue(num_classes=3, average=None) + >>> metric(preds, target) + tensor([0.6667, 1.0000, 1.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassNegativePredictiveValue + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassNegativePredictiveValue(num_classes=3) + >>> metric(preds, target) + tensor(0.8889) + >>> metric = MulticlassNegativePredictiveValue(num_classes=3, average=None) + >>> metric(preds, target) + tensor([0.6667, 1.0000, 1.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassNegativePredictiveValue + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassNegativePredictiveValue(num_classes=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.7833, 0.6556]) + >>> metric = MulticlassNegativePredictiveValue(num_classes=3, multidim_average='samplewise', average=None) + >>> metric(preds, target) + tensor([[1.0000, 0.6000, 0.7500], + [0.8000, 0.5000, 0.6667]]) + + """ + + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _negative_predictive_value_reduce( + tp, fp, tn, fn, average=self.average, multidim_average=self.multidim_average, top_k=self.top_k + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassNegativePredictiveValue + >>> metric = MulticlassNegativePredictiveValue(num_classes=3, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassNegativePredictiveValue + >>> metric = MulticlassNegativePredictiveValue(num_classes=3, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelNegativePredictiveValue(MultilabelStatScores): + r"""Compute `Negative Predictive Value`_ for multilabel tasks. + + .. math:: \text{Negative Predictive Value} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TN} + \text{FP} \neq 0`. If this case is + encountered for any label, the metric for that label will be set to 0 and the overall metric may therefore be + affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, C, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``npv`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global`` + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise`` + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelNegativePredictiveValue + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelNegativePredictiveValue(num_labels=3) + >>> metric(preds, target) + tensor(0.5000) + >>> mls = MultilabelNegativePredictiveValue(num_labels=3, average=None) + >>> mls(preds, target) + tensor([1.0000, 0.5000, 0.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelNegativePredictiveValue + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelNegativePredictiveValue(num_labels=3) + >>> metric(preds, target) + tensor(0.5000) + >>> mls = MultilabelNegativePredictiveValue(num_labels=3, average=None) + >>> mls(preds, target) + tensor([1.0000, 0.5000, 0.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelNegativePredictiveValue + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelNegativePredictiveValue(num_labels=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.0000, 0.1667]) + >>> mls = MultilabelNegativePredictiveValue(num_labels=3, multidim_average='samplewise', average=None) + >>> mls(preds, target) + tensor([[0.0000, 0.0000, 0.0000], + [0.0000, 0.0000, 0.5000]]) + + """ + + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _negative_predictive_value_reduce( + tp, fp, tn, fn, average=self.average, multidim_average=self.multidim_average, multilabel=True + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling ``metric.forward`` or ``metric.compute`` or a list of these + results. If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelNegativePredictiveValue + >>> metric = MultilabelNegativePredictiveValue(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelNegativePredictiveValue + >>> metric = MultilabelNegativePredictiveValue(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class NegativePredictiveValue(_ClassificationTaskWrapper): + r"""Compute `Negative Predictive Value`_. + + .. math:: \text{Negative Predictive Value} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0`. If this case is + encountered for any class/label, the metric for that class/label will be set to 0 and the overall metric may + therefore be affected in turn. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryNegativePredictiveValue`, + :class:`~torchmetrics.classification.MulticlassNegativePredictiveValue` + and :class:`~torchmetrics.classification.MultilabelNegativePredictiveValue` for the specific details of each + argument influence and examples. + + Legacy Example: + >>> from torch import tensor + >>> preds = tensor([2, 0, 2, 1]) + >>> target = tensor([1, 1, 2, 0]) + >>> nvp = NegativePredictiveValue(task="multiclass", average='macro', num_classes=3) + >>> nvp(preds, target) + tensor(0.6667) + >>> nvp = NegativePredictiveValue(task="multiclass", average='micro', num_classes=3) + >>> nvp(preds, target) + tensor(0.6250) + + """ + + def __new__( # type: ignore[misc] + cls: type["NegativePredictiveValue"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + }) + if task == ClassificationTask.BINARY: + return BinaryNegativePredictiveValue(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassNegativePredictiveValue(num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelNegativePredictiveValue(num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_fixed_recall.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_fixed_recall.py new file mode 100644 index 0000000000000000000000000000000000000000..7ec96603a4b6e888fe64cff1bcee368b0c2f477b --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_fixed_recall.py @@ -0,0 +1,515 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.precision_recall_curve import ( + BinaryPrecisionRecallCurve, + MulticlassPrecisionRecallCurve, + MultilabelPrecisionRecallCurve, +) +from torchmetrics.functional.classification.precision_fixed_recall import _precision_at_recall +from torchmetrics.functional.classification.recall_fixed_precision import ( + _binary_recall_at_fixed_precision_arg_validation, + _binary_recall_at_fixed_precision_compute, + _multiclass_recall_at_fixed_precision_arg_compute, + _multiclass_recall_at_fixed_precision_arg_validation, + _multilabel_recall_at_fixed_precision_arg_compute, + _multilabel_recall_at_fixed_precision_arg_validation, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryPrecisionAtFixedRecall.plot", + "MulticlassPrecisionAtFixedRecall.plot", + "MultilabelPrecisionAtFixedRecall.plot", + ] + + +class BinaryPrecisionAtFixedRecall(BinaryPrecisionRecallCurve): + r"""Compute the highest possible precision value given the minimum recall thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the precision for + a given recall level. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, ...)``. Preds should be a tensor containing + probabilities or logits for each observation. If preds has values outside [0,1] range we consider the input + to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). The value + 1 always encodes the positive class. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``precision`` (:class:`~torch.Tensor`): A scalar tensor with the maximum precision for the given recall level + - ``threshold`` (:class:`~torch.Tensor`): A scalar tensor with the corresponding threshold level + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a + binned version that is less accurate but more memory efficient. Setting the `thresholds` argument to ``None`` + will activate the non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting + the `thresholds` argument to either an integer, list or a 1d tensor will use a binned version that uses memory + of size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + Args: + min_recall: float value specifying minimum recall threshold. + thresholds: + Can be one of: + + - If set to ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d :class:`~torch.Tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryPrecisionAtFixedRecall + >>> preds = tensor([0, 0.5, 0.7, 0.8]) + >>> target = tensor([0, 1, 1, 0]) + >>> metric = BinaryPrecisionAtFixedRecall(min_recall=0.5, thresholds=None) + >>> metric(preds, target) + (tensor(0.6667), tensor(0.5000)) + >>> metric = BinaryPrecisionAtFixedRecall(min_recall=0.5, thresholds=5) + >>> metric(preds, target) + (tensor(0.6667), tensor(0.5000)) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + min_recall: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(thresholds, ignore_index, validate_args=False, **kwargs) + if validate_args: + _binary_recall_at_fixed_precision_arg_validation(min_recall, thresholds, ignore_index) + self.validate_args = validate_args + self.min_recall = min_recall + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _binary_recall_at_fixed_precision_compute( + state, self.thresholds, self.min_recall, reduce_fn=_precision_at_recall + ) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryPrecisionAtFixedRecall + >>> metric = BinaryPrecisionAtFixedRecall(min_recall=0.5) + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() # the returned plot only shows the maximum recall value by default + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryPrecisionAtFixedRecall + >>> metric = BinaryPrecisionAtFixedRecall(min_recall=0.5) + >>> values = [ ] + >>> for _ in range(10): + ... # we index by 0 such that only the maximum recall value is plotted + ... values.append(metric(rand(10), randint(2,(10,)))[0]) + >>> fig_, ax_ = metric.plot(values) + + """ + val = val or self.compute()[0] # by default we select the maximum recall value to plot + return self._plot(val, ax) + + +class MulticlassPrecisionAtFixedRecall(MulticlassPrecisionRecallCurve): + r"""Compute the highest possible precision value given the minimum recall thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the precision for + a given recall level. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply softmax per sample. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain values in the [0, n_classes-1] range (except if `ignore_index` + is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns a tuple of either 2 tensors or 2 lists containing: + + - ``precision`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the maximum precision for the + given recall level per class + - ``threshold`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the corresponding threshold + level per class + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to ``None`` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{classes})` (constant memory). + + Args: + num_classes: Integer specifying the number of classes + min_recall: float value specifying minimum recall threshold. + thresholds: + Can be one of: + + - If set to ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d :class:`~torch.Tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassPrecisionAtFixedRecall + >>> preds = tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = tensor([0, 1, 3, 2]) + >>> metric = MulticlassPrecisionAtFixedRecall(num_classes=5, min_recall=0.5, thresholds=None) + >>> metric(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([1.0000, 1.0000, 0.2500, 0.2500, 0.0000]), + tensor([7.5000e-01, 7.5000e-01, 5.0000e-02, 5.0000e-02, 1.0000e+06])) + >>> mcrafp = MulticlassPrecisionAtFixedRecall(num_classes=5, min_recall=0.5, thresholds=5) + >>> mcrafp(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([1.0000, 1.0000, 0.2500, 0.2500, 0.0000]), + tensor([7.5000e-01, 7.5000e-01, 0.0000e+00, 0.0000e+00, 1.0000e+06])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + min_recall: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_classes=num_classes, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multiclass_recall_at_fixed_precision_arg_validation(num_classes, min_recall, thresholds, ignore_index) + self.validate_args = validate_args + self.min_recall = min_recall + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _multiclass_recall_at_fixed_precision_arg_compute( + state, self.num_classes, self.thresholds, self.min_recall, reduce_fn=_precision_at_recall + ) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassPrecisionAtFixedRecall + >>> metric = MulticlassPrecisionAtFixedRecall(num_classes=3, min_recall=0.5) + >>> metric.update(rand(20, 3).softmax(dim=-1), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() # the returned plot only shows the maximum recall value by default + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassPrecisionAtFixedRecall + >>> metric = MulticlassPrecisionAtFixedRecall(num_classes=3, min_recall=0.5) + >>> values = [] + >>> for _ in range(20): + ... # we index by 0 such that only the maximum recall value is plotted + ... values.append(metric(rand(20, 3).softmax(dim=-1), randint(3, (20,)))[0]) + >>> fig_, ax_ = metric.plot(values) + + """ + val = val or self.compute()[0] # by default we select the maximum recall value to plot + return self._plot(val, ax) + + +class MultilabelPrecisionAtFixedRecall(MultilabelPrecisionRecallCurve): + r"""Compute the highest possible precision value given the minimum recall thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the precision for + a given recall level. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). The value + 1 always encodes the positive class. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns a tuple of either 2 tensors or 2 lists containing: + + - ``precision`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the maximum precision for the + given recall level per class + - ``threshold`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the corresponding threshold + level per class + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to ``None`` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{labels})` (constant memory). + + Args: + num_labels: Integer specifying the number of labels + min_recall: float value specifying minimum recall threshold. + thresholds: + Can be one of: + + - If set to ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d :class:`~torch.Tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelPrecisionAtFixedRecall + >>> preds = tensor([[0.75, 0.05, 0.35], + ... [0.45, 0.75, 0.05], + ... [0.05, 0.55, 0.75], + ... [0.05, 0.65, 0.05]]) + >>> target = tensor([[1, 0, 1], + ... [0, 0, 0], + ... [0, 1, 1], + ... [1, 1, 1]]) + >>> metric = MultilabelPrecisionAtFixedRecall(num_labels=3, min_recall=0.5, thresholds=None) + >>> metric(preds, target) + (tensor([1.0000, 0.6667, 1.0000]), tensor([0.7500, 0.5500, 0.3500])) + >>> mlrafp = MultilabelPrecisionAtFixedRecall(num_labels=3, min_recall=0.5, thresholds=5) + >>> mlrafp(preds, target) + (tensor([1.0000, 0.6667, 1.0000]), tensor([0.7500, 0.5000, 0.2500])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + min_recall: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_labels=num_labels, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multilabel_recall_at_fixed_precision_arg_validation(num_labels, min_recall, thresholds, ignore_index) + self.validate_args = validate_args + self.min_recall = min_recall + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _multilabel_recall_at_fixed_precision_arg_compute( + state, self.num_labels, self.thresholds, self.ignore_index, self.min_recall, reduce_fn=_precision_at_recall + ) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelPrecisionAtFixedRecall + >>> metric = MultilabelPrecisionAtFixedRecall(num_labels=3, min_recall=0.5) + >>> metric.update(rand(20, 3), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() # the returned plot only shows the maximum recall value by default + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelPrecisionAtFixedRecall + >>> metric = MultilabelPrecisionAtFixedRecall(num_labels=3, min_recall=0.5) + >>> values = [ ] + >>> for _ in range(10): + ... # we index by 0 such that only the maximum recall value is plotted + ... values.append(metric(rand(20, 3), randint(2, (20, 3)))[0]) + >>> fig_, ax_ = metric.plot(values) + + """ + val = val or self.compute()[0] # by default we select the maximum recall value to plot + return self._plot(val, ax) + + +class PrecisionAtFixedRecall(_ClassificationTaskWrapper): + r"""Compute the highest possible recall value given the minimum precision thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the recall for + a given precision level. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryPrecisionAtFixedRecall`, + :class:`~torchmetrics.classification.MulticlassPrecisionAtFixedRecall` and + :class:`~torchmetrics.classification.MultilabelPrecisionAtFixedRecall` for the specific details of each argument + influence and examples. + + """ + + def __new__( # type: ignore[misc] + cls: type["PrecisionAtFixedRecall"], + task: Literal["binary", "multiclass", "multilabel"], + min_recall: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + if task == ClassificationTask.BINARY: + return BinaryPrecisionAtFixedRecall(min_recall, thresholds, ignore_index, validate_args, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassPrecisionAtFixedRecall( + num_classes, min_recall, thresholds, ignore_index, validate_args, **kwargs + ) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelPrecisionAtFixedRecall( + num_labels, min_recall, thresholds, ignore_index, validate_args, **kwargs + ) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_recall.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_recall.py new file mode 100644 index 0000000000000000000000000000000000000000..aeb98fcf4632dd91292dae947d575bd2916f30a1 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_recall.py @@ -0,0 +1,1086 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.stat_scores import BinaryStatScores, MulticlassStatScores, MultilabelStatScores +from torchmetrics.functional.classification.precision_recall import ( + _precision_recall_reduce, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryPrecision.plot", + "MulticlassPrecision.plot", + "MultilabelPrecision.plot", + "BinaryRecall.plot", + "MulticlassRecall.plot", + "MultilabelRecall.plot", + ] + + +class BinaryPrecision(BinaryStatScores): + r"""Compute `Precision`_ for binary tasks. + + .. math:: \text{Precision} = \frac{\text{TP}}{\text{TP} + \text{FP}} + + Where :math:`\text{TP}` and :math:`\text{FP}` represent the number of true positives and false positives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0`. If this case is + encountered a score of `zero_division` (0 or 1, default is 0) is returned. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A int or float tensor of shape ``(N, ...)``. If preds is a floating point + tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid per + element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bp`` (:class:`~torch.Tensor`): If ``multidim_average`` is set to ``global``, the metric returns a scalar + value. If ``multidim_average`` is set to ``samplewise``, the metric returns ``(N,)`` vector consisting of a + scalar value per sample. + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when :math:`\text{TP} + \text{FP} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryPrecision + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinaryPrecision() + >>> metric(preds, target) + tensor(0.6667) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryPrecision + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinaryPrecision() + >>> metric(preds, target) + tensor(0.6667) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinaryPrecision + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinaryPrecision(multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.4000, 0.0000]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _precision_recall_reduce( + "precision", + tp, + fp, + tn, + fn, + average="binary", + multidim_average=self.multidim_average, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryPrecision + >>> metric = BinaryPrecision() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryPrecision + >>> metric = BinaryPrecision() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassPrecision(MulticlassStatScores): + r"""Compute `Precision`_ for multiclass tasks. + + .. math:: \text{Precision} = \frac{\text{TP}}{\text{TP} + \text{FP}} + + Where :math:`\text{TP}` and :math:`\text{FP}` represent the number of true positives and false positives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0`. If this case is + encountered for any class, the metric for that class will be set to `zero_division` (0 or 1, default is 0) and + the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. + + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcp`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + top_k: + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when :math:`\text{TP} + \text{FP} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassPrecision + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassPrecision(num_classes=3) + >>> metric(preds, target) + tensor(0.8333) + >>> mcp = MulticlassPrecision(num_classes=3, average=None) + >>> mcp(preds, target) + tensor([1.0000, 0.5000, 1.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassPrecision + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassPrecision(num_classes=3) + >>> metric(preds, target) + tensor(0.8333) + >>> mcp = MulticlassPrecision(num_classes=3, average=None) + >>> mcp(preds, target) + tensor([1.0000, 0.5000, 1.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassPrecision + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassPrecision(num_classes=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.3889, 0.2778]) + >>> mcp = MulticlassPrecision(num_classes=3, multidim_average='samplewise', average=None) + >>> mcp(preds, target) + tensor([[0.6667, 0.0000, 0.5000], + [0.0000, 0.5000, 0.3333]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _precision_recall_reduce( + "precision", + tp, + fp, + tn, + fn, + average=self.average, + multidim_average=self.multidim_average, + top_k=self.top_k, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassPrecision + >>> metric = MulticlassPrecision(num_classes=3, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassPrecision + >>> metric = MulticlassPrecision(num_classes=3, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelPrecision(MultilabelStatScores): + r"""Compute `Precision`_ for multilabel tasks. + + .. math:: \text{Precision} = \frac{\text{TP}}{\text{TP} + \text{FP}} + + Where :math:`\text{TP}` and :math:`\text{FP}` represent the number of true positives and false positives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0`. If this case is + encountered for any label, the metric for that label will be set to `zero_division` (0 or 1, default is 0) and + the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor or float tensor of shape ``(N, C, ...)``. + If preds is a floating point tensor with values outside [0,1] range we consider the input to be logits and + will auto apply sigmoid per element. Additionally, we convert to int tensor with thresholding using the value + in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlp`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when :math:`\text{TP} + \text{FP} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelPrecision + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelPrecision(num_labels=3) + >>> metric(preds, target) + tensor(0.5000) + >>> mlp = MultilabelPrecision(num_labels=3, average=None) + >>> mlp(preds, target) + tensor([1.0000, 0.0000, 0.5000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelPrecision + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelPrecision(num_labels=3) + >>> metric(preds, target) + tensor(0.5000) + >>> mlp = MultilabelPrecision(num_labels=3, average=None) + >>> mlp(preds, target) + tensor([1.0000, 0.0000, 0.5000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelPrecision + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelPrecision(num_labels=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.3333, 0.0000]) + >>> mlp = MultilabelPrecision(num_labels=3, multidim_average='samplewise', average=None) + >>> mlp(preds, target) + tensor([[0.5000, 0.5000, 0.0000], + [0.0000, 0.0000, 0.0000]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _precision_recall_reduce( + "precision", + tp, + fp, + tn, + fn, + average=self.average, + multidim_average=self.multidim_average, + multilabel=True, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelPrecision + >>> metric = MultilabelPrecision(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelPrecision + >>> metric = MultilabelPrecision(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class BinaryRecall(BinaryStatScores): + r"""Compute `Recall`_ for binary tasks. + + .. math:: \text{Recall} = \frac{\text{TP}}{\text{TP} + \text{FN}} + + Where :math:`\text{TP}` and :math:`\text{FN}` represent the number of true positives and false negatives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FN} \neq 0`. If this case is + encountered a score of `zero_division` (0 or 1, default is 0) is returned. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor or float tensor of shape ``(N, ...)``. If preds is a + floating point tensor with values outside [0,1] range we consider the input to be logits and will auto apply + sigmoid per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``br`` (:class:`~torch.Tensor`): If ``multidim_average`` is set to ``global``, the metric returns a scalar + value. If ``multidim_average`` is set to ``samplewise``, the metric returns ``(N,)`` vector consisting of + a scalar value per sample. + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when :math:`\text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryRecall + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinaryRecall() + >>> metric(preds, target) + tensor(0.6667) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryRecall + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinaryRecall() + >>> metric(preds, target) + tensor(0.6667) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinaryRecall + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinaryRecall(multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.6667, 0.0000]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _precision_recall_reduce( + "recall", + tp, + fp, + tn, + fn, + average="binary", + multidim_average=self.multidim_average, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryRecall + >>> metric = BinaryRecall() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryRecall + >>> metric = BinaryRecall() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassRecall(MulticlassStatScores): + r"""Compute `Recall`_ for multiclass tasks. + + .. math:: \text{Recall} = \frac{\text{TP}}{\text{TP} + \text{FN}} + + Where :math:`\text{TP}` and :math:`\text{FN}` represent the number of true positives and false negatives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FN} \neq 0`. If this case is + encountered for any class, the metric for that class will be set to `zero_division` (0 or 1, default is 0) and + the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)`` + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcr`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + top_k: + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when :math:`\text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassRecall + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassRecall(num_classes=3) + >>> metric(preds, target) + tensor(0.8333) + >>> mcr = MulticlassRecall(num_classes=3, average=None) + >>> mcr(preds, target) + tensor([0.5000, 1.0000, 1.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassRecall + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassRecall(num_classes=3) + >>> metric(preds, target) + tensor(0.8333) + >>> mcr = MulticlassRecall(num_classes=3, average=None) + >>> mcr(preds, target) + tensor([0.5000, 1.0000, 1.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassRecall + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassRecall(num_classes=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.5000, 0.2778]) + >>> mcr = MulticlassRecall(num_classes=3, multidim_average='samplewise', average=None) + >>> mcr(preds, target) + tensor([[1.0000, 0.0000, 0.5000], + [0.0000, 0.3333, 0.5000]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _precision_recall_reduce( + "recall", + tp, + fp, + tn, + fn, + average=self.average, + multidim_average=self.multidim_average, + top_k=self.top_k, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassRecall + >>> metric = MulticlassRecall(num_classes=3, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassRecall + >>> metric = MulticlassRecall(num_classes=3, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelRecall(MultilabelStatScores): + r"""Compute `Recall`_ for multilabel tasks. + + .. math:: \text{Recall} = \frac{\text{TP}}{\text{TP} + \text{FN}} + + Where :math:`\text{TP}` and :math:`\text{FN}` represent the number of true positives and false negatives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FN} \neq 0`. If this case is + encountered for any label, the metric for that label will be set to `zero_division` (0 or 1, default is 0) and + the overall metric may therefore be affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, C, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlr`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + zero_division: Should be `0` or `1`. The value returned when :math:`\text{TP} + \text{FN} = 0`. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelRecall + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelRecall(num_labels=3) + >>> metric(preds, target) + tensor(0.6667) + >>> mlr = MultilabelRecall(num_labels=3, average=None) + >>> mlr(preds, target) + tensor([1., 0., 1.]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelRecall + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelRecall(num_labels=3) + >>> metric(preds, target) + tensor(0.6667) + >>> mlr = MultilabelRecall(num_labels=3, average=None) + >>> mlr(preds, target) + tensor([1., 0., 1.]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelRecall + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelRecall(num_labels=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.6667, 0.0000]) + >>> mlr = MultilabelRecall(num_labels=3, multidim_average='samplewise', average=None) + >>> mlr(preds, target) + tensor([[1., 1., 0.], + [0., 0., 0.]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _precision_recall_reduce( + "recall", + tp, + fp, + tn, + fn, + average=self.average, + multidim_average=self.multidim_average, + multilabel=True, + zero_division=self.zero_division, + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelRecall + >>> metric = MultilabelRecall(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelRecall + >>> metric = MultilabelRecall(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class Precision(_ClassificationTaskWrapper): + r"""Compute `Precision`_. + + .. math:: \text{Precision} = \frac{\text{TP}}{\text{TP} + \text{FP}} + + Where :math:`\text{TP}` and :math:`\text{FP}` represent the number of true positives and false positives + respectively. The metric is only proper defined when :math:`\text{TP} + \text{FP} \neq 0`. If this case is + encountered for any class/label, the metric for that class/label will be set to 0 and the overall metric may + therefore be affected in turn. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryPrecision`, :class:`~torchmetrics.classification.MulticlassPrecision` and + :class:`~torchmetrics.classification.MultilabelPrecision` for the specific details of each argument influence and + examples. + + Legacy Example: + >>> from torch import tensor + >>> preds = tensor([2, 0, 2, 1]) + >>> target = tensor([1, 1, 2, 0]) + >>> precision = Precision(task="multiclass", average='macro', num_classes=3) + >>> precision(preds, target) + tensor(0.1667) + >>> precision = Precision(task="multiclass", average='micro', num_classes=3) + >>> precision(preds, target) + tensor(0.2500) + + """ + + def __new__( # type: ignore[misc] + cls: type["Precision"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + }) + task = ClassificationTask.from_str(task) + if task == ClassificationTask.BINARY: + return BinaryPrecision(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassPrecision(num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelPrecision(num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") + + +class Recall(_ClassificationTaskWrapper): + r"""Compute `Recall`_. + + .. math:: \text{Recall} = \frac{\text{TP}}{\text{TP} + \text{FN}} + + Where :math:`\text{TP}` and :math:`\text{FN}` represent the number of true positives and + false negatives respectively. The metric is only proper defined when :math:`\text{TP} + \text{FN} \neq 0`. If this + case is encountered for any class/label, the metric for that class/label will be set to 0 and the overall metric may + therefore be affected in turn. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryRecall`, + :class:`~torchmetrics.classification.MulticlassRecall` and :class:`~torchmetrics.classification.MultilabelRecall` + for the specific details of each argument influence and examples. + + Legacy Example: + >>> from torch import tensor + >>> preds = tensor([2, 0, 2, 1]) + >>> target = tensor([1, 1, 2, 0]) + >>> recall = Recall(task="multiclass", average='macro', num_classes=3) + >>> recall(preds, target) + tensor(0.3333) + >>> recall = Recall(task="multiclass", average='micro', num_classes=3) + >>> recall(preds, target) + tensor(0.2500) + + """ + + def __new__( # type: ignore[misc] + cls: type["Recall"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + }) + if task == ClassificationTask.BINARY: + return BinaryRecall(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassRecall(num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelRecall(num_labels, threshold, average, **kwargs) + return None # type: ignore[return-value] diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_recall_curve.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_recall_curve.py new file mode 100644 index 0000000000000000000000000000000000000000..351ecc890060dcacc5b9b3774713fdb3fff6c381 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/precision_recall_curve.py @@ -0,0 +1,692 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Any, List, Optional, Union + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.functional.classification.auroc import _reduce_auroc +from torchmetrics.functional.classification.precision_recall_curve import ( + _adjust_threshold_arg, + _binary_precision_recall_curve_arg_validation, + _binary_precision_recall_curve_compute, + _binary_precision_recall_curve_format, + _binary_precision_recall_curve_tensor_validation, + _binary_precision_recall_curve_update, + _multiclass_precision_recall_curve_arg_validation, + _multiclass_precision_recall_curve_compute, + _multiclass_precision_recall_curve_format, + _multiclass_precision_recall_curve_tensor_validation, + _multiclass_precision_recall_curve_update, + _multilabel_precision_recall_curve_arg_validation, + _multilabel_precision_recall_curve_compute, + _multilabel_precision_recall_curve_format, + _multilabel_precision_recall_curve_tensor_validation, + _multilabel_precision_recall_curve_update, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.compute import _auc_compute_without_check +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE, plot_curve + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryPrecisionRecallCurve.plot", + "MulticlassPrecisionRecallCurve.plot", + "MultilabelPrecisionRecallCurve.plot", + ] + + +class BinaryPrecisionRecallCurve(Metric): + r"""Compute the precision-recall curve for binary tasks. + + The curve consist of multiple pairs of precision and recall values evaluated at different thresholds, such that the + tradeoff between the two values can been seen. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, ...)``. Preds should be a tensor containing + probabilities or logits for each observation. If preds has values outside [0,1] range we consider the input + to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). The value + 1 always encodes the positive class. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``precision`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each class is returned with an 1d + tensor of size ``(n_thresholds+1, )`` with precision values (length may differ between classes). If `thresholds` + is set to something else, then a single 2d tensor of size ``(n_classes, n_thresholds+1)`` with precision values + is returned. + - ``recall`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each class is returned with an 1d tensor + of size ``(n_thresholds+1, )`` with recall values (length may differ between classes). If `thresholds` is set to + something else, then a single 2d tensor of size ``(n_classes, n_thresholds+1)`` with recall values is returned. + - ``thresholds`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each class is returned with an 1d + tensor of size ``(n_thresholds, )`` with increasing threshold values (length may differ between classes). If + `threshold` is set to something else, then a single 1d tensor of size ``(n_thresholds, )`` is returned with + shared threshold values for all classes. + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + Args: + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torchmetrics.classification import BinaryPrecisionRecallCurve + >>> preds = torch.tensor([0, 0.5, 0.7, 0.8]) + >>> target = torch.tensor([0, 1, 1, 0]) + >>> bprc = BinaryPrecisionRecallCurve(thresholds=None) + >>> bprc(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([0.5000, 0.6667, 0.5000, 0.0000, 1.0000]), + tensor([1.0000, 1.0000, 0.5000, 0.0000, 0.0000]), + tensor([0.0000, 0.5000, 0.7000, 0.8000])) + >>> bprc = BinaryPrecisionRecallCurve(thresholds=5) + >>> bprc(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([0.5000, 0.6667, 0.6667, 0.0000, 0.0000, 1.0000]), + tensor([1., 1., 1., 0., 0., 0.]), + tensor([0.0000, 0.2500, 0.5000, 0.7500, 1.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + + preds: List[Tensor] + target: List[Tensor] + confmat: Tensor + + def __init__( + self, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _binary_precision_recall_curve_arg_validation(thresholds, ignore_index) + + self.ignore_index = ignore_index + self.validate_args = validate_args + + thresholds = _adjust_threshold_arg(thresholds) + if thresholds is None: + self.thresholds = thresholds + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + else: + self.register_buffer("thresholds", thresholds, persistent=False) + self.add_state( + "confmat", default=torch.zeros(len(thresholds), 2, 2, dtype=torch.long), dist_reduce_fx="sum" + ) + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric states.""" + if self.validate_args: + _binary_precision_recall_curve_tensor_validation(preds, target, self.ignore_index) + preds, target, _ = _binary_precision_recall_curve_format(preds, target, self.thresholds, self.ignore_index) + state = _binary_precision_recall_curve_update(preds, target, self.thresholds) + if isinstance(state, Tensor): + self.confmat += state + else: + self.preds.append(state[0]) + self.target.append(state[1]) + + def compute(self) -> tuple[Tensor, Tensor, Tensor]: + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _binary_precision_recall_curve_compute(state, self.thresholds) + + def plot( + self, + curve: Optional[tuple[Tensor, Tensor, Tensor]] = None, + score: Optional[Union[Tensor, bool]] = None, + ax: Optional[_AX_TYPE] = None, + ) -> _PLOT_OUT_TYPE: + """Plot a single curve from the metric. + + Args: + curve: the output of either `metric.compute` or `metric.forward`. If no value is provided, will + automatically call `metric.compute` and plot that result. + score: Provide a area-under-the-curve score to be displayed on the plot. If `True` and no curve is provided, + will automatically compute the score. The score is computed by using the trapezoidal rule to compute the + area under the curve. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryPrecisionRecallCurve + >>> preds = rand(20) + >>> target = randint(2, (20,)) + >>> metric = BinaryPrecisionRecallCurve() + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot(score=True) + + """ + curve_computed = curve or self.compute() + # switch order as the standard way is recall along x-axis and precision along y-axis + curve_computed = (curve_computed[1], curve_computed[0], curve_computed[2]) + + score = ( + _auc_compute_without_check(curve_computed[0], curve_computed[1], direction=-1.0) + if not curve and score is True + else None + ) + return plot_curve( + curve_computed, score=score, ax=ax, label_names=("Recall", "Precision"), name=self.__class__.__name__ + ) + + +class MulticlassPrecisionRecallCurve(Metric): + r"""Compute the precision-recall curve for multiclass tasks. + + The curve consist of multiple pairs of precision and recall values evaluated at different thresholds, such that the + tradeoff between the two values can been seen. + + For multiclass the metric is calculated by iteratively treating each class as the positive class and all other + classes as the negative, which is referred to as the one-vs-rest approach. One-vs-one is currently not supported by + this metric. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor containing + probabilities or logits for each observation. If preds has values outside [0,1] range we consider the input to + be logits and will auto apply softmax per sample. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain values in the [0, n_classes-1] range (except if `ignore_index` + is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``precision`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_thresholds+1, )`` with precision values + - ``recall`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_thresholds+1, )`` with recall values + - ``thresholds`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_thresholds, )`` with increasing threshold values + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{classes})` (constant memory). + + Args: + num_classes: Integer specifying the number of classes + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to a 1D `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + average: + If aggregation of curves should be applied. By default, the curves are not aggregated and a curve for + each class is returned. If `average` is set to ``"micro"``, the metric will aggregate the curves by one hot + encoding the targets and flattening the predictions, considering all classes jointly as a binary problem. + If `average` is set to ``"macro"``, the metric will aggregate the curves by first interpolating the curves + from each class at a combined set of thresholds and then average over the classwise interpolated curves. + See `averaging curve objects`_ for more info on the different averaging methods. + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torchmetrics.classification import MulticlassPrecisionRecallCurve + >>> preds = torch.tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = torch.tensor([0, 1, 3, 2]) + >>> mcprc = MulticlassPrecisionRecallCurve(num_classes=5, thresholds=None) + >>> precision, recall, thresholds = mcprc(preds, target) + >>> precision # doctest: +NORMALIZE_WHITESPACE + [tensor([0.2500, 1.0000, 1.0000]), tensor([0.2500, 1.0000, 1.0000]), tensor([0.2500, 0.0000, 1.0000]), + tensor([0.2500, 0.0000, 1.0000]), tensor([0., 1.])] + >>> recall + [tensor([1., 1., 0.]), tensor([1., 1., 0.]), tensor([1., 0., 0.]), tensor([1., 0., 0.]), tensor([nan, 0.])] + >>> thresholds + [tensor([0.0500, 0.7500]), tensor([0.0500, 0.7500]), tensor([0.0500, 0.7500]), tensor([0.0500, 0.7500]), + tensor(0.0500)] + >>> mcprc = MulticlassPrecisionRecallCurve(num_classes=5, thresholds=5) + >>> mcprc(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([[0.2500, 1.0000, 1.0000, 1.0000, 0.0000, 1.0000], + [0.2500, 1.0000, 1.0000, 1.0000, 0.0000, 1.0000], + [0.2500, 0.0000, 0.0000, 0.0000, 0.0000, 1.0000], + [0.2500, 0.0000, 0.0000, 0.0000, 0.0000, 1.0000], + [0.0000, 0.0000, 0.0000, 0.0000, 0.0000, 1.0000]]), + tensor([[1., 1., 1., 1., 0., 0.], + [1., 1., 1., 1., 0., 0.], + [1., 0., 0., 0., 0., 0.], + [1., 0., 0., 0., 0., 0.], + [0., 0., 0., 0., 0., 0.]]), + tensor([0.0000, 0.2500, 0.5000, 0.7500, 1.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + + preds: List[Tensor] + target: List[Tensor] + confmat: Tensor + + def __init__( + self, + num_classes: int, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + average: Optional[Literal["micro", "macro"]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _multiclass_precision_recall_curve_arg_validation(num_classes, thresholds, ignore_index, average) + + self.num_classes = num_classes + self.average = average + self.ignore_index = ignore_index + self.validate_args = validate_args + + thresholds = _adjust_threshold_arg(thresholds) + if thresholds is None: + self.thresholds = thresholds + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + else: + self.register_buffer("thresholds", thresholds, persistent=False) + self.add_state( + "confmat", + default=torch.zeros(len(thresholds), num_classes, 2, 2, dtype=torch.long), + dist_reduce_fx="sum", + ) + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric states.""" + if self.validate_args: + _multiclass_precision_recall_curve_tensor_validation(preds, target, self.num_classes, self.ignore_index) + preds, target, _ = _multiclass_precision_recall_curve_format( + preds, target, self.num_classes, self.thresholds, self.ignore_index, self.average + ) + state = _multiclass_precision_recall_curve_update( + preds, target, self.num_classes, self.thresholds, self.average + ) + if isinstance(state, Tensor): + self.confmat += state + else: + self.preds.append(state[0]) + self.target.append(state[1]) + + def compute(self) -> Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]: + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _multiclass_precision_recall_curve_compute(state, self.num_classes, self.thresholds, self.average) + + def plot( + self, + curve: Optional[Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]] = None, + score: Optional[Union[Tensor, bool]] = None, + ax: Optional[_AX_TYPE] = None, + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + curve: the output of either `metric.compute` or `metric.forward`. If no value is provided, will + automatically call `metric.compute` and plot that result. + score: Provide a area-under-the-curve score to be displayed on the plot. If `True` and no curve is provided, + will automatically compute the score. The score is computed by using the trapezoidal rule to compute the + area under the curve. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randn, randint + >>> from torchmetrics.classification import MulticlassPrecisionRecallCurve + >>> preds = randn(20, 3).softmax(dim=-1) + >>> target = randint(3, (20,)) + >>> metric = MulticlassPrecisionRecallCurve(num_classes=3) + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot(score=True) + + """ + curve_computed = curve or self.compute() + # switch order as the standard way is recall along x-axis and precision along y-axis + curve_computed = (curve_computed[1], curve_computed[0], curve_computed[2]) + score = ( + _reduce_auroc(curve_computed[0], curve_computed[1], average=None, direction=-1.0) + if not curve and score is True + else None + ) + return plot_curve( + curve_computed, score=score, ax=ax, label_names=("Recall", "Precision"), name=self.__class__.__name__ + ) + + +class MultilabelPrecisionRecallCurve(Metric): + r"""Compute the precision-recall curve for multilabel tasks. + + The curve consist of multiple pairs of precision and recall values evaluated at different thresholds, such that the + tradeoff between the two values can been seen. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor containing + probabilities or logits for each observation. If preds has values outside [0,1] range we consider the input to + be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following a tuple of either 3 tensors or + 3 lists containing: + + - ``precision`` (:class:`~torch.Tensor` or :class:`~List`): if `thresholds=None` a list for each label is returned + with an 1d tensor of size ``(n_thresholds+1, )`` with precision values (length may differ between labels). If + `thresholds` is set to something else, then a single 2d tensor of size ``(n_labels, n_thresholds+1)`` with + precision values is returned. + - ``recall`` (:class:`~torch.Tensor` or :class:`~List`): if `thresholds=None` a list for each label is returned + with an 1d tensor of size ``(n_thresholds+1, )`` with recall values (length may differ between labels). If + `thresholds` is set to something else, then a single 2d tensor of size ``(n_labels, n_thresholds+1)`` with recall + values is returned. + - ``thresholds`` (:class:`~torch.Tensor` or :class:`~List`): if `thresholds=None` a list for each label is + returned with an 1d tensor of size ``(n_thresholds, )`` with increasing threshold values (length may differ + between labels). If `threshold` is set to something else, then a single 1d tensor of size ``(n_thresholds, )`` + is returned with shared threshold values for all labels. + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{labels})` (constant memory). + + Args: + preds: Tensor with predictions + target: Tensor with true labels + num_labels: Integer specifying the number of labels + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example: + >>> from torchmetrics.classification import MultilabelPrecisionRecallCurve + >>> preds = torch.tensor([[0.75, 0.05, 0.35], + ... [0.45, 0.75, 0.05], + ... [0.05, 0.55, 0.75], + ... [0.05, 0.65, 0.05]]) + >>> target = torch.tensor([[1, 0, 1], + ... [0, 0, 0], + ... [0, 1, 1], + ... [1, 1, 1]]) + >>> mlprc = MultilabelPrecisionRecallCurve(num_labels=3, thresholds=None) + >>> precision, recall, thresholds = mlprc(preds, target) + >>> precision # doctest: +NORMALIZE_WHITESPACE + [tensor([0.5000, 0.5000, 1.0000, 1.0000]), tensor([0.5000, 0.6667, 0.5000, 0.0000, 1.0000]), + tensor([0.7500, 1.0000, 1.0000, 1.0000])] + >>> recall # doctest: +NORMALIZE_WHITESPACE + [tensor([1.0000, 0.5000, 0.5000, 0.0000]), tensor([1.0000, 1.0000, 0.5000, 0.0000, 0.0000]), + tensor([1.0000, 0.6667, 0.3333, 0.0000])] + >>> thresholds # doctest: +NORMALIZE_WHITESPACE + [tensor([0.0500, 0.4500, 0.7500]), tensor([0.0500, 0.5500, 0.6500, 0.7500]), tensor([0.0500, 0.3500, 0.7500])] + >>> mlprc = MultilabelPrecisionRecallCurve(num_labels=3, thresholds=5) + >>> mlprc(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([[0.5000, 0.5000, 1.0000, 1.0000, 0.0000, 1.0000], + [0.5000, 0.6667, 0.6667, 0.0000, 0.0000, 1.0000], + [0.7500, 1.0000, 1.0000, 1.0000, 0.0000, 1.0000]]), + tensor([[1.0000, 0.5000, 0.5000, 0.5000, 0.0000, 0.0000], + [1.0000, 1.0000, 1.0000, 0.0000, 0.0000, 0.0000], + [1.0000, 0.6667, 0.3333, 0.3333, 0.0000, 0.0000]]), + tensor([0.0000, 0.2500, 0.5000, 0.7500, 1.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + + preds: List[Tensor] + target: List[Tensor] + confmat: Tensor + + def __init__( + self, + num_labels: int, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _multilabel_precision_recall_curve_arg_validation(num_labels, thresholds, ignore_index) + + self.num_labels = num_labels + self.ignore_index = ignore_index + self.validate_args = validate_args + + thresholds = _adjust_threshold_arg(thresholds) + if thresholds is None: + self.thresholds = thresholds + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + else: + self.register_buffer("thresholds", thresholds, persistent=False) + self.add_state( + "confmat", + default=torch.zeros(len(thresholds), num_labels, 2, 2, dtype=torch.long), + dist_reduce_fx="sum", + ) + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric states.""" + if self.validate_args: + _multilabel_precision_recall_curve_tensor_validation(preds, target, self.num_labels, self.ignore_index) + preds, target, _ = _multilabel_precision_recall_curve_format( + preds, target, self.num_labels, self.thresholds, self.ignore_index + ) + state = _multilabel_precision_recall_curve_update(preds, target, self.num_labels, self.thresholds) + if isinstance(state, Tensor): + self.confmat += state + else: + self.preds.append(state[0]) + self.target.append(state[1]) + + def compute(self) -> Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]: + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _multilabel_precision_recall_curve_compute(state, self.num_labels, self.thresholds, self.ignore_index) + + def plot( + self, + curve: Optional[Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]] = None, + score: Optional[Union[Tensor, bool]] = None, + ax: Optional[_AX_TYPE] = None, + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + curve: the output of either `metric.compute` or `metric.forward`. If no value is provided, will + automatically call `metric.compute` and plot that result. + score: Provide a area-under-the-curve score to be displayed on the plot. If `True` and no curve is provided, + will automatically compute the score. The score is computed by using the trapezoidal rule to compute the + area under the curve. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelPrecisionRecallCurve + >>> preds = rand(20, 3) + >>> target = randint(2, (20,3)) + >>> metric = MultilabelPrecisionRecallCurve(num_labels=3) + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot(score=True) + + """ + curve_computed = curve or self.compute() + # switch order as the standard way is recall along x-axis and precision along y-axis + curve_computed = (curve_computed[1], curve_computed[0], curve_computed[2]) + score = ( + _reduce_auroc(curve_computed[0], curve_computed[1], average=None, direction=-1.0) + if not curve and score is True + else None + ) + return plot_curve( + curve_computed, score=score, ax=ax, label_names=("Recall", "Precision"), name=self.__class__.__name__ + ) + + +class PrecisionRecallCurve(_ClassificationTaskWrapper): + r"""Compute the precision-recall curve. + + The curve consist of multiple pairs of precision and recall values evaluated at different thresholds, such that the + tradeoff between the two values can been seen. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryPrecisionRecallCurve`, + :class:`~torchmetrics.classification.MulticlassPrecisionRecallCurve` and + :class:`~torchmetrics.classification.MultilabelPrecisionRecallCurve` for the specific details of each argument + influence and examples. + + Legacy Example: + >>> pred = torch.tensor([0, 0.1, 0.8, 0.4]) + >>> target = torch.tensor([0, 1, 1, 0]) + >>> pr_curve = PrecisionRecallCurve(task="binary") + >>> precision, recall, thresholds = pr_curve(pred, target) + >>> precision + tensor([0.5000, 0.6667, 0.5000, 1.0000, 1.0000]) + >>> recall + tensor([1.0000, 1.0000, 0.5000, 0.5000, 0.0000]) + >>> thresholds + tensor([0.0000, 0.1000, 0.4000, 0.8000]) + + >>> pred = torch.tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = torch.tensor([0, 1, 3, 2]) + >>> pr_curve = PrecisionRecallCurve(task="multiclass", num_classes=5) + >>> precision, recall, thresholds = pr_curve(pred, target) + >>> precision + [tensor([0.2500, 1.0000, 1.0000]), tensor([0.2500, 1.0000, 1.0000]), tensor([0.2500, 0.0000, 1.0000]), + tensor([0.2500, 0.0000, 1.0000]), tensor([0., 1.])] + >>> recall + [tensor([1., 1., 0.]), tensor([1., 1., 0.]), tensor([1., 0., 0.]), tensor([1., 0., 0.]), tensor([nan, 0.])] + >>> thresholds + [tensor([0.0500, 0.7500]), tensor([0.0500, 0.7500]), tensor([0.0500, 0.7500]), tensor([0.0500, 0.7500]), + tensor(0.0500)] + + """ + + def __new__( # type: ignore[misc] + cls: type["PrecisionRecallCurve"], + task: Literal["binary", "multiclass", "multilabel"], + thresholds: Optional[Union[int, list[float], Tensor]] = None, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + kwargs.update({"thresholds": thresholds, "ignore_index": ignore_index, "validate_args": validate_args}) + if task == ClassificationTask.BINARY: + return BinaryPrecisionRecallCurve(**kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassPrecisionRecallCurve(num_classes, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelPrecisionRecallCurve(num_labels, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/ranking.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/ranking.py new file mode 100644 index 0000000000000000000000000000000000000000..9738386aae2486f52d90088346d9fe5f9f41ec65 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/ranking.py @@ -0,0 +1,431 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +import torch +from torch import Tensor + +from torchmetrics.functional.classification.ranking import ( + _multilabel_confusion_matrix_arg_validation, + _multilabel_confusion_matrix_format, + _multilabel_coverage_error_update, + _multilabel_ranking_average_precision_update, + _multilabel_ranking_loss_update, + _multilabel_ranking_tensor_validation, + _ranking_reduce, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "MultilabelCoverageError.plot", + "MultilabelRankingAveragePrecision.plot", + "MultilabelRankingLoss.plot", + ] + + +class MultilabelCoverageError(Metric): + """Compute `Multilabel coverage error`_. + + The score measure how far we need to go through the ranked scores to cover all true labels. The best value is equal + to the average number of labels in the target tensor per sample. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. Target should be a tensor + containing ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlce`` (:class:`~torch.Tensor`): A tensor containing the multilabel coverage error. + + Args: + num_labels: Integer specifying the number of labels + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example: + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelCoverageError + >>> preds = rand(10, 5) + >>> target = randint(2, (10, 5)) + >>> mlce = MultilabelCoverageError(num_labels=5) + >>> mlce(preds, target) + tensor(3.9000) + + """ + + higher_is_better: bool = False + is_differentiable: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _multilabel_confusion_matrix_arg_validation(num_labels, threshold=0.0, ignore_index=ignore_index) + self.validate_args = validate_args + self.num_labels = num_labels + self.ignore_index = ignore_index + self.add_state("measure", torch.tensor(0.0), dist_reduce_fx="sum") + self.add_state("total", torch.tensor(0.0), dist_reduce_fx="sum") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric states.""" + if self.validate_args: + _multilabel_ranking_tensor_validation(preds, target, self.num_labels, self.ignore_index) + preds, target = _multilabel_confusion_matrix_format( + preds, target, self.num_labels, threshold=0.0, ignore_index=self.ignore_index, should_threshold=False + ) + measure, num_elements = _multilabel_coverage_error_update(preds, target) + + if not isinstance(self.measure, Tensor): + raise TypeError(f"Expected 'self.measure' to be of type Tensor, but got {type(self.measure)}.") + if not isinstance(self.total, Tensor): + raise TypeError(f"Expected 'self.total' to be of type Tensor, but got {type(self.total)}.") + + self.measure += measure + self.total += num_elements + + def compute(self) -> Tensor: + """Compute metric.""" + if not isinstance(self.measure, Tensor): + raise TypeError(f"Expected 'self.measure' to be of type Tensor, but got {type(self.measure)}.") + if not isinstance(self.total, Tensor): + raise TypeError(f"Expected 'self.total' to be of type Tensor, but got {type(self.total)}.") + + return _ranking_reduce(self.measure, int(self.total.item())) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelCoverageError + >>> metric = MultilabelCoverageError(num_labels=3) + >>> metric.update(rand(20, 3), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelCoverageError + >>> metric = MultilabelCoverageError(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(20, 3), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelRankingAveragePrecision(Metric): + """Compute label ranking average precision score for multilabel data [1]. + + The score is the average over each ground truth label assigned to each sample of the ratio of true vs. total labels + with lower score. Best score is 1. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. Target should be a tensor + containing ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlrap`` (:class:`~torch.Tensor`): A tensor containing the multilabel ranking average precision. + + Args: + num_labels: Integer specifying the number of labels + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example: + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelRankingAveragePrecision + >>> preds = rand(10, 5) + >>> target = randint(2, (10, 5)) + >>> mlrap = MultilabelRankingAveragePrecision(num_labels=5) + >>> mlrap(preds, target) + tensor(0.7744) + + """ + + higher_is_better: bool = True + is_differentiable: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _multilabel_confusion_matrix_arg_validation(num_labels, threshold=0.0, ignore_index=ignore_index) + self.validate_args = validate_args + self.num_labels = num_labels + self.ignore_index = ignore_index + self.add_state("measure", torch.tensor(0.0), dist_reduce_fx="sum") + self.add_state("total", torch.tensor(0.0), dist_reduce_fx="sum") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric states.""" + if self.validate_args: + _multilabel_ranking_tensor_validation(preds, target, self.num_labels, self.ignore_index) + preds, target = _multilabel_confusion_matrix_format( + preds, target, self.num_labels, threshold=0.0, ignore_index=self.ignore_index, should_threshold=False + ) + if not isinstance(self.measure, Tensor): + raise TypeError(f"Expected 'self.measure' to be of type Tensor, but got {type(self.measure)}.") + if not isinstance(self.total, Tensor): + raise TypeError(f"Expected 'self.total' to be of type Tensor, but got {type(self.total)}.") + + measure, num_elements = _multilabel_ranking_average_precision_update(preds, target) + self.measure += measure + self.total += num_elements + + def compute(self) -> Tensor: + """Compute metric.""" + if not isinstance(self.measure, Tensor): + raise TypeError(f"Expected 'self.measure' to be of type Tensor, but got {type(self.measure)}.") + if not isinstance(self.total, Tensor): + raise TypeError(f"Expected 'self.total' to be of type Tensor, but got {type(self.total)}.") + + return _ranking_reduce(self.measure, int(self.total.item())) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelRankingAveragePrecision + >>> metric = MultilabelRankingAveragePrecision(num_labels=3) + >>> metric.update(rand(20, 3), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelRankingAveragePrecision + >>> metric = MultilabelRankingAveragePrecision(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(20, 3), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelRankingLoss(Metric): + """Compute the label ranking loss for multilabel data [1]. + + The score is corresponds to the average number of label pairs that are incorrectly ordered given some predictions + weighted by the size of the label set and the number of labels not in the label set. The best score is 0. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. Target should be a tensor + containing ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlrl`` (:class:`~torch.Tensor`): A tensor containing the multilabel ranking loss. + + Args: + preds: Tensor with predictions + target: Tensor with true labels + num_labels: Integer specifying the number of labels + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example: + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelRankingLoss + >>> preds = rand(10, 5) + >>> target = randint(2, (10, 5)) + >>> mlrl = MultilabelRankingLoss(num_labels=5) + >>> mlrl(preds, target) + tensor(0.4167) + + """ + + higher_is_better: bool = False + is_differentiable: bool = False + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if validate_args: + _multilabel_confusion_matrix_arg_validation(num_labels, threshold=0.0, ignore_index=ignore_index) + self.validate_args = validate_args + self.num_labels = num_labels + self.ignore_index = ignore_index + self.add_state("measure", torch.tensor(0.0), dist_reduce_fx="sum") + self.add_state("total", torch.tensor(0.0), dist_reduce_fx="sum") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update metric states.""" + if self.validate_args: + _multilabel_ranking_tensor_validation(preds, target, self.num_labels, self.ignore_index) + preds, target = _multilabel_confusion_matrix_format( + preds, target, self.num_labels, threshold=0.0, ignore_index=self.ignore_index, should_threshold=False + ) + if not isinstance(self.measure, Tensor): + raise TypeError(f"Expected 'self.measure' to be of type Tensor, but got {type(self.measure)}.") + if not isinstance(self.total, Tensor): + raise TypeError(f"Expected 'self.total' to be of type Tensor, but got {type(self.total)}.") + + measure, num_elements = _multilabel_ranking_loss_update(preds, target) + self.measure += measure + self.total += num_elements + + def compute(self) -> Tensor: + """Compute metric.""" + if not isinstance(self.measure, Tensor): + raise TypeError(f"Expected 'self.measure' to be of type Tensor, but got {type(self.measure)}.") + if not isinstance(self.total, Tensor): + raise TypeError(f"Expected 'self.total' to be of type Tensor, but got {type(self.total)}.") + + return _ranking_reduce(self.measure, int(self.total.item())) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelRankingLoss + >>> metric = MultilabelRankingLoss(num_labels=3) + >>> metric.update(rand(20, 3), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelRankingLoss + >>> metric = MultilabelRankingLoss(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(20, 3), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/recall_fixed_precision.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/recall_fixed_precision.py new file mode 100644 index 0000000000000000000000000000000000000000..196ec51e6b0382bcf4ce38a185a7b3b6395fc3de --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/recall_fixed_precision.py @@ -0,0 +1,514 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.precision_recall_curve import ( + BinaryPrecisionRecallCurve, + MulticlassPrecisionRecallCurve, + MultilabelPrecisionRecallCurve, +) +from torchmetrics.functional.classification.recall_fixed_precision import ( + _binary_recall_at_fixed_precision_arg_validation, + _binary_recall_at_fixed_precision_compute, + _multiclass_recall_at_fixed_precision_arg_compute, + _multiclass_recall_at_fixed_precision_arg_validation, + _multilabel_recall_at_fixed_precision_arg_compute, + _multilabel_recall_at_fixed_precision_arg_validation, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinaryRecallAtFixedPrecision.plot", + "MulticlassRecallAtFixedPrecision.plot", + "MultilabelRecallAtFixedPrecision.plot", + ] + + +class BinaryRecallAtFixedPrecision(BinaryPrecisionRecallCurve): + r"""Compute the highest possible recall value given the minimum precision thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the recall for + a given precision level. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, ...)``. Preds should be a tensor containing + probabilities or logits for each observation. If preds has values outside [0,1] range we consider the input + to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). The value + 1 always encodes the positive class. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``recall`` (:class:`~torch.Tensor`): A scalar tensor with the maximum recall for the given precision level + - ``threshold`` (:class:`~torch.Tensor`): A scalar tensor with the corresponding threshold level + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a + binned version that is less accurate but more memory efficient. Setting the `thresholds` argument to ``None`` + will activate the non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting + the `thresholds` argument to either an integer, list or a 1d tensor will use a binned version that uses memory + of size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + Args: + min_precision: float value specifying minimum precision threshold. + thresholds: + Can be one of: + + - If set to ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d :class:`~torch.Tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryRecallAtFixedPrecision + >>> preds = tensor([0, 0.5, 0.7, 0.8]) + >>> target = tensor([0, 1, 1, 0]) + >>> metric = BinaryRecallAtFixedPrecision(min_precision=0.5, thresholds=None) + >>> metric(preds, target) + (tensor(1.), tensor(0.5000)) + >>> metric = BinaryRecallAtFixedPrecision(min_precision=0.5, thresholds=5) + >>> metric(preds, target) + (tensor(1.), tensor(0.5000)) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + min_precision: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(thresholds, ignore_index, validate_args=False, **kwargs) + if validate_args: + _binary_recall_at_fixed_precision_arg_validation(min_precision, thresholds, ignore_index) + self.validate_args = validate_args + self.min_precision = min_precision + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _binary_recall_at_fixed_precision_compute(state, self.thresholds, self.min_precision) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinaryRecallAtFixedPrecision + >>> metric = BinaryRecallAtFixedPrecision(min_precision=0.5) + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() # the returned plot only shows the maximum recall value by default + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinaryRecallAtFixedPrecision + >>> metric = BinaryRecallAtFixedPrecision(min_precision=0.5) + >>> values = [ ] + >>> for _ in range(10): + ... # we index by 0 such that only the maximum recall value is plotted + ... values.append(metric(rand(10), randint(2,(10,)))[0]) + >>> fig_, ax_ = metric.plot(values) + + """ + val = val or self.compute()[0] # by default we select the maximum recall value to plot + return self._plot(val, ax) + + +class MulticlassRecallAtFixedPrecision(MulticlassPrecisionRecallCurve): + r"""Compute the highest possible recall value given the minimum precision thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the recall for + a given precision level. + + For multiclass the metric is calculated by iteratively treating each class as the positive class and all other + classes as the negative, which is referred to as the one-vs-rest approach. One-vs-one is currently not supported by + this metric. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply softmax per sample. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain values in the [0, n_classes-1] range (except if `ignore_index` + is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns a tuple of either 2 tensors or 2 lists containing: + + - ``recall`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the maximum recall for the + given precision level per class + - ``threshold`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the corresponding threshold + level per class + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to ``None`` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{classes})` (constant memory). + + Args: + num_classes: Integer specifying the number of classes + min_precision: float value specifying minimum precision threshold. + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d :class:`~torch.Tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassRecallAtFixedPrecision + >>> preds = tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = tensor([0, 1, 3, 2]) + >>> metric = MulticlassRecallAtFixedPrecision(num_classes=5, min_precision=0.5, thresholds=None) + >>> metric(preds, target) + (tensor([1., 1., 0., 0., 0.]), tensor([7.5000e-01, 7.5000e-01, 1.0000e+06, 1.0000e+06, 1.0000e+06])) + >>> mcrafp = MulticlassRecallAtFixedPrecision(num_classes=5, min_precision=0.5, thresholds=5) + >>> mcrafp(preds, target) + (tensor([1., 1., 0., 0., 0.]), tensor([7.5000e-01, 7.5000e-01, 1.0000e+06, 1.0000e+06, 1.0000e+06])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + min_precision: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_classes=num_classes, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multiclass_recall_at_fixed_precision_arg_validation(num_classes, min_precision, thresholds, ignore_index) + self.validate_args = validate_args + self.min_precision = min_precision + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _multiclass_recall_at_fixed_precision_arg_compute( + state, self.num_classes, self.thresholds, self.min_precision + ) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassRecallAtFixedPrecision + >>> metric = MulticlassRecallAtFixedPrecision(num_classes=3, min_precision=0.5) + >>> metric.update(rand(20, 3).softmax(dim=-1), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() # the returned plot only shows the maximum recall value by default + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassRecallAtFixedPrecision + >>> metric = MulticlassRecallAtFixedPrecision(num_classes=3, min_precision=0.5) + >>> values = [] + >>> for _ in range(20): + ... # we index by 0 such that only the maximum recall value is plotted + ... values.append(metric(rand(20, 3).softmax(dim=-1), randint(3, (20,)))[0]) + >>> fig_, ax_ = metric.plot(values) + + """ + val = val or self.compute()[0] # by default we select the maximum recall value to plot + return self._plot(val, ax) + + +class MultilabelRecallAtFixedPrecision(MultilabelPrecisionRecallCurve): + r"""Compute the highest possible recall value given the minimum precision thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the recall for + a given precision level. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). The value + 1 always encodes the positive class. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns a tuple of either 2 tensors or 2 lists containing: + + - ``recall`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the maximum recall for the + given precision level per class + - ``threshold`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_classes, )`` with the corresponding threshold + level per class + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to ```None``` will activate + the non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the + `thresholds` argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{labels})` (constant memory). + + Args: + num_labels: Integer specifying the number of labels + min_precision: float value specifying minimum precision threshold. + thresholds: + Can be one of: + + - If set to ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d :class:`~torch.Tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelRecallAtFixedPrecision + >>> preds = tensor([[0.75, 0.05, 0.35], + ... [0.45, 0.75, 0.05], + ... [0.05, 0.55, 0.75], + ... [0.05, 0.65, 0.05]]) + >>> target = tensor([[1, 0, 1], + ... [0, 0, 0], + ... [0, 1, 1], + ... [1, 1, 1]]) + >>> metric = MultilabelRecallAtFixedPrecision(num_labels=3, min_precision=0.5, thresholds=None) + >>> metric(preds, target) + (tensor([1., 1., 1.]), tensor([0.0500, 0.5500, 0.0500])) + >>> mlrafp = MultilabelRecallAtFixedPrecision(num_labels=3, min_precision=0.5, thresholds=5) + >>> mlrafp(preds, target) + (tensor([1., 1., 1.]), tensor([0.0000, 0.5000, 0.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + min_precision: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_labels=num_labels, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multilabel_recall_at_fixed_precision_arg_validation(num_labels, min_precision, thresholds, ignore_index) + self.validate_args = validate_args + self.min_precision = min_precision + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (dim_zero_cat(self.preds), dim_zero_cat(self.target)) if self.thresholds is None else self.confmat + return _multilabel_recall_at_fixed_precision_arg_compute( + state, self.num_labels, self.thresholds, self.ignore_index, self.min_precision + ) + + def plot( # type: ignore[override] + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelRecallAtFixedPrecision + >>> metric = MultilabelRecallAtFixedPrecision(num_labels=3, min_precision=0.5) + >>> metric.update(rand(20, 3), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() # the returned plot only shows the maximum recall value by default + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelRecallAtFixedPrecision + >>> metric = MultilabelRecallAtFixedPrecision(num_labels=3, min_precision=0.5) + >>> values = [ ] + >>> for _ in range(10): + ... # we index by 0 such that only the maximum recall value is plotted + ... values.append(metric(rand(20, 3), randint(2, (20, 3)))[0]) + >>> fig_, ax_ = metric.plot(values) + + """ + val = val or self.compute()[0] # by default we select the maximum recall value to plot + return self._plot(val, ax) + + +class RecallAtFixedPrecision(_ClassificationTaskWrapper): + r"""Compute the highest possible recall value given the minimum precision thresholds provided. + + This is done by first calculating the precision-recall curve for different thresholds and the find the recall for + a given precision level. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryRecallAtFixedPrecision`, + :class:`~torchmetrics.classification.MulticlassRecallAtFixedPrecision` and + :class:`~torchmetrics.classification.MultilabelRecallAtFixedPrecision` for the specific details of each argument + influence and examples. + + """ + + def __new__( # type: ignore[misc] + cls: type["RecallAtFixedPrecision"], + task: Literal["binary", "multiclass", "multilabel"], + min_precision: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + if task == ClassificationTask.BINARY: + return BinaryRecallAtFixedPrecision(min_precision, thresholds, ignore_index, validate_args, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassRecallAtFixedPrecision( + num_classes, min_precision, thresholds, ignore_index, validate_args, **kwargs + ) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelRecallAtFixedPrecision( + num_labels, min_precision, thresholds, ignore_index, validate_args, **kwargs + ) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/roc.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/roc.py new file mode 100644 index 0000000000000000000000000000000000000000..5bc1ad1cbfa1ba723656efda1a343f4bc3a253bc --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/roc.py @@ -0,0 +1,596 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Any, List, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.precision_recall_curve import ( + BinaryPrecisionRecallCurve, + MulticlassPrecisionRecallCurve, + MultilabelPrecisionRecallCurve, +) +from torchmetrics.functional.classification.auroc import _reduce_auroc +from torchmetrics.functional.classification.roc import ( + _binary_roc_compute, + _multiclass_roc_compute, + _multilabel_roc_compute, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.compute import _auc_compute_without_check +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE, plot_curve + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["BinaryROC.plot", "MulticlassROC.plot", "MultilabelROC.plot"] + + +class BinaryROC(BinaryPrecisionRecallCurve): + r"""Compute the Receiver Operating Characteristic (ROC) for binary tasks. + + The curve consist of multiple pairs of true positive rate (TPR) and false positive rate (FPR) values evaluated at + different thresholds, such that the tradeoff between the two values can be seen. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, ...)``. Preds should be a tensor containing + probabilities or logits for each observation. If preds has values outside [0,1] range we consider the input + to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). The value + 1 always encodes the positive class. + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns a tuple of 3 tensors containing: + + - ``fpr`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_thresholds+1, )`` with false positive rate values + - ``tpr`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_thresholds+1, )`` with true positive rate values + - ``thresholds`` (:class:`~torch.Tensor`): A 1d tensor of size ``(n_thresholds, )`` with decreasing threshold + values + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a + binned version that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will + activate the non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the + `thresholds` argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + .. attention:: + The outputted thresholds will be in reversed order to ensure that they correspond to both fpr and + tpr which are sorted in reversed order during their calculation, such that they are monotome increasing. + + Args: + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryROC + >>> preds = tensor([0, 0.5, 0.7, 0.8]) + >>> target = tensor([0, 1, 1, 0]) + >>> metric = BinaryROC(thresholds=None) + >>> metric(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([0.0000, 0.5000, 0.5000, 0.5000, 1.0000]), + tensor([0.0000, 0.0000, 0.5000, 1.0000, 1.0000]), + tensor([1.0000, 0.8000, 0.7000, 0.5000, 0.0000])) + >>> broc = BinaryROC(thresholds=5) + >>> broc(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([0.0000, 0.5000, 0.5000, 0.5000, 1.0000]), + tensor([0., 0., 1., 1., 1.]), + tensor([1.0000, 0.7500, 0.5000, 0.2500, 0.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def compute(self) -> tuple[Tensor, Tensor, Tensor]: + """Compute metric.""" + state = [dim_zero_cat(self.preds), dim_zero_cat(self.target)] if self.thresholds is None else self.confmat + return _binary_roc_compute(state, self.thresholds) # type: ignore[arg-type] + + def plot( + self, + curve: Optional[tuple[Tensor, Tensor, Tensor]] = None, + score: Optional[Union[Tensor, bool]] = None, + ax: Optional[_AX_TYPE] = None, + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + curve: the output of either `metric.compute` or `metric.forward`. If no value is provided, will + automatically call `metric.compute` and plot that result. + score: Provide a area-under-the-curve score to be displayed on the plot. If `True` and no curve is provided, + will automatically compute the score. The score is computed by using the trapezoidal rule to compute the + area under the curve. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> from torchmetrics.classification import BinaryROC + >>> preds = rand(20) + >>> target = randint(2, (20,)) + >>> metric = BinaryROC() + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot(score=True) + + """ + curve_computed = curve or self.compute() + score = ( + _auc_compute_without_check(curve_computed[0], curve_computed[1], 1.0) + if not curve and score is True + else None + ) + return plot_curve( + curve_computed, + score=score, + ax=ax, + label_names=("False positive rate", "True positive rate"), + name=self.__class__.__name__, + ) + + +class MulticlassROC(MulticlassPrecisionRecallCurve): + r"""Compute the Receiver Operating Characteristic (ROC) for binary tasks. + + The curve consist of multiple pairs of true positive rate (TPR) and false positive rate (FPR) values evaluated at + different thresholds, such that the tradeoff between the two values can be seen. + + For multiclass the metric is calculated by iteratively treating each class as the positive class and all other + classes as the negative, which is referred to as the one-vs-rest approach. One-vs-one is currently not supported by + this metric. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply softmax per sample. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)``. Target should be a tensor containing + ground truth labels, and therefore only contain values in the [0, n_classes-1] range (except if `ignore_index` + is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns a tuple of either 3 tensors or 3 lists containing + + - ``fpr`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each class is returned with an 1d tensor of + size ``(n_thresholds+1, )`` with false positive rate values (length may differ between classes). If `thresholds` + is set to something else, then a single 2d tensor of size ``(n_classes, n_thresholds+1)`` with false positive rate + values is returned. + - ``tpr`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each class is returned with an 1d tensor of + size ``(n_thresholds+1, )`` with true positive rate values (length may differ between classes). If `thresholds` is + set to something else, then a single 2d tensor of size ``(n_classes, n_thresholds+1)`` with true positive rate + values is returned. + - ``thresholds`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each class is returned with an 1d + tensor of size ``(n_thresholds, )`` with decreasing threshold values (length may differ between classes). If + `threshold` is set to something else, then a single 1d tensor of size ``(n_thresholds, )`` is returned with shared + threshold values for all classes. + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a + binned version that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will + activate the non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the + `thresholds` argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{classes})` (constant memory). + + .. attention:: + Note that outputted thresholds will be in reversed order to ensure that they correspond to both fpr + and tpr which are sorted in reversed order during their calculation, such that they are monotome increasing. + + Args: + num_classes: Integer specifying the number of classes + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + average: + If aggregation of curves should be applied. By default, the curves are not aggregated and a curve for + each class is returned. If `average` is set to ``"micro"``, the metric will aggregate the curves by one hot + encoding the targets and flattening the predictions, considering all classes jointly as a binary problem. + If `average` is set to ``"macro"``, the metric will aggregate the curves by first interpolating the curves + from each class at a combined set of thresholds and then average over the classwise interpolated curves. + See `averaging curve objects`_ for more info on the different averaging methods. + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassROC + >>> preds = tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = tensor([0, 1, 3, 2]) + >>> metric = MulticlassROC(num_classes=5, thresholds=None) + >>> fpr, tpr, thresholds = metric(preds, target) + >>> fpr # doctest: +NORMALIZE_WHITESPACE + [tensor([0., 0., 1.]), tensor([0., 0., 1.]), tensor([0.0000, 0.3333, 1.0000]), + tensor([0.0000, 0.3333, 1.0000]), tensor([0., 1.])] + >>> tpr + [tensor([0., 1., 1.]), tensor([0., 1., 1.]), tensor([0., 0., 1.]), tensor([0., 0., 1.]), tensor([0., 0.])] + >>> thresholds # doctest: +NORMALIZE_WHITESPACE + [tensor([1.0000, 0.7500, 0.0500]), tensor([1.0000, 0.7500, 0.0500]), + tensor([1.0000, 0.7500, 0.0500]), tensor([1.0000, 0.7500, 0.0500]), tensor([1.0000, 0.0500])] + >>> mcroc = MulticlassROC(num_classes=5, thresholds=5) + >>> mcroc(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([[0.0000, 0.0000, 0.0000, 0.0000, 1.0000], + [0.0000, 0.0000, 0.0000, 0.0000, 1.0000], + [0.0000, 0.3333, 0.3333, 0.3333, 1.0000], + [0.0000, 0.3333, 0.3333, 0.3333, 1.0000], + [0.0000, 0.0000, 0.0000, 0.0000, 1.0000]]), + tensor([[0., 1., 1., 1., 1.], + [0., 1., 1., 1., 1.], + [0., 0., 0., 0., 1.], + [0., 0., 0., 0., 1.], + [0., 0., 0., 0., 0.]]), + tensor([1.0000, 0.7500, 0.5000, 0.2500, 0.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def compute(self) -> Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]: + """Compute metric.""" + state = [dim_zero_cat(self.preds), dim_zero_cat(self.target)] if self.thresholds is None else self.confmat + return _multiclass_roc_compute(state, self.num_classes, self.thresholds, self.average) # type: ignore[arg-type] + + def plot( + self, + curve: Optional[Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]] = None, + score: Optional[Union[Tensor, bool]] = None, + ax: Optional[_AX_TYPE] = None, + labels: Optional[list[str]] = None, + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + curve: the output of either `metric.compute` or `metric.forward`. If no value is provided, will + automatically call `metric.compute` and plot that result. + score: Provide a area-under-the-curve score to be displayed on the plot. If `True` and no curve is provided, + will automatically compute the score. The score is computed by using the trapezoidal rule to compute the + area under the curve. + ax: An matplotlib axis object. If provided will add plot to that axis + labels: a list of strings, if provided will be added to the plot to indicate the different classes + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randn, randint + >>> from torchmetrics.classification import MulticlassROC + >>> preds = randn(20, 3).softmax(dim=-1) + >>> target = randint(3, (20,)) + >>> metric = MulticlassROC(num_classes=3) + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot(score=True) + + """ + curve_computed = curve or self.compute() + score = ( + _reduce_auroc(curve_computed[0], curve_computed[1], average=None) if not curve and score is True else None + ) + return plot_curve( + curve_computed, + score=score, + ax=ax, + label_names=("False positive rate", "True positive rate"), + name=self.__class__.__name__, + labels=labels, + ) + + +class MultilabelROC(MultilabelPrecisionRecallCurve): + r"""Compute the Receiver Operating Characteristic (ROC) for binary tasks. + + The curve consist of multiple pairs of true positive rate (TPR) and false positive rate (FPR) values evaluated at + different thresholds, such that the tradeoff between the two values can be seen. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): A float tensor of shape ``(N, C, ...)``. Preds should be a tensor + containing probabilities or logits for each observation. If preds has values outside [0,1] range we consider + the input to be logits and will auto apply sigmoid per element. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)``. Target should be a tensor + containing ground truth labels, and therefore only contain {0,1} values (except if `ignore_index` is specified). + + .. tip:: + Additional dimension ``...`` will be flattened into the batch dimension. + + As output to ``forward`` and ``compute`` the metric returns a tuple of either 3 tensors or 3 lists containing + + - ``fpr`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each label is returned with an 1d tensor of + size ``(n_thresholds+1, )`` with false positive rate values (length may differ between labels). If `thresholds` is + set to something else, then a single 2d tensor of size ``(n_labels, n_thresholds+1)`` with false positive rate + values is returned. + - ``tpr`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each label is returned with an 1d tensor of + size ``(n_thresholds+1, )`` with true positive rate values (length may differ between labels). If `thresholds` is + set to something else, then a single 2d tensor of size ``(n_labels, n_thresholds+1)`` with true positive rate + values is returned. + - ``thresholds`` (:class:`~torch.Tensor`): if `thresholds=None` a list for each label is returned with an 1d + tensor of size ``(n_thresholds, )`` with decreasing threshold values (length may differ between labels). If + `threshold` is set to something else, then a single 1d tensor of size ``(n_thresholds, )`` is returned with shared + threshold values for all labels. + + .. note:: + The implementation both supports calculating the metric in a non-binned but accurate version and a + binned version that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will + activate the non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the + `thresholds` argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{labels})` (constant memory). + + .. attention:: + The outputted thresholds will be in reversed order to ensure that they correspond to both fpr and tpr + which are sorted in reversed order during their calculation, such that they are monotome increasing. + + Args: + num_labels: Integer specifying the number of labels + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelROC + >>> preds = tensor([[0.75, 0.05, 0.35], + ... [0.45, 0.75, 0.05], + ... [0.05, 0.55, 0.75], + ... [0.05, 0.65, 0.05]]) + >>> target = tensor([[1, 0, 1], + ... [0, 0, 0], + ... [0, 1, 1], + ... [1, 1, 1]]) + >>> metric = MultilabelROC(num_labels=3, thresholds=None) + >>> fpr, tpr, thresholds = metric(preds, target) + >>> fpr # doctest: +NORMALIZE_WHITESPACE + [tensor([0.0000, 0.0000, 0.5000, 1.0000]), + tensor([0.0000, 0.5000, 0.5000, 0.5000, 1.0000]), + tensor([0., 0., 0., 1.])] + >>> tpr # doctest: +NORMALIZE_WHITESPACE + [tensor([0.0000, 0.5000, 0.5000, 1.0000]), + tensor([0.0000, 0.0000, 0.5000, 1.0000, 1.0000]), + tensor([0.0000, 0.3333, 0.6667, 1.0000])] + >>> thresholds # doctest: +NORMALIZE_WHITESPACE + [tensor([1.0000, 0.7500, 0.4500, 0.0500]), + tensor([1.0000, 0.7500, 0.6500, 0.5500, 0.0500]), + tensor([1.0000, 0.7500, 0.3500, 0.0500])] + >>> mlroc = MultilabelROC(num_labels=3, thresholds=5) + >>> mlroc(preds, target) # doctest: +NORMALIZE_WHITESPACE + (tensor([[0.0000, 0.0000, 0.0000, 0.5000, 1.0000], + [0.0000, 0.5000, 0.5000, 0.5000, 1.0000], + [0.0000, 0.0000, 0.0000, 0.0000, 1.0000]]), + tensor([[0.0000, 0.5000, 0.5000, 0.5000, 1.0000], + [0.0000, 0.0000, 1.0000, 1.0000, 1.0000], + [0.0000, 0.3333, 0.3333, 0.6667, 1.0000]]), + tensor([1.0000, 0.7500, 0.5000, 0.2500, 0.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def compute(self) -> Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]: + """Compute metric.""" + state = [dim_zero_cat(self.preds), dim_zero_cat(self.target)] if self.thresholds is None else self.confmat + return _multilabel_roc_compute(state, self.num_labels, self.thresholds, self.ignore_index) # type: ignore[arg-type] + + def plot( + self, + curve: Optional[Union[tuple[Tensor, Tensor, Tensor], tuple[List[Tensor], List[Tensor], List[Tensor]]]] = None, + score: Optional[Union[Tensor, bool]] = None, + ax: Optional[_AX_TYPE] = None, + labels: Optional[list[str]] = None, + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + curve: the output of either `metric.compute` or `metric.forward`. If no value is provided, will + automatically call `metric.compute` and plot that result. + score: Provide a area-under-the-curve score to be displayed on the plot. If `True` and no curve is provided, + will automatically compute the score. The score is computed by using the trapezoidal rule to compute the + area under the curve. + ax: An matplotlib axis object. If provided will add plot to that axis + labels: a list of strings, if provided will be added to the plot to indicate the different classes + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> from torchmetrics.classification import MultilabelROC + >>> preds = rand(20, 3) + >>> target = randint(2, (20,3)) + >>> metric = MultilabelROC(num_labels=3) + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot(score=True) + + """ + curve_computed = curve or self.compute() + score = ( + _reduce_auroc(curve_computed[0], curve_computed[1], average=None) if not curve and score is True else None + ) + return plot_curve( + curve_computed, + score=score, + ax=ax, + label_names=("False positive rate", "True positive rate"), + name=self.__class__.__name__, + labels=labels, + ) + + +class ROC(_ClassificationTaskWrapper): + r"""Compute the Receiver Operating Characteristic (ROC). + + The curve consist of multiple pairs of true positive rate (TPR) and false positive rate (FPR) values evaluated at + different thresholds, such that the tradeoff between the two values can be seen. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryROC`, + :class:`~torchmetrics.classification.MulticlassROC` and + :class:`~torchmetrics.classification.MultilabelROC` for the specific details of each argument + influence and examples. + + Legacy Example: + >>> from torch import tensor + >>> pred = tensor([0.0, 1.0, 2.0, 3.0]) + >>> target = tensor([0, 1, 1, 1]) + >>> roc = ROC(task="binary") + >>> fpr, tpr, thresholds = roc(pred, target) + >>> fpr + tensor([0., 0., 0., 0., 1.]) + >>> tpr + tensor([0.0000, 0.3333, 0.6667, 1.0000, 1.0000]) + >>> thresholds + tensor([1.0000, 0.9526, 0.8808, 0.7311, 0.5000]) + + >>> pred = tensor([[0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05], + ... [0.05, 0.05, 0.05, 0.75]]) + >>> target = tensor([0, 1, 3, 2]) + >>> roc = ROC(task="multiclass", num_classes=4) + >>> fpr, tpr, thresholds = roc(pred, target) + >>> fpr + [tensor([0., 0., 1.]), tensor([0., 0., 1.]), tensor([0.0000, 0.3333, 1.0000]), tensor([0.0000, 0.3333, 1.0000])] + >>> tpr + [tensor([0., 1., 1.]), tensor([0., 1., 1.]), tensor([0., 0., 1.]), tensor([0., 0., 1.])] + >>> thresholds # doctest: +NORMALIZE_WHITESPACE + [tensor([1.0000, 0.7500, 0.0500]), + tensor([1.0000, 0.7500, 0.0500]), + tensor([1.0000, 0.7500, 0.0500]), + tensor([1.0000, 0.7500, 0.0500])] + + >>> pred = tensor([[0.8191, 0.3680, 0.1138], + ... [0.3584, 0.7576, 0.1183], + ... [0.2286, 0.3468, 0.1338], + ... [0.8603, 0.0745, 0.1837]]) + >>> target = tensor([[1, 1, 0], [0, 1, 0], [0, 0, 0], [0, 1, 1]]) + >>> roc = ROC(task='multilabel', num_labels=3) + >>> fpr, tpr, thresholds = roc(pred, target) + >>> fpr + [tensor([0.0000, 0.3333, 0.3333, 0.6667, 1.0000]), + tensor([0., 0., 0., 1., 1.]), + tensor([0.0000, 0.0000, 0.3333, 0.6667, 1.0000])] + >>> tpr + [tensor([0., 0., 1., 1., 1.]), + tensor([0.0000, 0.3333, 0.6667, 0.6667, 1.0000]), + tensor([0., 1., 1., 1., 1.])] + >>> thresholds + [tensor([1.0000, 0.8603, 0.8191, 0.3584, 0.2286]), + tensor([1.0000, 0.7576, 0.3680, 0.3468, 0.0745]), + tensor([1.0000, 0.1837, 0.1338, 0.1183, 0.1138])] + + """ + + def __new__( # type: ignore[misc] + cls: type["ROC"], + task: Literal["binary", "multiclass", "multilabel"], + thresholds: Optional[Union[int, list[float], Tensor]] = None, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + kwargs.update({"thresholds": thresholds, "ignore_index": ignore_index, "validate_args": validate_args}) + if task == ClassificationTask.BINARY: + return BinaryROC(**kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassROC(num_classes, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelROC(num_labels, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/sensitivity_specificity.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/sensitivity_specificity.py new file mode 100644 index 0000000000000000000000000000000000000000..bb1afb87b21b72ac197730c0ebe488769885a843 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/sensitivity_specificity.py @@ -0,0 +1,375 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.precision_recall_curve import ( + BinaryPrecisionRecallCurve, + MulticlassPrecisionRecallCurve, + MultilabelPrecisionRecallCurve, +) +from torchmetrics.functional.classification.sensitivity_specificity import ( + _binary_sensitivity_at_specificity_arg_validation, + _binary_sensitivity_at_specificity_compute, + _multiclass_sensitivity_at_specificity_arg_validation, + _multiclass_sensitivity_at_specificity_compute, + _multilabel_sensitivity_at_specificity_arg_validation, + _multilabel_sensitivity_at_specificity_compute, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat as _cat +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinarySensitivityAtSpecificity.plot", + "MulticlassSensitivityAtSpecificity.plot", + "MultilabelSensitivityAtSpecificity.plot", + ] + + +class BinarySensitivityAtSpecificity(BinaryPrecisionRecallCurve): + r"""Compute the highest possible sensitivity value given the minimum specificity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the sensitivity for a given specificity level. + + Accepts the following input tensors: + + - ``preds`` (float tensor): ``(N, ...)``. Preds should be a tensor containing probabilities or logits for each + observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply + sigmoid per element. + - ``target`` (int tensor): ``(N, ...)``. Target should be a tensor containing ground truth labels, and therefore + only contain {0,1} values (except if `ignore_index` is specified). + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + Args: + min_specificity: float value specifying minimum specificity threshold. + thresholds: + Can be one of: + + - ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. It is the most accurate but also the most memory-consuming approach. + - ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - 1d ``tensor`` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + (tuple): a tuple of 2 tensors containing: + + - sensitivity: an scalar tensor with the maximum sensitivity for the given specificity level + - threshold: an scalar tensor with the corresponding threshold level + + Example: + >>> from torchmetrics.classification import BinarySensitivityAtSpecificity + >>> from torch import tensor + >>> preds = tensor([0, 0.5, 0.4, 0.1]) + >>> target = tensor([0, 1, 1, 1]) + >>> metric = BinarySensitivityAtSpecificity(min_specificity=0.5, thresholds=None) + >>> metric(preds, target) + (tensor(1.), tensor(0.1000)) + >>> metric = BinarySensitivityAtSpecificity(min_specificity=0.5, thresholds=5) + >>> metric(preds, target) + (tensor(0.6667), tensor(0.2500)) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + min_specificity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(thresholds, ignore_index, validate_args=False, **kwargs) + if validate_args: + _binary_sensitivity_at_specificity_arg_validation(min_specificity, thresholds, ignore_index) + self.validate_args = validate_args + self.min_specificity = min_specificity + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (_cat(self.preds), _cat(self.target)) if self.thresholds is None else self.confmat + return _binary_sensitivity_at_specificity_compute(state, self.thresholds, self.min_specificity) + + +class MulticlassSensitivityAtSpecificity(MulticlassPrecisionRecallCurve): + r"""Compute the highest possible sensitivity value given the minimum specificity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the sensitivity for a given specificity level. + + For multiclass the metric is calculated by iteratively treating each class as the positive class and all other + classes as the negative, which is referred to as the one-vs-rest approach. One-vs-one is currently not supported by + this metric. + + Accepts the following input tensors: + + - ``preds`` (float tensor): ``(N, C, ...)``. Preds should be a tensor containing probabilities or logits for each + observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply + softmax per sample. + - ``target`` (int tensor): ``(N, ...)``. Target should be a tensor containing ground truth labels, and therefore + only contain values in the [0, n_classes-1] range (except if `ignore_index` is specified). + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{classes})` (constant memory). + + Args: + num_classes: Integer specifying the number of classes + min_specificity: float value specifying minimum specificity threshold. + thresholds: + Can be one of: + + - ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. It is the most accurate but also the most memory-consuming approach. + - ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - 1d ``tensor`` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + (tuple): a tuple of either 2 tensors or 2 lists containing + + - sensitivity: an 1d tensor of size (n_classes, ) with the maximum sensitivity for the given + specificity level per class + - thresholds: an 1d tensor of size (n_classes, ) with the corresponding threshold level per class + + + Example: + >>> from torchmetrics.classification import MulticlassSensitivityAtSpecificity + >>> from torch import tensor + >>> preds = tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = tensor([0, 1, 3, 2]) + >>> metric = MulticlassSensitivityAtSpecificity(num_classes=5, min_specificity=0.5, thresholds=None) + >>> metric(preds, target) + (tensor([1., 1., 0., 0., 0.]), tensor([0.7500, 0.7500, 1.0000, 1.0000, 1.0000])) + >>> metric = MulticlassSensitivityAtSpecificity(num_classes=5, min_specificity=0.5, thresholds=5) + >>> metric(preds, target) + (tensor([1., 1., 0., 0., 0.]), tensor([0.7500, 0.7500, 1.0000, 1.0000, 1.0000])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + min_specificity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_classes=num_classes, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multiclass_sensitivity_at_specificity_arg_validation( + num_classes, min_specificity, thresholds, ignore_index + ) + self.validate_args = validate_args + self.min_specificity = min_specificity + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (_cat(self.preds), _cat(self.target)) if self.thresholds is None else self.confmat + return _multiclass_sensitivity_at_specificity_compute( + state, self.num_classes, self.thresholds, self.min_specificity + ) + + +class MultilabelSensitivityAtSpecificity(MultilabelPrecisionRecallCurve): + r"""Compute the highest possible sensitivity value given the minimum specificity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the sensitivity for a given specificity level. + + Accepts the following input tensors: + + - ``preds`` (float tensor): ``(N, C, ...)``. Preds should be a tensor containing probabilities or logits for each + observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply + sigmoid per element. + - ``target`` (int tensor): ``(N, C, ...)``. Target should be a tensor containing ground truth labels, and therefore + only contain {0,1} values (except if `ignore_index` is specified). + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{labels})` (constant memory). + + Args: + num_labels: Integer specifying the number of labels + min_specificity: float value specifying minimum specificity threshold. + thresholds: + Can be one of: + + - ``None``, will use a non-binned approach where thresholds are dynamically calculated from + all the data. It is the most accurate but also the most memory-consuming approach. + - ``int`` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - ``list`` of floats, will use the indicated thresholds in the list as bins for the calculation + - 1d ``tensor`` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + (tuple): a tuple of either 2 tensors or 2 lists containing + + - sensitivity: an 1d tensor of size ``(n_classes, )`` with the maximum sensitivity for the given + specificity level per class + - thresholds: an 1d tensor of size ``(n_classes, )`` with the corresponding threshold level per class + + Example: + >>> from torchmetrics.classification import MultilabelSensitivityAtSpecificity + >>> from torch import tensor + >>> preds = tensor([[0.75, 0.05, 0.35], + ... [0.45, 0.75, 0.05], + ... [0.05, 0.55, 0.75], + ... [0.05, 0.65, 0.05]]) + >>> target = tensor([[1, 0, 1], + ... [0, 0, 0], + ... [0, 1, 1], + ... [1, 1, 1]]) + >>> metric = MultilabelSensitivityAtSpecificity(num_labels=3, min_specificity=0.5, thresholds=None) + >>> metric(preds, target) + (tensor([0.5000, 1.0000, 0.6667]), tensor([0.7500, 0.5500, 0.3500])) + >>> metric = MultilabelSensitivityAtSpecificity(num_labels=3, min_specificity=0.5, thresholds=5) + >>> metric(preds, target) + (tensor([0.5000, 1.0000, 0.6667]), tensor([0.7500, 0.5000, 0.2500])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + min_specificity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_labels=num_labels, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multilabel_sensitivity_at_specificity_arg_validation(num_labels, min_specificity, thresholds, ignore_index) + self.validate_args = validate_args + self.min_specificity = min_specificity + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (_cat(self.preds), _cat(self.target)) if self.thresholds is None else self.confmat + return _multilabel_sensitivity_at_specificity_compute( + state, self.num_labels, self.thresholds, self.ignore_index, self.min_specificity + ) + + +class SensitivityAtSpecificity(_ClassificationTaskWrapper): + r"""Compute the highest possible sensitivity value given the minimum specificity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the sensitivity for a given specificity level. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinarySensitivityAtSpecificity`, + :class:`~torchmetrics.classification.MulticlassSensitivityAtSpecificity` and + :class:`~torchmetrics.classification.MultilabelSensitivityAtSpecificity` for the specific details of each argument + influence and examples. + + """ + + def __new__( # type: ignore[misc] + cls: type["SensitivityAtSpecificity"], + task: Literal["binary", "multiclass", "multilabel"], + min_specificity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + if task == ClassificationTask.BINARY: + return BinarySensitivityAtSpecificity(min_specificity, thresholds, ignore_index, validate_args, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassSensitivityAtSpecificity( + num_classes, min_specificity, thresholds, ignore_index, validate_args, **kwargs + ) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelSensitivityAtSpecificity( + num_labels, min_specificity, thresholds, ignore_index, validate_args, **kwargs + ) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/specificity.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/specificity.py new file mode 100644 index 0000000000000000000000000000000000000000..742ca10134db20faf5d8957ad938593525a7fe39 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/specificity.py @@ -0,0 +1,513 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.stat_scores import BinaryStatScores, MulticlassStatScores, MultilabelStatScores +from torchmetrics.functional.classification.specificity import _specificity_reduce +from torchmetrics.metric import Metric +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["BinarySpecificity.plot", "MulticlassSpecificity.plot", "MultilabelSpecificity.plot"] + + +class BinarySpecificity(BinaryStatScores): + r"""Compute `Specificity`_ for binary tasks. + + .. math:: \text{Specificity} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TN} + \text{FP} \neq 0`. If this case is + encountered a score of 0 is returned. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, ...)``. If preds is a floating point + tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid per + element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bs`` (:class:`~torch.Tensor`): If ``multidim_average`` is set to ``global``, the metric returns a scalar value. + If ``multidim_average`` is set to ``samplewise``, the metric returns ``(N,)`` vector consisting of a scalar value + per sample. + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinarySpecificity + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinarySpecificity() + >>> metric(preds, target) + tensor(0.6667) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinarySpecificity + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinarySpecificity() + >>> metric(preds, target) + tensor(0.6667) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinarySpecificity + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinarySpecificity(multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.0000, 0.3333]) + + """ + + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _specificity_reduce(tp, fp, tn, fn, average="binary", multidim_average=self.multidim_average) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import BinarySpecificity + >>> metric = BinarySpecificity() + >>> metric.update(rand(10), randint(2,(10,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import BinarySpecificity + >>> metric = BinarySpecificity() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(rand(10), randint(2,(10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MulticlassSpecificity(MulticlassStatScores): + r"""Compute `Specificity`_ for multiclass tasks. + + .. math:: \text{Specificity} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TN} + \text{FP} \neq 0`. If this case is + encountered for any class, the metric for that class will be set to 0 and the overall metric may therefore be + affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcs`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + top_k: + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassSpecificity + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassSpecificity(num_classes=3) + >>> metric(preds, target) + tensor(0.8889) + >>> mcs = MulticlassSpecificity(num_classes=3, average=None) + >>> mcs(preds, target) + tensor([1.0000, 0.6667, 1.0000]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassSpecificity + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassSpecificity(num_classes=3) + >>> metric(preds, target) + tensor(0.8889) + >>> mcs = MulticlassSpecificity(num_classes=3, average=None) + >>> mcs(preds, target) + tensor([1.0000, 0.6667, 1.0000]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassSpecificity + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassSpecificity(num_classes=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.7500, 0.6556]) + >>> mcs = MulticlassSpecificity(num_classes=3, multidim_average='samplewise', average=None) + >>> mcs(preds, target) + tensor([[0.7500, 0.7500, 0.7500], + [0.8000, 0.6667, 0.5000]]) + + """ + + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _specificity_reduce(tp, fp, tn, fn, average=self.average, multidim_average=self.multidim_average) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a single value per class + >>> from torchmetrics.classification import MulticlassSpecificity + >>> metric = MulticlassSpecificity(num_classes=3, average=None) + >>> metric.update(randint(3, (20,)), randint(3, (20,))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import randint + >>> # Example plotting a multiple values per class + >>> from torchmetrics.classification import MulticlassSpecificity + >>> metric = MulticlassSpecificity(num_classes=3, average=None) + >>> values = [] + >>> for _ in range(20): + ... values.append(metric(randint(3, (20,)), randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class MultilabelSpecificity(MultilabelStatScores): + r"""Compute `Specificity`_ for multilabel tasks. + + .. math:: \text{Specificity} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TN} + \text{FP} \neq 0`. If this case is + encountered for any label, the metric for that label will be set to 0 and the overall metric may therefore be + affected in turn. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, C, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mls`` (:class:`~torch.Tensor`): The returned shape depends on the ``average`` and ``multidim_average`` + arguments: + + - If ``multidim_average`` is set to ``global`` + + - If ``average='micro'/'macro'/'weighted'``, the output will be a scalar tensor + - If ``average=None/'none'``, the shape will be ``(C,)`` + + - If ``multidim_average`` is set to ``samplewise`` + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N,)`` + - If ``average=None/'none'``, the shape will be ``(N, C)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelSpecificity + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelSpecificity(num_labels=3) + >>> metric(preds, target) + tensor(0.6667) + >>> mls = MultilabelSpecificity(num_labels=3, average=None) + >>> mls(preds, target) + tensor([1., 1., 0.]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelSpecificity + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelSpecificity(num_labels=3) + >>> metric(preds, target) + tensor(0.6667) + >>> mls = MultilabelSpecificity(num_labels=3, average=None) + >>> mls(preds, target) + tensor([1., 1., 0.]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelSpecificity + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelSpecificity(num_labels=3, multidim_average='samplewise') + >>> metric(preds, target) + tensor([0.0000, 0.3333]) + >>> mls = MultilabelSpecificity(num_labels=3, multidim_average='samplewise', average=None) + >>> mls(preds, target) + tensor([[0., 0., 0.], + [0., 0., 1.]]) + + """ + + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def compute(self) -> Tensor: + """Compute metric.""" + tp, fp, tn, fn = self._final_state() + return _specificity_reduce( + tp, fp, tn, fn, average=self.average, multidim_average=self.multidim_average, multilabel=True + ) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting a single value + >>> from torchmetrics.classification import MultilabelSpecificity + >>> metric = MultilabelSpecificity(num_labels=3) + >>> metric.update(randint(2, (20, 3)), randint(2, (20, 3))) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> from torch import rand, randint + >>> # Example plotting multiple values + >>> from torchmetrics.classification import MultilabelSpecificity + >>> metric = MultilabelSpecificity(num_labels=3) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(randint(2, (20, 3)), randint(2, (20, 3)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class Specificity(_ClassificationTaskWrapper): + r"""Compute `Specificity`_. + + .. math:: \text{Specificity} = \frac{\text{TN}}{\text{TN} + \text{FP}} + + Where :math:`\text{TN}` and :math:`\text{FP}` represent the number of true negatives and false positives + respectively. The metric is only proper defined when :math:`\text{TN} + \text{FP} \neq 0`. If this case is + encountered for any class/label, the metric for that class/label will be set to 0 and the overall metric may + therefore be affected in turn. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinarySpecificity`, :class:`~torchmetrics.classification.MulticlassSpecificity` + and :class:`~torchmetrics.classification.MultilabelSpecificity` for the specific details of each argument influence + and examples. + + Legacy Example: + >>> from torch import tensor + >>> preds = tensor([2, 0, 2, 1]) + >>> target = tensor([1, 1, 2, 0]) + >>> specificity = Specificity(task="multiclass", average='macro', num_classes=3) + >>> specificity(preds, target) + tensor(0.6111) + >>> specificity = Specificity(task="multiclass", average='micro', num_classes=3) + >>> specificity(preds, target) + tensor(0.6250) + + """ + + def __new__( # type: ignore[misc] + cls: type["Specificity"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + }) + if task == ClassificationTask.BINARY: + return BinarySpecificity(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassSpecificity(num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelSpecificity(num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/specificity_sensitivity.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/specificity_sensitivity.py new file mode 100644 index 0000000000000000000000000000000000000000..2c4fe2709dbe063d9623f0963789ed675705e875 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/specificity_sensitivity.py @@ -0,0 +1,375 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Any, Optional, Union + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.classification.precision_recall_curve import ( + BinaryPrecisionRecallCurve, + MulticlassPrecisionRecallCurve, + MultilabelPrecisionRecallCurve, +) +from torchmetrics.functional.classification.specificity_sensitivity import ( + _binary_specificity_at_sensitivity_arg_validation, + _binary_specificity_at_sensitivity_compute, + _multiclass_specificity_at_sensitivity_arg_validation, + _multiclass_specificity_at_sensitivity_compute, + _multilabel_specificity_at_sensitivity_arg_validation, + _multilabel_specificity_at_sensitivity_compute, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat as _cat +from torchmetrics.utilities.enums import ClassificationTask +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = [ + "BinarySpecificityAtSensitivity.plot", + "MulticlassSpecificityAtSensitivity.plot", + "MultilabelSpecificityAtSensitivity.plot", + ] + + +class BinarySpecificityAtSensitivity(BinaryPrecisionRecallCurve): + r"""Compute the highest possible specificity value given the minimum sensitivity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the specificity for a given sensitivity level. + + Accepts the following input tensors: + + - ``preds`` (float tensor): ``(N, ...)``. Preds should be a tensor containing probabilities or logits for each + observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply + sigmoid per element. + - ``target`` (int tensor): ``(N, ...)``. Target should be a tensor containing ground truth labels, and therefore + only contain {0,1} values (except if `ignore_index` is specified). + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds})` (constant memory). + + Args: + min_sensitivity: float value specifying minimum sensitivity threshold. + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + (tuple): a tuple of 2 tensors containing: + + - specificity: an scalar tensor with the maximum specificity for the given sensitivity level + - threshold: an scalar tensor with the corresponding threshold level + + Example: + >>> from torchmetrics.classification import BinarySpecificityAtSensitivity + >>> from torch import tensor + >>> preds = tensor([0, 0.5, 0.4, 0.1]) + >>> target = tensor([0, 1, 1, 1]) + >>> metric = BinarySpecificityAtSensitivity(min_sensitivity=0.5, thresholds=None) + >>> metric(preds, target) + (tensor(1.), tensor(0.4000)) + >>> metric = BinarySpecificityAtSensitivity(min_sensitivity=0.5, thresholds=5) + >>> metric(preds, target) + (tensor(1.), tensor(0.2500)) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + def __init__( + self, + min_sensitivity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(thresholds, ignore_index, validate_args=False, **kwargs) + if validate_args: + _binary_specificity_at_sensitivity_arg_validation(min_sensitivity, thresholds, ignore_index) + self.validate_args = validate_args + self.min_sensitivity = min_sensitivity + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (_cat(self.preds), _cat(self.target)) if self.thresholds is None else self.confmat + return _binary_specificity_at_sensitivity_compute(state, self.thresholds, self.min_sensitivity) + + +class MulticlassSpecificityAtSensitivity(MulticlassPrecisionRecallCurve): + r"""Compute the highest possible specificity value given the minimum sensitivity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the specificity for a given sensitivity level. + + For multiclass the metric is calculated by iteratively treating each class as the positive class and all other + classes as the negative, which is referred to as the one-vs-rest approach. One-vs-one is currently not supported by + this metric. + + Accepts the following input tensors: + + - ``preds`` (float tensor): ``(N, C, ...)``. Preds should be a tensor containing probabilities or logits for each + observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply + softmax per sample. + - ``target`` (int tensor): ``(N, ...)``. Target should be a tensor containing ground truth labels, and therefore + only contain values in the [0, n_classes-1] range (except if `ignore_index` is specified). + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{classes})` (constant memory). + + Args: + num_classes: Integer specifying the number of classes + min_sensitivity: float value specifying minimum sensitivity threshold. + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + (tuple): a tuple of either 2 tensors or 2 lists containing + + - specificity: an 1d tensor of size (n_classes, ) with the maximum specificity for the given + sensitivity level per class + - thresholds: an 1d tensor of size (n_classes, ) with the corresponding threshold level per class + + + Example: + >>> from torchmetrics.classification import MulticlassSpecificityAtSensitivity + >>> from torch import tensor + >>> preds = tensor([[0.75, 0.05, 0.05, 0.05, 0.05], + ... [0.05, 0.75, 0.05, 0.05, 0.05], + ... [0.05, 0.05, 0.75, 0.05, 0.05], + ... [0.05, 0.05, 0.05, 0.75, 0.05]]) + >>> target = tensor([0, 1, 3, 2]) + >>> metric = MulticlassSpecificityAtSensitivity(num_classes=5, min_sensitivity=0.5, thresholds=None) + >>> metric(preds, target) + (tensor([1., 1., 0., 0., 0.]), tensor([7.5000e-01, 7.5000e-01, 5.0000e-02, 5.0000e-02, 1.0000e+06])) + >>> metric = MulticlassSpecificityAtSensitivity(num_classes=5, min_sensitivity=0.5, thresholds=5) + >>> metric(preds, target) + (tensor([1., 1., 0., 0., 0.]), tensor([7.5000e-01, 7.5000e-01, 0.0000e+00, 0.0000e+00, 1.0000e+06])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Class" + + def __init__( + self, + num_classes: int, + min_sensitivity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_classes=num_classes, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multiclass_specificity_at_sensitivity_arg_validation( + num_classes, min_sensitivity, thresholds, ignore_index + ) + self.validate_args = validate_args + self.min_sensitivity = min_sensitivity + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (_cat(self.preds), _cat(self.target)) if self.thresholds is None else self.confmat + return _multiclass_specificity_at_sensitivity_compute( + state, self.num_classes, self.thresholds, self.min_sensitivity + ) + + +class MultilabelSpecificityAtSensitivity(MultilabelPrecisionRecallCurve): + r"""Compute the highest possible specificity value given the minimum sensitivity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the specificity for a given sensitivity level. + + Accepts the following input tensors: + + - ``preds`` (float tensor): ``(N, C, ...)``. Preds should be a tensor containing probabilities or logits for each + observation. If preds has values outside [0,1] range we consider the input to be logits and will auto apply + sigmoid per element. + - ``target`` (int tensor): ``(N, C, ...)``. Target should be a tensor containing ground truth labels, and therefore + only contain {0,1} values (except if `ignore_index` is specified). + + Additional dimension ``...`` will be flattened into the batch dimension. + + The implementation both supports calculating the metric in a non-binned but accurate version and a binned version + that is less accurate but more memory efficient. Setting the `thresholds` argument to `None` will activate the + non-binned version that uses memory of size :math:`\mathcal{O}(n_{samples})` whereas setting the `thresholds` + argument to either an integer, list or a 1d tensor will use a binned version that uses memory of + size :math:`\mathcal{O}(n_{thresholds} \times n_{labels})` (constant memory). + + Args: + num_labels: Integer specifying the number of labels + min_sensitivity: float value specifying minimum sensitivity threshold. + thresholds: + Can be one of: + + - If set to `None`, will use a non-binned approach where thresholds are dynamically calculated from + all the data. Most accurate but also most memory consuming approach. + - If set to an `int` (larger than 1), will use that number of thresholds linearly spaced from + 0 to 1 as bins for the calculation. + - If set to an `list` of floats, will use the indicated thresholds in the list as bins for the calculation + - If set to an 1d `tensor` of floats, will use the indicated thresholds in the tensor as + bins for the calculation. + + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Returns: + (tuple): a tuple of either 2 tensors or 2 lists containing + + - specificity: an 1d tensor of size (n_classes, ) with the maximum specificity for the given + sensitivity level per class + - thresholds: an 1d tensor of size (n_classes, ) with the corresponding threshold level per class + + Example: + >>> from torchmetrics.classification import MultilabelSpecificityAtSensitivity + >>> from torch import tensor + >>> preds = tensor([[0.75, 0.05, 0.35], + ... [0.45, 0.75, 0.05], + ... [0.05, 0.55, 0.75], + ... [0.05, 0.65, 0.05]]) + >>> target = tensor([[1, 0, 1], + ... [0, 0, 0], + ... [0, 1, 1], + ... [1, 1, 1]]) + >>> metric = MultilabelSpecificityAtSensitivity(num_labels=3, min_sensitivity=0.5, thresholds=None) + >>> metric(preds, target) + (tensor([1.0000, 0.5000, 1.0000]), tensor([0.7500, 0.6500, 0.3500])) + >>> metric = MultilabelSpecificityAtSensitivity(num_labels=3, min_sensitivity=0.5, thresholds=5) + >>> metric(preds, target) + (tensor([1.0000, 0.5000, 1.0000]), tensor([0.7500, 0.5000, 0.2500])) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + plot_legend_name: str = "Label" + + def __init__( + self, + num_labels: int, + min_sensitivity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + super().__init__( + num_labels=num_labels, thresholds=thresholds, ignore_index=ignore_index, validate_args=False, **kwargs + ) + if validate_args: + _multilabel_specificity_at_sensitivity_arg_validation(num_labels, min_sensitivity, thresholds, ignore_index) + self.validate_args = validate_args + self.min_sensitivity = min_sensitivity + + def compute(self) -> tuple[Tensor, Tensor]: # type: ignore[override] + """Compute metric.""" + state = (_cat(self.preds), _cat(self.target)) if self.thresholds is None else self.confmat + return _multilabel_specificity_at_sensitivity_compute( + state, self.num_labels, self.thresholds, self.ignore_index, self.min_sensitivity + ) + + +class SpecificityAtSensitivity(_ClassificationTaskWrapper): + r"""Compute the highest possible specificity value given the minimum sensitivity thresholds provided. + + This is done by first calculating the Receiver Operating Characteristic (ROC) curve for different thresholds and the + find the specificity for a given sensitivity level. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinarySpecificityAtSensitivity`, + :class:`~torchmetrics.classification.MulticlassSpecificityAtSensitivity` and + :class:`~torchmetrics.classification.MultilabelSpecificityAtSensitivity` for the specific details of each argument + influence and examples. + + """ + + def __new__( # type: ignore[misc] + cls: type["SpecificityAtSensitivity"], + task: Literal["binary", "multiclass", "multilabel"], + min_sensitivity: float, + thresholds: Optional[Union[int, list[float], Tensor]] = None, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + if task == ClassificationTask.BINARY: + return BinarySpecificityAtSensitivity(min_sensitivity, thresholds, ignore_index, validate_args, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + return MulticlassSpecificityAtSensitivity( + num_classes, min_sensitivity, thresholds, ignore_index, validate_args, **kwargs + ) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelSpecificityAtSensitivity( + num_labels, min_sensitivity, thresholds, ignore_index, validate_args, **kwargs + ) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/classification/stat_scores.py b/rtme/lib/python3.10/site-packages/torchmetrics/classification/stat_scores.py new file mode 100644 index 0000000000000000000000000000000000000000..d54ae5edf89b61f5c5515f737c6728a4c0d64247 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/classification/stat_scores.py @@ -0,0 +1,562 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Any, Callable, List, Optional, Union + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.classification.base import _ClassificationTaskWrapper +from torchmetrics.functional.classification.stat_scores import ( + _binary_stat_scores_arg_validation, + _binary_stat_scores_compute, + _binary_stat_scores_format, + _binary_stat_scores_tensor_validation, + _binary_stat_scores_update, + _multiclass_stat_scores_arg_validation, + _multiclass_stat_scores_compute, + _multiclass_stat_scores_format, + _multiclass_stat_scores_tensor_validation, + _multiclass_stat_scores_update, + _multilabel_stat_scores_arg_validation, + _multilabel_stat_scores_compute, + _multilabel_stat_scores_format, + _multilabel_stat_scores_tensor_validation, + _multilabel_stat_scores_update, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.enums import ClassificationTask + + +class _AbstractStatScores(Metric): + tp: Union[List[Tensor], Tensor] + fp: Union[List[Tensor], Tensor] + tn: Union[List[Tensor], Tensor] + fn: Union[List[Tensor], Tensor] + + # define common functions + def _create_state( + self, + size: int, + multidim_average: Literal["global", "samplewise"] = "global", + ) -> None: + """Initialize the states for the different statistics.""" + default: Union[Callable[[], list], Callable[[], Tensor]] + if multidim_average == "samplewise": + default = list + dist_reduce_fx = "cat" + else: + default = lambda: torch.zeros(size, dtype=torch.long) + dist_reduce_fx = "sum" + + self.add_state("tp", default(), dist_reduce_fx=dist_reduce_fx) + self.add_state("fp", default(), dist_reduce_fx=dist_reduce_fx) + self.add_state("tn", default(), dist_reduce_fx=dist_reduce_fx) + self.add_state("fn", default(), dist_reduce_fx=dist_reduce_fx) + + def _update_state(self, tp: Tensor, fp: Tensor, tn: Tensor, fn: Tensor) -> None: + """Update states depending on multidim_average argument.""" + if self.multidim_average == "samplewise": + self.tp.append(tp) # type: ignore[union-attr] + self.fp.append(fp) # type: ignore[union-attr] + self.tn.append(tn) # type: ignore[union-attr] + self.fn.append(fn) # type: ignore[union-attr] + else: + self.tp = self.tp + tp if not isinstance(self.tp, list) else [*self.tp, tp] + self.fp = self.fp + fp if not isinstance(self.fp, list) else [*self.fp, fp] + self.tn = self.tn + tn if not isinstance(self.tn, list) else [*self.tn, tn] + self.fn = self.fn + fn if not isinstance(self.fn, list) else [*self.fn, fn] + + def _final_state(self) -> tuple[Tensor, Tensor, Tensor, Tensor]: + """Aggregate states that are lists and return final states.""" + tp = dim_zero_cat(self.tp) + fp = dim_zero_cat(self.fp) + tn = dim_zero_cat(self.tn) + fn = dim_zero_cat(self.fn) + return tp, fp, tn, fn + + +class BinaryStatScores(_AbstractStatScores): + r"""Compute true positives, false positives, true negatives, false negatives and the support for binary tasks. + + Related to `Type I and Type II errors`_. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``bss`` (:class:`~torch.Tensor`): A tensor of shape ``(..., 5)``, where the last dimension corresponds + to ``[tp, fp, tn, fn, sup]`` (``sup`` stands for support and equals ``tp + fn``). The shape + depends on the ``multidim_average`` parameter: + + - If ``multidim_average`` is set to ``global``, the shape will be ``(5,)`` + - If ``multidim_average`` is set to ``samplewise``, the shape will be ``(N, 5)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + threshold: Threshold for transforming probability to binary {0,1} predictions + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import BinaryStatScores + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0, 0, 1, 1, 0, 1]) + >>> metric = BinaryStatScores() + >>> metric(preds, target) + tensor([2, 1, 2, 1, 3]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import BinaryStatScores + >>> target = tensor([0, 1, 0, 1, 0, 1]) + >>> preds = tensor([0.11, 0.22, 0.84, 0.73, 0.33, 0.92]) + >>> metric = BinaryStatScores() + >>> metric(preds, target) + tensor([2, 1, 2, 1, 3]) + + Example (multidim tensors): + >>> from torchmetrics.classification import BinaryStatScores + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = BinaryStatScores(multidim_average='samplewise') + >>> metric(preds, target) + tensor([[2, 3, 0, 1, 3], + [0, 2, 1, 3, 3]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + + def __init__( + self, + threshold: float = 0.5, + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + zero_division = kwargs.pop("zero_division", 0) + super(_AbstractStatScores, self).__init__(**kwargs) + if validate_args: + _binary_stat_scores_arg_validation(threshold, multidim_average, ignore_index, zero_division) + self.threshold = threshold + self.multidim_average = multidim_average + self.ignore_index = ignore_index + self.validate_args = validate_args + self.zero_division = zero_division + + self._create_state(size=1, multidim_average=multidim_average) + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + if self.validate_args: + _binary_stat_scores_tensor_validation(preds, target, self.multidim_average, self.ignore_index) + preds, target = _binary_stat_scores_format(preds, target, self.threshold, self.ignore_index) + tp, fp, tn, fn = _binary_stat_scores_update(preds, target, self.multidim_average) + self._update_state(tp, fp, tn, fn) + + def compute(self) -> Tensor: + """Compute the final statistics.""" + tp, fp, tn, fn = self._final_state() + return _binary_stat_scores_compute(tp, fp, tn, fn, self.multidim_average) + + +class MulticlassStatScores(_AbstractStatScores): + r"""Computes true positives, false positives, true negatives, false negatives and the support for multiclass tasks. + + Related to `Type I and Type II errors`_. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` or float tensor of shape ``(N, C, ..)``. + If preds is a floating point we apply ``torch.argmax`` along the ``C`` dimension to automatically convert + probabilities/logits into an int tensor. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, ...)`` + + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mcss`` (:class:`~torch.Tensor`): A tensor of shape ``(..., 5)``, where the last dimension corresponds + to ``[tp, fp, tn, fn, sup]`` (``sup`` stands for support and equals ``tp + fn``). The shape + depends on ``average`` and ``multidim_average`` parameters: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(5,)`` + - If ``average=None/'none'``, the shape will be ``(C, 5)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N, 5)`` + - If ``average=None/'none'``, the shape will be ``(N, C, 5)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_classes: Integer specifying the number of classes + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + top_k: + Number of highest probability or logit score predictions considered to find the correct label. + Only works when ``preds`` contain probabilities/logits. + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MulticlassStatScores + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([2, 1, 0, 1]) + >>> metric = MulticlassStatScores(num_classes=3, average='micro') + >>> metric(preds, target) + tensor([3, 1, 7, 1, 4]) + >>> mcss = MulticlassStatScores(num_classes=3, average=None) + >>> mcss(preds, target) + tensor([[1, 0, 2, 1, 2], + [1, 1, 2, 0, 1], + [1, 0, 3, 0, 1]]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MulticlassStatScores + >>> target = tensor([2, 1, 0, 0]) + >>> preds = tensor([[0.16, 0.26, 0.58], + ... [0.22, 0.61, 0.17], + ... [0.71, 0.09, 0.20], + ... [0.05, 0.82, 0.13]]) + >>> metric = MulticlassStatScores(num_classes=3, average='micro') + >>> metric(preds, target) + tensor([3, 1, 7, 1, 4]) + >>> mcss = MulticlassStatScores(num_classes=3, average=None) + >>> mcss(preds, target) + tensor([[1, 0, 2, 1, 2], + [1, 1, 2, 0, 1], + [1, 0, 3, 0, 1]]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MulticlassStatScores + >>> target = tensor([[[0, 1], [2, 1], [0, 2]], [[1, 1], [2, 0], [1, 2]]]) + >>> preds = tensor([[[0, 2], [2, 0], [0, 1]], [[2, 2], [2, 1], [1, 0]]]) + >>> metric = MulticlassStatScores(num_classes=3, multidim_average="samplewise", average='micro') + >>> metric(preds, target) + tensor([[3, 3, 9, 3, 6], + [2, 4, 8, 4, 6]]) + >>> mcss = MulticlassStatScores(num_classes=3, multidim_average="samplewise", average=None) + >>> mcss(preds, target) + tensor([[[2, 1, 3, 0, 2], + [0, 1, 3, 2, 2], + [1, 1, 3, 1, 2]], + [[0, 1, 4, 1, 1], + [1, 1, 2, 2, 3], + [1, 2, 2, 1, 2]]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + + def __init__( + self, + num_classes: Optional[int] = None, + top_k: int = 1, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + zero_division = kwargs.pop("zero_division", 0) + super(_AbstractStatScores, self).__init__(**kwargs) + if validate_args: + _multiclass_stat_scores_arg_validation( + num_classes, top_k, average, multidim_average, ignore_index, zero_division + ) + self.num_classes = num_classes + self.top_k = top_k + self.average = average + self.multidim_average = multidim_average + self.ignore_index = ignore_index + self.validate_args = validate_args + self.zero_division = zero_division + + self._create_state( + size=1 if (average == "micro" and top_k == 1) else (num_classes or 1), multidim_average=multidim_average + ) + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + if self.validate_args: + _multiclass_stat_scores_tensor_validation( + preds, target, self.num_classes, self.multidim_average, self.ignore_index + ) + preds, target = _multiclass_stat_scores_format(preds, target, self.top_k) + num_classes = self.num_classes if self.num_classes is not None else 1 + tp, fp, tn, fn = _multiclass_stat_scores_update( + preds, target, num_classes, self.top_k, self.average, self.multidim_average, self.ignore_index + ) + self._update_state(tp, fp, tn, fn) + + def compute(self) -> Tensor: + """Compute the final statistics.""" + tp, fp, tn, fn = self._final_state() + return _multiclass_stat_scores_compute(tp, fp, tn, fn, self.average, self.multidim_average) + + +class MultilabelStatScores(_AbstractStatScores): + r"""Compute true positives, false positives, true negatives, false negatives and the support for multilabel tasks. + + Related to `Type I and Type II errors`_. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): An int or float tensor of shape ``(N, C, ...)``. If preds is a floating + point tensor with values outside [0,1] range we consider the input to be logits and will auto apply sigmoid + per element. Additionally, we convert to int tensor with thresholding using the value in ``threshold``. + - ``target`` (:class:`~torch.Tensor`): An int tensor of shape ``(N, C, ...)`` + + As output to ``forward`` and ``compute`` the metric returns the following output: + + - ``mlss`` (:class:`~torch.Tensor`): A tensor of shape ``(..., 5)``, where the last dimension corresponds + to ``[tp, fp, tn, fn, sup]`` (``sup`` stands for support and equals ``tp + fn``). The shape + depends on ``average`` and ``multidim_average`` parameters: + + - If ``multidim_average`` is set to ``global``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(5,)`` + - If ``average=None/'none'``, the shape will be ``(C, 5)`` + + - If ``multidim_average`` is set to ``samplewise``: + + - If ``average='micro'/'macro'/'weighted'``, the shape will be ``(N, 5)`` + - If ``average=None/'none'``, the shape will be ``(N, C, 5)`` + + If ``multidim_average`` is set to ``samplewise`` we expect at least one additional dimension ``...`` to be present, + which the reduction will then be applied over instead of the sample dimension ``N``. + + Args: + num_labels: Integer specifying the number of labels + threshold: Threshold for transforming probability to binary (0,1) predictions + average: + Defines the reduction that is applied over labels. Should be one of the following: + + - ``micro``: Sum statistics over all labels + - ``macro``: Calculate statistics for each label and average them + - ``weighted``: calculates statistics for each label and computes weighted average using their support + - ``"none"`` or ``None``: calculates statistic for each label and applies no reduction + + multidim_average: + Defines how additionally dimensions ``...`` should be handled. Should be one of the following: + + - ``global``: Additional dimensions are flatted along the batch dimension + - ``samplewise``: Statistic will be calculated independently for each sample on the ``N`` axis. + The statistics in this case are calculated over the additional dimensions. + + ignore_index: + Specifies a target value that is ignored and does not contribute to the metric calculation + validate_args: bool indicating if input arguments and tensors should be validated for correctness. + Set to ``False`` for faster computations. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example (preds is int tensor): + >>> from torch import tensor + >>> from torchmetrics.classification import MultilabelStatScores + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0, 0, 1], [1, 0, 1]]) + >>> metric = MultilabelStatScores(num_labels=3, average='micro') + >>> metric(preds, target) + tensor([2, 1, 2, 1, 3]) + >>> mlss = MultilabelStatScores(num_labels=3, average=None) + >>> mlss(preds, target) + tensor([[1, 0, 1, 0, 1], + [0, 0, 1, 1, 1], + [1, 1, 0, 0, 1]]) + + Example (preds is float tensor): + >>> from torchmetrics.classification import MultilabelStatScores + >>> target = tensor([[0, 1, 0], [1, 0, 1]]) + >>> preds = tensor([[0.11, 0.22, 0.84], [0.73, 0.33, 0.92]]) + >>> metric = MultilabelStatScores(num_labels=3, average='micro') + >>> metric(preds, target) + tensor([2, 1, 2, 1, 3]) + >>> mlss = MultilabelStatScores(num_labels=3, average=None) + >>> mlss(preds, target) + tensor([[1, 0, 1, 0, 1], + [0, 0, 1, 1, 1], + [1, 1, 0, 0, 1]]) + + Example (multidim tensors): + >>> from torchmetrics.classification import MultilabelStatScores + >>> target = tensor([[[0, 1], [1, 0], [0, 1]], [[1, 1], [0, 0], [1, 0]]]) + >>> preds = tensor([[[0.59, 0.91], [0.91, 0.99], [0.63, 0.04]], + ... [[0.38, 0.04], [0.86, 0.780], [0.45, 0.37]]]) + >>> metric = MultilabelStatScores(num_labels=3, multidim_average='samplewise', average='micro') + >>> metric(preds, target) + tensor([[2, 3, 0, 1, 3], + [0, 2, 1, 3, 3]]) + >>> mlss = MultilabelStatScores(num_labels=3, multidim_average='samplewise', average=None) + >>> mlss(preds, target) + tensor([[[1, 1, 0, 0, 1], + [1, 1, 0, 0, 1], + [0, 1, 0, 1, 1]], + [[0, 0, 0, 2, 2], + [0, 2, 0, 0, 0], + [0, 0, 1, 1, 1]]]) + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = None + full_state_update: bool = False + + def __init__( + self, + num_labels: int, + threshold: float = 0.5, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "macro", + multidim_average: Literal["global", "samplewise"] = "global", + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> None: + zero_division = kwargs.pop("zero_division", 0) + super(_AbstractStatScores, self).__init__(**kwargs) + if validate_args: + _multilabel_stat_scores_arg_validation( + num_labels, threshold, average, multidim_average, ignore_index, zero_division + ) + self.num_labels = num_labels + self.threshold = threshold + self.average = average + self.multidim_average = multidim_average + self.ignore_index = ignore_index + self.validate_args = validate_args + self.zero_division = zero_division + + self._create_state(size=num_labels, multidim_average=multidim_average) + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + if self.validate_args: + _multilabel_stat_scores_tensor_validation( + preds, target, self.num_labels, self.multidim_average, self.ignore_index + ) + preds, target = _multilabel_stat_scores_format( + preds, target, self.num_labels, self.threshold, self.ignore_index + ) + tp, fp, tn, fn = _multilabel_stat_scores_update(preds, target, self.multidim_average) + self._update_state(tp, fp, tn, fn) + + def compute(self) -> Tensor: + """Compute the final statistics.""" + tp, fp, tn, fn = self._final_state() + return _multilabel_stat_scores_compute(tp, fp, tn, fn, self.average, self.multidim_average) + + +class StatScores(_ClassificationTaskWrapper): + r"""Compute the number of true positives, false positives, true negatives, false negatives and the support. + + This function is a simple wrapper to get the task specific versions of this metric, which is done by setting the + ``task`` argument to either ``'binary'``, ``'multiclass'`` or ``'multilabel'``. See the documentation of + :class:`~torchmetrics.classification.BinaryStatScores`, :class:`~torchmetrics.classification.MulticlassStatScores` + and :class:`~torchmetrics.classification.MultilabelStatScores` for the specific details of each argument influence + and examples. + + Legacy Example: + >>> from torch import tensor + >>> preds = tensor([1, 0, 2, 1]) + >>> target = tensor([1, 1, 2, 0]) + >>> stat_scores = StatScores(task="multiclass", num_classes=3, average='micro') + >>> stat_scores(preds, target) + tensor([2, 2, 6, 2, 4]) + >>> stat_scores = StatScores(task="multiclass", num_classes=3, average=None) + >>> stat_scores(preds, target) + tensor([[0, 1, 2, 1, 1], + [1, 1, 1, 1, 2], + [1, 0, 3, 0, 1]]) + + """ + + def __new__( # type: ignore[misc] + cls: type["StatScores"], + task: Literal["binary", "multiclass", "multilabel"], + threshold: float = 0.5, + num_classes: Optional[int] = None, + num_labels: Optional[int] = None, + average: Optional[Literal["micro", "macro", "weighted", "none"]] = "micro", + multidim_average: Optional[Literal["global", "samplewise"]] = "global", + top_k: Optional[int] = 1, + ignore_index: Optional[int] = None, + validate_args: bool = True, + **kwargs: Any, + ) -> Metric: + """Initialize task metric.""" + task = ClassificationTask.from_str(task) + assert multidim_average is not None # noqa: S101 # needed for mypy + kwargs.update({ + "multidim_average": multidim_average, + "ignore_index": ignore_index, + "validate_args": validate_args, + }) + if task == ClassificationTask.BINARY: + return BinaryStatScores(threshold, **kwargs) + if task == ClassificationTask.MULTICLASS: + if not isinstance(num_classes, int): + raise ValueError(f"`num_classes` is expected to be `int` but `{type(num_classes)} was passed.`") + if not isinstance(top_k, int): + raise ValueError(f"`top_k` is expected to be `int` but `{type(top_k)} was passed.`") + return MulticlassStatScores(num_classes, top_k, average, **kwargs) + if task == ClassificationTask.MULTILABEL: + if not isinstance(num_labels, int): + raise ValueError(f"`num_labels` is expected to be `int` but `{type(num_labels)} was passed.`") + return MultilabelStatScores(num_labels, threshold, average, **kwargs) + raise ValueError(f"Task {task} not supported!") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/__init__.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a72414792300029904c11846cd97539a63bb0e88 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/__init__.py @@ -0,0 +1,44 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from torchmetrics.clustering.adjusted_mutual_info_score import AdjustedMutualInfoScore +from torchmetrics.clustering.adjusted_rand_score import AdjustedRandScore +from torchmetrics.clustering.calinski_harabasz_score import CalinskiHarabaszScore +from torchmetrics.clustering.cluster_accuracy import ClusterAccuracy +from torchmetrics.clustering.davies_bouldin_score import DaviesBouldinScore +from torchmetrics.clustering.dunn_index import DunnIndex +from torchmetrics.clustering.fowlkes_mallows_index import FowlkesMallowsIndex +from torchmetrics.clustering.homogeneity_completeness_v_measure import ( + CompletenessScore, + HomogeneityScore, + VMeasureScore, +) +from torchmetrics.clustering.mutual_info_score import MutualInfoScore +from torchmetrics.clustering.normalized_mutual_info_score import NormalizedMutualInfoScore +from torchmetrics.clustering.rand_score import RandScore + +__all__ = [ + "AdjustedMutualInfoScore", + "AdjustedRandScore", + "CalinskiHarabaszScore", + "ClusterAccuracy", + "CompletenessScore", + "DaviesBouldinScore", + "DunnIndex", + "FowlkesMallowsIndex", + "HomogeneityScore", + "MutualInfoScore", + "NormalizedMutualInfoScore", + "RandScore", + "VMeasureScore", +] diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/adjusted_rand_score.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/adjusted_rand_score.py new file mode 100644 index 0000000000000000000000000000000000000000..20278f74bc3e80ad8b7fd111505a2e0362add08c --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/adjusted_rand_score.py @@ -0,0 +1,127 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering.adjusted_rand_score import adjusted_rand_score +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["AdjustedRandScore.plot"] + + +class AdjustedRandScore(Metric): + r"""Compute `Adjusted Rand Score`_ (also known as Adjusted Rand Index). + + .. math:: + ARS(U, V) = (\text{RS} - \text{Expected RS}) / (\text{Max RS} - \text{Expected RS}) + + The adjusted rand score :math:`\text{ARS}` is in essence the :math:`\text{RS}` (rand score) adjusted for chance. + The score ensures that completely randomly cluster labels have a score close to zero and only a perfect match will + have a score of 1 (up to a permutation of the labels). The adjusted rand score is symmetric, therefore swapping + :math:`U` and :math:`V` yields the same adjusted rand score. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering is generally used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``adj_rand_score`` (:class:`~torch.Tensor`): Scalar tensor with the adjusted rand score + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> import torch + >>> from torchmetrics.clustering import AdjustedRandScore + >>> metric = AdjustedRandScore() + >>> metric(torch.tensor([0, 0, 1, 1]), torch.tensor([0, 0, 1, 1])) + tensor(1.) + >>> metric(torch.tensor([0, 0, 1, 1]), torch.tensor([0, 1, 0, 1])) + tensor(-0.5000) + + """ + + is_differentiable = True + higher_is_better = None + full_state_update: bool = False + plot_lower_bound: float = -0.5 + plot_upper_bound: float = 1.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + self.preds.append(preds) + self.target.append(target) + + def compute(self) -> Tensor: + """Compute mutual information over state.""" + return adjusted_rand_score(dim_zero_cat(self.preds), dim_zero_cat(self.target)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import AdjustedRandScore + >>> metric = AdjustedRandScore() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import AdjustedRandScore + >>> metric = AdjustedRandScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/calinski_harabasz_score.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/calinski_harabasz_score.py new file mode 100644 index 0000000000000000000000000000000000000000..c331fba7866e33925fadee5abfb9072b820d3fe7 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/calinski_harabasz_score.py @@ -0,0 +1,128 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering.calinski_harabasz_score import calinski_harabasz_score +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["CalinskiHarabaszScore.plot"] + + +class CalinskiHarabaszScore(Metric): + r"""Compute Calinski Harabasz Score (also known as variance ratio criterion) for clustering algorithms. + + .. math:: + CHS(X, L) = \frac{B(X, L) \cdot (n_\text{samples} - n_\text{labels})}{W(X, L) \cdot (n_\text{labels} - 1)} + + where :math:`B(X, L)` is the between-cluster dispersion, which is the squared distance between the cluster centers + and the dataset mean, weighted by the size of the clusters, :math:`n_\text{samples}` is the number of samples, + :math:`n_\text{labels}` is the number of labels, and :math:`W(X, L)` is the within-cluster dispersion e.g. the + sum of squared distances between each samples and its closest cluster center. + + This clustering metric is an intrinsic measure, because it does not rely on ground truth labels for the evaluation. + Instead it examines how well the clusters are separated from each other. The score is higher when clusters are dense + and well separated, which relates to a standard concept of a cluster. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``data`` (:class:`~torch.Tensor`): float tensor with shape ``(N,d)`` with the embedded data. ``d`` is the + dimensionality of the embedding space. + - ``labels`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``chs`` (:class:`~torch.Tensor`): A tensor with the Calinski Harabasz Score + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> from torch import randn, randint + >>> from torchmetrics.clustering import CalinskiHarabaszScore + >>> data = randn(20, 3) + >>> labels = randint(3, (20,)) + >>> metric = CalinskiHarabaszScore() + >>> metric(data, labels) + tensor(2.2128) + + """ + + is_differentiable: bool = True + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + data: List[Tensor] + labels: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("data", default=[], dist_reduce_fx="cat") + self.add_state("labels", default=[], dist_reduce_fx="cat") + + def update(self, data: Tensor, labels: Tensor) -> None: + """Update metric state with new data and labels.""" + self.data.append(data) + self.labels.append(labels) + + def compute(self) -> Tensor: + """Compute the Calinski Harabasz Score over all data and labels.""" + return calinski_harabasz_score(dim_zero_cat(self.data), dim_zero_cat(self.labels)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import CalinskiHarabaszScore + >>> metric = CalinskiHarabaszScore() + >>> metric.update(torch.randn(20, 3), torch.randint(3, (20,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import CalinskiHarabaszScore + >>> metric = CalinskiHarabaszScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randn(20, 3), torch.randint(3, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/cluster_accuracy.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/cluster_accuracy.py new file mode 100644 index 0000000000000000000000000000000000000000..bdb42f5e7c10e772127fbe2d8d5773681fdeeed2 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/cluster_accuracy.py @@ -0,0 +1,148 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Any, Optional, Sequence, Union + +import torch +from torch import Tensor + +from torchmetrics.functional.classification import multiclass_confusion_matrix +from torchmetrics.functional.clustering.cluster_accuracy import _cluster_accuracy_compute +from torchmetrics.metric import Metric +from torchmetrics.utilities.imports import ( + _MATPLOTLIB_AVAILABLE, + _TORCH_LINEAR_ASSIGNMENT_AVAILABLE, +) +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["ClusterAccuracy.plot"] + +if not _TORCH_LINEAR_ASSIGNMENT_AVAILABLE: + __doctest_skip__ = ["ClusterAccuracy", "ClusterAccuracy.plot"] + + +class ClusterAccuracy(Metric): + r"""Compute `Cluster Accuracy`_ between predicted and target clusters. + + .. math:: + + \text{Cluster Accuracy} = \max_g \frac{1}{N} \sum_{n=1}^N \mathbb{1}_{g(p_n) = t_n} + + Where :math:`g` is a function that maps predicted clusters :math:`p` to target clusters :math:`t`, :math:`N` is the + number of samples, :math:`p_n` is the predicted cluster for sample :math:`n`, :math:`t_n` is the target cluster for + sample :math:`n`, and :math:`\mathbb{1}` is the indicator function. The function :math:`g` is determined by solving + the linear sum assignment problem. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering in generally is used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``acc_score`` (:class:`~torch.Tensor`): A tensor with the Cluster Accuracy score + + Args: + num_classes: number of classes + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Raises: + RuntimeError: + If ``torch_linear_assignment`` is not installed. To install, run ``pip install torchmetrics[clustering]``. + ValueError + If ``num_classes`` is not a positive integer + + Example:: + >>> import torch + >>> from torchmetrics.clustering import ClusterAccuracy + >>> preds = torch.tensor([0, 0, 1, 1]) + >>> target = torch.tensor([1, 1, 0, 0]) + >>> metric = ClusterAccuracy(num_classes=2) + >>> metric(preds, target) + tensor(1.) + + """ + + is_differentiable: bool = False + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + confmat: Tensor + + def __init__(self, num_classes: int, **kwargs: Any) -> None: + super().__init__(**kwargs) + if not _TORCH_LINEAR_ASSIGNMENT_AVAILABLE: + raise RuntimeError( + "Missing `torch_linear_assignment`. Please install it with `pip install torchmetrics[clustering]`." + ) + + if not isinstance(num_classes, int) or num_classes <= 0: + raise ValueError("Argument `num_classes` should be a positive integer") + self.add_state( + "confmat", default=torch.zeros((num_classes, num_classes), dtype=torch.int64), dist_reduce_fx="sum" + ) + self.num_classes = num_classes + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update the confusion matrix with the new predictions and targets.""" + self.confmat += multiclass_confusion_matrix(preds, target, num_classes=self.num_classes) + + def compute(self) -> Tensor: + """Computes the clustering accuracy.""" + return _cluster_accuracy_compute(self.confmat) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling ``metric.forward`` or ``metric.compute`` + or a list of these results. If no value is provided, will automatically call `metric.compute` + and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import ClusterAccuracy + >>> metric = ClusterAccuracy(num_classes=4) + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import ClusterAccuracy + >>> metric = ClusterAccuracy(num_classes=4) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/davies_bouldin_score.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/davies_bouldin_score.py new file mode 100644 index 0000000000000000000000000000000000000000..5547c81796e30e801df27d2b4756ffd3944fa238 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/davies_bouldin_score.py @@ -0,0 +1,138 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering.davies_bouldin_score import davies_bouldin_score +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["DaviesBouldinScore.plot"] + + +class DaviesBouldinScore(Metric): + r"""Compute `Davies-Bouldin Score`_ for clustering algorithms. + + Given the following quantities: + + .. math:: + S_i = \left( \frac{1}{T_i} \sum_{j=1}^{T_i} ||X_j - A_i||^2_2 \right)^{1/2} + + where :math:`T_i` is the number of samples in cluster :math:`i`, :math:`X_j` is the :math:`j`-th sample in cluster + :math:`i`, and :math:`A_i` is the centroid of cluster :math:`i`. This quantity is the average distance between all + the samples in cluster :math:`i` and its centroid. Let + + .. math:: + M_{i,j} = ||A_i - A_j||_2 + + e.g. the distance between the centroids of cluster :math:`i` and cluster :math:`j`. Then the Davies-Bouldin score + is defined as: + + .. math:: + DB = \frac{1}{n_{clusters}} \sum_{i=1}^{n_{clusters}} \max_{j \neq i} \left( \frac{S_i + S_j}{M_{i,j}} \right) + + This clustering metric is an intrinsic measure, because it does not rely on ground truth labels for the evaluation. + Instead it examines how well the clusters are separated from each other. The score is higher when clusters are dense + and well separated, which relates to a standard concept of a cluster. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``data`` (:class:`~torch.Tensor`): float tensor with shape ``(N,d)`` with the embedded data. ``d`` is the + dimensionality of the embedding space. + - ``labels`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``chs`` (:class:`~torch.Tensor`): A tensor with the Calinski Harabasz Score + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> from torch import randn, randint + >>> from torchmetrics.clustering import DaviesBouldinScore + >>> data = randn(10, 3) + >>> labels = randint(3, (10,)) + >>> metric = DaviesBouldinScore() + >>> metric(data, labels) + tensor(1.2540) + + """ + + is_differentiable: bool = True + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + data: List[Tensor] + labels: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("data", default=[], dist_reduce_fx="cat") + self.add_state("labels", default=[], dist_reduce_fx="cat") + + def update(self, data: Tensor, labels: Tensor) -> None: + """Update metric state with new data and labels.""" + self.data.append(data) + self.labels.append(labels) + + def compute(self) -> Tensor: + """Compute the Davies Bouldin Score over all data and labels.""" + return davies_bouldin_score(dim_zero_cat(self.data), dim_zero_cat(self.labels)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import DaviesBouldinScore + >>> metric = DaviesBouldinScore() + >>> metric.update(torch.randn(20, 3), torch.randint(0, 2, (20,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import DaviesBouldinScore + >>> metric = DaviesBouldinScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randn(20, 3), torch.randint(0, 2, (20,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/dunn_index.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/dunn_index.py new file mode 100644 index 0000000000000000000000000000000000000000..65d1c0c9a9499f6c22e4d4b284e3bedad0186985 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/dunn_index.py @@ -0,0 +1,129 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering.dunn_index import dunn_index +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["DunnIndex.plot"] + + +class DunnIndex(Metric): + r"""Compute `Dunn Index`_. + + .. math:: + DI_m = \frac{\min_{1\leq i>> import torch + >>> from torchmetrics.clustering import DunnIndex + >>> data = torch.tensor([[0, 0], [0.5, 0], [1, 0], [0.5, 1]]) + >>> labels = torch.tensor([0, 0, 0, 1]) + >>> dunn_index = DunnIndex(p=2) + >>> dunn_index(data, labels) + tensor(2.) + + """ + + is_differentiable: bool = True + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + data: List[Tensor] + labels: List[Tensor] + + def __init__(self, p: float = 2, **kwargs: Any) -> None: + super().__init__(**kwargs) + self.p = p + + self.add_state("data", default=[], dist_reduce_fx="cat") + self.add_state("labels", default=[], dist_reduce_fx="cat") + + def update(self, data: Tensor, labels: Tensor) -> None: + """Update state with predictions and targets.""" + self.data.append(data) + self.labels.append(labels) + + def compute(self) -> Tensor: + """Compute mutual information over state.""" + return dunn_index(dim_zero_cat(self.data), dim_zero_cat(self.labels), self.p) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import DunnIndex + >>> data = torch.tensor([[0, 0], [0.5, 0], [1, 0], [0.5, 1]]) + >>> labels = torch.tensor([0, 0, 0, 1]) + >>> metric = DunnIndex(p=2) + >>> metric.update(data, labels) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import DunnIndex + >>> metric = DunnIndex(p=2) + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randn(50, 3), torch.randint(0, 2, (50,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/fowlkes_mallows_index.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/fowlkes_mallows_index.py new file mode 100644 index 0000000000000000000000000000000000000000..bc18a76f0f635f7e8efcc7b8fcceb4946feb29e2 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/fowlkes_mallows_index.py @@ -0,0 +1,122 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering import fowlkes_mallows_index +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["FowlkesMallowsIndex.plot"] + + +class FowlkesMallowsIndex(Metric): + r"""Compute `Fowlkes-Mallows Index`_. + + .. math:: + FMI(U,V) = \frac{TP}{\sqrt{(TP + FP) * (TP + FN)}} + + Where :math:`TP` is the number of true positives, :math:`FP` is the number of false positives, and :math:`FN` is + the number of false negatives. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``fmi`` (:class:`~torch.Tensor`): A tensor with the Fowlkes-Mallows index. + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> import torch + >>> from torchmetrics.clustering import FowlkesMallowsIndex + >>> preds = torch.tensor([2, 2, 0, 1, 0]) + >>> target = torch.tensor([2, 2, 1, 1, 0]) + >>> fmi = FowlkesMallowsIndex() + >>> fmi(preds, target) + tensor(0.5000) + + """ + + is_differentiable: bool = True + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + self.preds.append(preds) + self.target.append(target) + + def compute(self) -> Tensor: + """Compute Fowlkes-Mallows index over state.""" + return fowlkes_mallows_index(dim_zero_cat(self.preds), dim_zero_cat(self.target)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import FowlkesMallowsIndex + >>> metric = FowlkesMallowsIndex() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import FowlkesMallowsIndex + >>> metric = FowlkesMallowsIndex() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/homogeneity_completeness_v_measure.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/homogeneity_completeness_v_measure.py new file mode 100644 index 0000000000000000000000000000000000000000..260ab52224519d3a499c24019929db1a3552ef43 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/homogeneity_completeness_v_measure.py @@ -0,0 +1,329 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering.homogeneity_completeness_v_measure import ( + completeness_score, + homogeneity_score, + v_measure_score, +) +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["HomogeneityScore.plot", "CompletenessScore.plot", "VMeasureScore.plot"] + + +class HomogeneityScore(Metric): + r"""Compute `Homogeneity Score`_. + + The homogeneity score is a metric to measure the homogeneity of a clustering. A clustering result satisfies + homogeneity if all of its clusters contain only data points which are members of a single class. The metric is not + symmetric, therefore swapping ``preds`` and ``target`` yields a different score. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering in generally is used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``rand_score`` (:class:`~torch.Tensor`): A tensor with the Rand Score + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> import torch + >>> from torchmetrics.clustering import HomogeneityScore + >>> preds = torch.tensor([2, 1, 0, 1, 0]) + >>> target = torch.tensor([0, 2, 1, 1, 0]) + >>> metric = HomogeneityScore() + >>> metric(preds, target) + tensor(0.4744) + + """ + + is_differentiable: bool = True + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + self.preds.append(preds) + self.target.append(target) + + def compute(self) -> Tensor: + """Compute rand score over state.""" + return homogeneity_score(dim_zero_cat(self.preds), dim_zero_cat(self.target)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import HomogeneityScore + >>> metric = HomogeneityScore() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import HomogeneityScore + >>> metric = HomogeneityScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class CompletenessScore(Metric): + r"""Compute `Completeness Score`_. + + A clustering result satisfies completeness if all the data points that are members of a given class are elements of + the same cluster. The metric is not symmetric, therefore swapping ``preds`` and ``target`` yields a different score. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering in generally is used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``rand_score`` (:class:`~torch.Tensor`): A tensor with the Rand Score + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> import torch + >>> from torchmetrics.clustering import CompletenessScore + >>> preds = torch.tensor([2, 1, 0, 1, 0]) + >>> target = torch.tensor([0, 2, 1, 1, 0]) + >>> metric = CompletenessScore() + >>> metric(preds, target) + tensor(0.4744) + + """ + + is_differentiable: bool = True + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + self.preds.append(preds) + self.target.append(target) + + def compute(self) -> Tensor: + """Compute rand score over state.""" + return completeness_score(dim_zero_cat(self.preds), dim_zero_cat(self.target)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import CompletenessScore + >>> metric = CompletenessScore() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import CompletenessScore + >>> metric = CompletenessScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) + + +class VMeasureScore(Metric): + r"""Compute `V-Measure Score`_. + + The V-measure is the harmonic mean between homogeneity and completeness: + + ..math:: + v = \frac{(1 + \beta) * homogeneity * completeness}{\beta * homogeneity + completeness} + + where :math:`\beta` is a weight parameter that defines the weight of homogeneity in the harmonic mean, with the + default value :math:`\beta=1`. The V-measure is symmetric, which means that swapping ``preds`` and ``target`` does + not change the score. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering in generally is used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``rand_score`` (:class:`~torch.Tensor`): A tensor with the Rand Score + + Args: + beta: Weight parameter that defines the weight of homogeneity in the harmonic mean + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> import torch + >>> from torchmetrics.clustering import VMeasureScore + >>> preds = torch.tensor([2, 1, 0, 1, 0]) + >>> target = torch.tensor([0, 2, 1, 1, 0]) + >>> metric = VMeasureScore(beta=2.0) + >>> metric(preds, target) + tensor(0.4744) + + """ + + is_differentiable: bool = True + higher_is_better: bool = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__(self, beta: float = 1.0, **kwargs: Any) -> None: + super().__init__(**kwargs) + if not (isinstance(beta, float) and beta > 0): + raise ValueError(f"Argument `beta` should be a positive float. Got {beta}.") + self.beta = beta + + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + self.preds.append(preds) + self.target.append(target) + + def compute(self) -> Tensor: + """Compute rand score over state.""" + return v_measure_score(dim_zero_cat(self.preds), dim_zero_cat(self.target), beta=self.beta) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import VMeasureScore + >>> metric = VMeasureScore() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import VMeasureScore + >>> metric = VMeasureScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/mutual_info_score.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/mutual_info_score.py new file mode 100644 index 0000000000000000000000000000000000000000..13246d1e6c9aba91a9fc86f58e0c0fae79e82783 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/mutual_info_score.py @@ -0,0 +1,127 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering.mutual_info_score import mutual_info_score +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["MutualInfoScore.plot"] + + +class MutualInfoScore(Metric): + r"""Compute `Mutual Information Score`_. + + .. math:: + MI(U,V) = \sum_{i=1}^{|U|} \sum_{j=1}^{|V|} \frac{|U_i\cap V_j|}{N} + \log\frac{N|U_i\cap V_j|}{|U_i||V_j|} + + Where :math:`U` is a tensor of target values, :math:`V` is a tensor of predictions, + :math:`|U_i|` is the number of samples in cluster :math:`U_i`, and :math:`|V_i|` is the number of samples in + cluster :math:`V_i`. The metric is symmetric, therefore swapping :math:`U` and :math:`V` yields the same mutual + information score. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering in generally is used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``mi_score`` (:class:`~torch.Tensor`): A tensor with the Mutual Information Score + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> import torch + >>> from torchmetrics.clustering import MutualInfoScore + >>> preds = torch.tensor([2, 1, 0, 1, 0]) + >>> target = torch.tensor([0, 2, 1, 1, 0]) + >>> mi_score = MutualInfoScore() + >>> mi_score(preds, target) + tensor(0.5004) + + """ + + is_differentiable: bool = True + higher_is_better: Optional[bool] = True + full_state_update: bool = False + plot_lower_bound: float = 0.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + self.preds.append(preds) + self.target.append(target) + + def compute(self) -> Tensor: + """Compute mutual information over state.""" + return mutual_info_score(dim_zero_cat(self.preds), dim_zero_cat(self.target)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import MutualInfoScore + >>> metric = MutualInfoScore() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import MutualInfoScore + >>> metric = MutualInfoScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/normalized_mutual_info_score.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/normalized_mutual_info_score.py new file mode 100644 index 0000000000000000000000000000000000000000..373374fcda71095398a355265d3401d3b7a20787 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/normalized_mutual_info_score.py @@ -0,0 +1,127 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Literal, Optional, Union + +from torch import Tensor + +from torchmetrics.clustering.mutual_info_score import MutualInfoScore +from torchmetrics.functional.clustering.normalized_mutual_info_score import ( + _validate_average_method_arg, + normalized_mutual_info_score, +) +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["NormalizedMutualInfoScore.plot"] + + +class NormalizedMutualInfoScore(MutualInfoScore): + r"""Compute `Normalized Mutual Information Score`_. + + .. math:: + NMI(U,V) = \frac{MI(U,V)}{M_p(U,V)} + + Where :math:`U` is a tensor of target values, :math:`V` is a tensor of predictions, :math:`M_p(U,V)` is the + generalized mean of order :math:`p` of :math:`U` and :math:`V`, and :math:`MI(U,V)` is the mutual information score + between clusters :math:`U` and :math:`V`. The metric is symmetric, therefore swapping :math:`U` and :math:`V` yields + the same mutual information score. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering in generally is used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``nmi_score`` (:class:`~torch.Tensor`): A tensor with the Normalized Mutual Information Score + + Args: + average_method: Method used to calculate generalized mean for normalization + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> import torch + >>> from torchmetrics.clustering import NormalizedMutualInfoScore + >>> preds = torch.tensor([2, 1, 0, 1, 0]) + >>> target = torch.tensor([0, 2, 1, 1, 0]) + >>> nmi_score = NormalizedMutualInfoScore("arithmetic") + >>> nmi_score(preds, target) + tensor(0.4744) + + """ + + is_differentiable: bool = True + higher_is_better: Optional[bool] = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 0.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__( + self, average_method: Literal["min", "geometric", "arithmetic", "max"] = "arithmetic", **kwargs: Any + ) -> None: + super().__init__(**kwargs) + _validate_average_method_arg(average_method) + self.average_method = average_method + + def compute(self) -> Tensor: + """Compute normalized mutual information over state.""" + return normalized_mutual_info_score(dim_zero_cat(self.preds), dim_zero_cat(self.target), self.average_method) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import NormalizedMutualInfoScore + >>> metric = NormalizedMutualInfoScore() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import NormalizedMutualInfoScore + >>> metric = NormalizedMutualInfoScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/clustering/rand_score.py b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/rand_score.py new file mode 100644 index 0000000000000000000000000000000000000000..4ca73f28fe577f3db2b40d40c011894dae35e1ff --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/clustering/rand_score.py @@ -0,0 +1,125 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +from torch import Tensor + +from torchmetrics.functional.clustering.rand_score import rand_score +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["RandScore.plot"] + + +class RandScore(Metric): + r"""Compute `Rand Score`_ (alternatively known as Rand Index). + + .. math:: + RS(U, V) = \text{number of agreeing pairs} / \text{number of pairs} + + The number of agreeing pairs is every :math:`(i, j)` pair of samples where :math:`i \in U` and :math:`j \in V` + (the predicted and true clusterings, respectively) that are in the same cluster for both clusterings. The metric is + symmetric, therefore swapping :math:`U` and :math:`V` yields the same rand score. + + This clustering metric is an extrinsic measure, because it requires ground truth clustering labels, which may not + be available in practice since clustering in generally is used for unsupervised learning. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with predicted cluster labels + - ``target`` (:class:`~torch.Tensor`): single integer tensor with shape ``(N,)`` with ground truth cluster labels + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``rand_score`` (:class:`~torch.Tensor`): A tensor with the Rand Score + + Args: + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + >>> import torch + >>> from torchmetrics.clustering import RandScore + >>> preds = torch.tensor([2, 1, 0, 1, 0]) + >>> target = torch.tensor([0, 2, 1, 1, 0]) + >>> metric = RandScore() + >>> metric(preds, target) + tensor(0.6000) + + """ + + is_differentiable = True + higher_is_better = None + full_state_update: bool = False + plot_lower_bound: float = 0.0 + preds: List[Tensor] + target: List[Tensor] + + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + self.add_state("preds", default=[], dist_reduce_fx="cat") + self.add_state("target", default=[], dist_reduce_fx="cat") + + def update(self, preds: Tensor, target: Tensor) -> None: + """Update state with predictions and targets.""" + self.preds.append(preds) + self.target.append(target) + + def compute(self) -> Tensor: + """Compute rand score over state.""" + return rand_score(dim_zero_cat(self.preds), dim_zero_cat(self.target)) + + def plot(self, val: Union[Tensor, Sequence[Tensor], None] = None, ax: Optional[_AX_TYPE] = None) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting a single value + >>> import torch + >>> from torchmetrics.clustering import RandScore + >>> metric = RandScore() + >>> metric.update(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,))) + >>> fig_, ax_ = metric.plot(metric.compute()) + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.clustering import RandScore + >>> metric = RandScore() + >>> values = [ ] + >>> for _ in range(10): + ... values.append(metric(torch.randint(0, 4, (10,)), torch.randint(0, 4, (10,)))) + >>> fig_, ax_ = metric.plot(values) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/detection/__init__.py b/rtme/lib/python3.10/site-packages/torchmetrics/detection/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..968135b042844c062bd9117ba383b74a4f2389e3 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/detection/__init__.py @@ -0,0 +1,32 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from torchmetrics.detection.panoptic_qualities import ModifiedPanopticQuality, PanopticQuality +from torchmetrics.utilities.imports import _TORCHVISION_AVAILABLE + +__all__ = ["ModifiedPanopticQuality", "PanopticQuality"] + +if _TORCHVISION_AVAILABLE: + from torchmetrics.detection.ciou import CompleteIntersectionOverUnion + from torchmetrics.detection.diou import DistanceIntersectionOverUnion + from torchmetrics.detection.giou import GeneralizedIntersectionOverUnion + from torchmetrics.detection.iou import IntersectionOverUnion + from torchmetrics.detection.mean_ap import MeanAveragePrecision + + __all__ += [ + "CompleteIntersectionOverUnion", + "DistanceIntersectionOverUnion", + "GeneralizedIntersectionOverUnion", + "IntersectionOverUnion", + "MeanAveragePrecision", + ] diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/detection/_deprecated.py b/rtme/lib/python3.10/site-packages/torchmetrics/detection/_deprecated.py new file mode 100644 index 0000000000000000000000000000000000000000..f8acd23adb63177b0cc4790af0e496634a385faa --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/detection/_deprecated.py @@ -0,0 +1,63 @@ +from collections.abc import Collection +from typing import Any + +from torchmetrics.detection import ModifiedPanopticQuality, PanopticQuality +from torchmetrics.utilities.prints import _deprecated_root_import_class + + +class _ModifiedPanopticQuality(ModifiedPanopticQuality): + """Wrapper for deprecated import. + + >>> from torch import tensor + >>> preds = tensor([[[0, 0], [0, 1], [6, 0], [7, 0], [0, 2], [1, 0]]]) + >>> target = tensor([[[0, 1], [0, 0], [6, 0], [7, 0], [6, 0], [255, 0]]]) + >>> pq_modified = _ModifiedPanopticQuality(things = {0, 1}, stuffs = {6, 7}) + >>> pq_modified(preds, target) + tensor(0.7667, dtype=torch.float64) + + """ + + def __init__( + self, + things: Collection[int], + stuffs: Collection[int], + allow_unknown_preds_category: bool = False, + **kwargs: Any, + ) -> None: + _deprecated_root_import_class("ModifiedPanopticQuality", "detection") + super().__init__( + things=things, stuffs=stuffs, allow_unknown_preds_category=allow_unknown_preds_category, **kwargs + ) + + +class _PanopticQuality(PanopticQuality): + """Wrapper for deprecated import. + + >>> from torch import tensor + >>> preds = tensor([[[[6, 0], [0, 0], [6, 0], [6, 0]], + ... [[0, 0], [0, 0], [6, 0], [0, 1]], + ... [[0, 0], [0, 0], [6, 0], [0, 1]], + ... [[0, 0], [7, 0], [6, 0], [1, 0]], + ... [[0, 0], [7, 0], [7, 0], [7, 0]]]]) + >>> target = tensor([[[[6, 0], [0, 1], [6, 0], [0, 1]], + ... [[0, 1], [0, 1], [6, 0], [0, 1]], + ... [[0, 1], [0, 1], [6, 0], [1, 0]], + ... [[0, 1], [7, 0], [1, 0], [1, 0]], + ... [[0, 1], [7, 0], [7, 0], [7, 0]]]]) + >>> panoptic_quality = _PanopticQuality(things = {0, 1}, stuffs = {6, 7}) + >>> panoptic_quality(preds, target) + tensor(0.5463, dtype=torch.float64) + + """ + + def __init__( + self, + things: Collection[int], + stuffs: Collection[int], + allow_unknown_preds_category: bool = False, + **kwargs: Any, + ) -> None: + _deprecated_root_import_class("PanopticQuality", "detection") + super().__init__( + things=things, stuffs=stuffs, allow_unknown_preds_category=allow_unknown_preds_category, **kwargs + ) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/detection/_mean_ap.py b/rtme/lib/python3.10/site-packages/torchmetrics/detection/_mean_ap.py new file mode 100644 index 0000000000000000000000000000000000000000..857a87df9205d4911b2dd67622011da9b0afbb7a --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/detection/_mean_ap.py @@ -0,0 +1,988 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import logging +from collections.abc import Sequence +from typing import Any, Callable, List, Literal, Optional, Union + +import numpy as np +import torch +import torch.distributed as dist +from torch import IntTensor, Tensor + +from torchmetrics.detection.helpers import _fix_empty_tensors, _input_validator +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import _cumsum +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE, _PYCOCOTOOLS_AVAILABLE, _TORCHVISION_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["MeanAveragePrecision.plot"] + +if not _TORCHVISION_AVAILABLE or not _PYCOCOTOOLS_AVAILABLE: + __doctest_skip__ = ["MeanAveragePrecision.plot", "MeanAveragePrecision"] + +log = logging.getLogger(__name__) + + +def compute_area(inputs: list[Any], iou_type: Literal["bbox", "segm"] = "bbox") -> Tensor: + """Compute area of input depending on the specified iou_type. + + Default output for empty input is :class:`~torch.Tensor` + + """ + import pycocotools.mask as mask_utils + from torchvision.ops import box_area + + if len(inputs) == 0: + return Tensor([]) + + if iou_type == "bbox": + return box_area(torch.stack(inputs)) + if iou_type == "segm": + inputs = [{"size": i[0], "counts": i[1]} for i in inputs] + return torch.tensor(mask_utils.area(inputs).astype("float")) + + raise Exception(f"IOU type {iou_type} is not supported") + + +def compute_iou( + det: list[Any], + gt: list[Any], + iou_type: Literal["bbox", "segm"] = "bbox", +) -> Tensor: + """Compute IOU between detections and ground-truth using the specified iou_type.""" + from torchvision.ops import box_iou + + if iou_type == "bbox": + return box_iou(torch.stack(det), torch.stack(gt)) + if iou_type == "segm": + return _segm_iou(det, gt) + raise Exception(f"IOU type {iou_type} is not supported") + + +class BaseMetricResults(dict): + """Base metric class, that allows fields for pre-defined metrics.""" + + def __getattr__(self, key: str) -> Tensor: + """Get a specific metric attribute.""" + # Using this you get the correct error message, an AttributeError instead of a KeyError + if key in self: + return self[key] + raise AttributeError(f"No such attribute: {key}") + + def __setattr__(self, key: str, value: Tensor) -> None: + """Set a specific metric attribute.""" + self[key] = value + + def __delattr__(self, key: str) -> None: + """Delete a specific metric attribute.""" + if key in self: + del self[key] + raise AttributeError(f"No such attribute: {key}") + + +class MAPMetricResults(BaseMetricResults): + """Class to wrap the final mAP results.""" + + __slots__ = ("classes", "map", "map_50", "map_75", "map_large", "map_medium", "map_small") + + +class MARMetricResults(BaseMetricResults): + """Class to wrap the final mAR results.""" + + __slots__ = ("mar_1", "mar_10", "mar_100", "mar_large", "mar_medium", "mar_small") + + +class COCOMetricResults(BaseMetricResults): + """Class to wrap the final COCO metric results including various mAP/mAR values.""" + + __slots__ = ( + "map", + "map_50", + "map_75", + "map_large", + "map_medium", + "map_per_class", + "map_small", + "mar_1", + "mar_10", + "mar_100", + "mar_100_per_class", + "mar_large", + "mar_medium", + "mar_small", + ) + + +def _segm_iou(det: list[tuple[np.ndarray, np.ndarray]], gt: list[tuple[np.ndarray, np.ndarray]]) -> Tensor: + """Compute IOU between detections and ground-truths using mask-IOU. + + Implementation is based on pycocotools toolkit for mask_utils. + + Args: + det: A list of detection masks as ``[(RLE_SIZE, RLE_COUNTS)]``, where ``RLE_SIZE`` is (width, height) dimension + of the input and RLE_COUNTS is its RLE representation; + + gt: A list of ground-truth masks as ``[(RLE_SIZE, RLE_COUNTS)]``, where ``RLE_SIZE`` is (width, height) dimension + of the input and RLE_COUNTS is its RLE representation; + + """ + import pycocotools.mask as mask_utils + + det_coco_format = [{"size": i[0], "counts": i[1]} for i in det] + gt_coco_format = [{"size": i[0], "counts": i[1]} for i in gt] + + return torch.tensor(mask_utils.iou(det_coco_format, gt_coco_format, [False for _ in gt])) + + +class MeanAveragePrecision(Metric): + r"""Compute the `Mean-Average-Precision (mAP) and Mean-Average-Recall (mAR)`_ for object detection predictions. + + .. math:: + \text{mAP} = \frac{1}{n} \sum_{i=1}^{n} AP_i + + where :math:`AP_i` is the average precision for class :math:`i` and :math:`n` is the number of classes. The average + precision is defined as the area under the precision-recall curve. If argument `class_metrics` is set to ``True``, + the metric will also return the mAP/mAR per class. + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~List`): A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict + + - boxes: (:class:`~torch.FloatTensor`) of shape ``(num_boxes, 4)`` containing ``num_boxes`` detection + boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - scores: :class:`~torch.FloatTensor` of shape ``(num_boxes)`` containing detection scores for the boxes. + - labels: :class:`~torch.IntTensor` of shape ``(num_boxes)`` containing 0-indexed detection classes for + the boxes. + - masks: :class:`~torch.bool` of shape ``(num_boxes, image_height, image_width)`` containing boolean masks. + Only required when `iou_type="segm"`. + + - ``target`` (:class:`~List`) A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict: + + - boxes: :class:`~torch.FloatTensor` of shape ``(num_boxes, 4)`` containing ``num_boxes`` ground truth + boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - labels: :class:`~torch.IntTensor` of shape ``(num_boxes)`` containing 0-indexed ground truth + classes for the boxes. + - masks: :class:`~torch.bool` of shape ``(num_boxes, image_height, image_width)`` containing boolean masks. + Only required when `iou_type="segm"`. + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``map_dict``: A dictionary containing the following key-values: + + - map: (:class:`~torch.Tensor`) + - map_small: (:class:`~torch.Tensor`) + - map_medium:(:class:`~torch.Tensor`) + - map_large: (:class:`~torch.Tensor`) + - mar_1: (:class:`~torch.Tensor`) + - mar_10: (:class:`~torch.Tensor`) + - mar_100: (:class:`~torch.Tensor`) + - mar_small: (:class:`~torch.Tensor`) + - mar_medium: (:class:`~torch.Tensor`) + - mar_large: (:class:`~torch.Tensor`) + - map_50: (:class:`~torch.Tensor`) (-1 if 0.5 not in the list of iou thresholds) + - map_75: (:class:`~torch.Tensor`) (-1 if 0.75 not in the list of iou thresholds) + - map_per_class: (:class:`~torch.Tensor`) (-1 if class metrics are disabled) + - mar_100_per_class: (:class:`~torch.Tensor`) (-1 if class metrics are disabled) + - classes (:class:`~torch.Tensor`) + + For an example on how to use this metric check the `torchmetrics mAP example`_. + + .. attention:: + The ``map`` score is calculated with @[ IoU=self.iou_thresholds | area=all | max_dets=max_detection_thresholds ] + **Caution:** If the initialization parameters are changed, dictionary keys for mAR can change as well. + The default properties are also accessible via fields and will raise an ``AttributeError`` if not available. + + .. important:: + This metric is following the mAP implementation of `pycocotools`_ a standard implementation for the mAP metric + for object detection. + + .. hint:: + This metric requires you to have `torchvision` version 0.8.0 or newer installed + (with corresponding version 1.7.0 of torch or newer). This metric requires `pycocotools` + installed when iou_type is `segm`. Please install with ``pip install torchvision`` or + ``pip install torchmetrics[detection]``. + + Args: + box_format: + Input format of given boxes. Supported formats are ``[`xyxy`, `xywh`, `cxcywh`]``. + iou_type: + Type of input (either masks or bounding-boxes) used for computing IOU. + Supported IOU types are ``["bbox", "segm"]``. + If using ``"segm"``, masks should be provided (see :meth:`update`). + iou_thresholds: + IoU thresholds for evaluation. If set to ``None`` it corresponds to the stepped range ``[0.5,...,0.95]`` + with step ``0.05``. Else provide a list of floats. + rec_thresholds: + Recall thresholds for evaluation. If set to ``None`` it corresponds to the stepped range ``[0,...,1]`` + with step ``0.01``. Else provide a list of floats. + max_detection_thresholds: + Thresholds on max detections per image. If set to `None` will use thresholds ``[1, 10, 100]``. + Else, please provide a list of ints. + class_metrics: + Option to enable per-class metrics for mAP and mAR_100. Has a performance impact. + kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Raises: + ModuleNotFoundError: + If ``torchvision`` is not installed or version installed is lower than 0.8.0 + ModuleNotFoundError: + If ``iou_type`` is equal to ``segm`` and ``pycocotools`` is not installed + ValueError: + If ``class_metrics`` is not a boolean + ValueError: + If ``preds`` is not of type (:class:`~List[Dict[str, Tensor]]`) + ValueError: + If ``target`` is not of type ``List[Dict[str, Tensor]]`` + ValueError: + If ``preds`` and ``target`` are not of the same length + ValueError: + If any of ``preds.boxes``, ``preds.scores`` and ``preds.labels`` are not of the same length + ValueError: + If any of ``target.boxes`` and ``target.labels`` are not of the same length + ValueError: + If any box is not type float and of length 4 + ValueError: + If any class is not type int and of length 1 + ValueError: + If any score is not type float and of length 1 + + Example: + >>> from torch import tensor + >>> from torchmetrics.detection import MeanAveragePrecision + >>> preds = [ + ... dict( + ... boxes=tensor([[258.0, 41.0, 606.0, 285.0]]), + ... scores=tensor([0.536]), + ... labels=tensor([0]), + ... ) + ... ] + >>> target = [ + ... dict( + ... boxes=tensor([[214.0, 41.0, 562.0, 285.0]]), + ... labels=tensor([0]), + ... ) + ... ] + >>> metric = MeanAveragePrecision() + >>> metric.update(preds, target) + >>> from pprint import pprint + >>> pprint(metric.compute()) + {'classes': tensor(0, dtype=torch.int32), + 'map': tensor(0.6000), + 'map_50': tensor(1.), + 'map_75': tensor(1.), + 'map_large': tensor(0.6000), + 'map_medium': tensor(-1.), + 'map_per_class': tensor(-1.), + 'map_small': tensor(-1.), + 'mar_1': tensor(0.6000), + 'mar_10': tensor(0.6000), + 'mar_100': tensor(0.6000), + 'mar_100_per_class': tensor(-1.), + 'mar_large': tensor(0.6000), + 'mar_medium': tensor(-1.), + 'mar_small': tensor(-1.)} + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = True + plot_lower_bound: float = 0.0 + plot_upper_bound: float = 1.0 + + detections: List[Tensor] + detection_scores: List[Tensor] + detection_labels: List[Tensor] + groundtruths: List[Tensor] + groundtruth_labels: List[Tensor] + + def __init__( + self, + box_format: str = "xyxy", + iou_type: Literal["bbox", "segm"] = "bbox", + iou_thresholds: Optional[list[float]] = None, + rec_thresholds: Optional[list[float]] = None, + max_detection_thresholds: Optional[list[int]] = None, + class_metrics: bool = False, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + if not _PYCOCOTOOLS_AVAILABLE: + raise ModuleNotFoundError( + "`MAP` metric requires that `pycocotools` installed." + " Please install with `pip install pycocotools` or `pip install torchmetrics[detection]`" + ) + if not _TORCHVISION_AVAILABLE: + raise ModuleNotFoundError( + "`MeanAveragePrecision` metric requires that `torchvision` is installed." + " Please install with `pip install torchmetrics[detection]`." + ) + + allowed_box_formats = ("xyxy", "xywh", "cxcywh") + allowed_iou_types = ("segm", "bbox") + if box_format not in allowed_box_formats: + raise ValueError(f"Expected argument `box_format` to be one of {allowed_box_formats} but got {box_format}") + self.box_format = box_format + self.iou_thresholds = iou_thresholds or torch.linspace(0.5, 0.95, round((0.95 - 0.5) / 0.05) + 1).tolist() + self.rec_thresholds = rec_thresholds or torch.linspace(0.0, 1.00, round(1.00 / 0.01) + 1).tolist() + max_det_threshold, _ = torch.sort(IntTensor(max_detection_thresholds or [1, 10, 100])) + self.max_detection_thresholds = max_det_threshold.tolist() + if iou_type not in allowed_iou_types: + raise ValueError(f"Expected argument `iou_type` to be one of {allowed_iou_types} but got {iou_type}") + if iou_type == "segm" and not _PYCOCOTOOLS_AVAILABLE: + raise ModuleNotFoundError("When `iou_type` is set to 'segm', pycocotools need to be installed") + self.iou_type = iou_type + self.bbox_area_ranges = { + "all": (float(0**2), float(1e5**2)), + "small": (float(0**2), float(32**2)), + "medium": (float(32**2), float(96**2)), + "large": (float(96**2), float(1e5**2)), + } + + if not isinstance(class_metrics, bool): + raise ValueError("Expected argument `class_metrics` to be a boolean") + + self.class_metrics = class_metrics + self.add_state("detections", default=[], dist_reduce_fx=None) + self.add_state("detection_scores", default=[], dist_reduce_fx=None) + self.add_state("detection_labels", default=[], dist_reduce_fx=None) + self.add_state("groundtruths", default=[], dist_reduce_fx=None) + self.add_state("groundtruth_labels", default=[], dist_reduce_fx=None) + + def update(self, preds: list[dict[str, Tensor]], target: list[dict[str, Tensor]]) -> None: + """Update state with predictions and targets.""" + _input_validator(preds, target, iou_type=self.iou_type) + + for item in preds: + detections = self._get_safe_item_values(item) + + self.detections.append(detections) # type: ignore[arg-type] + self.detection_labels.append(item["labels"]) + self.detection_scores.append(item["scores"]) + + for item in target: + groundtruths = self._get_safe_item_values(item) + self.groundtruths.append(groundtruths) # type: ignore[arg-type] + self.groundtruth_labels.append(item["labels"]) + + def _move_list_states_to_cpu(self) -> None: + """Move list states to cpu to save GPU memory.""" + for key in self._defaults: + current_val = getattr(self, key) + current_to_cpu = [] + if isinstance(current_val, Sequence): + for cur_v in current_val: + # Cannot handle RLE as Tensor + if not isinstance(cur_v, tuple): + cur_v = cur_v.to("cpu") + current_to_cpu.append(cur_v) + setattr(self, key, current_to_cpu) + + def _get_safe_item_values(self, item: dict[str, Any]) -> Union[Tensor, tuple]: + import pycocotools.mask as mask_utils + from torchvision.ops import box_convert + + if self.iou_type == "bbox": + boxes = _fix_empty_tensors(item["boxes"]) + if boxes.numel() > 0: + boxes = box_convert(boxes, in_fmt=self.box_format, out_fmt="xyxy") + return boxes + if self.iou_type == "segm": + masks = [] + for i in item["masks"].cpu().numpy(): + rle = mask_utils.encode(np.asfortranarray(i)) + masks.append((tuple(rle["size"]), rle["counts"])) + return tuple(masks) + raise Exception(f"IOU type {self.iou_type} is not supported") + + def _get_classes(self) -> list: + """Return a list of unique classes found in ground truth and detection data.""" + if len(self.detection_labels) > 0 or len(self.groundtruth_labels) > 0: + return torch.cat(self.detection_labels + self.groundtruth_labels).unique().tolist() + return [] + + def _compute_iou(self, idx: int, class_id: int, max_det: int) -> Tensor: + """Compute the Intersection over Union (IoU) between bounding boxes for the given image and class. + + Args: + idx: + Image Id, equivalent to the index of supplied samples + class_id: + Class Id of the supplied ground truth and detection labels + max_det: + Maximum number of evaluated detection bounding boxes + + """ + # if self.iou_type == "bbox": + gt = self.groundtruths[idx] + det = self.detections[idx] + + gt_label_mask = (self.groundtruth_labels[idx] == class_id).nonzero().squeeze(1) + det_label_mask = (self.detection_labels[idx] == class_id).nonzero().squeeze(1) + + if len(gt_label_mask) == 0 or len(det_label_mask) == 0: + return Tensor([]) + + gt = [gt[i] for i in gt_label_mask] + det = [det[i] for i in det_label_mask] + + if len(gt) == 0 or len(det) == 0: + return Tensor([]) + + # Sort by scores and use only max detections + scores = self.detection_scores[idx] + scores_filtered = scores[self.detection_labels[idx] == class_id] + inds = torch.argsort(scores_filtered, descending=True) + + # TODO Fix (only for masks is necessary) + det = [det[i] for i in inds] + if len(det) > max_det: + det = det[:max_det] + + return compute_iou(det, gt, self.iou_type).to(self.device) + + def __evaluate_image_gt_no_preds( + self, gt: Tensor, gt_label_mask: Tensor, area_range: tuple[int, int], num_iou_thrs: int + ) -> dict[str, Any]: + """Evaluate images with a ground truth but no predictions.""" + # GTs + gt = [gt[i] for i in gt_label_mask] + num_gt = len(gt) + areas = compute_area(gt, iou_type=self.iou_type).to(self.device) + ignore_area = (areas < area_range[0]) | (areas > area_range[1]) + gt_ignore, _ = torch.sort(ignore_area.to(torch.uint8)) + gt_ignore = gt_ignore.to(torch.bool) + + # Detections + num_det = 0 + det_ignore = torch.zeros((num_iou_thrs, num_det), dtype=torch.bool, device=self.device) + + return { + "dtMatches": torch.zeros((num_iou_thrs, num_det), dtype=torch.bool, device=self.device), + "gtMatches": torch.zeros((num_iou_thrs, num_gt), dtype=torch.bool, device=self.device), + "dtScores": torch.zeros(num_det, dtype=torch.float32, device=self.device), + "gtIgnore": gt_ignore, + "dtIgnore": det_ignore, + } + + def __evaluate_image_preds_no_gt( + self, + det: Tensor, + idx: int, + det_label_mask: Tensor, + max_det: int, + area_range: tuple[int, int], + num_iou_thrs: int, + ) -> dict[str, Any]: + """Evaluate images with a prediction but no ground truth.""" + # GTs + num_gt = 0 + + gt_ignore = torch.zeros(num_gt, dtype=torch.bool, device=self.device) + + # Detections + + det = [det[i] for i in det_label_mask] + scores = self.detection_scores[idx] + scores_filtered = scores[det_label_mask] + scores_sorted, dtind = torch.sort(scores_filtered, descending=True) + + det = [det[i] for i in dtind] + if len(det) > max_det: + det = det[:max_det] + num_det = len(det) + det_areas = compute_area(det, iou_type=self.iou_type).to(self.device) + det_ignore_area = (det_areas < area_range[0]) | (det_areas > area_range[1]) + ar = det_ignore_area.reshape((1, num_det)) + det_ignore = torch.repeat_interleave(ar, num_iou_thrs, 0) + + return { + "dtMatches": torch.zeros((num_iou_thrs, num_det), dtype=torch.bool, device=self.device), + "gtMatches": torch.zeros((num_iou_thrs, num_gt), dtype=torch.bool, device=self.device), + "dtScores": scores_sorted.to(self.device), + "gtIgnore": gt_ignore.to(self.device), + "dtIgnore": det_ignore.to(self.device), + } + + def _evaluate_image( + self, idx: int, class_id: int, area_range: tuple[int, int], max_det: int, ious: dict + ) -> Optional[dict]: + """Perform evaluation for single class and image. + + Args: + idx: + Image Id, equivalent to the index of supplied samples. + class_id: + Class Id of the supplied ground truth and detection labels. + area_range: + List of lower and upper bounding box area threshold. + max_det: + Maximum number of evaluated detection bounding boxes. + ious: + IoU results for image and class. + + """ + gt = self.groundtruths[idx] + det = self.detections[idx] + gt_label_mask = (self.groundtruth_labels[idx] == class_id).nonzero().squeeze(1) + det_label_mask = (self.detection_labels[idx] == class_id).nonzero().squeeze(1) + + # No Gt and No predictions --> ignore image + if len(gt_label_mask) == 0 and len(det_label_mask) == 0: + return None + + num_iou_thrs = len(self.iou_thresholds) + + # Some GT but no predictions + if len(gt_label_mask) > 0 and len(det_label_mask) == 0: + return self.__evaluate_image_gt_no_preds(gt, gt_label_mask, area_range, num_iou_thrs) + + # Some predictions but no GT + if len(gt_label_mask) == 0 and len(det_label_mask) > 0: + return self.__evaluate_image_preds_no_gt(det, idx, det_label_mask, max_det, area_range, num_iou_thrs) + + gt = [gt[i] for i in gt_label_mask] + det = [det[i] for i in det_label_mask] + if len(gt) == 0 and len(det) == 0: + return None + if isinstance(det, dict): + det = [det] + if isinstance(gt, dict): + gt = [gt] + + areas = compute_area(gt, iou_type=self.iou_type).to(self.device) + + ignore_area = torch.logical_or(areas < area_range[0], areas > area_range[1]) + + # sort dt highest score first, sort gt ignore last + ignore_area_sorted, gtind = torch.sort(ignore_area.to(torch.uint8)) + # Convert to uint8 temporarily and back to bool, because "Sort currently does not support bool dtype on CUDA" + + ignore_area_sorted = ignore_area_sorted.to(torch.bool).to(self.device) + + gt = [gt[i] for i in gtind] + scores = self.detection_scores[idx] + scores_filtered = scores[det_label_mask] + scores_sorted, dtind = torch.sort(scores_filtered, descending=True) + det = [det[i] for i in dtind] + if len(det) > max_det: + det = det[:max_det] + # load computed ious + ious = ious[idx, class_id][:, gtind] if len(ious[idx, class_id]) > 0 else ious[idx, class_id] + + num_iou_thrs = len(self.iou_thresholds) + num_gt = len(gt) + num_det = len(det) + gt_matches = torch.zeros((num_iou_thrs, num_gt), dtype=torch.bool, device=self.device) + det_matches = torch.zeros((num_iou_thrs, num_det), dtype=torch.bool, device=self.device) + gt_ignore = ignore_area_sorted + det_ignore = torch.zeros((num_iou_thrs, num_det), dtype=torch.bool, device=self.device) + + if torch.numel(ious) > 0: + for idx_iou, t in enumerate(self.iou_thresholds): + for idx_det, _ in enumerate(det): + m = MeanAveragePrecision._find_best_gt_match(t, gt_matches, idx_iou, gt_ignore, ious, idx_det) + if m == -1: + continue + det_ignore[idx_iou, idx_det] = gt_ignore[m] + det_matches[idx_iou, idx_det] = 1 + gt_matches[idx_iou, m] = 1 + + # set unmatched detections outside of area range to ignore + det_areas = compute_area(det, iou_type=self.iou_type).to(self.device) + det_ignore_area = (det_areas < area_range[0]) | (det_areas > area_range[1]) + ar = det_ignore_area.reshape((1, num_det)) + det_ignore = torch.logical_or( + det_ignore, torch.logical_and(det_matches == 0, torch.repeat_interleave(ar, num_iou_thrs, 0)) + ) + + return { + "dtMatches": det_matches.to(self.device), + "gtMatches": gt_matches.to(self.device), + "dtScores": scores_sorted.to(self.device), + "gtIgnore": gt_ignore.to(self.device), + "dtIgnore": det_ignore.to(self.device), + } + + @staticmethod + def _find_best_gt_match( + threshold: int, gt_matches: Tensor, idx_iou: float, gt_ignore: Tensor, ious: Tensor, idx_det: int + ) -> int: + """Return id of best ground truth match with current detection. + + Args: + threshold: + Current threshold value. + gt_matches: + Tensor showing if a ground truth matches for threshold ``t`` exists. + idx_iou: + Id of threshold ``t``. + gt_ignore: + Tensor showing if ground truth should be ignored. + ious: + IoUs for all combinations of detection and ground truth. + idx_det: + Id of current detection. + + """ + previously_matched = gt_matches[idx_iou] # type: ignore[index] + # Remove previously matched or ignored gts + remove_mask = previously_matched | gt_ignore + gt_ious = ious[idx_det] * ~remove_mask + match_idx = gt_ious.argmax().item() + if gt_ious[match_idx] > threshold: # type: ignore[index] + return match_idx # type: ignore[return-value] + return -1 + + def _summarize( + self, + results: dict, + avg_prec: bool = True, + iou_threshold: Optional[float] = None, + area_range: str = "all", + max_dets: int = 100, + ) -> Tensor: + """Perform evaluation for single class and image. + + Args: + results: + Dictionary including precision, recall and scores for all combinations. + avg_prec: + Calculate average precision. Else calculate average recall. + iou_threshold: + IoU threshold. If set to ``None`` it all values are used. Else results are filtered. + area_range: + Bounding box area range key. + max_dets: + Maximum detections. + + """ + area_inds = [i for i, k in enumerate(self.bbox_area_ranges.keys()) if k == area_range] + mdet_inds = [i for i, k in enumerate(self.max_detection_thresholds) if k == max_dets] + if avg_prec: + # dimension of precision: [TxRxKxAxM] + prec = results["precision"] + # IoU + if iou_threshold is not None: + threshold = self.iou_thresholds.index(iou_threshold) + prec = prec[threshold, :, :, area_inds, mdet_inds] + else: + prec = prec[:, :, :, area_inds, mdet_inds] + else: + # dimension of recall: [TxKxAxM] + prec = results["recall"] + if iou_threshold is not None: + threshold = self.iou_thresholds.index(iou_threshold) + prec = prec[threshold, :, :, area_inds, mdet_inds] + else: + prec = prec[:, :, area_inds, mdet_inds] + + return torch.tensor([-1.0]) if len(prec[prec > -1]) == 0 else torch.mean(prec[prec > -1]) + + def _calculate(self, class_ids: list) -> tuple[MAPMetricResults, MARMetricResults]: + """Calculate the precision and recall for all supplied classes to calculate mAP/mAR. + + Args: + class_ids: + List of label class Ids. + + """ + img_ids = range(len(self.groundtruths)) + max_detections = self.max_detection_thresholds[-1] + area_ranges = self.bbox_area_ranges.values() + + ious = { + (idx, class_id): self._compute_iou(idx, class_id, max_detections) + for idx in img_ids + for class_id in class_ids + } + + eval_imgs = [ + self._evaluate_image(img_id, class_id, area, max_detections, ious) # type: ignore[arg-type] + for class_id in class_ids + for area in area_ranges + for img_id in img_ids + ] + + num_iou_thrs = len(self.iou_thresholds) + num_rec_thrs = len(self.rec_thresholds) + num_classes = len(class_ids) + num_bbox_areas = len(self.bbox_area_ranges) + num_max_det_thresholds = len(self.max_detection_thresholds) + num_imgs = len(img_ids) + precision = -torch.ones((num_iou_thrs, num_rec_thrs, num_classes, num_bbox_areas, num_max_det_thresholds)) + recall = -torch.ones((num_iou_thrs, num_classes, num_bbox_areas, num_max_det_thresholds)) + scores = -torch.ones((num_iou_thrs, num_rec_thrs, num_classes, num_bbox_areas, num_max_det_thresholds)) + + # move tensors if necessary + rec_thresholds_tensor = torch.tensor(self.rec_thresholds) + + # retrieve E at each category, area range, and max number of detections + for idx_cls, _ in enumerate(class_ids): + for idx_bbox_area, _ in enumerate(self.bbox_area_ranges): + for idx_max_det_thresholds, max_det in enumerate(self.max_detection_thresholds): + recall, precision, scores = MeanAveragePrecision.__calculate_recall_precision_scores( + recall, + precision, + scores, + idx_cls=idx_cls, + idx_bbox_area=idx_bbox_area, + idx_max_det_thresholds=idx_max_det_thresholds, + eval_imgs=eval_imgs, + rec_thresholds=rec_thresholds_tensor, + max_det=max_det, + num_imgs=num_imgs, + num_bbox_areas=num_bbox_areas, + ) + + return precision, recall # type: ignore[return-value] + + def _summarize_results(self, precisions: Tensor, recalls: Tensor) -> tuple[MAPMetricResults, MARMetricResults]: + """Summarizes the precision and recall values to calculate mAP/mAR. + + Args: + precisions: + Precision values for different thresholds + recalls: + Recall values for different thresholds + + """ + results = {"precision": precisions, "recall": recalls} + map_metrics = MAPMetricResults() + last_max_det_threshold = self.max_detection_thresholds[-1] + map_metrics.map = self._summarize(results, True, max_dets=last_max_det_threshold) + if 0.5 in self.iou_thresholds: + map_metrics.map_50 = self._summarize(results, True, iou_threshold=0.5, max_dets=last_max_det_threshold) + else: + map_metrics.map_50 = torch.tensor([-1]) + if 0.75 in self.iou_thresholds: + map_metrics.map_75 = self._summarize(results, True, iou_threshold=0.75, max_dets=last_max_det_threshold) + else: + map_metrics.map_75 = torch.tensor([-1]) + map_metrics.map_small = self._summarize(results, True, area_range="small", max_dets=last_max_det_threshold) + map_metrics.map_medium = self._summarize(results, True, area_range="medium", max_dets=last_max_det_threshold) + map_metrics.map_large = self._summarize(results, True, area_range="large", max_dets=last_max_det_threshold) + + mar_metrics = MARMetricResults() + for max_det in self.max_detection_thresholds: + mar_metrics[f"mar_{max_det}"] = self._summarize(results, False, max_dets=max_det) + mar_metrics.mar_small = self._summarize(results, False, area_range="small", max_dets=last_max_det_threshold) + mar_metrics.mar_medium = self._summarize(results, False, area_range="medium", max_dets=last_max_det_threshold) + mar_metrics.mar_large = self._summarize(results, False, area_range="large", max_dets=last_max_det_threshold) + + return map_metrics, mar_metrics + + @staticmethod + def __calculate_recall_precision_scores( + recall: Tensor, + precision: Tensor, + scores: Tensor, + idx_cls: int, + idx_bbox_area: int, + idx_max_det_thresholds: int, + eval_imgs: list, + rec_thresholds: Tensor, + max_det: int, + num_imgs: int, + num_bbox_areas: int, + ) -> tuple[Tensor, Tensor, Tensor]: + num_rec_thrs = len(rec_thresholds) + idx_cls_pointer = idx_cls * num_bbox_areas * num_imgs + idx_bbox_area_pointer = idx_bbox_area * num_imgs + # Load all image evals for current class_id and area_range + img_eval_cls_bbox = [eval_imgs[idx_cls_pointer + idx_bbox_area_pointer + i] for i in range(num_imgs)] + img_eval_cls_bbox = [e for e in img_eval_cls_bbox if e is not None] + if not img_eval_cls_bbox: + return recall, precision, scores + + det_scores = torch.cat([e["dtScores"][:max_det] for e in img_eval_cls_bbox]) + + # different sorting method generates slightly different results. + # mergesort is used to be consistent as Matlab implementation. + # Sort in PyTorch does not support bool types on CUDA (yet, 1.11.0) + dtype = torch.uint8 if det_scores.is_cuda and det_scores.dtype is torch.bool else det_scores.dtype + # Explicitly cast to uint8 to avoid error for bool inputs on CUDA to argsort + inds = torch.argsort(det_scores.to(dtype), descending=True) + det_scores_sorted = det_scores[inds] + + det_matches = torch.cat([e["dtMatches"][:, :max_det] for e in img_eval_cls_bbox], axis=1)[:, inds] # type: ignore[call-overload] + det_ignore = torch.cat([e["dtIgnore"][:, :max_det] for e in img_eval_cls_bbox], axis=1)[:, inds] # type: ignore[call-overload] + gt_ignore = torch.cat([e["gtIgnore"] for e in img_eval_cls_bbox]) + npig = torch.count_nonzero(gt_ignore == False) # noqa: E712 + if npig == 0: + return recall, precision, scores + tps = torch.logical_and(det_matches, torch.logical_not(det_ignore)) + fps = torch.logical_and(torch.logical_not(det_matches), torch.logical_not(det_ignore)) + + tp_sum = _cumsum(tps, dim=1, dtype=torch.float) + fp_sum = _cumsum(fps, dim=1, dtype=torch.float) + for idx, (tp, fp) in enumerate(zip(tp_sum, fp_sum)): + tp_len = len(tp) + rc = tp / npig + pr = tp / (fp + tp + torch.finfo(torch.float64).eps) + prec = torch.zeros((num_rec_thrs,)) + score = torch.zeros((num_rec_thrs,)) + + recall[idx, idx_cls, idx_bbox_area, idx_max_det_thresholds] = rc[-1] if tp_len else 0 + + # Remove zigzags for AUC + diff_zero = torch.zeros((1,), device=pr.device) + diff = torch.ones((1,), device=pr.device) + while not torch.all(diff == 0): + diff = torch.clamp(torch.cat(((pr[1:] - pr[:-1]), diff_zero), 0), min=0) + pr += diff + + inds = torch.searchsorted(rc, rec_thresholds.to(rc.device), right=False) + num_inds = inds.argmax() if inds.max() >= tp_len else num_rec_thrs + inds = inds[:num_inds] + prec[:num_inds] = pr[inds] + score[:num_inds] = det_scores_sorted[inds] + precision[idx, :, idx_cls, idx_bbox_area, idx_max_det_thresholds] = prec + scores[idx, :, idx_cls, idx_bbox_area, idx_max_det_thresholds] = score + + return recall, precision, scores + + def compute(self) -> dict: + """Compute metric.""" + classes = self._get_classes() + precisions, recalls = self._calculate(classes) + map_val, mar_val = self._summarize_results(precisions, recalls) # type: ignore[arg-type] + + # if class mode is enabled, evaluate metrics per class + map_per_class_values: Tensor = torch.tensor([-1.0]) + mar_max_dets_per_class_values: Tensor = torch.tensor([-1.0]) + if self.class_metrics: + map_per_class_list = [] + mar_max_dets_per_class_list = [] + + for class_idx, _ in enumerate(classes): + cls_precisions = precisions[:, :, class_idx].unsqueeze(dim=2) + cls_recalls = recalls[:, class_idx].unsqueeze(dim=1) + cls_map, cls_mar = self._summarize_results(cls_precisions, cls_recalls) + map_per_class_list.append(cls_map.map) + mar_max_dets_per_class_list.append(cls_mar[f"mar_{self.max_detection_thresholds[-1]}"]) + + map_per_class_values = torch.tensor(map_per_class_list, dtype=torch.float) + mar_max_dets_per_class_values = torch.tensor(mar_max_dets_per_class_list, dtype=torch.float) + + metrics = COCOMetricResults() + metrics.update(map_val) + metrics.update(mar_val) + metrics.map_per_class = map_per_class_values + metrics[f"mar_{self.max_detection_thresholds[-1]}_per_class"] = mar_max_dets_per_class_values + metrics.classes = torch.tensor(classes, dtype=torch.int) + return metrics + + def _apply(self, fn: Callable) -> torch.nn.Module: # type: ignore[override] + """Custom apply function. + + Excludes the detections and groundtruths from the casting when the iou_type is set to `segm` as the state is + no longer a tensor but a tuple. + + """ + if self.iou_type == "segm": + this = super()._apply(fn, exclude_state=("detections", "groundtruths")) + else: + this = super()._apply(fn) + return this + + def _sync_dist(self, dist_sync_fn: Optional[Callable] = None, process_group: Optional[Any] = None) -> None: + """Custom sync function. + + For the iou_type `segm` the detections and groundtruths are no longer tensors but tuples. Therefore, we need + to gather the list of tuples and then convert it back to a list of tuples. + + """ + super()._sync_dist(dist_sync_fn=dist_sync_fn, process_group=process_group) # type: ignore[arg-type] + + if self.iou_type == "segm": + self.detections = self._gather_tuple_list(self.detections, process_group) # type: ignore[arg-type] + self.groundtruths = self._gather_tuple_list(self.groundtruths, process_group) # type: ignore[arg-type] + + @staticmethod + def _gather_tuple_list( + list_to_gather: list[Union[tuple, Tensor]], process_group: Optional[Any] = None + ) -> list[Any]: + """Gather a list of tuples over multiple devices.""" + world_size = dist.get_world_size(group=process_group) + dist.barrier(group=process_group) + + list_gathered = [None for _ in range(world_size)] + dist.all_gather_object(list_gathered, list_to_gather, group=process_group) + + return [list_gathered[rank][idx] for idx in range(len(list_gathered[0])) for rank in range(world_size)] # type: ignore[arg-type,index] + + def plot( + self, val: Optional[Union[dict[str, Tensor], Sequence[dict[str, Tensor]]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> from torch import tensor + >>> from torchmetrics.detection.mean_ap import MeanAveragePrecision + >>> preds = [dict( + ... boxes=tensor([[258.0, 41.0, 606.0, 285.0]]), + ... scores=tensor([0.536]), + ... labels=tensor([0]), + ... )] + >>> target = [dict( + ... boxes=tensor([[214.0, 41.0, 562.0, 285.0]]), + ... labels=tensor([0]), + ... )] + >>> metric = MeanAveragePrecision() + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.detection.mean_ap import MeanAveragePrecision + >>> preds = lambda: [dict( + ... boxes=torch.tensor([[258.0, 41.0, 606.0, 285.0]]) + torch.randint(10, (1,4)), + ... scores=torch.tensor([0.536]) + 0.1*torch.rand(1), + ... labels=torch.tensor([0]), + ... )] + >>> target = [dict( + ... boxes=torch.tensor([[214.0, 41.0, 562.0, 285.0]]), + ... labels=torch.tensor([0]), + ... )] + >>> metric = MeanAveragePrecision() + >>> vals = [] + >>> for _ in range(20): + ... vals.append(metric(preds(), target)) + >>> fig_, ax_ = metric.plot(vals) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/detection/diou.py b/rtme/lib/python3.10/site-packages/torchmetrics/detection/diou.py new file mode 100644 index 0000000000000000000000000000000000000000..edfebb3818471df0c993b302affcd8eff585a5da --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/detection/diou.py @@ -0,0 +1,195 @@ +# Copyright The PyTorch Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor + +from torchmetrics.detection.iou import IntersectionOverUnion +from torchmetrics.functional.detection.diou import _diou_compute, _diou_update +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE, _TORCHVISION_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _TORCHVISION_AVAILABLE: + __doctest_skip__ = ["DistanceIntersectionOverUnion", "DistanceIntersectionOverUnion.plot"] +elif not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["DistanceIntersectionOverUnion.plot"] + + +class DistanceIntersectionOverUnion(IntersectionOverUnion): + r"""Computes Distance Intersection Over Union (`DIoU`_). + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~List`): A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict: + + - ``boxes`` (:class:`~torch.Tensor`): float tensor of shape ``(num_boxes, 4)`` containing ``num_boxes`` + detection boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - ``labels`` (:class:`~torch.Tensor`): integer tensor of shape ``(num_boxes)`` containing 0-indexed detection + classes for the boxes. + + - ``target`` (:class:`~List`): A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict: + + - ``boxes`` (:class:`~torch.Tensor`): float tensor of shape ``(num_boxes, 4)`` containing ``num_boxes`` ground + truth boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - ``labels`` (:class:`~torch.Tensor`): integer tensor of shape ``(num_boxes)`` containing 0-indexed ground truth + classes for the boxes. + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``diou_dict``: A dictionary containing the following key-values: + + - diou: (:class:`~torch.Tensor`) with overall diou value over all classes and samples. + - diou/cl_{cl}: (:class:`~torch.Tensor`), if argument ``class_metrics=True`` + + Args: + box_format: + Input format of given boxes. Supported formats are ``['xyxy', 'xywh', 'cxcywh']``. + iou_thresholds: + Optional IoU thresholds for evaluation. If set to `None` the threshold is ignored. + class_metrics: + Option to enable per-class metrics for IoU. Has a performance impact. + respect_labels: + Ignore values from boxes that do not have the same label as the ground truth box. Else will compute Iou + between all pairs of boxes. + kwargs: + Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> import torch + >>> from torchmetrics.detection import DistanceIntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = DistanceIntersectionOverUnion() + >>> metric(preds, target) + {'diou': tensor(0.8611)} + + Raises: + ModuleNotFoundError: + If torchvision is not installed with version 0.13.0 or newer. + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = True + + _iou_type: str = "diou" + _invalid_val: float = -1.0 + + def __init__( + self, + box_format: str = "xyxy", + iou_threshold: Optional[float] = None, + class_metrics: bool = False, + respect_labels: bool = True, + **kwargs: Any, + ) -> None: + if not _TORCHVISION_AVAILABLE: + raise ModuleNotFoundError( + f"Metric `{self._iou_type.upper()}` requires that `torchvision` is installed." + " Please install with `pip install torchmetrics[detection]`." + ) + super().__init__(box_format, iou_threshold, class_metrics, respect_labels, **kwargs) + + @staticmethod + def _iou_update_fn(*args: Any, **kwargs: Any) -> Tensor: + return _diou_update(*args, **kwargs) + + @staticmethod + def _iou_compute_fn(*args: Any, **kwargs: Any) -> Tensor: + return _diou_compute(*args, **kwargs) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting single value + >>> import torch + >>> from torchmetrics.detection import DistanceIntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = DistanceIntersectionOverUnion() + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.detection import DistanceIntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = lambda : [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]) + torch.randint(-10, 10, (1, 4)), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = DistanceIntersectionOverUnion() + >>> vals = [] + >>> for _ in range(20): + ... vals.append(metric(preds, target())) + >>> fig_, ax_ = metric.plot(vals) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/detection/giou.py b/rtme/lib/python3.10/site-packages/torchmetrics/detection/giou.py new file mode 100644 index 0000000000000000000000000000000000000000..cf03a73c65d7e321bc9ffc7057e1aef11fa50f52 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/detection/giou.py @@ -0,0 +1,190 @@ +# Copyright The PyTorch Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, Optional, Union + +from torch import Tensor + +from torchmetrics.detection.iou import IntersectionOverUnion +from torchmetrics.functional.detection.giou import _giou_compute, _giou_update +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE, _TORCHVISION_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _TORCHVISION_AVAILABLE: + __doctest_skip__ = ["GeneralizedIntersectionOverUnion", "GeneralizedIntersectionOverUnion.plot"] +elif not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["GeneralizedIntersectionOverUnion.plot"] + + +class GeneralizedIntersectionOverUnion(IntersectionOverUnion): + r"""Compute Generalized Intersection Over Union (`GIoU`_). + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~List`): A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict: + + - ``boxes`` (:class:`~torch.Tensor`): float tensor of shape ``(num_boxes, 4)`` containing ``num_boxes`` + detection boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - ``labels`` (:class:`~torch.Tensor`): integer tensor of shape ``(num_boxes)`` containing 0-indexed detection + classes for the boxes. + + - ``target`` (:class:`~List`): A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict: + + - ``boxes`` (:class:`~torch.Tensor`): float tensor of shape ``(num_boxes, 4)`` containing ``num_boxes`` ground + truth boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - ``labels`` (:class:`~torch.Tensor`): integer tensor of shape ``(num_boxes)`` containing 0-indexed ground truth + classes for the boxes. + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``giou_dict``: A dictionary containing the following key-values: + + - giou: (:class:`~torch.Tensor`) with overall giou value over all classes and samples. + - giou/cl_{cl}: (:class:`~torch.Tensor`), if argument ``class metrics=True`` + + Args: + box_format: + Input format of given boxes. Supported formats are ``[`xyxy`, `xywh`, `cxcywh`]``. + iou_thresholds: + Optional IoU thresholds for evaluation. If set to `None` the threshold is ignored. + class_metrics: + Option to enable per-class metrics for IoU. Has a performance impact. + respect_labels: + Ignore values from boxes that do not have the same label as the ground truth box. Else will compute Iou + between all pairs of boxes. + kwargs: + Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example: + >>> import torch + >>> from torchmetrics.detection import GeneralizedIntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = GeneralizedIntersectionOverUnion() + >>> metric(preds, target) + {'giou': tensor(0.8613)} + + Raises: + ModuleNotFoundError: + If torchvision is not installed with version 0.8.0 or newer. + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = True + + _iou_type: str = "giou" + _invalid_val: float = -1.0 + + def __init__( + self, + box_format: str = "xyxy", + iou_threshold: Optional[float] = None, + class_metrics: bool = False, + respect_labels: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(box_format, iou_threshold, class_metrics, respect_labels, **kwargs) + + @staticmethod + def _iou_update_fn(*args: Any, **kwargs: Any) -> Tensor: + return _giou_update(*args, **kwargs) + + @staticmethod + def _iou_compute_fn(*args: Any, **kwargs: Any) -> Tensor: + return _giou_compute(*args, **kwargs) + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> # Example plotting single value + >>> import torch + >>> from torchmetrics.detection import GeneralizedIntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = GeneralizedIntersectionOverUnion() + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.detection import GeneralizedIntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = lambda : [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 335.00, 150.00]]) + torch.randint(-10, 10, (1, 4)), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = GeneralizedIntersectionOverUnion() + >>> vals = [] + >>> for _ in range(20): + ... vals.append(metric(preds, target())) + >>> fig_, ax_ = metric.plot(vals) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/detection/iou.py b/rtme/lib/python3.10/site-packages/torchmetrics/detection/iou.py new file mode 100644 index 0000000000000000000000000000000000000000..22d7e5225d42cedaa44825f41e24642d09469cd3 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/detection/iou.py @@ -0,0 +1,297 @@ +# Copyright The PyTorch Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Any, List, Optional, Union + +import torch +from torch import Tensor + +from torchmetrics.detection.helpers import _fix_empty_tensors, _input_validator +from torchmetrics.functional.detection.iou import _iou_compute, _iou_update +from torchmetrics.metric import Metric +from torchmetrics.utilities.data import dim_zero_cat +from torchmetrics.utilities.imports import _MATPLOTLIB_AVAILABLE, _TORCHVISION_AVAILABLE +from torchmetrics.utilities.plot import _AX_TYPE, _PLOT_OUT_TYPE + +if not _TORCHVISION_AVAILABLE: + __doctest_skip__ = ["IntersectionOverUnion", "IntersectionOverUnion.plot"] +elif not _MATPLOTLIB_AVAILABLE: + __doctest_skip__ = ["IntersectionOverUnion.plot"] + + +class IntersectionOverUnion(Metric): + r"""Computes Intersection Over Union (IoU). + + As input to ``forward`` and ``update`` the metric accepts the following input: + + - ``preds`` (:class:`~List`): A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict: + + - ``boxes`` (:class:`~torch.Tensor`): float tensor of shape ``(num_boxes, 4)`` containing ``num_boxes`` + detection boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - labels: ``IntTensor`` of shape ``(num_boxes)`` containing 0-indexed detection classes for + the boxes. + + - ``target`` (:class:`~List`): A list consisting of dictionaries each containing the key-values + (each dictionary corresponds to a single image). Parameters that should be provided per dict: + + - ``boxes`` (:class:`~torch.Tensor`): float tensor of shape ``(num_boxes, 4)`` containing ``num_boxes`` ground + truth boxes of the format specified in the constructor. + By default, this method expects ``(xmin, ymin, xmax, ymax)`` in absolute image coordinates. + - ``labels`` (:class:`~torch.Tensor`): integer tensor of shape ``(num_boxes)`` containing 0-indexed ground truth + classes for the boxes. + + As output of ``forward`` and ``compute`` the metric returns the following output: + + - ``iou_dict``: A dictionary containing the following key-values: + + - iou: (:class:`~torch.Tensor`) + - iou/cl_{cl}: (:class:`~torch.Tensor`), if argument ``class metrics=True`` + + Args: + box_format: + Input format of given boxes. Supported formats are ``[`xyxy`, `xywh`, `cxcywh`]``. + iou_thresholds: + Optional IoU thresholds for evaluation. If set to `None` the threshold is ignored. + class_metrics: + Option to enable per-class metrics for IoU. Has a performance impact. + respect_labels: + Ignore values from boxes that do not have the same label as the ground truth box. Else will compute Iou + between all pairs of boxes. + kwargs: + Additional keyword arguments, see :ref:`Metric kwargs` for more info. + + Example:: + + >>> import torch + >>> from torchmetrics.detection import IntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([ + ... [296.55, 93.96, 314.97, 152.79], + ... [298.55, 98.96, 314.97, 151.79]]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = IntersectionOverUnion() + >>> metric(preds, target) + {'iou': tensor(0.8614)} + + Example:: + + The metric can also return the score per class: + + >>> import torch + >>> from torchmetrics.detection import IntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([ + ... [296.55, 93.96, 314.97, 152.79], + ... [298.55, 98.96, 314.97, 151.79]]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = [ + ... { + ... "boxes": torch.tensor([ + ... [300.00, 100.00, 315.00, 150.00], + ... [300.00, 100.00, 315.00, 150.00] + ... ]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> metric = IntersectionOverUnion(class_metrics=True) + >>> metric(preds, target) + {'iou': tensor(0.7756), 'iou/cl_4': tensor(0.6898), 'iou/cl_5': tensor(0.8614)} + + Raises: + ModuleNotFoundError: + If torchvision is not installed with version 0.8.0 or newer. + + """ + + is_differentiable: bool = False + higher_is_better: Optional[bool] = True + full_state_update: bool = True + + groundtruth_labels: List[Tensor] + iou_matrix: List[Tensor] + _iou_type: str = "iou" + _invalid_val: float = -1.0 + + def __init__( + self, + box_format: str = "xyxy", + iou_threshold: Optional[float] = None, + class_metrics: bool = False, + respect_labels: bool = True, + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + + if not _TORCHVISION_AVAILABLE: + raise ModuleNotFoundError( + f"Metric `{self._iou_type.upper()}` requires that `torchvision` is installed." + " Please install with `pip install torchmetrics[detection]`." + ) + + allowed_box_formats = ("xyxy", "xywh", "cxcywh") + if box_format not in allowed_box_formats: + raise ValueError(f"Expected argument `box_format` to be one of {allowed_box_formats} but got {box_format}") + + self.box_format = box_format + self.iou_threshold = iou_threshold + + if not isinstance(class_metrics, bool): + raise ValueError("Expected argument `class_metrics` to be a boolean") + self.class_metrics = class_metrics + + if not isinstance(respect_labels, bool): + raise ValueError("Expected argument `respect_labels` to be a boolean") + self.respect_labels = respect_labels + + self.add_state("groundtruth_labels", default=[], dist_reduce_fx=None) + self.add_state("iou_matrix", default=[], dist_reduce_fx=None) + + @staticmethod + def _iou_update_fn(*args: Any, **kwargs: Any) -> Tensor: + return _iou_update(*args, **kwargs) + + @staticmethod + def _iou_compute_fn(*args: Any, **kwargs: Any) -> Tensor: + return _iou_compute(*args, **kwargs) + + def update(self, preds: list[dict[str, Tensor]], target: list[dict[str, Tensor]]) -> None: + """Update state with predictions and targets.""" + _input_validator(preds, target, ignore_score=True) + + for p_i, t_i in zip(preds, target): + det_boxes = self._get_safe_item_values(p_i["boxes"]) + gt_boxes = self._get_safe_item_values(t_i["boxes"]) + self.groundtruth_labels.append(t_i["labels"]) + + iou_matrix = self._iou_update_fn(det_boxes, gt_boxes, self.iou_threshold, self._invalid_val) # N x M + if self.respect_labels: + if det_boxes.numel() > 0 and gt_boxes.numel() > 0: + label_eq = p_i["labels"].unsqueeze(1) == t_i["labels"].unsqueeze(0) # N x M + else: + label_eq = torch.eye(iou_matrix.shape[0], dtype=bool, device=iou_matrix.device) # type: ignore[call-overload] + iou_matrix[~label_eq] = self._invalid_val + self.iou_matrix.append(iou_matrix) + + def _get_safe_item_values(self, boxes: Tensor) -> Tensor: + from torchvision.ops import box_convert + + boxes = _fix_empty_tensors(boxes) + if boxes.numel() > 0: + boxes = box_convert(boxes, in_fmt=self.box_format, out_fmt="xyxy") + return boxes + + def _get_gt_classes(self) -> list: + """Returns a list of unique classes found in ground truth and detection data.""" + if len(self.groundtruth_labels) > 0: + return torch.cat(self.groundtruth_labels).unique().tolist() + return [] + + def compute(self) -> dict: + """Computes IoU based on inputs passed in to ``update`` previously.""" + score = torch.cat([mat[mat != self._invalid_val] for mat in self.iou_matrix], 0).mean() + results: dict[str, Tensor] = {f"{self._iou_type}": score} + if torch.isnan(score): # if no valid boxes are found + results[f"{self._iou_type}"] = torch.tensor(0.0, device=score.device) + if self.class_metrics: + gt_labels = dim_zero_cat(self.groundtruth_labels) + classes = gt_labels.unique().tolist() if len(gt_labels) > 0 else [] + for cl in classes: + masked_iou, observed = torch.zeros_like(score), torch.zeros_like(score) + for mat, gt_lab in zip(self.iou_matrix, self.groundtruth_labels): + scores = mat[:, gt_lab == cl] + masked_iou += scores[scores != self._invalid_val].sum() + observed += scores[scores != self._invalid_val].numel() + results.update({f"{self._iou_type}/cl_{cl}": masked_iou / observed}) + return results + + def plot( + self, val: Optional[Union[Tensor, Sequence[Tensor]]] = None, ax: Optional[_AX_TYPE] = None + ) -> _PLOT_OUT_TYPE: + """Plot a single or multiple values from the metric. + + Args: + val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results. + If no value is provided, will automatically call `metric.compute` and plot that result. + ax: An matplotlib axis object. If provided will add plot to that axis + + Returns: + Figure object and Axes object + + Raises: + ModuleNotFoundError: + If `matplotlib` is not installed + + .. plot:: + :scale: 75 + + >>> import torch + >>> from torchmetrics.detection import IntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = IntersectionOverUnion() + >>> metric.update(preds, target) + >>> fig_, ax_ = metric.plot() + + .. plot:: + :scale: 75 + + >>> # Example plotting multiple values + >>> import torch + >>> from torchmetrics.detection import IntersectionOverUnion + >>> preds = [ + ... { + ... "boxes": torch.tensor([[296.55, 93.96, 314.97, 152.79], [298.55, 98.96, 314.97, 151.79]]), + ... "scores": torch.tensor([0.236, 0.56]), + ... "labels": torch.tensor([4, 5]), + ... } + ... ] + >>> target = lambda : [ + ... { + ... "boxes": torch.tensor([[300.00, 100.00, 315.00, 150.00]]) + torch.randint(-10, 10, (1, 4)), + ... "labels": torch.tensor([5]), + ... } + ... ] + >>> metric = IntersectionOverUnion() + >>> vals = [] + >>> for _ in range(20): + ... vals.append(metric(preds, target())) + >>> fig_, ax_ = metric.plot(vals) + + """ + return self._plot(val, ax) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/d_lambda.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/d_lambda.py new file mode 100644 index 0000000000000000000000000000000000000000..478455c0a685dd0c97e33cd3049d0997efc2ab74 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/d_lambda.py @@ -0,0 +1,152 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.image.uqi import universal_image_quality_index +from torchmetrics.utilities.distributed import reduce + + +def _spectral_distortion_index_update(preds: Tensor, target: Tensor) -> tuple[Tensor, Tensor]: + """Update and returns variables required to compute Spectral Distortion Index. + + Args: + preds: Low resolution multispectral image + target: High resolution fused image + + """ + if preds.dtype != target.dtype: + raise TypeError( + f"Expected `ms` and `fused` to have the same data type. Got ms: {preds.dtype} and fused: {target.dtype}." + ) + if len(preds.shape) != 4: + raise ValueError( + f"Expected `preds` and `target` to have BxCxHxW shape. Got preds: {preds.shape} and target: {target.shape}." + ) + if preds.shape[:2] != target.shape[:2]: + raise ValueError( + "Expected `preds` and `target` to have same batch and channel sizes." + f"Got preds: {preds.shape} and target: {target.shape}." + ) + return preds, target + + +def _spectral_distortion_index_compute( + preds: Tensor, + target: Tensor, + p: int = 1, + reduction: Literal["elementwise_mean", "sum", "none"] = "elementwise_mean", +) -> Tensor: + """Compute Spectral Distortion Index (SpectralDistortionIndex_). + + Args: + preds: Low resolution multispectral image + target: High resolution fused image + p: a parameter to emphasize large spectral difference + reduction: a method to reduce metric score over labels. + + - ``'elementwise_mean'``: takes the mean (default) + - ``'sum'``: takes the sum + - ``'none'``: no reduction will be applied + + Example: + >>> from torch import rand + >>> preds = rand([16, 3, 16, 16]) + >>> target = rand([16, 3, 16, 16]) + >>> preds, target = _spectral_distortion_index_update(preds, target) + >>> _spectral_distortion_index_compute(preds, target) + tensor(0.0234) + + """ + length = preds.shape[1] + + m1 = torch.zeros((length, length), device=preds.device) + m2 = torch.zeros((length, length), device=preds.device) + + for k in range(length): + num = length - (k + 1) + if num == 0: + continue + stack1 = target[:, k : k + 1, :, :].repeat(num, 1, 1, 1) + stack2 = torch.cat([target[:, r : r + 1, :, :] for r in range(k + 1, length)], dim=0) + score = [ + s.mean() for s in universal_image_quality_index(stack1, stack2, reduction="none").split(preds.shape[0]) + ] + m1[k, k + 1 :] = torch.stack(score, 0) + + stack1 = preds[:, k : k + 1, :, :].repeat(num, 1, 1, 1) + stack2 = torch.cat([preds[:, r : r + 1, :, :] for r in range(k + 1, length)], dim=0) + score = [ + s.mean() for s in universal_image_quality_index(stack1, stack2, reduction="none").split(preds.shape[0]) + ] + m2[k, k + 1 :] = torch.stack(score, 0) + m1 = m1 + m1.T + m2 = m2 + m2.T + + diff = torch.pow(torch.abs(m1 - m2), p) + # Special case: when number of channels (L) is 1, there will be only one element in M1 and M2. Hence no need to sum. + if length == 1: + output = torch.pow(diff, (1.0 / p)) + else: + output = torch.pow(1.0 / (length * (length - 1)) * torch.sum(diff), (1.0 / p)) + return reduce(output, reduction) + + +def spectral_distortion_index( + preds: Tensor, + target: Tensor, + p: int = 1, + reduction: Literal["elementwise_mean", "sum", "none"] = "elementwise_mean", +) -> Tensor: + """Calculate `Spectral Distortion Index`_ (SpectralDistortionIndex_) also known as D_lambda. + + Metric is used to compare the spectral distortion between two images. + + Args: + preds: Low resolution multispectral image + target: High resolution fused image + p: Large spectral differences + reduction: a method to reduce metric score over labels. + + - ``'elementwise_mean'``: takes the mean (default) + - ``'sum'``: takes the sum + - ``'none'``: no reduction will be applied + + Return: + Tensor with SpectralDistortionIndex score + + Raises: + TypeError: + If ``preds`` and ``target`` don't have the same data type. + ValueError: + If ``preds`` and ``target`` don't have ``BxCxHxW shape``. + ValueError: + If ``p`` is not a positive integer. + + Example: + >>> from torch import rand + >>> from torchmetrics.functional.image import spectral_distortion_index + >>> preds = rand([16, 3, 16, 16]) + >>> target = rand([16, 3, 16, 16]) + >>> spectral_distortion_index(preds, target) + tensor(0.0234) + + """ + if not isinstance(p, int) or p <= 0: + raise ValueError(f"Expected `p` to be a positive integer. Got p: {p}.") + preds, target = _spectral_distortion_index_update(preds, target) + return _spectral_distortion_index_compute(preds, target, p, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/dists.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/dists.py new file mode 100644 index 0000000000000000000000000000000000000000..c215714006bc9db055b381c5d20638bb100bb87e --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/dists.py @@ -0,0 +1,215 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# Below is a derivative work based on the original work: +# https://github.com/dingkeyan93/DISTS +# with the following license: +# +# MIT License +# Copyright (c) 2020 Keyan Ding +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +from pathlib import Path +from typing import List, Optional + +import numpy as np +import torch +import torch.nn as nn +from torch import Tensor +from torch.nn.functional import conv2d +from typing_extensions import Literal + +from torchmetrics.utilities.imports import _TORCHVISION_AVAILABLE + +if not _TORCHVISION_AVAILABLE: + __doctest_skip__ = ["deep_image_structure_and_texture_similarity"] + +_PATH_WEIGHT_DISTS = Path(__file__).resolve().parent / "dists_models" / "weights.pt" + + +class L2pooling(nn.Module): + """L2 pooling layer.""" + + filter: Tensor + + def __init__(self, filter_size: int = 5, stride: int = 2, channels: int = 3) -> None: + super().__init__() + self.padding = (filter_size - 2) // 2 + self.stride = stride + self.channels = channels + a = np.hanning(filter_size)[1:-1] + g = torch.Tensor(a[:, None] * a[None, :]) + g = g / torch.sum(g) + self.register_buffer("filter", g[None, None, :, :].repeat(self.channels, 1, 1, 1)) + + def forward(self, tensor: Tensor) -> Tensor: + """Forward pass of the layer.""" + tensor = tensor**2 + out = conv2d(tensor, self.filter, stride=self.stride, padding=self.padding, groups=tensor.shape[1]) + return (out + 1e-12).sqrt() + + +class DISTSNetwork(torch.nn.Module): + """DISTS network.""" + + alpha: Tensor + beta: Tensor + mean: Tensor + std: Tensor + + def __init__(self, load_weights: bool = True) -> None: + super().__init__() + + if _TORCHVISION_AVAILABLE: + from torchvision import models + else: + raise ModuleNotFoundError( + "DISTS requires torchvision to be installed. Please install it with `pip install torchvision`." + ) + + vgg_pretrained_features = models.vgg16(pretrained=True).features + self.stage1 = torch.nn.Sequential() + self.stage2 = torch.nn.Sequential() + self.stage3 = torch.nn.Sequential() + self.stage4 = torch.nn.Sequential() + self.stage5 = torch.nn.Sequential() + for x in range(4): + self.stage1.add_module(str(x), vgg_pretrained_features[x]) + self.stage2.add_module(str(4), L2pooling(channels=64)) + for x in range(5, 9): + self.stage2.add_module(str(x), vgg_pretrained_features[x]) + self.stage3.add_module(str(9), L2pooling(channels=128)) + for x in range(10, 16): + self.stage3.add_module(str(x), vgg_pretrained_features[x]) + self.stage4.add_module(str(16), L2pooling(channels=256)) + for x in range(17, 23): + self.stage4.add_module(str(x), vgg_pretrained_features[x]) + self.stage5.add_module(str(23), L2pooling(channels=512)) + for x in range(24, 30): + self.stage5.add_module(str(x), vgg_pretrained_features[x]) + + for param in self.parameters(): + param.requires_grad = False + + self.register_buffer("mean", torch.tensor([0.485, 0.456, 0.406]).view(1, -1, 1, 1)) + self.register_buffer("std", torch.tensor([0.229, 0.224, 0.225]).view(1, -1, 1, 1)) + + self.chns = [3, 64, 128, 256, 512, 512] + self.register_parameter("alpha", nn.Parameter(torch.randn(1, sum(self.chns), 1, 1))) + self.register_parameter("beta", nn.Parameter(torch.randn(1, sum(self.chns), 1, 1))) + self.alpha.data.normal_(0.1, 0.01) + self.beta.data.normal_(0.1, 0.01) + if load_weights: + if not _PATH_WEIGHT_DISTS.exists(): + raise FileNotFoundError(f"The weights file is not found in {_PATH_WEIGHT_DISTS}") + weights = torch.load(str(_PATH_WEIGHT_DISTS)) + self.alpha.data = weights["alpha"] + self.beta.data = weights["beta"] + + def forward_once(self, x: Tensor) -> List[Tensor]: + """Forward pass of the network.""" + h = (x - self.mean) / self.std + h = self.stage1(h) + h_relu1_2 = h + h = self.stage2(h) + h_relu2_2 = h + h = self.stage3(h) + h_relu3_3 = h + h = self.stage4(h) + h_relu4_3 = h + h = self.stage5(h) + h_relu5_3 = h + return [x, h_relu1_2, h_relu2_2, h_relu3_3, h_relu4_3, h_relu5_3] + + def forward(self, x: Tensor, y: Tensor, require_grad: bool = False) -> Tensor: + """Computes DISTS score between two images.""" + if require_grad: + feats0 = self.forward_once(x) + feats1 = self.forward_once(y) + else: + with torch.inference_mode(): + feats0 = self.forward_once(x) + feats1 = self.forward_once(y) + dist1, dist2, c1, c2 = 0, 0, 1e-6, 1e-6 + w_sum = self.alpha.sum() + self.beta.sum() + alpha = torch.split(self.alpha / w_sum, self.chns, dim=1) + beta = torch.split(self.beta / w_sum, self.chns, dim=1) + for k in range(len(self.chns)): + x_mean = feats0[k].mean([2, 3], keepdim=True) + y_mean = feats1[k].mean([2, 3], keepdim=True) + s1 = (2 * x_mean * y_mean + c1) / (x_mean**2 + y_mean**2 + c1) + dist1 = dist1 + (alpha[k] * s1).sum(1, keepdim=True) + + x_var = ((feats0[k] - x_mean) ** 2).mean([2, 3], keepdim=True) + y_var = ((feats1[k] - y_mean) ** 2).mean([2, 3], keepdim=True) + xy_cov = (feats0[k] * feats1[k]).mean([2, 3], keepdim=True) - x_mean * y_mean + s2 = (2 * xy_cov + c2) / (x_var + y_var + c2) + dist2 = dist2 + (beta[k] * s2).sum(1, keepdim=True) + + return 1 - (dist1 + dist2).squeeze() + + +def _dists_update(preds: Tensor, target: Tensor) -> Tensor: + dists = DISTSNetwork().to(preds.device) + return dists(preds, target, require_grad=preds.requires_grad) + + +def _dists_compute(scores: Tensor, reduction: Optional[Literal["sum", "mean", "none"]]) -> Tensor: + if reduction == "sum": + return scores.sum() + if reduction == "mean": + return scores.mean() + if reduction is None or reduction == "none": + return scores + raise ValueError(f"Argument {reduction} is not valid. Choose 'sum', 'mean' or 'none'., but got {reduction}") + + +def deep_image_structure_and_texture_similarity( + preds: Tensor, target: Tensor, reduction: Optional[Literal["sum", "mean", "none"]] = None +) -> Tensor: + """Calculates `Deep Image Structure and Texture Similarity`_ (DISTS) score. + + Args: + preds: Predicted image tensor. + target: Target image tensor. + reduction: Reduction method for the output. + + Returns: + DISTS Similarity score between the two images. + + Example: + >>> from torch import rand + >>> preds = rand(5, 3, 256, 256) + >>> target = rand(5, 3, 256, 256) + >>> deep_image_structure_and_texture_similarity(preds, target) + tensor([0.1285, 0.1344, 0.1356, 0.1277, 0.1276], grad_fn=) + >>> deep_image_structure_and_texture_similarity(preds, target, reduction='mean') + tensor(0.1308, grad_fn=) + + """ + scores = _dists_update(preds, target) + return _dists_compute(scores, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/gradients.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/gradients.py new file mode 100644 index 0000000000000000000000000000000000000000..68045663fa9ebe40c03a09469c8a0b0bbb93eb19 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/gradients.py @@ -0,0 +1,80 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import torch +from torch import Tensor + + +def _image_gradients_validate(img: Tensor) -> None: + """Validate whether img is a 4D torch Tensor.""" + if not isinstance(img, Tensor): + raise TypeError(f"The `img` expects a value of type but got {type(img)}") + if img.ndim != 4: + raise RuntimeError(f"The `img` expects a 4D tensor but got {img.ndim}D tensor") + + +def _compute_image_gradients(img: Tensor) -> tuple[Tensor, Tensor]: + """Compute image gradients (dy/dx) for a given image.""" + batch_size, channels, height, width = img.shape + + dy = img[..., 1:, :] - img[..., :-1, :] + dx = img[..., :, 1:] - img[..., :, :-1] + + shapey = [batch_size, channels, 1, width] + dy = torch.cat([dy, torch.zeros(shapey, device=img.device, dtype=img.dtype)], dim=2) + dy = dy.view(img.shape) + + shapex = [batch_size, channels, height, 1] + dx = torch.cat([dx, torch.zeros(shapex, device=img.device, dtype=img.dtype)], dim=3) + dx = dx.view(img.shape) + + return dy, dx + + +def image_gradients(img: Tensor) -> tuple[Tensor, Tensor]: + """Compute `Gradient Computation of Image`_ of a given image using finite difference. + + Args: + img: An ``(N, C, H, W)`` input tensor where ``C`` is the number of image channels + + Return: + Tuple of ``(dy, dx)`` with each gradient of shape ``[N, C, H, W]`` + + Raises: + TypeError: + If ``img`` is not of the type :class:`~torch.Tensor`. + RuntimeError: + If ``img`` is not a 4D tensor. + + Example: + >>> from torchmetrics.functional.image import image_gradients + >>> image = torch.arange(0, 1*1*5*5, dtype=torch.float32) + >>> image = torch.reshape(image, (1, 1, 5, 5)) + >>> dy, dx = image_gradients(image) + >>> dy[0, 0, :, :] + tensor([[5., 5., 5., 5., 5.], + [5., 5., 5., 5., 5.], + [5., 5., 5., 5., 5.], + [5., 5., 5., 5., 5.], + [0., 0., 0., 0., 0.]]) + + .. note:: + The implementation follows the 1-step finite difference method as followed + by the TF implementation. The values are organized such that the gradient of + [I(x+1, y)-[I(x, y)]] are at the (x, y) location + + """ + _image_gradients_validate(img) + + return _compute_image_gradients(img) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/perceptual_path_length.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/perceptual_path_length.py new file mode 100644 index 0000000000000000000000000000000000000000..035425539d200ee6b8491f4122c76844d9ec9730 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/perceptual_path_length.py @@ -0,0 +1,283 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import math +from typing import Literal, Optional, Union + +import torch +from torch import Tensor, nn + +from torchmetrics.functional.image.lpips import _LPIPS +from torchmetrics.utilities.imports import _TORCHVISION_AVAILABLE + +if not _TORCHVISION_AVAILABLE: + __doctest_skip__ = ["perceptual_path_length"] + + +class GeneratorType(nn.Module): + """Basic interface for a generator model. + + Users can inherit from this class and implement their own generator model. The requirements are that the ``sample`` + method is implemented and that the ``num_classes`` attribute is present when ``conditional=True`` metric. + + """ + + @property + def num_classes(self) -> int: + """Return the number of classes for conditional generation.""" + raise NotImplementedError + + def sample(self, num_samples: int) -> Tensor: + """Sample from the generator. + + Args: + num_samples: Number of samples to generate. + + """ + raise NotImplementedError + + +def _validate_generator_model(generator: GeneratorType, conditional: bool = False) -> None: + """Validate that the user provided generator has the right methods and attributes. + + Args: + generator: Generator model + conditional: Whether the generator is conditional or not (i.e. whether it takes labels as input). + + """ + if not hasattr(generator, "sample"): + raise NotImplementedError( + "The generator must have a `sample` method with signature `sample(num_samples: int) -> Tensor` where the" + " returned tensor has shape `(num_samples, z_size)`." + ) + if not callable(generator.sample): + raise ValueError("The generator's `sample` method must be callable.") + if conditional and not hasattr(generator, "num_classes"): + raise AttributeError("The generator must have a `num_classes` attribute when `conditional=True`.") + if conditional and not isinstance(generator.num_classes, int): + raise ValueError("The generator's `num_classes` attribute must be an integer when `conditional=True`.") + + +def _perceptual_path_length_validate_arguments( + num_samples: int = 10_000, + conditional: bool = False, + batch_size: int = 128, + interpolation_method: Literal["lerp", "slerp_any", "slerp_unit"] = "lerp", + epsilon: float = 1e-4, + resize: Optional[int] = 64, + lower_discard: Optional[float] = 0.01, + upper_discard: Optional[float] = 0.99, +) -> None: + """Validate arguments for perceptual path length.""" + if not (isinstance(num_samples, int) and num_samples > 0): + raise ValueError(f"Argument `num_samples` must be a positive integer, but got {num_samples}.") + if not isinstance(conditional, bool): + raise ValueError(f"Argument `conditional` must be a boolean, but got {conditional}.") + if not (isinstance(batch_size, int) and batch_size > 0): + raise ValueError(f"Argument `batch_size` must be a positive integer, but got {batch_size}.") + if interpolation_method not in ["lerp", "slerp_any", "slerp_unit"]: + raise ValueError( + f"Argument `interpolation_method` must be one of 'lerp', 'slerp_any', 'slerp_unit'," + f"got {interpolation_method}." + ) + if not (isinstance(epsilon, float) and epsilon > 0): + raise ValueError(f"Argument `epsilon` must be a positive float, but got {epsilon}.") + if resize is not None and not (isinstance(resize, int) and resize > 0): + raise ValueError(f"Argument `resize` must be a positive integer or `None`, but got {resize}.") + if lower_discard is not None and not (isinstance(lower_discard, float) and 0 <= lower_discard <= 1): + raise ValueError( + f"Argument `lower_discard` must be a float between 0 and 1 or `None`, but got {lower_discard}." + ) + if upper_discard is not None and not (isinstance(upper_discard, float) and 0 <= upper_discard <= 1): + raise ValueError( + f"Argument `upper_discard` must be a float between 0 and 1 or `None`, but got {upper_discard}." + ) + + +def _interpolate( + latents1: Tensor, + latents2: Tensor, + epsilon: float = 1e-4, + interpolation_method: Literal["lerp", "slerp_any", "slerp_unit"] = "lerp", +) -> Tensor: + """Interpolate between two sets of latents. + + Inspired by: https://github.com/toshas/torch-fidelity/blob/master/torch_fidelity/noise.py + + Args: + latents1: First set of latents. + latents2: Second set of latents. + epsilon: Spacing between the points on the path between latent points. + interpolation_method: Interpolation method to use. Choose from 'lerp', 'slerp_any', 'slerp_unit'. + + """ + eps = 1e-7 + if latents1.shape != latents2.shape: + raise ValueError("Latents must have the same shape.") + if interpolation_method == "lerp": + return latents1 + (latents2 - latents1) * epsilon + if interpolation_method == "slerp_any": + ndims = latents1.dim() - 1 + z_size = latents1.shape[-1] + latents1_norm = latents1 / (latents1**2).sum(dim=-1, keepdim=True).sqrt().clamp_min(eps) + latents2_norm = latents2 / (latents2**2).sum(dim=-1, keepdim=True).sqrt().clamp_min(eps) + d = (latents1_norm * latents2_norm).sum(dim=-1, keepdim=True) + mask_zero = (latents1_norm.norm(dim=-1, keepdim=True) < eps) | (latents2_norm.norm(dim=-1, keepdim=True) < eps) + mask_collinear = (d > 1 - eps) | (d < -1 + eps) + mask_lerp = (mask_zero | mask_collinear).repeat([1 for _ in range(ndims)] + [z_size]) + omega = d.acos() + denom = omega.sin().clamp_min(eps) + coef_latents1 = ((1 - epsilon) * omega).sin() / denom + coef_latents2 = (epsilon * omega).sin() / denom + out = coef_latents1 * latents1 + coef_latents2 * latents2 + out[mask_lerp] = _interpolate(latents1, latents2, epsilon, interpolation_method="lerp")[mask_lerp] + return out + if interpolation_method == "slerp_unit": + out = _interpolate(latents1=latents1, latents2=latents2, epsilon=epsilon, interpolation_method="slerp_any") + return out / (out**2).sum(dim=-1, keepdim=True).sqrt().clamp_min(eps) + raise ValueError( + f"Interpolation method {interpolation_method} not supported. Choose from 'lerp', 'slerp_any', 'slerp_unit'." + ) + + +def perceptual_path_length( + generator: GeneratorType, + num_samples: int = 10_000, + conditional: bool = False, + batch_size: int = 64, + interpolation_method: Literal["lerp", "slerp_any", "slerp_unit"] = "lerp", + epsilon: float = 1e-4, + resize: Optional[int] = 64, + lower_discard: Optional[float] = 0.01, + upper_discard: Optional[float] = 0.99, + sim_net: Union[nn.Module, Literal["alex", "vgg", "squeeze"]] = "vgg", + device: Union[str, torch.device] = "cpu", +) -> tuple[Tensor, Tensor, Tensor]: + r"""Computes the perceptual path length (`PPL`_) of a generator model. + + The perceptual path length can be used to measure the consistency of interpolation in latent-space models. It is + defined as + + .. math:: + PPL = \mathbb{E}\left[\frac{1}{\epsilon^2} D(G(I(z_1, z_2, t)), G(I(z_1, z_2, t+\epsilon)))\right] + + where :math:`G` is the generator, :math:`I` is the interpolation function, :math:`D` is a similarity metric, + :math:`z_1` and :math:`z_2` are two sets of latent points, and :math:`t` is a parameter between 0 and 1. The metric + thus works by interpolating between two sets of latent points, and measuring the similarity between the generated + images. The expectation is approximated by sampling :math:`z_1` and :math:`z_2` from the generator, and averaging + the calculated distanced. The similarity metric :math:`D` is by default the `LPIPS`_ metric, but can be changed by + setting the `sim_net` argument. + + The provided generator model must have a `sample` method with signature `sample(num_samples: int) -> Tensor` where + the returned tensor has shape `(num_samples, z_size)`. If the generator is conditional, it must also have a + `num_classes` attribute. The `forward` method of the generator must have signature `forward(z: Tensor) -> Tensor` + if `conditional=False`, and `forward(z: Tensor, labels: Tensor) -> Tensor` if `conditional=True`. The returned + tensor should have shape `(num_samples, C, H, W)` and be scaled to the range [0, 255]. + + Args: + generator: Generator model, with specific requirements. See above. + num_samples: Number of samples to use for the PPL computation. + conditional: Whether the generator is conditional or not (i.e. whether it takes labels as input). + batch_size: Batch size to use for the PPL computation. + interpolation_method: Interpolation method to use. Choose from 'lerp', 'slerp_any', 'slerp_unit'. + epsilon: Spacing between the points on the path between latent points. + resize: Resize images to this size before computing the similarity between generated images. + lower_discard: Lower quantile to discard from the distances, before computing the mean and standard deviation. + upper_discard: Upper quantile to discard from the distances, before computing the mean and standard deviation. + sim_net: Similarity network to use. Can be a `nn.Module` or one of 'alex', 'vgg', 'squeeze', where the three + latter options correspond to the pretrained networks from the `LPIPS`_ paper. + device: Device to use for the computation. + + Returns: + A tuple containing the mean, standard deviation and all distances. + + Example:: + >>> import torch + >>> from torchmetrics.functional.image import perceptual_path_length + >>> class DummyGenerator(torch.nn.Module): + ... def __init__(self, z_size) -> None: + ... super().__init__() + ... self.z_size = z_size + ... self.model = torch.nn.Sequential(torch.nn.Linear(z_size, 3*128*128), torch.nn.Sigmoid()) + ... def forward(self, z): + ... return 255 * (self.model(z).reshape(-1, 3, 128, 128) + 1) + ... def sample(self, num_samples): + ... return torch.randn(num_samples, self.z_size) + >>> generator = DummyGenerator(2) + >>> perceptual_path_length(generator, num_samples=10) # doctest: +SKIP + (tensor(0.1945), + tensor(0.1222), + tensor([0.0990, 0.4173, 0.1628, 0.3573, 0.1875, 0.0335, 0.1095, 0.1887, 0.1953])) + + """ + if not _TORCHVISION_AVAILABLE: + raise ModuleNotFoundError( + "Metric `perceptual_path_length` requires torchvision which is not installed." + "Install with `pip install torchvision` or `pip install torchmetrics[image]`" + ) + _perceptual_path_length_validate_arguments( + num_samples, conditional, batch_size, interpolation_method, epsilon, resize, lower_discard, upper_discard + ) + _validate_generator_model(generator, conditional) + generator = generator.to(device) + + latent1 = generator.sample(num_samples).to(device) + latent2 = generator.sample(num_samples).to(device) + latent2 = _interpolate(latent1, latent2, epsilon, interpolation_method=interpolation_method) + + if conditional: + labels = torch.randint(0, generator.num_classes, (num_samples,)).to(device) + + if isinstance(sim_net, nn.Module): + net = sim_net.to(device) + elif sim_net in ["alex", "vgg", "squeeze"]: + net = _LPIPS(pretrained=True, net=sim_net, resize=resize).to(device) + else: + raise ValueError(f"sim_net must be a nn.Module or one of 'alex', 'vgg', 'squeeze', got {sim_net}") + + with torch.inference_mode(): + distances = [] + num_batches = math.ceil(num_samples / batch_size) + for batch_idx in range(num_batches): + batch_latent1 = latent1[batch_idx * batch_size : (batch_idx + 1) * batch_size].to(device) + batch_latent2 = latent2[batch_idx * batch_size : (batch_idx + 1) * batch_size].to(device) + + if conditional: + batch_labels = labels[batch_idx * batch_size : (batch_idx + 1) * batch_size].to(device) + outputs = generator( + torch.cat((batch_latent1, batch_latent2), dim=0), torch.cat((batch_labels, batch_labels), dim=0) + ) + else: + outputs = generator(torch.cat((batch_latent1, batch_latent2), dim=0)) + + out1, out2 = outputs.chunk(2, dim=0) + + # rescale to lpips expected domain: [0, 255] -> [0, 1] -> [-1, 1] + out1_rescale = 2 * (out1 / 255) - 1 + out2_rescale = 2 * (out2 / 255) - 1 + + similarity = net(out1_rescale, out2_rescale) + dist = similarity / epsilon**2 + distances.append(dist.detach()) + + distances = torch.cat(distances) + + lower = torch.quantile(distances, lower_discard, interpolation="lower") if lower_discard is not None else 0.0 + upper = ( + torch.quantile(distances, upper_discard, interpolation="lower") + if upper_discard is not None + else max(distances) + ) + distances = distances[(distances >= lower) & (distances <= upper)] + + return distances.mean(), distances.std(), distances diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/psnr.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/psnr.py new file mode 100644 index 0000000000000000000000000000000000000000..e058d34e7b65a909a103c31ed4d1098d222eea92 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/psnr.py @@ -0,0 +1,159 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional, Union + +import torch +from torch import Tensor, tensor +from typing_extensions import Literal + +from torchmetrics.utilities import rank_zero_warn, reduce + + +def _psnr_compute( + sum_squared_error: Tensor, + num_obs: Tensor, + data_range: Tensor, + base: float = 10.0, + reduction: Literal["elementwise_mean", "sum", "none", None] = "elementwise_mean", +) -> Tensor: + """Compute peak signal-to-noise ratio. + + Args: + sum_squared_error: Sum of square of errors over all observations + num_obs: Number of predictions or observations + data_range: the range of the data. If None, it is determined from the data (max - min). + ``data_range`` must be given when ``dim`` is not None. + base: a base of a logarithm to use + reduction: a method to reduce metric score over labels. + + - ``'elementwise_mean'``: takes the mean (default) + - ``'sum'``: takes the sum + - ``'none'`` or ``None``: no reduction will be applied + + Example: + >>> preds = torch.tensor([[0.0, 1.0], [2.0, 3.0]]) + >>> target = torch.tensor([[3.0, 2.0], [1.0, 0.0]]) + >>> data_range = target.max() - target.min() + >>> sum_squared_error, num_obs = _psnr_update(preds, target) + >>> _psnr_compute(sum_squared_error, num_obs, data_range) + tensor(2.5527) + + """ + psnr_base_e = 2 * torch.log(data_range) - torch.log(sum_squared_error / num_obs) + psnr_vals = psnr_base_e * (10 / torch.log(tensor(base))) + return reduce(psnr_vals, reduction=reduction) + + +def _psnr_update( + preds: Tensor, + target: Tensor, + dim: Optional[Union[int, tuple[int, ...]]] = None, +) -> tuple[Tensor, Tensor]: + """Update and return variables required to compute peak signal-to-noise ratio. + + Args: + preds: Predicted tensor + target: Ground truth tensor + dim: Dimensions to reduce PSNR scores over provided as either an integer or a list of integers. + Default is None meaning scores will be reduced across all dimensions. + + """ + if not preds.is_floating_point(): + preds = preds.to(torch.float32) + if not target.is_floating_point(): + target = target.to(torch.float32) + + if dim is None: + sum_squared_error = torch.sum(torch.pow(preds - target, 2)) + num_obs = tensor(target.numel(), device=target.device) + return sum_squared_error, num_obs + + diff = preds - target + sum_squared_error = torch.sum(diff * diff, dim=dim) + + dim_list = [dim] if isinstance(dim, int) else list(dim) + if not dim_list: + num_obs = tensor(target.numel(), device=target.device) + else: + num_obs = tensor(target.size(), device=target.device)[dim_list].prod() + num_obs = num_obs.expand_as(sum_squared_error) + + return sum_squared_error, num_obs + + +def peak_signal_noise_ratio( + preds: Tensor, + target: Tensor, + data_range: Optional[Union[float, tuple[float, float]]] = None, + base: float = 10.0, + reduction: Literal["elementwise_mean", "sum", "none", None] = "elementwise_mean", + dim: Optional[Union[int, tuple[int, ...]]] = None, +) -> Tensor: + """Compute the peak signal-to-noise ratio. + + Args: + preds: estimated signal + target: groun truth signal + data_range: + the range of the data. If None, it is determined from the data (max - min). If a tuple is provided then + the range is calculated as the difference and input is clamped between the values. + The ``data_range`` must be given when ``dim`` is not None. + base: a base of a logarithm to use + reduction: a method to reduce metric score over labels. + + - ``'elementwise_mean'``: takes the mean (default) + - ``'sum'``: takes the sum + - ``'none'`` or None``: no reduction will be applied + + dim: + Dimensions to reduce PSNR scores over provided as either an integer or a list of integers. Default is + None meaning scores will be reduced across all dimensions. + + Return: + Tensor with PSNR score + + Raises: + ValueError: + If ``dim`` is not ``None`` and ``data_range`` is not provided. + + Example: + >>> from torchmetrics.functional.image import peak_signal_noise_ratio + >>> pred = torch.tensor([[0.0, 1.0], [2.0, 3.0]]) + >>> target = torch.tensor([[3.0, 2.0], [1.0, 0.0]]) + >>> peak_signal_noise_ratio(pred, target) + tensor(2.5527) + + .. attention:: + Half precision is only support on GPU for this metric. + + """ + if dim is None and reduction != "elementwise_mean": + rank_zero_warn(f"The `reduction={reduction}` will not have any effect when `dim` is None.") + + if data_range is None: + if dim is not None: + # Maybe we could use `torch.amax(target, dim=dim) - torch.amin(target, dim=dim)` in PyTorch 1.7 to calculate + # `data_range` in the future. + raise ValueError("The `data_range` must be given when `dim` is not None.") + + data_range = target.max() - target.min() # type: ignore[assignment] + elif isinstance(data_range, tuple): + preds = torch.clamp(preds, min=data_range[0], max=data_range[1]) + target = torch.clamp(target, min=data_range[0], max=data_range[1]) + data_range = tensor(data_range[1] - data_range[0]) # type: ignore[assignment] + else: + data_range = tensor(float(data_range)) # type: ignore[assignment] + + sum_squared_error, num_obs = _psnr_update(preds, target, dim=dim) + return _psnr_compute(sum_squared_error, num_obs, data_range, base=base, reduction=reduction) # type: ignore[arg-type] diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/psnrb.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/psnrb.py new file mode 100644 index 0000000000000000000000000000000000000000..4a8469e19fe0edf7f6679933d81712c8b5cef0d0 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/psnrb.py @@ -0,0 +1,134 @@ +# Copyright The PyTorch Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import math + +import torch +from torch import Tensor, tensor + + +def _compute_bef(x: Tensor, block_size: int = 8) -> Tensor: + """Compute block effect. + + Args: + x: input image + block_size: integer indication the block size + + Returns: + Computed block effect + + Raises: + ValueError: + If the image is not a grayscale image + + """ + ( + _, + channels, + height, + width, + ) = x.shape + if channels > 1: + raise ValueError(f"`psnrb` metric expects grayscale images, but got images with {channels} channels.") + + h = torch.arange(width - 1) + h_b = torch.tensor(range(block_size - 1, width - 1, block_size)) + h_bc = torch.tensor(list(set(h.tolist()).symmetric_difference(h_b.tolist()))) + + v = torch.arange(height - 1) + v_b = torch.tensor(range(block_size - 1, height - 1, block_size)) + v_bc = torch.tensor(list(set(v.tolist()).symmetric_difference(v_b.tolist()))) + + d_b = (x[:, :, :, h_b] - x[:, :, :, h_b + 1]).pow(2.0).sum() + d_bc = (x[:, :, :, h_bc] - x[:, :, :, h_bc + 1]).pow(2.0).sum() + d_b += (x[:, :, v_b, :] - x[:, :, v_b + 1, :]).pow(2.0).sum() + d_bc += (x[:, :, v_bc, :] - x[:, :, v_bc + 1, :]).pow(2.0).sum() + + n_hb = height * (width / block_size) - 1 + n_hbc = (height * (width - 1)) - n_hb + n_vb = width * (height / block_size) - 1 + n_vbc = (width * (height - 1)) - n_vb + d_b /= n_hb + n_vb + d_bc /= n_hbc + n_vbc + t = math.log2(block_size) / math.log2(min(height, width)) if d_b > d_bc else 0 + return t * (d_b - d_bc) + + +def _psnrb_compute( + sum_squared_error: Tensor, + bef: Tensor, + num_obs: Tensor, + data_range: Tensor, +) -> Tensor: + """Computes peak signal-to-noise ratio. + + Args: + sum_squared_error: Sum of square of errors over all observations + bef: block effect + num_obs: Number of predictions or observations + data_range: the range of the data. If None, it is determined from the data (max - min). + + """ + sum_squared_error = sum_squared_error / num_obs + bef + if data_range > 2: + return 10 * torch.log10(data_range**2 / sum_squared_error) + return 10 * torch.log10(1.0 / sum_squared_error) + + +def _psnrb_update(preds: Tensor, target: Tensor, block_size: int = 8) -> tuple[Tensor, Tensor, Tensor]: + """Updates and returns variables required to compute peak signal-to-noise ratio. + + Args: + preds: Predicted tensor + target: Ground truth tensor + block_size: Integer indication the block size + + """ + sum_squared_error = torch.sum(torch.pow(preds - target, 2)) + num_obs = tensor(target.numel(), device=target.device) + bef = _compute_bef(preds, block_size=block_size) + return sum_squared_error, bef, num_obs + + +def peak_signal_noise_ratio_with_blocked_effect( + preds: Tensor, + target: Tensor, + block_size: int = 8, +) -> Tensor: + r"""Computes `Peak Signal to Noise Ratio With Blocked Effect` (PSNRB) metrics. + + .. math:: + \text{PSNRB}(I, J) = 10 * \log_{10} \left(\frac{\max(I)^2}{\text{MSE}(I, J)-\text{B}(I, J)}\right) + + Where :math:`\text{MSE}` denotes the `mean-squared-error`_ function. + + Args: + preds: estimated signal + target: groun truth signal + block_size: integer indication the block size + + Return: + Tensor with PSNRB score + + Example: + >>> from torch import rand + >>> from torchmetrics.functional.image import peak_signal_noise_ratio_with_blocked_effect + >>> preds = rand(1, 1, 28, 28) + >>> target = rand(1, 1, 28, 28) + >>> peak_signal_noise_ratio_with_blocked_effect(preds, target) + tensor(7.8402) + + """ + data_range = target.max() - target.min() + sum_squared_error, bef, num_obs = _psnrb_update(preds, target, block_size=block_size) + return _psnrb_compute(sum_squared_error, bef, num_obs, data_range) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/qnr.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/qnr.py new file mode 100644 index 0000000000000000000000000000000000000000..e34e6bc2473c41fb283e48802d13b5b8d1d74806 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/qnr.py @@ -0,0 +1,81 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from typing import Optional + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.image.d_lambda import spectral_distortion_index +from torchmetrics.functional.image.d_s import spatial_distortion_index +from torchmetrics.utilities.imports import _TORCHVISION_AVAILABLE + +if not _TORCHVISION_AVAILABLE: + __doctest_skip__ = ["quality_with_no_reference"] + + +def quality_with_no_reference( + preds: Tensor, + ms: Tensor, + pan: Tensor, + pan_lr: Optional[Tensor] = None, + alpha: float = 1, + beta: float = 1, + norm_order: int = 1, + window_size: int = 7, + reduction: Literal["elementwise_mean", "sum", "none"] = "elementwise_mean", +) -> Tensor: + """Calculate `Quality with No Reference`_ (QualityWithNoReference_) also known as QNR. + + Metric is used to compare the joint spectral and spatial distortion between two images. + + Args: + preds: High resolution multispectral image. + ms: Low resolution multispectral image. + pan: High resolution panchromatic image. + pan_lr: Low resolution panchromatic image. + alpha: Relevance of spectral distortion. + beta: Relevance of spatial distortion. + norm_order: Order of the norm applied on the difference. + window_size: Window size of the filter applied to degrade the high resolution panchromatic image. + reduction: A method to reduce metric score over labels. + + - ``'elementwise_mean'``: takes the mean (default) + - ``'sum'``: takes the sum + - ``'none'``: no reduction will be applied + + Return: + Tensor with QualityWithNoReference score + + Raises: + ValueError: + If ``alpha`` or ``beta`` is not a non-negative real number. + + Example: + >>> from torch import rand + >>> from torchmetrics.functional.image import quality_with_no_reference + >>> preds = rand([16, 3, 32, 32]) + >>> ms = rand([16, 3, 16, 16]) + >>> pan = rand([16, 3, 32, 32]) + >>> quality_with_no_reference(preds, ms, pan) + tensor(0.9694) + + """ + if not isinstance(alpha, (int, float)) or alpha < 0: + raise ValueError(f"Expected `alpha` to be a non-negative real number. Got alpha: {alpha}.") + if not isinstance(beta, (int, float)) or beta < 0: + raise ValueError(f"Expected `beta` to be a non-negative real number. Got beta: {beta}.") + d_lambda = spectral_distortion_index(preds, ms, norm_order, reduction) + d_s = spatial_distortion_index(preds, ms, pan, pan_lr, norm_order, window_size, reduction) + return (1 - d_lambda) ** alpha * (1 - d_s) ** beta diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/rase.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/rase.py new file mode 100644 index 0000000000000000000000000000000000000000..51181852aa640851ba91ce8373e6cabdd33113b4 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/rase.py @@ -0,0 +1,102 @@ +# Copyright The PyTorch Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +import torch +from torch import Tensor + +from torchmetrics.functional.image.rmse_sw import _rmse_sw_compute, _rmse_sw_update +from torchmetrics.functional.image.utils import _uniform_filter + + +def _rase_update( + preds: Tensor, target: Tensor, window_size: int, rmse_map: Tensor, target_sum: Tensor, total_images: Tensor +) -> tuple[Tensor, Tensor, Tensor]: + """Calculate the sum of RMSE map values for the batch of examples and update intermediate states. + + Args: + preds: Deformed image + target: Ground truth image + window_size: Sliding window used for RMSE calculation + rmse_map: Sum of RMSE map values over all examples + target_sum: target... + total_images: Total number of images + + Return: + Intermediate state of RMSE map + Updated total number of already processed images + + """ + _, rmse_map, total_images = _rmse_sw_update( + preds, target, window_size, rmse_val_sum=None, rmse_map=rmse_map, total_images=total_images + ) + target_sum += torch.sum(_uniform_filter(target, window_size) / (window_size**2), dim=0) + return rmse_map, target_sum, total_images + + +def _rase_compute(rmse_map: Tensor, target_sum: Tensor, total_images: Tensor, window_size: int) -> Tensor: + """Compute RASE. + + Args: + rmse_map: Sum of RMSE map values over all examples + target_sum: target... + total_images: Total number of images. + window_size: Sliding window used for rmse calculation + + Return: + Relative Average Spectral Error (RASE) + + """ + _, rmse_map = _rmse_sw_compute(rmse_val_sum=None, rmse_map=rmse_map, total_images=total_images) + target_mean = target_sum / total_images + target_mean = target_mean.mean(0) # mean over image channels + rase_map = 100 / target_mean * torch.sqrt(torch.mean(rmse_map**2, 0)) + crop_slide = round(window_size / 2) + + return torch.mean(rase_map[crop_slide:-crop_slide, crop_slide:-crop_slide]) + + +def relative_average_spectral_error(preds: Tensor, target: Tensor, window_size: int = 8) -> Tensor: + """Compute Relative Average Spectral Error (RASE) (RelativeAverageSpectralError_). + + Args: + preds: Deformed image + target: Ground truth image + window_size: Sliding window used for rmse calculation + + Return: + Relative Average Spectral Error (RASE) + + Example: + >>> from torch import rand + >>> from torchmetrics.functional.image import relative_average_spectral_error + >>> preds = rand(4, 3, 16, 16) + >>> target = rand(4, 3, 16, 16) + >>> relative_average_spectral_error(preds, target) + tensor(5326.40...) + + Raises: + ValueError: If ``window_size`` is not a positive integer. + + """ + if not isinstance(window_size, int) or (isinstance(window_size, int) and window_size < 1): + raise ValueError("Argument `window_size` is expected to be a positive integer.") + + img_shape = target.shape[1:] # [num_channels, width, height] + rmse_map = torch.zeros(img_shape, dtype=target.dtype, device=target.device) + target_sum = torch.zeros(img_shape, dtype=target.dtype, device=target.device) + total_images = torch.tensor(0.0, device=target.device) + + rmse_map, target_sum, total_images = _rase_update(preds, target, window_size, rmse_map, target_sum, total_images) + return _rase_compute(rmse_map, target_sum, total_images, window_size) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/scc.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/scc.py new file mode 100644 index 0000000000000000000000000000000000000000..d680a22b67685d1947ea7e850215b854571f87ea --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/scc.py @@ -0,0 +1,220 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import math +from typing import Optional, Union + +import torch +from torch import Tensor, tensor +from torch.nn.functional import conv2d, pad +from typing_extensions import Literal + +from torchmetrics.utilities.checks import _check_same_shape +from torchmetrics.utilities.distributed import reduce + + +def _scc_update(preds: Tensor, target: Tensor, hp_filter: Tensor, window_size: int) -> tuple[Tensor, Tensor, Tensor]: + """Update and returns variables required to compute Spatial Correlation Coefficient. + + Args: + preds: Predicted tensor + target: Ground truth tensor + hp_filter: High-pass filter tensor + window_size: Local window size integer + + Return: + Tuple of (preds, target, hp_filter) tensors + + Raises: + ValueError: + If ``preds`` and ``target`` have different number of channels + If ``preds`` and ``target`` have different shapes + If ``preds`` and ``target`` have invalid shapes + If ``window_size`` is not a positive integer + If ``window_size`` is greater than the size of the image + + """ + if preds.dtype != target.dtype: + target = target.to(preds.dtype) + _check_same_shape(preds, target) + if preds.ndim not in (3, 4): + raise ValueError( + "Expected `preds` and `target` to have batch of colored images with BxCxHxW shape" + " or batch of grayscale images of BxHxW shape." + f" Got preds: {preds.shape} and target: {target.shape}." + ) + + if len(preds.shape) == 3: + preds = preds.unsqueeze(1) + target = target.unsqueeze(1) + + if not window_size > 0: + raise ValueError(f"Expected `window_size` to be a positive integer. Got {window_size}.") + + if window_size > preds.size(2) or window_size > preds.size(3): + raise ValueError( + f"Expected `window_size` to be less than or equal to the size of the image." + f" Got window_size: {window_size} and image size: {preds.size(2)}x{preds.size(3)}." + ) + + preds = preds.to(torch.float32) + target = target.to(torch.float32) + hp_filter = hp_filter[None, None, :].to(dtype=preds.dtype, device=preds.device) + return preds, target, hp_filter + + +def _symmetric_reflect_pad_2d(input_img: Tensor, pad: Union[int, tuple[int, ...]]) -> Tensor: + """Applies symmetric padding to the 2D image tensor input using ``reflect`` mode (d c b a | a b c d | d c b a).""" + if isinstance(pad, int): + pad = (pad, pad, pad, pad) + if len(pad) != 4: + raise ValueError(f"Expected padding to have length 4, but got {len(pad)}") + + left_pad = input_img[:, :, :, 0 : pad[0]].flip(dims=[3]) + right_pad = input_img[:, :, :, -pad[1] :].flip(dims=[3]) + padded = torch.cat([left_pad, input_img, right_pad], dim=3) + + top_pad = padded[:, :, 0 : pad[2], :].flip(dims=[2]) + bottom_pad = padded[:, :, -pad[3] :, :].flip(dims=[2]) + return torch.cat([top_pad, padded, bottom_pad], dim=2) + + +def _signal_convolve_2d(input_img: Tensor, kernel: Tensor) -> Tensor: + """Applies 2D signal convolution to the input tensor with the given kernel.""" + left_padding = int(math.floor((kernel.size(3) - 1) / 2)) + right_padding = int(math.ceil((kernel.size(3) - 1) / 2)) + top_padding = int(math.floor((kernel.size(2) - 1) / 2)) + bottom_padding = int(math.ceil((kernel.size(2) - 1) / 2)) + + padded = _symmetric_reflect_pad_2d(input_img, pad=(left_padding, right_padding, top_padding, bottom_padding)) + kernel = kernel.flip([2, 3]) + return conv2d(padded, kernel, stride=1, padding=0) + + +def _hp_2d_laplacian(input_img: Tensor, kernel: Tensor) -> Tensor: + """Applies 2-D Laplace filter to the input tensor with the given high pass filter.""" + return _signal_convolve_2d(input_img, kernel) * 2.0 + + +def _local_variance_covariance(preds: Tensor, target: Tensor, window: Tensor) -> tuple[Tensor, Tensor, Tensor]: + """Computes local variance and covariance of the input tensors.""" + # This code is inspired by + # https://github.com/andrewekhalel/sewar/blob/master/sewar/full_ref.py#L187. + + left_padding = int(math.ceil((window.size(3) - 1) / 2)) + right_padding = int(math.floor((window.size(3) - 1) / 2)) + + preds = pad(preds, (left_padding, right_padding, left_padding, right_padding)) + target = pad(target, (left_padding, right_padding, left_padding, right_padding)) + + preds_mean = conv2d(preds, window, stride=1, padding=0) + target_mean = conv2d(target, window, stride=1, padding=0) + + preds_var = conv2d(preds**2, window, stride=1, padding=0) - preds_mean**2 + target_var = conv2d(target**2, window, stride=1, padding=0) - target_mean**2 + target_preds_cov = conv2d(target * preds, window, stride=1, padding=0) - target_mean * preds_mean + + return preds_var, target_var, target_preds_cov + + +def _scc_per_channel_compute(preds: Tensor, target: Tensor, hp_filter: Tensor, window_size: int) -> Tensor: + """Computes per channel Spatial Correlation Coefficient. + + Args: + preds: estimated image of Bx1xHxW shape. + target: ground truth image of Bx1xHxW shape. + hp_filter: 2D high-pass filter. + window_size: size of window for local mean calculation. + + Return: + Tensor with Spatial Correlation Coefficient score + + """ + dtype = preds.dtype + device = preds.device + + # This code is inspired by + # https://github.com/andrewekhalel/sewar/blob/master/sewar/full_ref.py#L187. + + window = torch.ones(size=(1, 1, window_size, window_size), dtype=dtype, device=device) / (window_size**2) + + preds_hp = _hp_2d_laplacian(preds, hp_filter) + target_hp = _hp_2d_laplacian(target, hp_filter) + + preds_var, target_var, target_preds_cov = _local_variance_covariance(preds_hp, target_hp, window) + + preds_var[preds_var < 0] = 0 + target_var[target_var < 0] = 0 + + den = torch.sqrt(target_var) * torch.sqrt(preds_var) + idx = den == 0 + den[den == 0] = 1 + scc = target_preds_cov / den + scc[idx] = 0 + return scc + + +def spatial_correlation_coefficient( + preds: Tensor, + target: Tensor, + hp_filter: Optional[Tensor] = None, + window_size: int = 8, + reduction: Optional[Literal["mean", "none", None]] = "mean", +) -> Tensor: + """Compute Spatial Correlation Coefficient (SCC_). + + Args: + preds: predicted images of shape ``(N,C,H,W)`` or ``(N,H,W)``. + target: ground truth images of shape ``(N,C,H,W)`` or ``(N,H,W)``. + hp_filter: High-pass filter tensor. default: tensor([[-1,-1,-1],[-1,8,-1],[-1,-1,-1]]) + window_size: Local window size integer. default: 8, + reduction: Reduction method for output tensor. If ``None`` or ``"none"``, + returns a tensor with the per sample results. default: ``"mean"``. + + Return: + Tensor with scc score + + Example: + >>> from torch import randn + >>> from torchmetrics.functional.image import spatial_correlation_coefficient as scc + >>> x = randn(5, 3, 16, 16) + >>> scc(x, x) + tensor(1.) + >>> x = randn(5, 16, 16) + >>> scc(x, x) + tensor(1.) + >>> x = randn(5, 3, 16, 16) + >>> y = randn(5, 3, 16, 16) + >>> scc(x, y, reduction="none") + tensor([0.0223, 0.0256, 0.0616, 0.0159, 0.0170]) + + """ + if hp_filter is None: + hp_filter = tensor([[-1, -1, -1], [-1, 8, -1], [-1, -1, -1]]) + if reduction is None: + reduction = "none" + if reduction not in ("mean", "none"): + raise ValueError(f"Expected reduction to be 'mean' or 'none', but got {reduction}") + preds, target, hp_filter = _scc_update(preds, target, hp_filter, window_size) + + per_channel = [ + _scc_per_channel_compute( + preds[:, i, :, :].unsqueeze(1), target[:, i, :, :].unsqueeze(1), hp_filter, window_size + ) + for i in range(preds.size(1)) + ] + if reduction == "none": + return torch.mean(torch.cat(per_channel, dim=1), dim=[1, 2, 3]) + if reduction == "mean": + return reduce(torch.cat(per_channel, dim=1), reduction="elementwise_mean") + return None diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/uqi.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/uqi.py new file mode 100644 index 0000000000000000000000000000000000000000..2b8d6f3cfa77ac62d38592b23f30fbf72d62c5bf --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/image/uqi.py @@ -0,0 +1,171 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from collections.abc import Sequence +from typing import Optional + +import torch +from torch import Tensor, nn +from typing_extensions import Literal + +from torchmetrics.functional.image.utils import _gaussian_kernel_2d +from torchmetrics.utilities.checks import _check_same_shape +from torchmetrics.utilities.distributed import reduce + + +def _uqi_update(preds: Tensor, target: Tensor) -> tuple[Tensor, Tensor]: + """Update and returns variables required to compute Universal Image Quality Index. + + Args: + preds: Predicted tensor + target: Ground truth tensor + + """ + if preds.dtype != target.dtype: + raise TypeError( + "Expected `preds` and `target` to have the same data type." + f" Got preds: {preds.dtype} and target: {target.dtype}." + ) + _check_same_shape(preds, target) + if len(preds.shape) != 4: + raise ValueError( + f"Expected `preds` and `target` to have BxCxHxW shape. Got preds: {preds.shape} and target: {target.shape}." + ) + return preds, target + + +def _uqi_compute( + preds: Tensor, + target: Tensor, + kernel_size: Sequence[int] = (11, 11), + sigma: Sequence[float] = (1.5, 1.5), + reduction: Optional[Literal["elementwise_mean", "sum", "none"]] = "elementwise_mean", +) -> Tensor: + """Compute Universal Image Quality Index. + + Args: + preds: estimated image + target: ground truth image + kernel_size: size of the gaussian kernel + sigma: Standard deviation of the gaussian kernel + reduction: a method to reduce metric score over labels. + + - ``'elementwise_mean'``: takes the mean (default) + - ``'sum'``: takes the sum + - ``'none'`` or ``None``: no reduction will be applied + + Example: + >>> preds = torch.rand([16, 1, 16, 16]) + >>> target = preds * 0.75 + >>> preds, target = _uqi_update(preds, target) + >>> _uqi_compute(preds, target) + tensor(0.9216) + + """ + if len(kernel_size) != 2 or len(sigma) != 2: + raise ValueError( + "Expected `kernel_size` and `sigma` to have the length of two." + f" Got kernel_size: {len(kernel_size)} and sigma: {len(sigma)}." + ) + + if any(x % 2 == 0 or x <= 0 for x in kernel_size): + raise ValueError(f"Expected `kernel_size` to have odd positive number. Got {kernel_size}.") + + if any(y <= 0 for y in sigma): + raise ValueError(f"Expected `sigma` to have positive number. Got {sigma}.") + + device = preds.device + channel = preds.size(1) + dtype = preds.dtype + kernel = _gaussian_kernel_2d(channel, kernel_size, sigma, dtype, device) + pad_h = (kernel_size[0] - 1) // 2 + pad_w = (kernel_size[1] - 1) // 2 + + preds = nn.functional.pad(preds, (pad_h, pad_h, pad_w, pad_w), mode="reflect") + target = nn.functional.pad(target, (pad_h, pad_h, pad_w, pad_w), mode="reflect") + + input_list = torch.cat((preds, target, preds * preds, target * target, preds * target)) # (5 * B, C, H, W) + outputs = nn.functional.conv2d(input_list, kernel, groups=channel) + output_list = outputs.split(preds.shape[0]) + + mu_pred_sq = output_list[0].pow(2) + mu_target_sq = output_list[1].pow(2) + mu_pred_target = output_list[0] * output_list[1] + + # Calculate the variance of the predicted and target images, should be non-negative + sigma_pred_sq = torch.clamp(output_list[2] - mu_pred_sq, min=0.0) + sigma_target_sq = torch.clamp(output_list[3] - mu_target_sq, min=0.0) + sigma_pred_target = output_list[4] - mu_pred_target + + upper = 2 * sigma_pred_target + lower = sigma_pred_sq + sigma_target_sq + eps = torch.finfo(sigma_pred_sq.dtype).eps + uqi_idx = ((2 * mu_pred_target) * upper) / ((mu_pred_sq + mu_target_sq) * lower + eps) + uqi_idx = uqi_idx[..., pad_h:-pad_h, pad_w:-pad_w] + + return reduce(uqi_idx, reduction) + + +def universal_image_quality_index( + preds: Tensor, + target: Tensor, + kernel_size: Sequence[int] = (11, 11), + sigma: Sequence[float] = (1.5, 1.5), + reduction: Optional[Literal["elementwise_mean", "sum", "none"]] = "elementwise_mean", +) -> Tensor: + """Universal Image Quality Index. + + Args: + preds: estimated image + target: ground truth image + kernel_size: size of the gaussian kernel + sigma: Standard deviation of the gaussian kernel + reduction: a method to reduce metric score over labels. + + - ``'elementwise_mean'``: takes the mean (default) + - ``'sum'``: takes the sum + - ``'none'`` or ``None``: no reduction will be applied + + Return: + Tensor with UniversalImageQualityIndex score + + Raises: + TypeError: + If ``preds`` and ``target`` don't have the same data type. + ValueError: + If ``preds`` and ``target`` don't have ``BxCxHxW shape``. + ValueError: + If the length of ``kernel_size`` or ``sigma`` is not ``2``. + ValueError: + If one of the elements of ``kernel_size`` is not an ``odd positive number``. + ValueError: + If one of the elements of ``sigma`` is not a ``positive number``. + + Example: + >>> from torchmetrics.functional.image import universal_image_quality_index + >>> preds = torch.rand([16, 1, 16, 16]) + >>> target = preds * 0.75 + >>> universal_image_quality_index(preds, target) + tensor(0.9216) + + References: + [1] Zhou Wang and A. C. Bovik, "A universal image quality index," in IEEE Signal Processing Letters, vol. 9, + no. 3, pp. 81-84, March 2002, doi: 10.1109/97.995823. + + [2] Zhou Wang, A. C. Bovik, H. R. Sheikh and E. P. Simoncelli, "Image quality assessment: from error visibility + to structural similarity," in IEEE Transactions on Image Processing, vol. 13, no. 4, pp. 600-612, April 2004, + doi: 10.1109/TIP.2003.819861. + + """ + preds, target = _uqi_update(preds, target) + return _uqi_compute(preds, target, kernel_size, sigma, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/__init__.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..772cb39589595b698e44a1b6ad59e88830ac2079 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/__init__.py @@ -0,0 +1,34 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from torchmetrics.functional.nominal.cramers import cramers_v, cramers_v_matrix +from torchmetrics.functional.nominal.fleiss_kappa import fleiss_kappa +from torchmetrics.functional.nominal.pearson import ( + pearsons_contingency_coefficient, + pearsons_contingency_coefficient_matrix, +) +from torchmetrics.functional.nominal.theils_u import theils_u, theils_u_matrix +from torchmetrics.functional.nominal.tschuprows import tschuprows_t, tschuprows_t_matrix + +__all__ = [ + "cramers_v", + "cramers_v_matrix", + "fleiss_kappa", + "pearsons_contingency_coefficient", + "pearsons_contingency_coefficient_matrix", + "theils_u", + "theils_u_matrix", + "tschuprows_t", + "tschuprows_t_matrix", +] diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/cramers.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/cramers.py new file mode 100644 index 0000000000000000000000000000000000000000..33b89b92014604dddf462699aff643556276c1b2 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/cramers.py @@ -0,0 +1,183 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import itertools +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.classification.confusion_matrix import _multiclass_confusion_matrix_update +from torchmetrics.functional.nominal.utils import ( + _compute_bias_corrected_values, + _compute_chi_squared, + _drop_empty_rows_and_cols, + _handle_nan_in_data, + _nominal_input_validation, + _unable_to_use_bias_correction_warning, +) + + +def _cramers_v_update( + preds: Tensor, + target: Tensor, + num_classes: int, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + """Compute the bins to update the confusion matrix with for Cramer's V calculation. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data + target: 1D or 2D tensor of categorical (nominal) data + num_classes: Integer specifying the number of classes + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN`s when ``nan_strategy = 'replace``` + + Returns: + Non-reduced confusion matrix + + """ + preds = preds.argmax(1) if preds.ndim == 2 else preds + target = target.argmax(1) if target.ndim == 2 else target + preds, target = _handle_nan_in_data(preds, target, nan_strategy, nan_replace_value) + return _multiclass_confusion_matrix_update(preds, target, num_classes) + + +def _cramers_v_compute(confmat: Tensor, bias_correction: bool) -> Tensor: + """Compute Cramers' V statistic based on a pre-computed confusion matrix. + + Args: + confmat: Confusion matrix for observed data + bias_correction: Indication of whether to use bias correction. + + Returns: + Cramer's V statistic + + """ + confmat = _drop_empty_rows_and_cols(confmat) + cm_sum = confmat.sum() + chi_squared = _compute_chi_squared(confmat, bias_correction) + phi_squared = chi_squared / cm_sum + num_rows, num_cols = confmat.shape + + if bias_correction: + phi_squared_corrected, rows_corrected, cols_corrected = _compute_bias_corrected_values( + phi_squared, num_rows, num_cols, cm_sum + ) + if torch.min(rows_corrected, cols_corrected) == 1: + _unable_to_use_bias_correction_warning(metric_name="Cramer's V") + return torch.tensor(float("nan"), device=confmat.device) + cramers_v_value = torch.sqrt(phi_squared_corrected / torch.min(rows_corrected - 1, cols_corrected - 1)) + else: + cramers_v_value = torch.sqrt(phi_squared / min(num_rows - 1, num_cols - 1)) + return cramers_v_value.clamp(0.0, 1.0) + + +def cramers_v( + preds: Tensor, + target: Tensor, + bias_correction: bool = True, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Cramer's V`_ statistic measuring the association between two categorical (nominal) data series. + + .. math:: + V = \sqrt{\frac{\chi^2 / n}{\min(r - 1, k - 1)}} + + where + + .. math:: + \chi^2 = \sum_{i,j} \ frac{\left(n_{ij} - \frac{n_{i.} n_{.j}}{n}\right)^2}{\frac{n_{i.} n_{.j}}{n}} + + where :math:`n_{ij}` denotes the number of times the values :math:`(A_i, B_j)` are observed with :math:`A_i, B_j` + represent frequencies of values in ``preds`` and ``target``, respectively. + + Cramer's V is a symmetric coefficient, i.e. :math:`V(preds, target) = V(target, preds)`. + + The output values lies in [0, 1] with 1 meaning the perfect association. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + target: 1D or 2D tensor of categorical (nominal) data + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + bias_correction: Indication of whether to use bias correction. + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Cramer's V statistic + + Example: + >>> from torch import randint, round + >>> from torchmetrics.functional.nominal import cramers_v + >>> preds = randint(0, 4, (100,)) + >>> target = round(preds + torch.randn(100)).clamp(0, 4) + >>> cramers_v(preds, target) + tensor(0.5284) + + """ + _nominal_input_validation(nan_strategy, nan_replace_value) + num_classes = len(torch.cat([preds, target]).unique()) + confmat = _cramers_v_update(preds, target, num_classes, nan_strategy, nan_replace_value) + return _cramers_v_compute(confmat, bias_correction) + + +def cramers_v_matrix( + matrix: Tensor, + bias_correction: bool = True, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Cramer's V`_ statistic between a set of multiple variables. + + This can serve as a convenient tool to compute Cramer's V statistic for analyses of correlation between categorical + variables in your dataset. + + Args: + matrix: A tensor of categorical (nominal) data, where: + - rows represent a number of data points + - columns represent a number of categorical (nominal) features + bias_correction: Indication of whether to use bias correction. + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Cramer's V statistic for a dataset of categorical variables + + Example: + >>> from torch import randint + >>> from torchmetrics.functional.nominal import cramers_v_matrix + >>> matrix = randint(0, 4, (200, 5)) + >>> cramers_v_matrix(matrix) + tensor([[1.0000, 0.0637, 0.0000, 0.0542, 0.1337], + [0.0637, 1.0000, 0.0000, 0.0000, 0.0000], + [0.0000, 0.0000, 1.0000, 0.0000, 0.0649], + [0.0542, 0.0000, 0.0000, 1.0000, 0.1100], + [0.1337, 0.0000, 0.0649, 0.1100, 1.0000]]) + + """ + _nominal_input_validation(nan_strategy, nan_replace_value) + num_variables = matrix.shape[1] + cramers_v_matrix_value = torch.ones(num_variables, num_variables, device=matrix.device) + for i, j in itertools.combinations(range(num_variables), 2): + x, y = matrix[:, i], matrix[:, j] + num_classes = len(torch.cat([x, y]).unique()) + confmat = _cramers_v_update(x, y, num_classes, nan_strategy, nan_replace_value) + cramers_v_matrix_value[i, j] = cramers_v_matrix_value[j, i] = _cramers_v_compute(confmat, bias_correction) + return cramers_v_matrix_value diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/fleiss_kappa.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/fleiss_kappa.py new file mode 100644 index 0000000000000000000000000000000000000000..69990f552d80bfb2babb16bfd804b6d78f53fc7d --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/fleiss_kappa.py @@ -0,0 +1,99 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import torch +from torch import Tensor +from typing_extensions import Literal + + +def _fleiss_kappa_update(ratings: Tensor, mode: Literal["counts", "probs"] = "counts") -> Tensor: + """Updates the counts for fleiss kappa metric. + + Args: + ratings: ratings matrix + mode: whether ratings are provided as counts or probabilities + + """ + if mode == "probs": + if ratings.ndim != 3 or not ratings.is_floating_point(): + raise ValueError( + "If argument ``mode`` is 'probs', ratings must have 3 dimensions with the format" + " [n_samples, n_categories, n_raters] and be floating point." + ) + ratings = ratings.argmax(dim=1) + one_hot = torch.nn.functional.one_hot(ratings, num_classes=ratings.shape[1]).permute(0, 2, 1) + ratings = one_hot.sum(dim=-1) + elif mode == "counts" and (ratings.ndim != 2 or ratings.is_floating_point()): + raise ValueError( + "If argument ``mode`` is `counts`, ratings must have 2 dimensions with the format" + " [n_samples, n_categories] and be none floating point." + ) + return ratings + + +def _fleiss_kappa_compute(counts: Tensor) -> Tensor: + """Computes fleiss kappa from counts matrix. + + Args: + counts: counts matrix of shape [n_samples, n_categories] + + """ + total = counts.shape[0] + num_raters = counts.sum(1).max() + + p_i = counts.sum(dim=0) / (total * num_raters) + p_j = ((counts**2).sum(dim=1) - num_raters) / (num_raters * (num_raters - 1)) + p_bar = p_j.mean() + pe_bar = (p_i**2).sum() + return (p_bar - pe_bar) / (1 - pe_bar + 1e-5) + + +def fleiss_kappa(ratings: Tensor, mode: Literal["counts", "probs"] = "counts") -> Tensor: + r"""Calculatees `Fleiss kappa`_ a statistical measure for inter agreement between raters. + + .. math:: + \kappa = \frac{\bar{p} - \bar{p_e}}{1 - \bar{p_e}} + + where :math:`\bar{p}` is the mean of the agreement probability over all raters and :math:`\bar{p_e}` is the mean + agreement probability over all raters if they were randomly assigned. If the raters are in complete agreement then + the score 1 is returned, if there is no agreement among the raters (other than what would be expected by chance) + then a score smaller than 0 is returned. + + Args: + ratings: Ratings of shape [n_samples, n_categories] or [n_samples, n_categories, n_raters] depedenent on `mode`. + If `mode` is `counts`, `ratings` must be integer and contain the number of raters that chose each category. + If `mode` is `probs`, `ratings` must be floating point and contain the probability/logits that each rater + chose each category. + mode: Whether `ratings` will be provided as counts or probabilities. + + Example: + >>> # Ratings are provided as counts + >>> from torch import randint + >>> from torchmetrics.functional.nominal import fleiss_kappa + >>> ratings = randint(0, 10, size=(100, 5)).long() # 100 samples, 5 categories, 10 raters + >>> fleiss_kappa(ratings) + tensor(0.0089) + + Example: + >>> # Ratings are provided as probabilities + >>> from torch import randn + >>> from torchmetrics.functional.nominal import fleiss_kappa + >>> ratings = randn(100, 5, 10).softmax(dim=1) # 100 samples, 5 categories, 10 raters + >>> fleiss_kappa(ratings, mode='probs') + tensor(-0.0075) + + """ + if mode not in ["counts", "probs"]: + raise ValueError("Argument ``mode`` must be one of ['counts', 'probs'].") + counts = _fleiss_kappa_update(ratings, mode) + return _fleiss_kappa_compute(counts) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/pearson.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/pearson.py new file mode 100644 index 0000000000000000000000000000000000000000..55fe1681bf754654e863aff3c572ea69837212c8 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/pearson.py @@ -0,0 +1,174 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import itertools +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.classification.confusion_matrix import _multiclass_confusion_matrix_update +from torchmetrics.functional.nominal.utils import ( + _compute_chi_squared, + _drop_empty_rows_and_cols, + _handle_nan_in_data, + _nominal_input_validation, +) + + +def _pearsons_contingency_coefficient_update( + preds: Tensor, + target: Tensor, + num_classes: int, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + """Compute the bins to update the confusion matrix with for Pearson's Contingency Coefficient calculation. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data + target: 1D or 2D tensor of categorical (nominal) data + num_classes: Integer specifying the number of classes + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN`s when ``nan_strategy = 'replace``` + + Returns: + Non-reduced confusion matrix + + """ + preds = preds.argmax(1) if preds.ndim == 2 else preds + target = target.argmax(1) if target.ndim == 2 else target + preds, target = _handle_nan_in_data(preds, target, nan_strategy, nan_replace_value) + return _multiclass_confusion_matrix_update(preds, target, num_classes) + + +def _pearsons_contingency_coefficient_compute(confmat: Tensor) -> Tensor: + """Compute Pearson's Contingency Coefficient based on a pre-computed confusion matrix. + + Args: + confmat: Confusion matrix for observed data + + Returns: + Pearson's Contingency Coefficient + + """ + confmat = _drop_empty_rows_and_cols(confmat) + cm_sum = confmat.sum() + chi_squared = _compute_chi_squared(confmat, bias_correction=False) + phi_squared = chi_squared / cm_sum + + tschuprows_t_value = torch.sqrt(phi_squared / (1 + phi_squared)) + return tschuprows_t_value.clamp(0.0, 1.0) + + +def pearsons_contingency_coefficient( + preds: Tensor, + target: Tensor, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Pearson's Contingency Coefficient`_ for measuring the association between two categorical data series. + + .. math:: + Pearson = \sqrt{\frac{\chi^2 / n}{1 + \chi^2 / n}} + + where + + .. math:: + \chi^2 = \sum_{i,j} \ frac{\left(n_{ij} - \frac{n_{i.} n_{.j}}{n}\right)^2}{\frac{n_{i.} n_{.j}}{n}} + + where :math:`n_{ij}` denotes the number of times the values :math:`(A_i, B_j)` are observed with :math:`A_i, B_j` + represent frequencies of values in ``preds`` and ``target``, respectively. + + Pearson's Contingency Coefficient is a symmetric coefficient, i.e. + :math:`Pearson(preds, target) = Pearson(target, preds)`. + + The output values lies in [0, 1] with 1 meaning the perfect association. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data: + + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + + target: 1D or 2D tensor of categorical (nominal) data: + + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Pearson's Contingency Coefficient + + Example: + >>> from torch import randint, round + >>> from torchmetrics.functional.nominal import pearsons_contingency_coefficient + >>> preds = randint(0, 4, (100,)) + >>> target = round(preds + torch.randn(100)).clamp(0, 4) + >>> pearsons_contingency_coefficient(preds, target) + tensor(0.6948) + + """ + _nominal_input_validation(nan_strategy, nan_replace_value) + num_classes = len(torch.cat([preds, target]).unique()) + confmat = _pearsons_contingency_coefficient_update(preds, target, num_classes, nan_strategy, nan_replace_value) + return _pearsons_contingency_coefficient_compute(confmat) + + +def pearsons_contingency_coefficient_matrix( + matrix: Tensor, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Pearson's Contingency Coefficient`_ statistic between a set of multiple variables. + + This can serve as a convenient tool to compute Pearson's Contingency Coefficient for analyses + of correlation between categorical variables in your dataset. + + Args: + matrix: A tensor of categorical (nominal) data, where: + + - rows represent a number of data points + - columns represent a number of categorical (nominal) features + + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Pearson's Contingency Coefficient statistic for a dataset of categorical variables + + Example: + >>> from torch import randint + >>> from torchmetrics.functional.nominal import pearsons_contingency_coefficient_matrix + >>> matrix = randint(0, 4, (200, 5)) + >>> pearsons_contingency_coefficient_matrix(matrix) + tensor([[1.0000, 0.2326, 0.1959, 0.2262, 0.2989], + [0.2326, 1.0000, 0.1386, 0.1895, 0.1329], + [0.1959, 0.1386, 1.0000, 0.1840, 0.2335], + [0.2262, 0.1895, 0.1840, 1.0000, 0.2737], + [0.2989, 0.1329, 0.2335, 0.2737, 1.0000]]) + + """ + _nominal_input_validation(nan_strategy, nan_replace_value) + num_variables = matrix.shape[1] + pearsons_cont_coef_matrix_value = torch.ones(num_variables, num_variables, device=matrix.device) + for i, j in itertools.combinations(range(num_variables), 2): + x, y = matrix[:, i], matrix[:, j] + num_classes = len(torch.cat([x, y]).unique()) + confmat = _pearsons_contingency_coefficient_update(x, y, num_classes, nan_strategy, nan_replace_value) + val = _pearsons_contingency_coefficient_compute(confmat) + pearsons_cont_coef_matrix_value[i, j] = pearsons_cont_coef_matrix_value[j, i] = val + return pearsons_cont_coef_matrix_value diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/theils_u.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/theils_u.py new file mode 100644 index 0000000000000000000000000000000000000000..f356dbfd03d728a5cc69b8c4dca69162ebcdd791 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/theils_u.py @@ -0,0 +1,195 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import itertools +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.classification.confusion_matrix import _multiclass_confusion_matrix_update +from torchmetrics.functional.nominal.utils import ( + _drop_empty_rows_and_cols, + _handle_nan_in_data, + _nominal_input_validation, +) + + +def _conditional_entropy_compute(confmat: Tensor) -> Tensor: + r"""Compute Conditional Entropy Statistic based on a pre-computed confusion matrix. + + .. math:: + H(X|Y) = \sum_{x, y ~ (X, Y)} p(x, y)\frac{p(y)}{p(x, y)} + + Args: + confmat: Confusion matrix for observed data + + Returns: + Conditional Entropy Value + + """ + confmat = _drop_empty_rows_and_cols(confmat) + total_occurrences = confmat.sum() + # iterate over all i, j combinations + p_xy_m = confmat / total_occurrences + # get p_y by summing over x dim (=1) + p_y = confmat.sum(1) / total_occurrences + # repeat over rows (shape = p_xy_m.shape[1]) for tensor multiplication + p_y_m = p_y.unsqueeze(1).repeat(1, p_xy_m.shape[1]) + + # entropy calculated as p_xy * log (p_xy / p_y) + return torch.nansum(p_xy_m * torch.log(p_y_m / p_xy_m)) + + +def _theils_u_update( + preds: Tensor, + target: Tensor, + num_classes: int, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + """Compute the bins to update the confusion matrix with for Theil's U calculation. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data + target: 1D or 2D tensor of categorical (nominal) data + num_classes: Integer specifying the number of classes + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN`s when ``nan_strategy = 'replace``` + + Returns: + Non-reduced confusion matrix + + """ + preds = preds.argmax(1) if preds.ndim == 2 else preds + target = target.argmax(1) if target.ndim == 2 else target + preds, target = _handle_nan_in_data(preds, target, nan_strategy, nan_replace_value) + return _multiclass_confusion_matrix_update(preds, target, num_classes) + + +def _theils_u_compute(confmat: Tensor) -> Tensor: + """Compute Theil's U statistic based on a pre-computed confusion matrix. + + Args: + confmat: Confusion matrix for observed data + + Returns: + Theil's U statistic + + """ + confmat = _drop_empty_rows_and_cols(confmat) + + # compute conditional entropy + s_xy = _conditional_entropy_compute(confmat) + + # compute H(x) + total_occurrences = confmat.sum() + p_x = confmat.sum(0) / total_occurrences + s_x = -torch.sum(p_x * torch.log(p_x)) + + # compute u statistic + if s_x == 0: + return torch.tensor(0, device=confmat.device) + + return (s_x - s_xy) / s_x + + +def theils_u( + preds: Tensor, + target: Tensor, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Theils Uncertainty coefficient`_ statistic measuring the association between two nominal data series. + + .. math:: + U(X|Y) = \frac{H(X) - H(X|Y)}{H(X)} + + where :math:`H(X)` is entropy of variable :math:`X` while :math:`H(X|Y)` is the conditional entropy of :math:`X` + given :math:`Y`. + + Theils's U is an asymmetric coefficient, i.e. :math:`TheilsU(preds, target) \neq TheilsU(target, preds)`. + + The output values lies in [0, 1]. 0 means y has no information about x while value 1 means y has complete + information about x. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + target: 1D or 2D tensor of categorical (nominal) data + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Tensor containing Theil's U statistic + + Example: + >>> from torch import randint + >>> from torchmetrics.functional.nominal import theils_u + >>> preds = randint(10, (10,)) + >>> target = randint(10, (10,)) + >>> theils_u(preds, target) + tensor(0.8530) + + """ + num_classes = len(torch.cat([preds, target]).unique()) + confmat = _theils_u_update(preds, target, num_classes, nan_strategy, nan_replace_value) + return _theils_u_compute(confmat) + + +def theils_u_matrix( + matrix: Tensor, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Theil's U`_ statistic between a set of multiple variables. + + This can serve as a convenient tool to compute Theil's U statistic for analyses of correlation between categorical + variables in your dataset. + + Args: + matrix: A tensor of categorical (nominal) data, where: + - rows represent a number of data points + - columns represent a number of categorical (nominal) features + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Theil's U statistic for a dataset of categorical variables + + Example: + >>> from torch import randint + >>> from torchmetrics.functional.nominal import theils_u_matrix + >>> matrix = randint(0, 4, (200, 5)) + >>> theils_u_matrix(matrix) + tensor([[1.0000, 0.0202, 0.0142, 0.0196, 0.0353], + [0.0202, 1.0000, 0.0070, 0.0136, 0.0065], + [0.0143, 0.0070, 1.0000, 0.0125, 0.0206], + [0.0198, 0.0137, 0.0125, 1.0000, 0.0312], + [0.0352, 0.0065, 0.0204, 0.0308, 1.0000]]) + + """ + _nominal_input_validation(nan_strategy, nan_replace_value) + num_variables = matrix.shape[1] + theils_u_matrix_value = torch.ones(num_variables, num_variables, device=matrix.device) + for i, j in itertools.combinations(range(num_variables), 2): + x, y = matrix[:, i], matrix[:, j] + num_classes = len(torch.cat([x, y]).unique()) + confmat = _theils_u_update(x, y, num_classes, nan_strategy, nan_replace_value) + theils_u_matrix_value[i, j] = _theils_u_compute(confmat) + theils_u_matrix_value[j, i] = _theils_u_compute(confmat.T) + return theils_u_matrix_value diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/tschuprows.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/tschuprows.py new file mode 100644 index 0000000000000000000000000000000000000000..22d256d33d12c288aca7627e51c54b72b9594ff3 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/tschuprows.py @@ -0,0 +1,193 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import itertools +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.classification.confusion_matrix import _multiclass_confusion_matrix_update +from torchmetrics.functional.nominal.utils import ( + _compute_bias_corrected_values, + _compute_chi_squared, + _drop_empty_rows_and_cols, + _handle_nan_in_data, + _nominal_input_validation, + _unable_to_use_bias_correction_warning, +) + + +def _tschuprows_t_update( + preds: Tensor, + target: Tensor, + num_classes: int, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + """Compute the bins to update the confusion matrix with for Tschuprow's T calculation. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data + target: 1D or 2D tensor of categorical (nominal) data + num_classes: Integer specifying the number of classes + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN`s when ``nan_strategy = 'replace``` + + Returns: + Non-reduced confusion matrix + + """ + preds = preds.argmax(1) if preds.ndim == 2 else preds + target = target.argmax(1) if target.ndim == 2 else target + preds, target = _handle_nan_in_data(preds, target, nan_strategy, nan_replace_value) + return _multiclass_confusion_matrix_update(preds, target, num_classes) + + +def _tschuprows_t_compute(confmat: Tensor, bias_correction: bool) -> Tensor: + """Compute Tschuprow's T statistic based on a pre-computed confusion matrix. + + Args: + confmat: Confusion matrix for observed data + bias_correction: Indication of whether to use bias correction. + + Returns: + Tschuprow's T statistic + + """ + confmat = _drop_empty_rows_and_cols(confmat) + cm_sum = confmat.sum() + chi_squared = _compute_chi_squared(confmat, bias_correction) + phi_squared = chi_squared / cm_sum + num_rows, num_cols = confmat.shape + + if bias_correction: + phi_squared_corrected, rows_corrected, cols_corrected = _compute_bias_corrected_values( + phi_squared, num_rows, num_cols, cm_sum + ) + if torch.min(rows_corrected, cols_corrected) == 1: + _unable_to_use_bias_correction_warning(metric_name="Tschuprow's T") + return torch.tensor(float("nan"), device=confmat.device) + tschuprows_t_value = torch.sqrt(phi_squared_corrected / torch.sqrt((rows_corrected - 1) * (cols_corrected - 1))) + else: + n_rows_tensor = torch.tensor(num_rows, device=phi_squared.device) + n_cols_tensor = torch.tensor(num_cols, device=phi_squared.device) + tschuprows_t_value = torch.sqrt(phi_squared / torch.sqrt((n_rows_tensor - 1) * (n_cols_tensor - 1))) + return tschuprows_t_value.clamp(0.0, 1.0) + + +def tschuprows_t( + preds: Tensor, + target: Tensor, + bias_correction: bool = True, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Tschuprow's T`_ statistic measuring the association between two categorical (nominal) data series. + + .. math:: + T = \sqrt{\frac{\chi^2 / n}{\sqrt{(r - 1) * (k - 1)}}} + + where + + .. math:: + \chi^2 = \sum_{i,j} \ frac{\left(n_{ij} - \frac{n_{i.} n_{.j}}{n}\right)^2}{\frac{n_{i.} n_{.j}}{n}} + + where :math:`n_{ij}` denotes the number of times the values :math:`(A_i, B_j)` are observed with :math:`A_i, B_j` + represent frequencies of values in ``preds`` and ``target``, respectively. + + Tschuprow's T is a symmetric coefficient, i.e. :math:`T(preds, target) = T(target, preds)`. + + The output values lies in [0, 1] with 1 meaning the perfect association. + + Args: + preds: 1D or 2D tensor of categorical (nominal) data: + + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + + target: 1D or 2D tensor of categorical (nominal) data: + + - 1D shape: (batch_size,) + - 2D shape: (batch_size, num_classes) + + bias_correction: Indication of whether to use bias correction. + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Tschuprow's T statistic + + Example: + >>> from torch import randint, round + >>> from torchmetrics.functional.nominal import tschuprows_t + >>> preds = randint(0, 4, (100,)) + >>> target = round(preds + torch.randn(100)).clamp(0, 4) + >>> tschuprows_t(preds, target) + tensor(0.4930) + + """ + _nominal_input_validation(nan_strategy, nan_replace_value) + num_classes = len(torch.cat([preds, target]).unique()) + confmat = _tschuprows_t_update(preds, target, num_classes, nan_strategy, nan_replace_value) + return _tschuprows_t_compute(confmat, bias_correction) + + +def tschuprows_t_matrix( + matrix: Tensor, + bias_correction: bool = True, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> Tensor: + r"""Compute `Tschuprow's T`_ statistic between a set of multiple variables. + + This can serve as a convenient tool to compute Tschuprow's T statistic for analyses of correlation between + categorical variables in your dataset. + + Args: + matrix: A tensor of categorical (nominal) data, where: + + - rows represent a number of data points + - columns represent a number of categorical (nominal) features + + bias_correction: Indication of whether to use bias correction. + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN``s when ``nan_strategy = 'replace'`` + + Returns: + Tschuprow's T statistic for a dataset of categorical variables + + Example: + >>> from torch import randint + >>> from torchmetrics.functional.nominal import tschuprows_t_matrix + >>> matrix = randint(0, 4, (200, 5)) + >>> tschuprows_t_matrix(matrix) + tensor([[1.0000, 0.0637, 0.0000, 0.0542, 0.1337], + [0.0637, 1.0000, 0.0000, 0.0000, 0.0000], + [0.0000, 0.0000, 1.0000, 0.0000, 0.0649], + [0.0542, 0.0000, 0.0000, 1.0000, 0.1100], + [0.1337, 0.0000, 0.0649, 0.1100, 1.0000]]) + + """ + _nominal_input_validation(nan_strategy, nan_replace_value) + num_variables = matrix.shape[1] + tschuprows_t_matrix_value = torch.ones(num_variables, num_variables, device=matrix.device) + for i, j in itertools.combinations(range(num_variables), 2): + x, y = matrix[:, i], matrix[:, j] + num_classes = len(torch.cat([x, y]).unique()) + confmat = _tschuprows_t_update(x, y, num_classes, nan_strategy, nan_replace_value) + tschuprows_t_matrix_value[i, j] = tschuprows_t_matrix_value[j, i] = _tschuprows_t_compute( + confmat, bias_correction + ) + return tschuprows_t_matrix_value diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/utils.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..9d8dd8dc4afdb7ef2028170fd74c3b289fa311d5 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/nominal/utils.py @@ -0,0 +1,146 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.utilities.prints import rank_zero_warn + + +def _nominal_input_validation(nan_strategy: str, nan_replace_value: Optional[float]) -> None: + if nan_strategy not in ["replace", "drop"]: + raise ValueError( + f"Argument `nan_strategy` is expected to be one of `['replace', 'drop']`, but got {nan_strategy}" + ) + if nan_strategy == "replace" and not isinstance(nan_replace_value, (float, int)): + raise ValueError( + "Argument `nan_replace` is expected to be of a type `int` or `float` when `nan_strategy = 'replace`, " + f"but got {nan_replace_value}" + ) + + +def _compute_expected_freqs(confmat: Tensor) -> Tensor: + """Compute the expected frequenceis from the provided confusion matrix.""" + margin_sum_rows, margin_sum_cols = confmat.sum(1), confmat.sum(0) + return torch.einsum("r, c -> rc", margin_sum_rows, margin_sum_cols) / confmat.sum() + + +def _compute_chi_squared(confmat: Tensor, bias_correction: bool) -> Tensor: + """Chi-square test of independenc of variables in a confusion matrix table. + + Adapted from: https://github.com/scipy/scipy/blob/v1.9.2/scipy/stats/contingency.py. + + """ + expected_freqs = _compute_expected_freqs(confmat) + # Get degrees of freedom + df = expected_freqs.numel() - sum(expected_freqs.shape) + expected_freqs.ndim - 1 + if df == 0: + return torch.tensor(0.0, device=confmat.device) + + if df == 1 and bias_correction: + diff = expected_freqs - confmat + direction = diff.sign() + confmat += direction * torch.minimum(0.5 * torch.ones_like(direction), direction.abs()) + + return torch.sum((confmat - expected_freqs) ** 2 / expected_freqs) + + +def _drop_empty_rows_and_cols(confmat: Tensor) -> Tensor: + """Drop all rows and columns containing only zeros. + + Example: + >>> from torch import randint + >>> from torchmetrics.functional.nominal.utils import _drop_empty_rows_and_cols + >>> matrix = randint(10, size=(4, 3)) + >>> matrix[1, :] = matrix[:, 1] = 0 + >>> matrix + tensor([[2, 0, 6], + [0, 0, 0], + [0, 0, 0], + [3, 0, 4]]) + >>> _drop_empty_rows_and_cols(matrix) + tensor([[2, 6], + [3, 4]]) + + """ + confmat = confmat[confmat.sum(1) != 0] + return confmat[:, confmat.sum(0) != 0] + + +def _compute_phi_squared_corrected( + phi_squared: Tensor, + num_rows: int, + num_cols: int, + confmat_sum: Tensor, +) -> Tensor: + """Compute bias-corrected Phi Squared.""" + return torch.max( + torch.tensor(0.0, device=phi_squared.device), + phi_squared - ((num_rows - 1) * (num_cols - 1)) / (confmat_sum - 1), + ) + + +def _compute_rows_and_cols_corrected(num_rows: int, num_cols: int, confmat_sum: Tensor) -> tuple[Tensor, Tensor]: + """Compute bias-corrected number of rows and columns.""" + rows_corrected = num_rows - (num_rows - 1) ** 2 / (confmat_sum - 1) + cols_corrected = num_cols - (num_cols - 1) ** 2 / (confmat_sum - 1) + return rows_corrected, cols_corrected + + +def _compute_bias_corrected_values( + phi_squared: Tensor, num_rows: int, num_cols: int, confmat_sum: Tensor +) -> tuple[Tensor, Tensor, Tensor]: + """Compute bias-corrected Phi Squared and number of rows and columns.""" + phi_squared_corrected = _compute_phi_squared_corrected(phi_squared, num_rows, num_cols, confmat_sum) + rows_corrected, cols_corrected = _compute_rows_and_cols_corrected(num_rows, num_cols, confmat_sum) + return phi_squared_corrected, rows_corrected, cols_corrected + + +def _handle_nan_in_data( + preds: Tensor, + target: Tensor, + nan_strategy: Literal["replace", "drop"] = "replace", + nan_replace_value: Optional[float] = 0.0, +) -> tuple[Tensor, Tensor]: + """Handle ``NaN`` values in input data. + + If ``nan_strategy = 'replace'``, all ``NaN`` values are replaced with ``nan_replace_value``. + If ``nan_strategy = 'drop'``, all rows containing ``NaN`` in any of two vectors are dropped. + + Args: + preds: 1D tensor of categorical (nominal) data + target: 1D tensor of categorical (nominal) data + nan_strategy: Indication of whether to replace or drop ``NaN`` values + nan_replace_value: Value to replace ``NaN`s when ``nan_strategy = 'replace``` + + Returns: + Updated ``preds`` and ``target`` tensors which contain no ``Nan`` + + Raises: + ValueError: If ``nan_strategy`` is not from ``['replace', 'drop']``. + ValueError: If ``nan_strategy = replace`` and ``nan_replace_value`` is not of a type ``int`` or ``float``. + + """ + if nan_strategy == "replace": + return preds.nan_to_num(nan_replace_value), target.nan_to_num(nan_replace_value) + rows_contain_nan = torch.logical_or(preds.isnan(), target.isnan()) + return preds[~rows_contain_nan], target[~rows_contain_nan] + + +def _unable_to_use_bias_correction_warning(metric_name: str) -> None: + rank_zero_warn( + f"Unable to compute {metric_name} using bias correction. Please consider to set `bias_correction=False`." + ) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/__init__.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..c226169521d6a2ea8b3ea607f05e304b3b4bf99f --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/__init__.py @@ -0,0 +1,26 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from torchmetrics.functional.pairwise.cosine import pairwise_cosine_similarity +from torchmetrics.functional.pairwise.euclidean import pairwise_euclidean_distance +from torchmetrics.functional.pairwise.linear import pairwise_linear_similarity +from torchmetrics.functional.pairwise.manhattan import pairwise_manhattan_distance +from torchmetrics.functional.pairwise.minkowski import pairwise_minkowski_distance + +__all__ = [ + "pairwise_cosine_similarity", + "pairwise_euclidean_distance", + "pairwise_linear_similarity", + "pairwise_manhattan_distance", + "pairwise_minkowski_distance", +] diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/cosine.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/cosine.py new file mode 100644 index 0000000000000000000000000000000000000000..246b9adf5af05c08f778431a2a65d077768467c7 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/cosine.py @@ -0,0 +1,91 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.pairwise.helpers import _check_input, _reduce_distance_matrix +from torchmetrics.utilities.compute import _safe_matmul + + +def _pairwise_cosine_similarity_update( + x: Tensor, y: Optional[Tensor] = None, zero_diagonal: Optional[bool] = None +) -> Tensor: + """Calculate the pairwise cosine similarity matrix. + + Args: + x: tensor of shape ``[N,d]`` + y: tensor of shape ``[M,d]`` + zero_diagonal: determines if the diagonal of the distance matrix should be set to zero + + """ + x, y, zero_diagonal = _check_input(x, y, zero_diagonal) + + norm = torch.norm(x, p=2, dim=1) + x = x / norm.unsqueeze(1) + norm = torch.norm(y, p=2, dim=1) + y = y / norm.unsqueeze(1) + + distance = _safe_matmul(x, y) + if zero_diagonal: + distance.fill_diagonal_(0) + return distance + + +def pairwise_cosine_similarity( + x: Tensor, + y: Optional[Tensor] = None, + reduction: Literal["mean", "sum", "none", None] = None, + zero_diagonal: Optional[bool] = None, +) -> Tensor: + r"""Calculate pairwise cosine similarity. + + .. math:: + s_{cos}(x,y) = \frac{}{||x|| \cdot ||y||} + = \frac{\sum_{d=1}^D x_d \cdot y_d }{\sqrt{\sum_{d=1}^D x_i^2} \cdot \sqrt{\sum_{d=1}^D y_i^2}} + + If both :math:`x` and :math:`y` are passed in, the calculation will be performed pairwise + between the rows of :math:`x` and :math:`y`. + If only :math:`x` is passed in, the calculation will be performed between the rows of :math:`x`. + + Args: + x: Tensor with shape ``[N, d]`` + y: Tensor with shape ``[M, d]``, optional + reduction: reduction to apply along the last dimension. Choose between `'mean'`, `'sum'` + (applied along column dimension) or `'none'`, `None` for no reduction + zero_diagonal: if the diagonal of the distance matrix should be set to 0. If only :math:`x` is given + this defaults to ``True`` else if :math:`y` is also given it defaults to ``False`` + + Returns: + A ``[N,N]`` matrix of distances if only ``x`` is given, else a ``[N,M]`` matrix + + Example: + >>> import torch + >>> from torchmetrics.functional.pairwise import pairwise_cosine_similarity + >>> x = torch.tensor([[2, 3], [3, 5], [5, 8]], dtype=torch.float32) + >>> y = torch.tensor([[1, 0], [2, 1]], dtype=torch.float32) + >>> pairwise_cosine_similarity(x, y) + tensor([[0.5547, 0.8682], + [0.5145, 0.8437], + [0.5300, 0.8533]]) + >>> pairwise_cosine_similarity(x) + tensor([[0.0000, 0.9989, 0.9996], + [0.9989, 0.0000, 0.9998], + [0.9996, 0.9998, 0.0000]]) + + """ + distance = _pairwise_cosine_similarity_update(x, y, zero_diagonal) + return _reduce_distance_matrix(distance, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/euclidean.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/euclidean.py new file mode 100644 index 0000000000000000000000000000000000000000..7dc1e4b5b24a4042eaf99d844fea1d01e8694586 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/euclidean.py @@ -0,0 +1,89 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.pairwise.helpers import _check_input, _reduce_distance_matrix + + +def _pairwise_euclidean_distance_update( + x: Tensor, y: Optional[Tensor] = None, zero_diagonal: Optional[bool] = None +) -> Tensor: + """Calculate the pairwise euclidean distance matrix. + + Args: + x: tensor of shape ``[N,d]`` + y: tensor of shape ``[M,d]`` + zero_diagonal: determines if the diagonal of the distance matrix should be set to zero + + """ + x, y, zero_diagonal = _check_input(x, y, zero_diagonal) + # upcast to float64 to prevent precision issues + _orig_dtype = x.dtype + x = x.to(torch.float64) + y = y.to(torch.float64) + x_norm = (x * x).sum(dim=1, keepdim=True) + y_norm = (y * y).sum(dim=1) + distance = (x_norm + y_norm - 2 * x.mm(y.T)).to(_orig_dtype) + if zero_diagonal: + distance.fill_diagonal_(0) + return distance.sqrt() + + +def pairwise_euclidean_distance( + x: Tensor, + y: Optional[Tensor] = None, + reduction: Literal["mean", "sum", "none", None] = None, + zero_diagonal: Optional[bool] = None, +) -> Tensor: + r"""Calculate pairwise euclidean distances. + + .. math:: + d_{euc}(x,y) = ||x - y||_2 = \sqrt{\sum_{d=1}^D (x_d - y_d)^2} + + If both :math:`x` and :math:`y` are passed in, the calculation will be performed pairwise between + the rows of :math:`x` and :math:`y`. + If only :math:`x` is passed in, the calculation will be performed between the rows of :math:`x`. + + Args: + x: Tensor with shape ``[N, d]`` + y: Tensor with shape ``[M, d]``, optional + reduction: reduction to apply along the last dimension. Choose between `'mean'`, `'sum'` + (applied along column dimension) or `'none'`, `None` for no reduction + zero_diagonal: if the diagonal of the distance matrix should be set to 0. If only `x` is given + this defaults to `True` else if `y` is also given it defaults to `False` + + Returns: + A ``[N,N]`` matrix of distances if only ``x`` is given, else a ``[N,M]`` matrix + + Example: + >>> import torch + >>> from torchmetrics.functional.pairwise import pairwise_euclidean_distance + >>> x = torch.tensor([[2, 3], [3, 5], [5, 8]], dtype=torch.float32) + >>> y = torch.tensor([[1, 0], [2, 1]], dtype=torch.float32) + >>> pairwise_euclidean_distance(x, y) + tensor([[3.1623, 2.0000], + [5.3852, 4.1231], + [8.9443, 7.6158]]) + >>> pairwise_euclidean_distance(x) + tensor([[0.0000, 2.2361, 5.8310], + [2.2361, 0.0000, 3.6056], + [5.8310, 3.6056, 0.0000]]) + + """ + distance = _pairwise_euclidean_distance_update(x, y, zero_diagonal) + return _reduce_distance_matrix(distance, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/helpers.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/helpers.py new file mode 100644 index 0000000000000000000000000000000000000000..703b5ddb083004a29b1178d6d53342adc326b340 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/helpers.py @@ -0,0 +1,60 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +from torch import Tensor + + +def _check_input( + x: Tensor, y: Optional[Tensor] = None, zero_diagonal: Optional[bool] = None +) -> tuple[Tensor, Tensor, bool]: + """Check that input has the right dimensionality and sets the ``zero_diagonal`` argument if user has not set it. + + Args: + x: tensor of shape ``[N,d]`` + y: if provided, a tensor of shape ``[M,d]`` + zero_diagonal: determines if the diagonal of the distance matrix should be set to zero + + """ + if x.ndim != 2: + raise ValueError(f"Expected argument `x` to be a 2D tensor of shape `[N, d]` but got {x.shape}") + + if y is not None: + if y.ndim != 2 or y.shape[1] != x.shape[1]: + raise ValueError( + "Expected argument `y` to be a 2D tensor of shape `[M, d]` where" + " `d` should be same as the last dimension of `x`" + ) + zero_diagonal = False if zero_diagonal is None else zero_diagonal + else: + y = x.clone() + zero_diagonal = True if zero_diagonal is None else zero_diagonal + return x, y, zero_diagonal + + +def _reduce_distance_matrix(distmat: Tensor, reduction: Optional[str] = None) -> Tensor: + """Reduction of distance matrix. + + Args: + distmat: a ``[N,M]`` matrix + reduction: string determining how to reduce along last dimension + + """ + if reduction == "mean": + return distmat.mean(dim=-1) + if reduction == "sum": + return distmat.sum(dim=-1) + if reduction is None or reduction == "none": + return distmat + raise ValueError(f"Expected reduction to be one of `['mean', 'sum', None]` but got {reduction}") diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/linear.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/linear.py new file mode 100644 index 0000000000000000000000000000000000000000..67bebbae1eaa2967738cc1d6bd8fc38cbacb2e76 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/linear.py @@ -0,0 +1,84 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.pairwise.helpers import _check_input, _reduce_distance_matrix +from torchmetrics.utilities.compute import _safe_matmul + + +def _pairwise_linear_similarity_update( + x: Tensor, y: Optional[Tensor] = None, zero_diagonal: Optional[bool] = None +) -> Tensor: + """Calculate the pairwise linear similarity matrix. + + Args: + x: tensor of shape ``[N,d]`` + y: tensor of shape ``[M,d]`` + zero_diagonal: determines if the diagonal of the distance matrix should be set to zero + + """ + x, y, zero_diagonal = _check_input(x, y, zero_diagonal) + + distance = _safe_matmul(x, y) + if zero_diagonal: + distance.fill_diagonal_(0) + return distance + + +def pairwise_linear_similarity( + x: Tensor, + y: Optional[Tensor] = None, + reduction: Literal["mean", "sum", "none", None] = None, + zero_diagonal: Optional[bool] = None, +) -> Tensor: + r"""Calculate pairwise linear similarity. + + .. math:: + s_{lin}(x,y) = = \sum_{d=1}^D x_d \cdot y_d + + If both :math:`x` and :math:`y` are passed in, the calculation will be performed pairwise between + the rows of :math:`x` and :math:`y`. + If only :math:`x` is passed in, the calculation will be performed between the rows of :math:`x`. + + Args: + x: Tensor with shape ``[N, d]`` + y: Tensor with shape ``[M, d]``, optional + reduction: reduction to apply along the last dimension. Choose between `'mean'`, `'sum'` + (applied along column dimension) or `'none'`, `None` for no reduction + zero_diagonal: if the diagonal of the distance matrix should be set to 0. If only `x` is given + this defaults to `True` else if `y` is also given it defaults to `False` + + Returns: + A ``[N,N]`` matrix of distances if only ``x`` is given, else a ``[N,M]`` matrix + + Example: + >>> import torch + >>> from torchmetrics.functional.pairwise import pairwise_linear_similarity + >>> x = torch.tensor([[2, 3], [3, 5], [5, 8]], dtype=torch.float32) + >>> y = torch.tensor([[1, 0], [2, 1]], dtype=torch.float32) + >>> pairwise_linear_similarity(x, y) + tensor([[ 2., 7.], + [ 3., 11.], + [ 5., 18.]]) + >>> pairwise_linear_similarity(x) + tensor([[ 0., 21., 34.], + [21., 0., 55.], + [34., 55., 0.]]) + + """ + distance = _pairwise_linear_similarity_update(x, y, zero_diagonal) + return _reduce_distance_matrix(distance, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/manhattan.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/manhattan.py new file mode 100644 index 0000000000000000000000000000000000000000..3eda0c07a3820b56d92bc6497e667a05f63233a0 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/manhattan.py @@ -0,0 +1,83 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.pairwise.helpers import _check_input, _reduce_distance_matrix + + +def _pairwise_manhattan_distance_update( + x: Tensor, y: Optional[Tensor] = None, zero_diagonal: Optional[bool] = None +) -> Tensor: + """Calculate the pairwise manhattan similarity matrix. + + Args: + x: tensor of shape ``[N,d]`` + y: if provided, a tensor of shape ``[M,d]`` + zero_diagonal: determines if the diagonal of the distance matrix should be set to zero + + """ + x, y, zero_diagonal = _check_input(x, y, zero_diagonal) + + distance = (x.unsqueeze(1) - y.unsqueeze(0).repeat(x.shape[0], 1, 1)).abs().sum(dim=-1) + if zero_diagonal: + distance.fill_diagonal_(0) + return distance + + +def pairwise_manhattan_distance( + x: Tensor, + y: Optional[Tensor] = None, + reduction: Literal["mean", "sum", "none", None] = None, + zero_diagonal: Optional[bool] = None, +) -> Tensor: + r"""Calculate pairwise manhattan distance. + + .. math:: + d_{man}(x,y) = ||x-y||_1 = \sum_{d=1}^D |x_d - y_d| + + If both :math:`x` and :math:`y` are passed in, the calculation will be performed pairwise between + the rows of :math:`x` and :math:`y`. + If only :math:`x` is passed in, the calculation will be performed between the rows of :math:`x`. + + Args: + x: Tensor with shape ``[N, d]`` + y: Tensor with shape ``[M, d]``, optional + reduction: reduction to apply along the last dimension. Choose between `'mean'`, `'sum'` + (applied along column dimension) or `'none'`, `None` for no reduction + zero_diagonal: if the diagonal of the distance matrix should be set to 0. If only `x` is given + this defaults to `True` else if `y` is also given it defaults to `False` + + Returns: + A ``[N,N]`` matrix of distances if only ``x`` is given, else a ``[N,M]`` matrix + + Example: + >>> import torch + >>> from torchmetrics.functional.pairwise import pairwise_manhattan_distance + >>> x = torch.tensor([[2, 3], [3, 5], [5, 8]], dtype=torch.float32) + >>> y = torch.tensor([[1, 0], [2, 1]], dtype=torch.float32) + >>> pairwise_manhattan_distance(x, y) + tensor([[ 4., 2.], + [ 7., 5.], + [12., 10.]]) + >>> pairwise_manhattan_distance(x) + tensor([[0., 3., 8.], + [3., 0., 5.], + [8., 5., 0.]]) + + """ + distance = _pairwise_manhattan_distance_update(x, y, zero_diagonal) + return _reduce_distance_matrix(distance, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/minkowski.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/minkowski.py new file mode 100644 index 0000000000000000000000000000000000000000..298cedd14862511c9304699f2a531cf19d8a60ce --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/pairwise/minkowski.py @@ -0,0 +1,93 @@ +# Copyright The PyTorch Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.pairwise.helpers import _check_input, _reduce_distance_matrix +from torchmetrics.utilities.exceptions import TorchMetricsUserError + + +def _pairwise_minkowski_distance_update( + x: Tensor, y: Optional[Tensor] = None, exponent: float = 2, zero_diagonal: Optional[bool] = None +) -> Tensor: + """Calculate the pairwise minkowski distance matrix. + + Args: + x: tensor of shape ``[N,d]`` + y: tensor of shape ``[M,d]`` + exponent: int or float larger than 1, exponent to which the difference between preds and target is to be raised + zero_diagonal: determines if the diagonal of the distance matrix should be set to zero + + """ + x, y, zero_diagonal = _check_input(x, y, zero_diagonal) + if not (isinstance(exponent, (float, int)) and exponent >= 1): + raise TorchMetricsUserError(f"Argument ``p`` must be a float or int greater than 1, but got {exponent}") + # upcast to float64 to prevent precision issues + _orig_dtype = x.dtype + x = x.to(torch.float64) + y = y.to(torch.float64) + distance = (x.unsqueeze(1) - y.unsqueeze(0)).abs().pow(exponent).sum(-1).pow(1.0 / exponent) + if zero_diagonal: + distance.fill_diagonal_(0) + return distance.to(_orig_dtype) + + +def pairwise_minkowski_distance( + x: Tensor, + y: Optional[Tensor] = None, + exponent: float = 2, + reduction: Literal["mean", "sum", "none", None] = None, + zero_diagonal: Optional[bool] = None, +) -> Tensor: + r"""Calculate pairwise minkowski distances. + + .. math:: + d_{minkowski}(x,y,p) = ||x - y||_p = \sqrt[p]{\sum_{d=1}^D (x_d - y_d)^p} + + If both :math:`x` and :math:`y` are passed in, the calculation will be performed pairwise between the rows of + :math:`x` and :math:`y`. If only :math:`x` is passed in, the calculation will be performed between the rows + of :math:`x`. + + Args: + x: Tensor with shape ``[N, d]`` + y: Tensor with shape ``[M, d]``, optional + exponent: int or float larger than 1, exponent to which the difference between preds and target is to be raised + reduction: reduction to apply along the last dimension. Choose between `'mean'`, `'sum'` + (applied along column dimension) or `'none'`, `None` for no reduction + zero_diagonal: if the diagonal of the distance matrix should be set to 0. If only `x` is given + this defaults to `True` else if `y` is also given it defaults to `False` + + Returns: + A ``[N,N]`` matrix of distances if only ``x`` is given, else a ``[N,M]`` matrix + + Example: + >>> import torch + >>> from torchmetrics.functional.pairwise import pairwise_minkowski_distance + >>> x = torch.tensor([[2, 3], [3, 5], [5, 8]], dtype=torch.float32) + >>> y = torch.tensor([[1, 0], [2, 1]], dtype=torch.float32) + >>> pairwise_minkowski_distance(x, y, exponent=4) + tensor([[3.0092, 2.0000], + [5.0317, 4.0039], + [8.1222, 7.0583]]) + >>> pairwise_minkowski_distance(x, exponent=4) + tensor([[0.0000, 2.0305, 5.1547], + [2.0305, 0.0000, 3.1383], + [5.1547, 3.1383, 0.0000]]) + + """ + distance = _pairwise_minkowski_distance_update(x, y, exponent, zero_diagonal) + return _reduce_distance_matrix(distance, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/__init__.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..cd99da1133223c05b30dc7df57fe0b61d8694708 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/__init__.py @@ -0,0 +1,59 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from torchmetrics.functional.regression.concordance import concordance_corrcoef +from torchmetrics.functional.regression.cosine_similarity import cosine_similarity +from torchmetrics.functional.regression.csi import critical_success_index +from torchmetrics.functional.regression.explained_variance import explained_variance +from torchmetrics.functional.regression.js_divergence import jensen_shannon_divergence +from torchmetrics.functional.regression.kendall import kendall_rank_corrcoef +from torchmetrics.functional.regression.kl_divergence import kl_divergence +from torchmetrics.functional.regression.log_cosh import log_cosh_error +from torchmetrics.functional.regression.log_mse import mean_squared_log_error +from torchmetrics.functional.regression.mae import mean_absolute_error +from torchmetrics.functional.regression.mape import mean_absolute_percentage_error +from torchmetrics.functional.regression.minkowski import minkowski_distance +from torchmetrics.functional.regression.mse import mean_squared_error +from torchmetrics.functional.regression.nrmse import normalized_root_mean_squared_error +from torchmetrics.functional.regression.pearson import pearson_corrcoef +from torchmetrics.functional.regression.r2 import r2_score +from torchmetrics.functional.regression.rse import relative_squared_error +from torchmetrics.functional.regression.spearman import spearman_corrcoef +from torchmetrics.functional.regression.symmetric_mape import symmetric_mean_absolute_percentage_error +from torchmetrics.functional.regression.tweedie_deviance import tweedie_deviance_score +from torchmetrics.functional.regression.wmape import weighted_mean_absolute_percentage_error + +__all__ = [ + "concordance_corrcoef", + "cosine_similarity", + "critical_success_index", + "explained_variance", + "jensen_shannon_divergence", + "kendall_rank_corrcoef", + "kl_divergence", + "log_cosh_error", + "mean_absolute_error", + "mean_absolute_percentage_error", + "mean_absolute_percentage_error", + "mean_squared_error", + "mean_squared_log_error", + "minkowski_distance", + "normalized_root_mean_squared_error", + "pearson_corrcoef", + "r2_score", + "relative_squared_error", + "spearman_corrcoef", + "symmetric_mean_absolute_percentage_error", + "tweedie_deviance_score", + "weighted_mean_absolute_percentage_error", +] diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/concordance.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/concordance.py new file mode 100644 index 0000000000000000000000000000000000000000..501cf8da05472af02f4303c074f3706037f8c4e2 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/concordance.py @@ -0,0 +1,70 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import torch +from torch import Tensor + +from torchmetrics.functional.regression.pearson import _pearson_corrcoef_compute, _pearson_corrcoef_update + + +def _concordance_corrcoef_compute( + mean_x: Tensor, + mean_y: Tensor, + var_x: Tensor, + var_y: Tensor, + corr_xy: Tensor, + nb: Tensor, +) -> Tensor: + """Compute the final concordance correlation coefficient based on accumulated statistics.""" + pearson = _pearson_corrcoef_compute(var_x, var_y, corr_xy, nb) + var_x = var_x / (nb - 1) + var_y = var_y / (nb - 1) + return 2.0 * pearson * var_x.sqrt() * var_y.sqrt() / (var_x + var_y + (mean_x - mean_y) ** 2) + + +def concordance_corrcoef(preds: Tensor, target: Tensor) -> Tensor: + r"""Compute concordance correlation coefficient that measures the agreement between two variables. + + .. math:: + \rho_c = \frac{2 \rho \sigma_x \sigma_y}{\sigma_x^2 + \sigma_y^2 + (\mu_x - \mu_y)^2} + + where :math:`\mu_x, \mu_y` is the means for the two variables, :math:`\sigma_x^2, \sigma_y^2` are the corresponding + variances and \rho is the pearson correlation coefficient between the two variables. + + Args: + preds: estimated scores + target: ground truth scores + + Example (single output regression): + >>> from torchmetrics.functional.regression import concordance_corrcoef + >>> target = torch.tensor([3, -0.5, 2, 7]) + >>> preds = torch.tensor([2.5, 0.0, 2, 8]) + >>> concordance_corrcoef(preds, target) + tensor([0.9777]) + + Example (multi output regression): + >>> from torchmetrics.functional.regression import concordance_corrcoef + >>> target = torch.tensor([[3, -0.5], [2, 7]]) + >>> preds = torch.tensor([[2.5, 0.0], [2, 8]]) + >>> concordance_corrcoef(preds, target) + tensor([0.7273, 0.9887]) + + """ + d = preds.shape[1] if preds.ndim == 2 else 1 + _temp = torch.zeros(d, dtype=preds.dtype, device=preds.device) + mean_x, mean_y, var_x = _temp.clone(), _temp.clone(), _temp.clone() + var_y, corr_xy, nb = _temp.clone(), _temp.clone(), _temp.clone() + mean_x, mean_y, var_x, var_y, corr_xy, nb = _pearson_corrcoef_update( + preds, target, mean_x, mean_y, var_x, var_y, corr_xy, nb, num_outputs=1 if preds.ndim == 1 else preds.shape[-1] + ) + return _concordance_corrcoef_compute(mean_x, mean_y, var_x, var_y, corr_xy, nb) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/cosine_similarity.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/cosine_similarity.py new file mode 100644 index 0000000000000000000000000000000000000000..c57623931a4ef3202633b636d21a6e37d945acf2 --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/cosine_similarity.py @@ -0,0 +1,101 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +import torch +from torch import Tensor + +from torchmetrics.utilities.checks import _check_same_shape + + +def _cosine_similarity_update( + preds: Tensor, + target: Tensor, +) -> tuple[Tensor, Tensor]: + """Update and returns variables required to compute Cosine Similarity. Checks for same shape of input tensors. + + Args: + preds: Predicted tensor + target: Ground truth tensor + + """ + _check_same_shape(preds, target) + if preds.ndim != 2: + raise ValueError( + "Expected input to cosine similarity to be 2D tensors of shape `[N,D]` where `N` is the number of samples" + f" and `D` is the number of dimensions, but got tensor of shape {preds.shape}" + ) + preds = preds.float() + target = target.float() + + return preds, target + + +def _cosine_similarity_compute(preds: Tensor, target: Tensor, reduction: Optional[str] = "sum") -> Tensor: + """Compute Cosine Similarity. + + Args: + preds: Predicted tensor + target: Ground truth tensor + reduction: + The method of reducing along the batch dimension using sum, mean or taking the individual scores + + Example: + >>> target = torch.tensor([[1, 2, 3, 4], [1, 2, 3, 4]]) + >>> preds = torch.tensor([[1, 2, 3, 4], [-1, -2, -3, -4]]) + >>> preds, target = _cosine_similarity_update(preds, target) + >>> _cosine_similarity_compute(preds, target, 'none') + tensor([ 1.0000, -1.0000]) + + """ + dot_product = (preds * target).sum(dim=-1) + preds_norm = preds.norm(dim=-1) + target_norm = target.norm(dim=-1) + similarity = dot_product / (preds_norm * target_norm) + reduction_mapping = { + "sum": torch.sum, + "mean": torch.mean, + "none": lambda x: x, + None: lambda x: x, + } + return reduction_mapping[reduction](similarity) # type: ignore[operator] + + +def cosine_similarity(preds: Tensor, target: Tensor, reduction: Optional[str] = "sum") -> Tensor: + r"""Compute the `Cosine Similarity`_. + + .. math:: + cos_{sim}(x,y) = \frac{x \cdot y}{||x|| \cdot ||y||} = + \frac{\sum_{i=1}^n x_i y_i}{\sqrt{\sum_{i=1}^n x_i^2}\sqrt{\sum_{i=1}^n y_i^2}} + + where :math:`y` is a tensor of target values, and :math:`x` is a tensor of predictions. + + Args: + preds: Predicted tensor with shape ``(N,d)`` + target: Ground truth tensor with shape ``(N,d)`` + reduction: + The method of reducing along the batch dimension using sum, mean or taking the individual scores + + Example: + >>> from torchmetrics.functional.regression import cosine_similarity + >>> target = torch.tensor([[1, 2, 3, 4], + ... [1, 2, 3, 4]]) + >>> preds = torch.tensor([[1, 2, 3, 4], + ... [-1, -2, -3, -4]]) + >>> cosine_similarity(preds, target, 'none') + tensor([ 1.0000, -1.0000]) + + """ + preds, target = _cosine_similarity_update(preds, target) + return _cosine_similarity_compute(preds, target, reduction) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/csi.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/csi.py new file mode 100644 index 0000000000000000000000000000000000000000..65d38e6f573e92f749a4c8790503bdaeeb618b7b --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/csi.py @@ -0,0 +1,112 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import Optional + +import torch +from torch import Tensor + +from torchmetrics.utilities.checks import _check_same_shape +from torchmetrics.utilities.compute import _safe_divide + + +def _critical_success_index_update( + preds: Tensor, target: Tensor, threshold: float, keep_sequence_dim: Optional[int] = None +) -> tuple[Tensor, Tensor, Tensor]: + """Update and return variables required to compute Critical Success Index. Checks for same shape of tensors. + + Args: + preds: Predicted tensor + target: Ground truth tensor + threshold: Values above or equal to threshold are replaced with 1, below by 0 + keep_sequence_dim: Index of the sequence dimension if the inputs are sequences of images. If specified, + the score will be calculated separately for each image in the sequence. If ``None``, the score will be + calculated across all dimensions. + + """ + _check_same_shape(preds, target) + + if keep_sequence_dim is None: + sum_dims = None + elif not 0 <= keep_sequence_dim < preds.ndim: + raise ValueError(f"Expected keep_sequence dim to be in range [0, {preds.ndim}] but got {keep_sequence_dim}") + else: + sum_dims = tuple(i for i in range(preds.ndim) if i != keep_sequence_dim) + + # binarize the tensors with the threshold + preds_bin = (preds >= threshold).bool() + target_bin = (target >= threshold).bool() + + if keep_sequence_dim is None: + hits = torch.sum(preds_bin & target_bin).int() + misses = torch.sum((preds_bin ^ target_bin) & target_bin).int() + false_alarms = torch.sum((preds_bin ^ target_bin) & preds_bin).int() + else: + hits = torch.sum(preds_bin & target_bin, dim=sum_dims).int() + misses = torch.sum((preds_bin ^ target_bin) & target_bin, dim=sum_dims).int() + false_alarms = torch.sum((preds_bin ^ target_bin) & preds_bin, dim=sum_dims).int() + return hits, misses, false_alarms + + +def _critical_success_index_compute(hits: Tensor, misses: Tensor, false_alarms: Tensor) -> Tensor: + """Compute critical success index. + + Args: + hits: Number of true positives after binarization + misses: Number of false negatives after binarization + false_alarms: Number of false positives after binarization + + Returns: + If input tensors are 5-dimensional and ``keep_sequence_dim=True``, the metric returns a ``(S,)`` vector + with CSI scores for each image in the sequence. Otherwise, it returns a scalar tensor with the CSI score. + + """ + return _safe_divide(hits, hits + misses + false_alarms) + + +def critical_success_index( + preds: Tensor, target: Tensor, threshold: float, keep_sequence_dim: Optional[int] = None +) -> Tensor: + """Compute critical success index. + + Args: + preds: Predicted tensor + target: Ground truth tensor + threshold: Values above or equal to threshold are replaced with 1, below by 0 + keep_sequence_dim: Index of the sequence dimension if the inputs are sequences of images. If specified, + the score will be calculated separately for each image in the sequence. If ``None``, the score will be + calculated across all dimensions. + + Returns: + If ``keep_sequence_dim`` is specified, the metric returns a vector of with CSI scores for each image + in the sequence. Otherwise, it returns a scalar tensor with the CSI score. + + Example: + >>> import torch + >>> from torchmetrics.functional.regression import critical_success_index + >>> x = torch.Tensor([[0.2, 0.7], [0.9, 0.3]]) + >>> y = torch.Tensor([[0.4, 0.2], [0.8, 0.6]]) + >>> critical_success_index(x, y, 0.5) + tensor(0.3333) + + Example: + >>> import torch + >>> from torchmetrics.functional.regression import critical_success_index + >>> x = torch.Tensor([[[0.2, 0.7], [0.9, 0.3]], [[0.2, 0.7], [0.9, 0.3]]]) + >>> y = torch.Tensor([[[0.4, 0.2], [0.8, 0.6]], [[0.4, 0.2], [0.8, 0.6]]]) + >>> critical_success_index(x, y, 0.5, keep_sequence_dim=0) + tensor([0.3333, 0.3333]) + + """ + hits, misses, false_alarms = _critical_success_index_update(preds, target, threshold, keep_sequence_dim) + return _critical_success_index_compute(hits, misses, false_alarms) diff --git a/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/kendall.py b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/kendall.py new file mode 100644 index 0000000000000000000000000000000000000000..a34f032943a3fd99d07e1fad710cba93e73acffc --- /dev/null +++ b/rtme/lib/python3.10/site-packages/torchmetrics/functional/regression/kendall.py @@ -0,0 +1,430 @@ +# Copyright The Lightning team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import List, Optional, Union + +import torch +from torch import Tensor +from typing_extensions import Literal + +from torchmetrics.functional.regression.utils import _check_data_shape_to_num_outputs +from torchmetrics.utilities.checks import _check_same_shape +from torchmetrics.utilities.data import _bincount, _cumsum, dim_zero_cat +from torchmetrics.utilities.enums import EnumStr + + +class _MetricVariant(EnumStr): + """Enumerate for metric variants.""" + + A = "a" + B = "b" + C = "c" + + @staticmethod + def _name() -> str: + return "variant" + + +class _TestAlternative(EnumStr): + """Enumerate for test alternative options.""" + + TWO_SIDED = "two-sided" + LESS = "less" + GREATER = "greater" + + @staticmethod + def _name() -> str: + return "alternative" + + +def _sort_on_first_sequence(x: Tensor, y: Tensor) -> tuple[Tensor, Tensor]: + """Sort sequences in an ascent order according to the sequence ``x``.""" + # We need to clone `y` tensor not to change an object in memory + y = torch.clone(y) + x, y = x.T, y.T + x, perm = x.sort() + for i in range(x.shape[0]): + y[i] = y[i][perm[i]] + return x.T, y.T + + +def _concordant_element_sum(x: Tensor, y: Tensor, i: int) -> Tensor: + """Count a total number of concordant pairs in a single sequence.""" + return torch.logical_and(x[i] < x[(i + 1) :], y[i] < y[(i + 1) :]).sum(0).unsqueeze(0) + + +def _count_concordant_pairs(preds: Tensor, target: Tensor) -> Tensor: + """Count a total number of concordant pairs in given sequences.""" + return torch.cat([_concordant_element_sum(preds, target, i) for i in range(preds.shape[0])]).sum(0) + + +def _discordant_element_sum(x: Tensor, y: Tensor, i: int) -> Tensor: + """Count a total number of discordant pairs in a single sequences.""" + return ( + torch.logical_or( + torch.logical_and(x[i] > x[(i + 1) :], y[i] < y[(i + 1) :]), + torch.logical_and(x[i] < x[(i + 1) :], y[i] > y[(i + 1) :]), + ) + .sum(0) + .unsqueeze(0) + ) + + +def _count_discordant_pairs(preds: Tensor, target: Tensor) -> Tensor: + """Count a total number of discordant pairs in given sequences.""" + return torch.cat([_discordant_element_sum(preds, target, i) for i in range(preds.shape[0])]).sum(0) + + +def _convert_sequence_to_dense_rank(x: Tensor, sort: bool = False) -> Tensor: + """Convert a sequence to the rank tensor.""" + # Sort if a sequence has not been sorted before + if sort: + x = x.sort(dim=0).values + _ones = torch.zeros(1, x.shape[1], dtype=torch.int32, device=x.device) + return _cumsum(torch.cat([_ones, (x[1:] != x[:-1]).int()], dim=0), dim=0) + + +def _get_ties(x: Tensor) -> tuple[Tensor, Tensor, Tensor]: + """Get a total number of ties and staistics for p-value calculation for a given sequence.""" + ties = torch.zeros(x.shape[1], dtype=x.dtype, device=x.device) + ties_p1 = torch.zeros(x.shape[1], dtype=x.dtype, device=x.device) + ties_p2 = torch.zeros(x.shape[1], dtype=x.dtype, device=x.device) + for dim in range(x.shape[1]): + n_ties = _bincount(x[:, dim]) + n_ties = n_ties[n_ties > 1] + ties[dim] = (n_ties * (n_ties - 1) // 2).sum() + ties_p1[dim] = (n_ties * (n_ties - 1.0) * (n_ties - 2)).sum() + ties_p2[dim] = (n_ties * (n_ties - 1.0) * (2 * n_ties + 5)).sum() + + return ties, ties_p1, ties_p2 + + +def _get_metric_metadata( + preds: Tensor, target: Tensor, variant: _MetricVariant +) -> tuple[ + Tensor, + Tensor, + Optional[Tensor], + Optional[Tensor], + Optional[Tensor], + Optional[Tensor], + Optional[Tensor], + Optional[Tensor], + Tensor, +]: + """Obtain statistics to calculate metric value.""" + preds, target = _sort_on_first_sequence(preds, target) + + concordant_pairs = _count_concordant_pairs(preds, target) + discordant_pairs = _count_discordant_pairs(preds, target) + + n_total = torch.tensor(preds.shape[0], device=preds.device) + preds_ties = target_ties = None + preds_ties_p1 = preds_ties_p2 = target_ties_p1 = target_ties_p2 = None + if variant != _MetricVariant.A: + preds = _convert_sequence_to_dense_rank(preds) + target = _convert_sequence_to_dense_rank(target, sort=True) + preds_ties, preds_ties_p1, preds_ties_p2 = _get_ties(preds) + target_ties, target_ties_p1, target_ties_p2 = _get_ties(target) + return ( + concordant_pairs, + discordant_pairs, + preds_ties, + preds_ties_p1, + preds_ties_p2, + target_ties, + target_ties_p1, + target_ties_p2, + n_total, + ) + + +def _calculate_tau( + preds: Tensor, + target: Tensor, + concordant_pairs: Tensor, + discordant_pairs: Tensor, + con_min_dis_pairs: Tensor, + n_total: Tensor, + preds_ties: Optional[Tensor], + target_ties: Optional[Tensor], + variant: _MetricVariant, +) -> Tensor: + """Calculate Kendall's tau from metric metadata.""" + if variant == _MetricVariant.A: + return con_min_dis_pairs / (concordant_pairs + discordant_pairs) + if variant == _MetricVariant.B: + total_combinations: Tensor = n_total * (n_total - 1) // 2 + if preds_ties is None: + preds_ties = torch.tensor(0.0, dtype=total_combinations.dtype, device=total_combinations.device) + if target_ties is None: + target_ties = torch.tensor(0.0, dtype=total_combinations.dtype, device=total_combinations.device) + denominator = (total_combinations - preds_ties) * (total_combinations - target_ties) + return con_min_dis_pairs / torch.sqrt(denominator) + + preds_unique = torch.tensor([len(p.unique()) for p in preds.T], dtype=preds.dtype, device=preds.device) + target_unique = torch.tensor([len(t.unique()) for t in target.T], dtype=target.dtype, device=target.device) + min_classes = torch.minimum(preds_unique, target_unique) + return 2 * con_min_dis_pairs / ((min_classes - 1) / min_classes * n_total**2) + + +def _get_p_value_for_t_value_from_dist(t_value: Tensor) -> Tensor: + """Obtain p-value for a given Tensor of t-values. Handle ``nan`` which cannot be passed into torch distributions. + + When t-value is ``nan``, a resulted p-value should be alson ``nan``. + + """ + device = t_value + normal_dist = torch.distributions.normal.Normal(torch.tensor([0.0]).to(device), torch.tensor([1.0]).to(device)) + + is_nan = t_value.isnan() + t_value = t_value.nan_to_num() + p_value = normal_dist.cdf(t_value) + return p_value.where(~is_nan, torch.tensor(float("nan"), dtype=p_value.dtype, device=p_value.device)) + + +def _calculate_p_value( + con_min_dis_pairs: Tensor, + n_total: Tensor, + preds_ties: Optional[Tensor], + preds_ties_p1: Optional[Tensor], + preds_ties_p2: Optional[Tensor], + target_ties: Optional[Tensor], + target_ties_p1: Optional[Tensor], + target_ties_p2: Optional[Tensor], + variant: _MetricVariant, + alternative: Optional[_TestAlternative], +) -> Tensor: + """Calculate p-value for Kendall's tau from metric metadata.""" + t_value_denominator_base = n_total * (n_total - 1) * (2 * n_total + 5) + if variant == _MetricVariant.A: + t_value = 3 * con_min_dis_pairs / torch.sqrt(t_value_denominator_base / 2) + else: + m = n_total * (n_total - 1) + t_value_denominator: Tensor = ( + t_value_denominator_base + - (preds_ties_p2 if preds_ties_p2 is not None else 0) + - (target_ties_p2 if target_ties_p2 is not None else 0) + ) / 18 + t_value_denominator += ( + 2 * (preds_ties if preds_ties is not None else 0) * (target_ties if target_ties is not None else 0) + ) / m + t_value_denominator += ( + (preds_ties_p1 if preds_ties_p1 is not None else 0) + * (target_ties_p1 if target_ties_p1 is not None else 0) + / (9 * m * (n_total - 2)) + ) + t_value = con_min_dis_pairs / torch.sqrt(t_value_denominator) + + if alternative == _TestAlternative.TWO_SIDED: + t_value = torch.abs(t_value) + if alternative in [_TestAlternative.TWO_SIDED, _TestAlternative.GREATER]: + t_value *= -1 + p_value = _get_p_value_for_t_value_from_dist(t_value) + if alternative == _TestAlternative.TWO_SIDED: + p_value *= 2 + return p_value + + +def _kendall_corrcoef_update( + preds: Tensor, + target: Tensor, + concat_preds: Optional[List[Tensor]] = None, + concat_target: Optional[List[Tensor]] = None, + num_outputs: int = 1, +) -> tuple[List[Tensor], List[Tensor]]: + """Update variables required to compute Kendall rank correlation coefficient. + + Args: + preds: Sequence of data + target: Sequence of data + concat_preds: List of batches of preds sequence to be concatenated + concat_target: List of batches of target sequence to be concatenated + num_outputs: Number of outputs in multioutput setting + + Raises: + RuntimeError: If ``preds`` and ``target`` do not have the same shape + + """ + concat_preds = concat_preds or [] + concat_target = concat_target or [] + # Data checking + _check_same_shape(preds, target) + _check_data_shape_to_num_outputs(preds, target, num_outputs) + + if num_outputs == 1: + preds = preds.unsqueeze(1) + target = target.unsqueeze(1) + + concat_preds.append(preds) + concat_target.append(target) + + return concat_preds, concat_target + + +def _kendall_corrcoef_compute( + preds: Tensor, + target: Tensor, + variant: _MetricVariant, + alternative: Optional[_TestAlternative] = None, +) -> tuple[Tensor, Optional[Tensor]]: + """Compute Kendall rank correlation coefficient, and optionally p-value of corresponding statistical test. + + Args: + Args: + preds: Sequence of data + target: Sequence of data + variant: Indication of which variant of Kendall's tau to be used + alternative: Alternative hypothesis for for t-test. Possible values: + - 'two-sided': the rank correlation is nonzero + - 'less': the rank correlation is negative (less than zero) + - 'greater': the rank correlation is positive (greater than zero) + + """ + ( + concordant_pairs, + discordant_pairs, + preds_ties, + preds_ties_p1, + preds_ties_p2, + target_ties, + target_ties_p1, + target_ties_p2, + n_total, + ) = _get_metric_metadata(preds, target, variant) + con_min_dis_pairs = concordant_pairs - discordant_pairs + + tau = _calculate_tau( + preds, target, concordant_pairs, discordant_pairs, con_min_dis_pairs, n_total, preds_ties, target_ties, variant + ) + p_value = ( + _calculate_p_value( + con_min_dis_pairs, + n_total, + preds_ties, + preds_ties_p1, + preds_ties_p2, + target_ties, + target_ties_p1, + target_ties_p2, + variant, + alternative, + ) + if alternative + else None + ) + + # Squeeze tensor if num_outputs=1 + if tau.shape[0] == 1: + tau = tau.squeeze() + p_value = p_value.squeeze() if p_value is not None else None + + return tau.clamp(-1, 1), p_value + + +def kendall_rank_corrcoef( + preds: Tensor, + target: Tensor, + variant: Literal["a", "b", "c"] = "b", + t_test: bool = False, + alternative: Optional[Literal["two-sided", "less", "greater"]] = "two-sided", +) -> Union[Tensor, tuple[Tensor, Tensor]]: + r"""Compute `Kendall Rank Correlation Coefficient`_. + + .. math:: + tau_a = \frac{C - D}{C + D} + + where :math:`C` represents concordant pairs, :math:`D` stands for discordant pairs. + + .. math:: + tau_b = \frac{C - D}{\sqrt{(C + D + T_{preds}) * (C + D + T_{target})}} + + where :math:`C` represents concordant pairs, :math:`D` stands for discordant pairs and :math:`T` represents + a total number of ties. + + .. math:: + tau_c = 2 * \frac{C - D}{n^2 * \frac{m - 1}{m}} + + where :math:`C` represents concordant pairs, :math:`D` stands for discordant pairs, :math:`n` is a total number + of observations and :math:`m` is a ``min`` of unique values in ``preds`` and ``target`` sequence. + + Definitions according to Definition according to `The Treatment of Ties in Ranking Problems`_. + + Args: + preds: Sequence of data of either shape ``(N,)`` or ``(N,d)`` + target: Sequence of data of either shape ``(N,)`` or ``(N,d)`` + variant: Indication of which variant of Kendall's tau to be used + t_test: Indication whether to run t-test + alternative: Alternative hypothesis for t-test. Possible values: + - 'two-sided': the rank correlation is nonzero + - 'less': the rank correlation is negative (less than zero) + - 'greater': the rank correlation is positive (greater than zero) + + Return: + Correlation tau statistic + (Optional) p-value of corresponding statistical test (asymptotic) + + Raises: + ValueError: If ``t_test`` is not of a type bool + ValueError: If ``t_test=True`` and ``alternative=None`` + + Example (single output regression): + >>> from torchmetrics.functional.regression import kendall_rank_corrcoef + >>> preds = torch.tensor([2.5, 0.0, 2, 8]) + >>> target = torch.tensor([3, -0.5, 2, 1]) + >>> kendall_rank_corrcoef(preds, target) + tensor(0.3333) + + Example (multi output regression): + >>> from torchmetrics.functional.regression import kendall_rank_corrcoef + >>> preds = torch.tensor([[2.5, 0.0], [2, 8]]) + >>> target = torch.tensor([[3, -0.5], [2, 1]]) + >>> kendall_rank_corrcoef(preds, target) + tensor([1., 1.]) + + Example (single output regression with t-test) + >>> from torchmetrics.functional.regression import kendall_rank_corrcoef + >>> preds = torch.tensor([2.5, 0.0, 2, 8]) + >>> target = torch.tensor([3, -0.5, 2, 1]) + >>> kendall_rank_corrcoef(preds, target, t_test=True, alternative='two-sided') + (tensor(0.3333), tensor(0.4969)) + + Example (multi output regression with t-test): + >>> from torchmetrics.functional.regression import kendall_rank_corrcoef + >>> preds = torch.tensor([[2.5, 0.0], [2, 8]]) + >>> target = torch.tensor([[3, -0.5], [2, 1]]) + >>> kendall_rank_corrcoef(preds, target, t_test=True, alternative='two-sided') + (tensor([1., 1.]), tensor([nan, nan])) + + """ + if not isinstance(t_test, bool): + raise ValueError(f"Argument `t_test` is expected to be of a type `bool`, but got {type(t_test)}.") + if t_test and alternative is None: + raise ValueError("Argument `alternative` is required if `t_test=True` but got `None`.") + + _variant = _MetricVariant.from_str(str(variant)) + _alternative = _TestAlternative.from_str(str(alternative)) if t_test else None + + _preds, _target = _kendall_corrcoef_update( + preds, target, [], [], num_outputs=1 if preds.ndim == 1 else preds.shape[-1] + ) + tau, p_value = _kendall_corrcoef_compute( + dim_zero_cat(_preds), + dim_zero_cat(_target), + _variant, # type: ignore[arg-type] # todo + _alternative, # type: ignore[arg-type] # todo + ) + + if p_value is not None: + return tau, p_value + return tau