Compare commits
52 Commits
7f698b511e
...
7fe3cfcdee
| Author | SHA1 | Date |
|---|---|---|
|
|
7fe3cfcdee | |
|
|
ee803a339b | |
|
|
125d4bd756 | |
|
|
0965d99135 | |
|
|
6e83bbef67 | |
|
|
48a5240a6e | |
|
|
195552ef55 | |
|
|
2cc2237165 | |
|
|
8b0327f00a | |
|
|
24a971304d | |
|
|
4bfa212459 | |
|
|
574fc61032 | |
|
|
fd6373651e | |
|
|
c59ee76d9d | |
|
|
ed02be2c8d | |
|
|
b29966797a | |
|
|
f6c822ca8f | |
|
|
4916919833 | |
|
|
a1dff57db0 | |
|
|
e44056d860 | |
|
|
1907c040ad | |
|
|
5374ef3cae | |
|
|
c1ae85950a | |
|
|
f4ac6da7a4 | |
|
|
be23a321a7 | |
|
|
bb423c6856 | |
|
|
51190c2c7d | |
|
|
7562e7c85f | |
|
|
bd59b5d190 | |
|
|
6db659e3c4 | |
|
|
db49cd31be | |
|
|
9da4fd57db | |
|
|
5f6a151263 | |
|
|
047cb9e533 | |
|
|
6388d9b549 | |
|
|
d9cf6cc11d | |
|
|
46e7246f3a | |
|
|
c6492a0f20 | |
|
|
e4c07e1835 | |
|
|
c2044cb200 | |
|
|
3268e9fada | |
|
|
a868d2d561 | |
|
|
3c09b1c98a | |
|
|
f227ed2f60 | |
|
|
41baeb78ca | |
|
|
088ebcdd31 | |
|
|
c11b99e08a | |
|
|
854b3ba3f9 | |
|
|
eafe486893 | |
|
|
1fb8500e87 | |
|
|
d597820a59 | |
|
|
010676df12 |
|
|
@ -1,3 +1,38 @@
|
|||
Change Log 0.2.1
|
||||
-----------------
|
||||
|
||||
- Improved documentation of confidence regions. Added QuaPy logo :')
|
||||
|
||||
- Added Bayesian KDEy and Bayesian MAPLS quantifiers.
|
||||
|
||||
- Added temperature calibration utilities for Bayesian confidence-aware methods.
|
||||
|
||||
- Added compositional CLR and ILR transformations.
|
||||
|
||||
- Extended KDEy with Aitchison/ILR kernels, shrinkage, and improved numerical stability.
|
||||
|
||||
- Added image-embedding-based datasets including CIFAR10, CIFAR100, CIFAR100coarse, VSHN, FashionMNIST, MNIST.
|
||||
|
||||
- Added TemperatureScalingFromLogits for calibrating pretrained logits.
|
||||
|
||||
- Added DirichletProtocol for prevalence sampling from Dirichlet priors.
|
||||
|
||||
- Added ReadMe method by Daniel Hopkins and Gary King.
|
||||
|
||||
- Internal index in LabelledCollection is now "lazy", and is only constructed if required.
|
||||
|
||||
- Improved unit testing and separated integration tests.
|
||||
|
||||
- Added RLLS (Regularized Learning for Domain Adaptation under Label Shifts) method.
|
||||
|
||||
- Added visualization tools for 3-class problems in the simplex, see also the new example no.19
|
||||
|
||||
- Deep code revision and improved codebase
|
||||
|
||||
- Added EDx/EDy from quantificationlib (thanks to Pablo and Juanjo!)
|
||||
|
||||
|
||||
|
||||
Change Log 0.2.0
|
||||
-----------------
|
||||
|
||||
|
|
@ -14,6 +49,7 @@ Change Log 0.2.0
|
|||
in which case the_data is to be used for validation purposes. However, the val_split could be set as a fraction
|
||||
indicating only part of the_data must be used for validation, and the rest wasted... it was certainly confusing.
|
||||
- This change imposes a versioning constrain with qunfold, which now must be >= 0.1.6
|
||||
|
||||
- EMQ has been modified, so that the representation function "classify" now only provides posterior
|
||||
probabilities and, if required, these are recalibrated (e.g., by "bcts") during the aggregation function.
|
||||
- A new parameter "on_calib_error" is passed to the constructor, which informs of the policy to follow
|
||||
|
|
@ -21,13 +57,16 @@ Change Log 0.2.0
|
|||
- 'raise': raises a RuntimeException (default)
|
||||
- 'backup': reruns by silently avoiding calibration
|
||||
- Parameter "recalib" has been renamed "calib"
|
||||
|
||||
- Added aggregative bootstrap for deriving confidence regions (confidence intervals, ellipses in the simplex, or
|
||||
ellipses in the CLR space). This method is efficient as it leverages the two-phases of the aggregative quantifiers.
|
||||
This method applies resampling only to the aggregation phase, thus avoiding to train many quantifiers, or
|
||||
classify multiple times the instances of a sample. See:
|
||||
- quapy/method/confidence.py (new)
|
||||
- the new example no. 16.confidence_regions.py
|
||||
|
||||
- BayesianCC moved to confidence.py, where methods having to do with confidence intervals belong.
|
||||
|
||||
- Improved documentation of qp.plot module.
|
||||
|
||||
|
||||
|
|
@ -129,6 +168,7 @@ Change Log 0.1.8
|
|||
|
||||
- New API documentation template.
|
||||
|
||||
|
||||
Change Log 0.1.7
|
||||
----------------
|
||||
|
||||
|
|
@ -182,7 +222,7 @@ Change Log 0.1.7
|
|||
- hyperparameters yielding to inconsistent runs raise a ValueError exception, while hyperparameter combinations
|
||||
yielding to internal errors of surrogate functions are reported and skipped, without stopping the grid search.
|
||||
|
||||
- DistributionMatching methods added. This is a general framework for distribution matching methods that catters for
|
||||
- DistributionMatching methods added. This is a general framework for distribution matching methods that caters for
|
||||
multiclass quantification. That is to say, one could get a multiclass variant of the (originally binary) HDy
|
||||
method aligned with the Firat's formulation.
|
||||
|
||||
|
|
@ -207,4 +247,3 @@ Change Log 0.1.7
|
|||
any instance of BaseQuantifier), and a subclass of it called OneVsAllAggregative which implements the
|
||||
classify / aggregate interface. Both are instances of OneVsAll. There is a method getOneVsAll that returns the
|
||||
best instance based on the type of quantifier.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
# QuaPy
|
||||
|
||||
## version 0.2.1
|
||||
|
||||
QuaPy is an open source framework for quantification (a.k.a. supervised prevalence estimation, or learning to quantify)
|
||||
written in Python.
|
||||
|
||||
|
|
@ -13,7 +15,7 @@ for facilitating the analysis and interpretation of the experimental results.
|
|||
|
||||
### Last updates:
|
||||
|
||||
* Version 0.2.0 is released! major changes can be consulted [here](CHANGE_LOG.txt).
|
||||
* Version 0.2.1 is released! major changes can be consulted [here](CHANGE_LOG.txt).
|
||||
* The developer API documentation is available [here](https://hlt-isti.github.io/QuaPy/index.html)
|
||||
* Manuals are available [here](https://hlt-isti.github.io/QuaPy/manuals.html)
|
||||
|
||||
|
|
@ -74,6 +76,7 @@ See the [documentation](https://hlt-isti.github.io/QuaPy/manuals.html) for detai
|
|||
|
||||
* Implementation of many popular quantification methods (Classify-&-Count and its variants, Expectation Maximization,
|
||||
quantification methods based on structured output learning, HDy, QuaNet, quantification ensembles, among others).
|
||||
* Support for uncertainty quantification via bootstrap-based and Bayesian methods, including confidence intervals and simplex-aware confidence regions.
|
||||
* Versatile functionality for performing evaluation based on sampling generation protocols (e.g., APP, NPP, etc.).
|
||||
* Implementation of most commonly used evaluation metrics (e.g., AE, RAE, NAE, NRAE, SE, KLD, NKLD, etc.).
|
||||
* Datasets frequently used in quantification (textual and numeric), including:
|
||||
|
|
|
|||
1
TODO.txt
|
|
@ -16,7 +16,6 @@ scale each value by per-class thresholds, i.e., [0.33*0.1, 0.33*1, 0.33*1]/sum."
|
|||
|
||||
- [TODO] document confidence in manuals
|
||||
- [TODO] Test the return_type="index" in protocols and finish the "distributing_samples.py" example
|
||||
- [TODO] Add EDy (an implementation is available at quantificationlib)
|
||||
- [TODO] add ensemble methods SC-MQ, MC-SQ, MC-MQ
|
||||
- [TODO] add HistNetQ
|
||||
- [TODO] add CDE-iteration and Bayes-CDE methods
|
||||
|
|
|
|||
|
|
@ -1,78 +1,374 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>Overview: module code — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>Overview: module code — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../" id="documentation_options" src="../_static/documentation_options.js"></script>
|
||||
<script src="../_static/jquery.js"></script>
|
||||
<script src="../_static/underscore.js"></script>
|
||||
<script src="../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../_static/doctools.js"></script>
|
||||
<script src="../_static/sphinx_highlight.js"></script>
|
||||
<script src="../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/index';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../genindex.html" />
|
||||
<link rel="search" title="Search" href="../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../_static/searchtools.js"></script>
|
||||
<script src="../_static/language_data.js"></script>
|
||||
<script src="../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item active">Overview: module code</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">Overview: module code</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>All modules for which code is available</h1>
|
||||
<ul><li><a href="quapy/classification/calibration.html">quapy.classification.calibration</a></li>
|
||||
<li><a href="quapy/classification/methods.html">quapy.classification.methods</a></li>
|
||||
<li><a href="quapy/classification/neural.html">quapy.classification.neural</a></li>
|
||||
<li><a href="quapy/classification/svmperf.html">quapy.classification.svmperf</a></li>
|
||||
<li><a href="quapy/data/base.html">quapy.data.base</a></li>
|
||||
<li><a href="quapy/data/datasets.html">quapy.data.datasets</a></li>
|
||||
|
|
@ -82,10 +378,10 @@
|
|||
<li><a href="quapy/evaluation.html">quapy.evaluation</a></li>
|
||||
<li><a href="quapy/functional.html">quapy.functional</a></li>
|
||||
<li><a href="quapy/method/_kdey.html">quapy.method._kdey</a></li>
|
||||
<li><a href="quapy/method/_neural.html">quapy.method._neural</a></li>
|
||||
<li><a href="quapy/method/_threshold_optim.html">quapy.method._threshold_optim</a></li>
|
||||
<li><a href="quapy/method/aggregative.html">quapy.method.aggregative</a></li>
|
||||
<li><a href="quapy/method/base.html">quapy.method.base</a></li>
|
||||
<li><a href="quapy/method/confidence.html">quapy.method.confidence</a></li>
|
||||
<li><a href="quapy/method/meta.html">quapy.method.meta</a></li>
|
||||
<li><a href="quapy/method/non_aggregative.html">quapy.method.non_aggregative</a></li>
|
||||
<li><a href="quapy/model_selection.html">quapy.model_selection</a></li>
|
||||
|
|
@ -94,31 +390,75 @@
|
|||
<li><a href="quapy/util.html">quapy.util</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,82 +1,382 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.calibration — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.calibration — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../../" id="documentation_options" src="../../../_static/documentation_options.js"></script>
|
||||
<script src="../../../_static/jquery.js"></script>
|
||||
<script src="../../../_static/underscore.js"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../../_static/doctools.js"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/classification/calibration';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.classification.calibration</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
<h1>Source code for quapy.classification.calibration</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">from</span> <span class="nn">copy</span> <span class="kn">import</span> <span class="n">deepcopy</span>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<span class="kn">from</span> <span class="nn">abstention.calibration</span> <span class="kn">import</span> <span class="n">NoBiasVectorScaling</span><span class="p">,</span> <span class="n">TempScaling</span><span class="p">,</span> <span class="n">VectorScaling</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.base</span> <span class="kn">import</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">clone</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.model_selection</span> <span class="kn">import</span> <span class="n">cross_val_predict</span><span class="p">,</span> <span class="n">train_test_split</span>
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.classification.calibration</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.classification.calibration</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">from</span><span class="w"> </span><span class="nn">copy</span><span class="w"> </span><span class="kn">import</span> <span class="n">deepcopy</span>
|
||||
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">clone</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.model_selection</span><span class="w"> </span><span class="kn">import</span> <span class="n">cross_val_predict</span><span class="p">,</span> <span class="n">train_test_split</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.preprocessing</span><span class="w"> </span><span class="kn">import</span> <span class="n">LabelEncoder</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.utils.validation</span><span class="w"> </span><span class="kn">import</span> <span class="n">check_X_y</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
|
||||
|
||||
<span class="c1"># Wrappers of calibration defined by Alexandari et al. in paper <http://proceedings.mlr.press/v119/alexandari20a.html></span>
|
||||
|
|
@ -84,7 +384,20 @@
|
|||
<span class="c1"># see https://github.com/kundajelab/abstention</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifier"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifier">[docs]</a><span class="k">class</span> <span class="nc">RecalibratedProbabilisticClassifier</span><span class="p">:</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_require_abstention_calibration</span><span class="p">():</span>
|
||||
<span class="k">try</span><span class="p">:</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">abstention.calibration</span><span class="w"> </span><span class="kn">import</span> <span class="n">NoBiasVectorScaling</span><span class="p">,</span> <span class="n">TempScaling</span><span class="p">,</span> <span class="n">VectorScaling</span>
|
||||
<span class="k">except</span> <span class="ne">ImportError</span> <span class="k">as</span> <span class="n">exc</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="ne">ImportError</span><span class="p">(</span>
|
||||
<span class="s2">"Calibration methods in quapy.classification.calibration require the optional "</span>
|
||||
<span class="s2">"'abstention' package."</span>
|
||||
<span class="p">)</span> <span class="kn">from</span><span class="w"> </span><span class="nn">exc</span>
|
||||
<span class="k">return</span> <span class="n">NoBiasVectorScaling</span><span class="p">,</span> <span class="n">TempScaling</span><span class="p">,</span> <span class="n">VectorScaling</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifier">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifier">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">RecalibratedProbabilisticClassifier</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Abstract class for (re)calibration method from `abstention.calibration`, as defined in</span>
|
||||
<span class="sd"> `Alexandari, A., Kundaje, A., & Shrikumar, A. (2020, November). Maximum likelihood with bias-corrected calibration</span>
|
||||
|
|
@ -94,7 +407,10 @@
|
|||
<span class="k">pass</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase">[docs]</a><span class="k">class</span> <span class="nc">RecalibratedProbabilisticClassifierBase</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">RecalibratedProbabilisticClassifier</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">RecalibratedProbabilisticClassifierBase</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">RecalibratedProbabilisticClassifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Applies a (re)calibration method from `abstention.calibration`, as defined in</span>
|
||||
<span class="sd"> `Alexandari et al. paper <http://proceedings.mlr.press/v119/alexandari20a.html>`_.</span>
|
||||
|
|
@ -110,14 +426,16 @@
|
|||
<span class="sd"> :param verbose: whether or not to display information in the standard output</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">calibrator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">calibrator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">calibrator</span> <span class="o">=</span> <span class="n">calibrator</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">val_split</span> <span class="o">=</span> <span class="n">val_split</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span> <span class="o">=</span> <span class="n">n_jobs</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.fit"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.fit">[docs]</a> <span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fits the calibration for the probabilistic classifier.</span>
|
||||
|
||||
|
|
@ -135,7 +453,10 @@
|
|||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="s1">'wrong value for val_split: the proportion of validation documents must be in (0,1)'</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">fit_tr_val</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span></div>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.fit_cv"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.fit_cv">[docs]</a> <span class="k">def</span> <span class="nf">fit_cv</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.fit_cv">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.fit_cv">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit_cv</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fits the calibration in a cross-validation manner, i.e., it generates posterior probabilities for all</span>
|
||||
<span class="sd"> training instances via cross-validation, and then retrains the classifier on all training instances.</span>
|
||||
|
|
@ -153,7 +474,10 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">calibration_function</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">calibrator</span><span class="p">(</span><span class="n">posteriors</span><span class="p">,</span> <span class="n">np</span><span class="o">.</span><span class="n">eye</span><span class="p">(</span><span class="n">nclasses</span><span class="p">)[</span><span class="n">y</span><span class="p">],</span> <span class="n">posterior_supplied</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.fit_tr_val"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.fit_tr_val">[docs]</a> <span class="k">def</span> <span class="nf">fit_tr_val</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.fit_tr_val">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.fit_tr_val">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit_tr_val</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fits the calibration in a train/val-split manner, i.e.t, it partitions the training instances into a</span>
|
||||
<span class="sd"> training and a validation set, and then uses the training samples to learn classifier which is then used</span>
|
||||
|
|
@ -171,7 +495,10 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">calibration_function</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">calibrator</span><span class="p">(</span><span class="n">posteriors</span><span class="p">,</span> <span class="n">np</span><span class="o">.</span><span class="n">eye</span><span class="p">(</span><span class="n">nclasses</span><span class="p">)[</span><span class="n">yva</span><span class="p">],</span> <span class="n">posterior_supplied</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.predict"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.predict">[docs]</a> <span class="k">def</span> <span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts class labels for the data instances in `X`</span>
|
||||
|
||||
|
|
@ -180,7 +507,10 @@
|
|||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">predict</span><span class="p">(</span><span class="n">X</span><span class="p">)</span></div>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.predict_proba"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.predict_proba">[docs]</a> <span class="k">def</span> <span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="RecalibratedProbabilisticClassifierBase.predict_proba">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.RecalibratedProbabilisticClassifierBase.predict_proba">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Generates posterior probabilities for the data instances in `X`</span>
|
||||
|
||||
|
|
@ -190,8 +520,9 @@
|
|||
<span class="n">posteriors</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">predict_proba</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">calibration_function</span><span class="p">(</span><span class="n">posteriors</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">classes_</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">classes_</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns the classes on which the classifier has been trained on</span>
|
||||
|
||||
|
|
@ -200,7 +531,10 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">classes_</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="NBVSCalibration"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.NBVSCalibration">[docs]</a><span class="k">class</span> <span class="nc">NBVSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="NBVSCalibration">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.NBVSCalibration">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">NBVSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Applies the No-Bias Vector Scaling (NBVS) calibration method from `abstention.calibration`, as defined in</span>
|
||||
<span class="sd"> `Alexandari et al. paper <http://proceedings.mlr.press/v119/alexandari20a.html>`_:</span>
|
||||
|
|
@ -214,7 +548,8 @@
|
|||
<span class="sd"> :param verbose: whether or not to display information in the standard output</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="n">NoBiasVectorScaling</span><span class="p">,</span> <span class="n">_</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="n">_require_abstention_calibration</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">calibrator</span> <span class="o">=</span> <span class="n">NoBiasVectorScaling</span><span class="p">(</span><span class="n">verbose</span><span class="o">=</span><span class="n">verbose</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">val_split</span> <span class="o">=</span> <span class="n">val_split</span>
|
||||
|
|
@ -222,7 +557,10 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="BCTSCalibration"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.BCTSCalibration">[docs]</a><span class="k">class</span> <span class="nc">BCTSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="BCTSCalibration">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.BCTSCalibration">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">BCTSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Applies the Bias-Corrected Temperature Scaling (BCTS) calibration method from `abstention.calibration`, as defined in</span>
|
||||
<span class="sd"> `Alexandari et al. paper <http://proceedings.mlr.press/v119/alexandari20a.html>`_:</span>
|
||||
|
|
@ -236,7 +574,8 @@
|
|||
<span class="sd"> :param verbose: whether or not to display information in the standard output</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="n">_</span><span class="p">,</span> <span class="n">TempScaling</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="n">_require_abstention_calibration</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">calibrator</span> <span class="o">=</span> <span class="n">TempScaling</span><span class="p">(</span><span class="n">verbose</span><span class="o">=</span><span class="n">verbose</span><span class="p">,</span> <span class="n">bias_positions</span><span class="o">=</span><span class="s1">'all'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">val_split</span> <span class="o">=</span> <span class="n">val_split</span>
|
||||
|
|
@ -244,7 +583,10 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="TSCalibration"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.TSCalibration">[docs]</a><span class="k">class</span> <span class="nc">TSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="TSCalibration">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.TSCalibration">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">TSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Applies the Temperature Scaling (TS) calibration method from `abstention.calibration`, as defined in</span>
|
||||
<span class="sd"> `Alexandari et al. paper <http://proceedings.mlr.press/v119/alexandari20a.html>`_:</span>
|
||||
|
|
@ -258,7 +600,8 @@
|
|||
<span class="sd"> :param verbose: whether or not to display information in the standard output</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="n">_</span><span class="p">,</span> <span class="n">TempScaling</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="n">_require_abstention_calibration</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">calibrator</span> <span class="o">=</span> <span class="n">TempScaling</span><span class="p">(</span><span class="n">verbose</span><span class="o">=</span><span class="n">verbose</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">val_split</span> <span class="o">=</span> <span class="n">val_split</span>
|
||||
|
|
@ -266,7 +609,10 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="VSCalibration"><a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.VSCalibration">[docs]</a><span class="k">class</span> <span class="nc">VSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="VSCalibration">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.VSCalibration">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">VSCalibration</span><span class="p">(</span><span class="n">RecalibratedProbabilisticClassifierBase</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Applies the Vector Scaling (VS) calibration method from `abstention.calibration`, as defined in</span>
|
||||
<span class="sd"> `Alexandari et al. paper <http://proceedings.mlr.press/v119/alexandari20a.html>`_:</span>
|
||||
|
|
@ -280,40 +626,172 @@
|
|||
<span class="sd"> :param verbose: whether or not to display information in the standard output</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="n">_</span><span class="p">,</span> <span class="n">_</span><span class="p">,</span> <span class="n">VectorScaling</span> <span class="o">=</span> <span class="n">_require_abstention_calibration</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">calibrator</span> <span class="o">=</span> <span class="n">VectorScaling</span><span class="p">(</span><span class="n">verbose</span><span class="o">=</span><span class="n">verbose</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">val_split</span> <span class="o">=</span> <span class="n">val_split</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span> <span class="o">=</span> <span class="n">n_jobs</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span></div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="TemperatureScalingFromLogits">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.TemperatureScalingFromLogits">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">TemperatureScalingFromLogits</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Calibrates a matrix of logits by learning a temperature-scaling mapping</span>
|
||||
<span class="sd"> with the calibration methods from `abstention.calibration`.</span>
|
||||
|
||||
<span class="sd"> This estimator is useful when the inputs are already logits produced by a</span>
|
||||
<span class="sd"> pretrained classifier, and the goal is to transform them directly into</span>
|
||||
<span class="sd"> calibrated posterior probabilities without retraining the underlying model.</span>
|
||||
|
||||
<span class="sd"> :param bias_corrected: if True, uses Bias-Corrected Temperature Scaling</span>
|
||||
<span class="sd"> (BCTS); otherwise, uses standard Temperature Scaling (TS)</span>
|
||||
<span class="sd"> :param verbose: whether the underlying calibrator should display progress</span>
|
||||
<span class="sd"> information</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">bias_corrected</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">bias_corrected</span> <span class="o">=</span> <span class="n">bias_corrected</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span>
|
||||
|
||||
<div class="viewcode-block" id="TemperatureScalingFromLogits.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.TemperatureScalingFromLogits.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fits the logits calibrator.</span>
|
||||
|
||||
<span class="sd"> :param X: array-like of shape `(n_samples, n_classes)` containing</span>
|
||||
<span class="sd"> logits</span>
|
||||
<span class="sd"> :param y: array-like of shape `(n_samples,)` containing class labels</span>
|
||||
<span class="sd"> :return: self</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">X</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">check_X_y</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">label_encoder_</span> <span class="o">=</span> <span class="n">LabelEncoder</span><span class="p">()</span>
|
||||
<span class="n">y_enc</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">label_encoder_</span><span class="o">.</span><span class="n">fit_transform</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">label_encoder_</span><span class="o">.</span><span class="n">classes_</span>
|
||||
|
||||
<span class="n">n_classes</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">logits_dim</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="k">if</span> <span class="n">n_classes</span> <span class="o">!=</span> <span class="n">logits_dim</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span>
|
||||
<span class="sa">f</span><span class="s1">'mismatch between the number of classes (</span><span class="si">{</span><span class="n">n_classes</span><span class="si">}</span><span class="s1">) and the '</span>
|
||||
<span class="sa">f</span><span class="s1">'dimensionality of the logits (</span><span class="si">{</span><span class="n">logits_dim</span><span class="si">}</span><span class="s1">)'</span>
|
||||
<span class="p">)</span>
|
||||
|
||||
<span class="n">_</span><span class="p">,</span> <span class="n">TempScaling</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="n">_require_abstention_calibration</span><span class="p">()</span>
|
||||
<span class="n">calibrator</span> <span class="o">=</span> <span class="n">TempScaling</span><span class="p">(</span>
|
||||
<span class="n">verbose</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">verbose</span><span class="p">,</span>
|
||||
<span class="n">bias_positions</span><span class="o">=</span><span class="s1">'all'</span> <span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">bias_corrected</span> <span class="k">else</span> <span class="p">[],</span>
|
||||
<span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">calibrator_</span> <span class="o">=</span> <span class="n">calibrator</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">calibration_function_</span> <span class="o">=</span> <span class="n">calibrator</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">np</span><span class="o">.</span><span class="n">eye</span><span class="p">(</span><span class="n">n_classes</span><span class="p">)[</span><span class="n">y_enc</span><span class="p">])</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="TemperatureScalingFromLogits.predict_proba">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.TemperatureScalingFromLogits.predict_proba">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Converts logits into calibrated posterior probabilities.</span>
|
||||
|
||||
<span class="sd"> :param X: array-like of shape `(n_samples, n_classes)` containing</span>
|
||||
<span class="sd"> logits</span>
|
||||
<span class="sd"> :return: array-like of shape `(n_samples, n_classes)` with calibrated</span>
|
||||
<span class="sd"> posterior probabilities</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">calibration_function_</span><span class="p">(</span><span class="n">X</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="TemperatureScalingFromLogits.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.calibration.TemperatureScalingFromLogits.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts class labels after calibration.</span>
|
||||
|
||||
<span class="sd"> :param X: array-like of shape `(n_samples, n_classes)` containing</span>
|
||||
<span class="sd"> logits</span>
|
||||
<span class="sd"> :return: array-like of shape `(n_samples,)` with class label</span>
|
||||
<span class="sd"> predictions</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">posteriors</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">predict_proba</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">label_encoder_</span><span class="o">.</span><span class="n">inverse_transform</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">posteriors</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">))</span></div>
|
||||
</div>
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,83 +1,384 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en" data-content_root="../../../">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.methods — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=92fd9be5" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css?v=19f00094" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.methods — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script src="../../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../../_static/documentation_options.js?v=22607128"></script>
|
||||
<script src="../../../_static/doctools.js?v=9a2dae69"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=dc90522c"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/classification/methods';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.classification.methods</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.classification.methods</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.classification.methods</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">from</span> <span class="nn">sklearn.base</span> <span class="kn">import</span> <span class="n">BaseEstimator</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.decomposition</span> <span class="kn">import</span> <span class="n">TruncatedSVD</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.linear_model</span> <span class="kn">import</span> <span class="n">LogisticRegression</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseEstimator</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.decomposition</span><span class="w"> </span><span class="kn">import</span> <span class="n">TruncatedSVD</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.linear_model</span><span class="w"> </span><span class="kn">import</span> <span class="n">LogisticRegression</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="LowRankLogisticRegression">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.LowRankLogisticRegression">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">LowRankLogisticRegression</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">):</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">LowRankLogisticRegression</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> An example of a classification method (i.e., an object that implements `fit`, `predict`, and `predict_proba`)</span>
|
||||
<span class="sd"> that also generates embedded inputs (i.e., that implements `transform`), as those required for</span>
|
||||
|
|
@ -91,13 +392,13 @@
|
|||
<span class="sd"> `Logistic Regression <https://scikit-learn.org/stable/modules/generated/sklearn.linear_model.LogisticRegression.html>`__ classifier</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">n_components</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">n_components</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_components</span> <span class="o">=</span> <span class="n">n_components</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">LogisticRegression</span><span class="p">(</span><span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="LowRankLogisticRegression.get_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.LowRankLogisticRegression.get_params">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Get hyper-parameters for this estimator.</span>
|
||||
|
||||
|
|
@ -110,7 +411,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LowRankLogisticRegression.set_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.LowRankLogisticRegression.set_params">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">params</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">params</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Set the parameters of this estimator.</span>
|
||||
|
||||
|
|
@ -127,7 +428,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LowRankLogisticRegression.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.LowRankLogisticRegression.fit">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fit the model according to the given training data. The fit consists of</span>
|
||||
<span class="sd"> fitting `TruncatedSVD` and then `LogisticRegression` on the low-rank representation.</span>
|
||||
|
|
@ -148,7 +449,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LowRankLogisticRegression.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.LowRankLogisticRegression.predict">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts labels for the instances `X` embedded into the low-rank space.</span>
|
||||
|
||||
|
|
@ -162,7 +463,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LowRankLogisticRegression.predict_proba">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.LowRankLogisticRegression.predict_proba">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts posterior probabilities for the instances `X` embedded into the low-rank space.</span>
|
||||
|
||||
|
|
@ -175,7 +476,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LowRankLogisticRegression.transform">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.LowRankLogisticRegression.transform">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns the low-rank approximation of `X` with `n_components` dimensions, or `X` unaltered if</span>
|
||||
<span class="sd"> `n_components` >= `X.shape[1]`.</span>
|
||||
|
|
@ -188,33 +489,109 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">pca</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">X</span><span class="p">)</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="MockClassifierFromPosteriors">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.MockClassifierFromPosteriors">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">MockClassifierFromPosteriors</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Mock classifier that bypasses classifier training when the input instances</span>
|
||||
<span class="sd"> are already posterior probabilities produced by a pretrained probabilistic</span>
|
||||
<span class="sd"> classifier.</span>
|
||||
|
||||
<span class="sd"> :param X: arrays of shape `(n_samples, n_classes)` are interpreted as posterior probabilities</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<div class="viewcode-block" id="MockClassifierFromPosteriors.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.MockClassifierFromPosteriors.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">sort</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">y</span><span class="p">))</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="MockClassifierFromPosteriors.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.MockClassifierFromPosteriors.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="n">np</span><span class="o">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="MockClassifierFromPosteriors.predict_proba">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.methods.MockClassifierFromPosteriors.predict_proba">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="n">X</span></div>
|
||||
</div>
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,95 +1,396 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en" data-content_root="../../../">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.neural — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=92fd9be5" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css?v=19f00094" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.neural — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script src="../../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../../_static/documentation_options.js?v=22607128"></script>
|
||||
<script src="../../../_static/doctools.js?v=9a2dae69"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=dc90522c"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=a95f357e85573c9b56d5" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=a95f357e85573c9b56d5" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=a95f357e85573c9b56d5"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=a95f357e85573c9b56d5" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=a95f357e85573c9b56d5" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/classification/neural';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.classification.neural</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.classification.neural</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.classification.neural</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">os</span>
|
||||
<span class="kn">from</span> <span class="nn">abc</span> <span class="kn">import</span> <span class="n">ABCMeta</span><span class="p">,</span> <span class="n">abstractmethod</span>
|
||||
<span class="kn">from</span> <span class="nn">pathlib</span> <span class="kn">import</span> <span class="n">Path</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">logging</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">os</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">abc</span><span class="w"> </span><span class="kn">import</span> <span class="n">ABCMeta</span><span class="p">,</span> <span class="n">abstractmethod</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">pathlib</span><span class="w"> </span><span class="kn">import</span> <span class="n">Path</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">import</span> <span class="nn">torch</span>
|
||||
<span class="kn">import</span> <span class="nn">torch.nn</span> <span class="k">as</span> <span class="nn">nn</span>
|
||||
<span class="kn">import</span> <span class="nn">torch.nn.functional</span> <span class="k">as</span> <span class="nn">F</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.metrics</span> <span class="kn">import</span> <span class="n">accuracy_score</span><span class="p">,</span> <span class="n">f1_score</span>
|
||||
<span class="kn">from</span> <span class="nn">torch.nn.utils.rnn</span> <span class="kn">import</span> <span class="n">pad_sequence</span>
|
||||
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">torch</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">torch.nn</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">nn</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">torch.nn.functional</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">F</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.metrics</span><span class="w"> </span><span class="kn">import</span> <span class="n">accuracy_score</span><span class="p">,</span> <span class="n">f1_score</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">torch.nn.utils.rnn</span><span class="w"> </span><span class="kn">import</span> <span class="n">pad_sequence</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">tqdm</span><span class="w"> </span><span class="kn">import</span> <span class="n">tqdm</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">quapy</span> <span class="k">as</span> <span class="nn">qp</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.data</span> <span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.util</span> <span class="kn">import</span> <span class="n">EarlyStop</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.data</span><span class="w"> </span><span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.util</span><span class="w"> </span><span class="kn">import</span> <span class="n">EarlyStop</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">NeuralClassifierTrainer</span><span class="p">:</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">NeuralClassifierTrainer</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Trains a neural network for text classification.</span>
|
||||
|
||||
|
|
@ -107,7 +408,7 @@
|
|||
<span class="sd"> according to the evaluation in the held-out validation split (default '../checkpoint/classifier_net.dat')</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="n">net</span><span class="p">:</span> <span class="s1">'TextClassifierNet'</span><span class="p">,</span>
|
||||
<span class="n">lr</span><span class="o">=</span><span class="mf">1e-3</span><span class="p">,</span>
|
||||
<span class="n">weight_decay</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span>
|
||||
|
|
@ -116,7 +417,7 @@
|
|||
<span class="n">batch_size</span><span class="o">=</span><span class="mi">64</span><span class="p">,</span>
|
||||
<span class="n">batch_size_test</span><span class="o">=</span><span class="mi">512</span><span class="p">,</span>
|
||||
<span class="n">padding_length</span><span class="o">=</span><span class="mi">300</span><span class="p">,</span>
|
||||
<span class="n">device</span><span class="o">=</span><span class="s1">'cuda'</span><span class="p">,</span>
|
||||
<span class="n">device</span><span class="o">=</span><span class="s1">'cpu'</span><span class="p">,</span>
|
||||
<span class="n">checkpointpath</span><span class="o">=</span><span class="s1">'../checkpoint/classifier_net.dat'</span><span class="p">):</span>
|
||||
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">()</span>
|
||||
|
|
@ -137,12 +438,12 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">learner_hyperparams</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">net</span><span class="o">.</span><span class="n">get_params</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">checkpointpath</span> <span class="o">=</span> <span class="n">checkpointpath</span>
|
||||
|
||||
<span class="nb">print</span><span class="p">(</span><span class="sa">f</span><span class="s1">'[NeuralNetwork running on </span><span class="si">{</span><span class="n">device</span><span class="si">}</span><span class="s1">]'</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="sa">f</span><span class="s1">'NeuralNetwork running on </span><span class="si">{</span><span class="n">device</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">os</span><span class="o">.</span><span class="n">makedirs</span><span class="p">(</span><span class="n">Path</span><span class="p">(</span><span class="n">checkpointpath</span><span class="p">)</span><span class="o">.</span><span class="n">parent</span><span class="p">,</span> <span class="n">exist_ok</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer.reset_net_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer.reset_net_params">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">reset_net_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">vocab_size</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">reset_net_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">vocab_size</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Reinitialize the network parameters</span>
|
||||
|
||||
<span class="sd"> :param vocab_size: the size of the vocabulary</span>
|
||||
|
|
@ -155,7 +456,7 @@
|
|||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer.get_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer.get_params">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Get hyper-parameters for this estimator</span>
|
||||
|
||||
<span class="sd"> :return: a dictionary with parameter names mapped to their values</span>
|
||||
|
|
@ -165,7 +466,7 @@
|
|||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer.set_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer.set_params">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">params</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">params</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Set the parameters of this trainer and the learner it is training.</span>
|
||||
<span class="sd"> In this current version, parameter names for the trainer and learner should</span>
|
||||
<span class="sd"> be disjoint.</span>
|
||||
|
|
@ -191,14 +492,14 @@
|
|||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">device</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">device</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">""" Gets the device in which the network is allocated</span>
|
||||
|
||||
<span class="sd"> :return: device</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="nb">next</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">net</span><span class="o">.</span><span class="n">parameters</span><span class="p">())</span><span class="o">.</span><span class="n">device</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_train_epoch</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_train_epoch</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">net</span><span class="o">.</span><span class="n">train</span><span class="p">()</span>
|
||||
<span class="n">criterion</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">CrossEntropyLoss</span><span class="p">()</span>
|
||||
<span class="n">losses</span><span class="p">,</span> <span class="n">predictions</span><span class="p">,</span> <span class="n">true_labels</span> <span class="o">=</span> <span class="p">[],</span> <span class="p">[],</span> <span class="p">[]</span>
|
||||
|
|
@ -218,7 +519,7 @@
|
|||
<span class="n">status</span><span class="p">[</span><span class="s2">"f1"</span><span class="p">]</span> <span class="o">=</span> <span class="n">f1_score</span><span class="p">(</span><span class="n">true_labels</span><span class="p">,</span> <span class="n">predictions</span><span class="p">,</span> <span class="n">average</span><span class="o">=</span><span class="s1">'macro'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">__update_progress_bar</span><span class="p">(</span><span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_test_epoch</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_test_epoch</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">net</span><span class="o">.</span><span class="n">eval</span><span class="p">()</span>
|
||||
<span class="n">criterion</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">CrossEntropyLoss</span><span class="p">()</span>
|
||||
<span class="n">losses</span><span class="p">,</span> <span class="n">predictions</span><span class="p">,</span> <span class="n">true_labels</span> <span class="o">=</span> <span class="p">[],</span> <span class="p">[],</span> <span class="p">[]</span>
|
||||
|
|
@ -236,7 +537,7 @@
|
|||
<span class="n">status</span><span class="p">[</span><span class="s2">"f1"</span><span class="p">]</span> <span class="o">=</span> <span class="n">f1_score</span><span class="p">(</span><span class="n">true_labels</span><span class="p">,</span> <span class="n">predictions</span><span class="p">,</span> <span class="n">average</span><span class="o">=</span><span class="s1">'macro'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">__update_progress_bar</span><span class="p">(</span><span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">__update_progress_bar</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__update_progress_bar</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="n">pbar</span><span class="o">.</span><span class="n">set_description</span><span class="p">(</span><span class="sa">f</span><span class="s1">'[</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">net</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">] training epoch=</span><span class="si">{</span><span class="n">epoch</span><span class="si">}</span><span class="s1"> '</span>
|
||||
<span class="sa">f</span><span class="s1">'tr-loss=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">status</span><span class="p">[</span><span class="s2">"tr"</span><span class="p">][</span><span class="s2">"loss"</span><span class="p">]</span><span class="si">:</span><span class="s1">.5f</span><span class="si">}</span><span class="s1"> '</span>
|
||||
<span class="sa">f</span><span class="s1">'tr-acc=</span><span class="si">{</span><span class="mi">100</span><span class="w"> </span><span class="o">*</span><span class="w"> </span><span class="bp">self</span><span class="o">.</span><span class="n">status</span><span class="p">[</span><span class="s2">"tr"</span><span class="p">][</span><span class="s2">"acc"</span><span class="p">]</span><span class="si">:</span><span class="s1">.2f</span><span class="si">}</span><span class="s1">% '</span>
|
||||
|
|
@ -248,7 +549,7 @@
|
|||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer.fit">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mf">0.3</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mf">0.3</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fits the model according to the given training data.</span>
|
||||
|
||||
|
|
@ -283,21 +584,22 @@
|
|||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">early_stop</span><span class="o">.</span><span class="n">IMPROVED</span><span class="p">:</span>
|
||||
<span class="n">torch</span><span class="o">.</span><span class="n">save</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">net</span><span class="o">.</span><span class="n">state_dict</span><span class="p">(),</span> <span class="n">checkpoint</span><span class="p">)</span>
|
||||
<span class="k">elif</span> <span class="bp">self</span><span class="o">.</span><span class="n">early_stop</span><span class="o">.</span><span class="n">STOP</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="sa">f</span><span class="s1">'training ended by patience exhasted; loading best model parameters in </span><span class="si">{</span><span class="n">checkpoint</span><span class="si">}</span><span class="s1"> '</span>
|
||||
<span class="sa">f</span><span class="s1">'for epoch </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">early_stop</span><span class="o">.</span><span class="n">best_epoch</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span>
|
||||
<span class="sa">f</span><span class="s1">'training ended by patience exhausted; loading best model parameters in </span><span class="si">{</span><span class="n">checkpoint</span><span class="si">}</span><span class="s1"> '</span>
|
||||
<span class="sa">f</span><span class="s1">'for epoch </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">early_stop</span><span class="o">.</span><span class="n">best_epoch</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">net</span><span class="o">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">load</span><span class="p">(</span><span class="n">checkpoint</span><span class="p">))</span>
|
||||
<span class="k">break</span>
|
||||
|
||||
<span class="nb">print</span><span class="p">(</span><span class="s1">'performing one training pass over the validation set...'</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="s1">'performing one training pass over the validation set...'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_train_epoch</span><span class="p">(</span><span class="n">valid_generator</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">status</span><span class="p">[</span><span class="s1">'tr'</span><span class="p">],</span> <span class="n">pbar</span><span class="p">,</span> <span class="n">epoch</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="s1">'[done]'</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="s1">'done'</span><span class="p">)</span>
|
||||
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer.predict">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts labels for the instances</span>
|
||||
|
||||
|
|
@ -310,7 +612,7 @@
|
|||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer.predict_proba">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer.predict_proba">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts posterior probabilities for the instances</span>
|
||||
|
||||
|
|
@ -329,7 +631,7 @@
|
|||
|
||||
<div class="viewcode-block" id="NeuralClassifierTrainer.transform">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.NeuralClassifierTrainer.transform">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns the embeddings of the instances</span>
|
||||
|
||||
|
|
@ -351,7 +653,7 @@
|
|||
|
||||
<div class="viewcode-block" id="TorchDataset">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TorchDataset">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">TorchDataset</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">Dataset</span><span class="p">):</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">TorchDataset</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">Dataset</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Transforms labelled instances into a Torch's :class:`torch.utils.data.DataLoader` object</span>
|
||||
|
||||
|
|
@ -359,19 +661,19 @@
|
|||
<span class="sd"> :param labels: array-like of shape `(n_samples, n_classes)` with the class labels</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">,</span> <span class="n">labels</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">,</span> <span class="n">labels</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">instances</span> <span class="o">=</span> <span class="n">instances</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">labels</span> <span class="o">=</span> <span class="n">labels</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__len__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__len__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">instances</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__getitem__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">index</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__getitem__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">index</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="p">{</span><span class="s1">'doc'</span><span class="p">:</span> <span class="bp">self</span><span class="o">.</span><span class="n">instances</span><span class="p">[</span><span class="n">index</span><span class="p">],</span> <span class="s1">'label'</span><span class="p">:</span> <span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">[</span><span class="n">index</span><span class="p">]</span> <span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">labels</span> <span class="ow">is</span> <span class="ow">not</span> <span class="kc">None</span> <span class="k">else</span> <span class="kc">None</span><span class="p">}</span>
|
||||
|
||||
<div class="viewcode-block" id="TorchDataset.asDataloader">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TorchDataset.asDataloader">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">asDataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch_size</span><span class="p">,</span> <span class="n">shuffle</span><span class="p">,</span> <span class="n">pad_length</span><span class="p">,</span> <span class="n">device</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">asDataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch_size</span><span class="p">,</span> <span class="n">shuffle</span><span class="p">,</span> <span class="n">pad_length</span><span class="p">,</span> <span class="n">device</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Converts the labelled collection into a Torch DataLoader with dynamic padding for</span>
|
||||
<span class="sd"> the batch</span>
|
||||
|
|
@ -384,7 +686,7 @@
|
|||
<span class="sd"> :param device: whether to allocate tensors in cpu or in cuda</span>
|
||||
<span class="sd"> :return: a :class:`torch.utils.data.DataLoader` object</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">def</span> <span class="nf">collate</span><span class="p">(</span><span class="n">batch</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">collate</span><span class="p">(</span><span class="n">batch</span><span class="p">):</span>
|
||||
<span class="n">data</span> <span class="o">=</span> <span class="p">[</span><span class="n">torch</span><span class="o">.</span><span class="n">LongTensor</span><span class="p">(</span><span class="n">item</span><span class="p">[</span><span class="s1">'doc'</span><span class="p">][:</span><span class="n">pad_length</span><span class="p">])</span> <span class="k">for</span> <span class="n">item</span> <span class="ow">in</span> <span class="n">batch</span><span class="p">]</span>
|
||||
<span class="n">data</span> <span class="o">=</span> <span class="n">pad_sequence</span><span class="p">(</span><span class="n">data</span><span class="p">,</span> <span class="n">batch_first</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">padding_value</span><span class="o">=</span><span class="n">qp</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">'PAD_INDEX'</span><span class="p">])</span><span class="o">.</span><span class="n">to</span><span class="p">(</span><span class="n">device</span><span class="p">)</span>
|
||||
<span class="n">targets</span> <span class="o">=</span> <span class="p">[</span><span class="n">item</span><span class="p">[</span><span class="s1">'label'</span><span class="p">]</span> <span class="k">for</span> <span class="n">item</span> <span class="ow">in</span> <span class="n">batch</span><span class="p">]</span>
|
||||
|
|
@ -402,7 +704,7 @@
|
|||
|
||||
<div class="viewcode-block" id="TextClassifierNet">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TextClassifierNet">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">TextClassifierNet</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Module</span><span class="p">,</span> <span class="n">metaclass</span><span class="o">=</span><span class="n">ABCMeta</span><span class="p">):</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">TextClassifierNet</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Module</span><span class="p">,</span> <span class="n">metaclass</span><span class="o">=</span><span class="n">ABCMeta</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Abstract Text classifier (`torch.nn.Module`)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
|
@ -410,7 +712,7 @@
|
|||
<div class="viewcode-block" id="TextClassifierNet.document_embedding">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TextClassifierNet.document_embedding">[docs]</a>
|
||||
<span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span> <span class="nf">document_embedding</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">document_embedding</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Embeds documents (i.e., performs the forward pass up to the</span>
|
||||
<span class="sd"> next-to-last layer).</span>
|
||||
|
||||
|
|
@ -425,7 +727,7 @@
|
|||
|
||||
<div class="viewcode-block" id="TextClassifierNet.forward">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TextClassifierNet.forward">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Performs the forward pass.</span>
|
||||
|
||||
<span class="sd"> :param x: a batch of instances, typically generated by a torch's `DataLoader`</span>
|
||||
|
|
@ -439,7 +741,7 @@
|
|||
|
||||
<div class="viewcode-block" id="TextClassifierNet.dimensions">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TextClassifierNet.dimensions">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">dimensions</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">dimensions</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Gets the number of dimensions of the embedding space</span>
|
||||
|
||||
<span class="sd"> :return: integer</span>
|
||||
|
|
@ -449,7 +751,7 @@
|
|||
|
||||
<div class="viewcode-block" id="TextClassifierNet.predict_proba">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TextClassifierNet.predict_proba">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict_proba</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts posterior probabilities for the instances in `x`</span>
|
||||
|
||||
|
|
@ -464,7 +766,7 @@
|
|||
|
||||
<div class="viewcode-block" id="TextClassifierNet.xavier_uniform">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TextClassifierNet.xavier_uniform">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">xavier_uniform</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">xavier_uniform</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Performs Xavier initialization of the network parameters</span>
|
||||
<span class="sd"> """</span>
|
||||
|
|
@ -476,7 +778,7 @@
|
|||
<div class="viewcode-block" id="TextClassifierNet.get_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.TextClassifierNet.get_params">[docs]</a>
|
||||
<span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span> <span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Get hyper-parameters for this estimator</span>
|
||||
|
||||
|
|
@ -486,7 +788,7 @@
|
|||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Return the size of the vocabulary</span>
|
||||
|
||||
|
|
@ -498,7 +800,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LSTMnet">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.LSTMnet">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">LSTMnet</span><span class="p">(</span><span class="n">TextClassifierNet</span><span class="p">):</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">LSTMnet</span><span class="p">(</span><span class="n">TextClassifierNet</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> An implementation of :class:`quapy.classification.neural.TextClassifierNet` based on</span>
|
||||
<span class="sd"> Long Short Term Memory networks.</span>
|
||||
|
|
@ -512,7 +814,7 @@
|
|||
<span class="sd"> :param drop_p: drop probability for dropout (default 0.5)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">vocabulary_size</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">,</span> <span class="n">embedding_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">hidden_size</span><span class="o">=</span><span class="mi">256</span><span class="p">,</span> <span class="n">repr_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">lstm_class_nlayers</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">vocabulary_size</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">,</span> <span class="n">embedding_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">hidden_size</span><span class="o">=</span><span class="mi">256</span><span class="p">,</span> <span class="n">repr_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">lstm_class_nlayers</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
|
||||
<span class="n">drop_p</span><span class="o">=</span><span class="mf">0.5</span><span class="p">):</span>
|
||||
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">()</span>
|
||||
|
|
@ -534,7 +836,7 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">doc_embedder</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">hidden_size</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">dim</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">output</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Linear</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">dim</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">__init_hidden</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">set_size</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__init_hidden</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">set_size</span><span class="p">):</span>
|
||||
<span class="n">opt</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">hyperparams</span>
|
||||
<span class="n">var_hidden</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">opt</span><span class="p">[</span><span class="s1">'lstm_class_nlayers'</span><span class="p">],</span> <span class="n">set_size</span><span class="p">,</span> <span class="n">opt</span><span class="p">[</span><span class="s1">'hidden_size'</span><span class="p">])</span>
|
||||
<span class="n">var_cell</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">opt</span><span class="p">[</span><span class="s1">'lstm_class_nlayers'</span><span class="p">],</span> <span class="n">set_size</span><span class="p">,</span> <span class="n">opt</span><span class="p">[</span><span class="s1">'hidden_size'</span><span class="p">])</span>
|
||||
|
|
@ -544,7 +846,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LSTMnet.document_embedding">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.LSTMnet.document_embedding">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">document_embedding</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">document_embedding</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Embeds documents (i.e., performs the forward pass up to the</span>
|
||||
<span class="sd"> next-to-last layer).</span>
|
||||
|
||||
|
|
@ -563,7 +865,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LSTMnet.get_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.LSTMnet.get_params">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Get hyper-parameters for this estimator</span>
|
||||
|
||||
|
|
@ -573,7 +875,7 @@
|
|||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Return the size of the vocabulary</span>
|
||||
|
||||
|
|
@ -585,7 +887,7 @@
|
|||
|
||||
<div class="viewcode-block" id="CNNnet">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.CNNnet">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">CNNnet</span><span class="p">(</span><span class="n">TextClassifierNet</span><span class="p">):</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">CNNnet</span><span class="p">(</span><span class="n">TextClassifierNet</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> An implementation of :class:`quapy.classification.neural.TextClassifierNet` based on</span>
|
||||
<span class="sd"> Convolutional Neural Networks.</span>
|
||||
|
|
@ -602,7 +904,7 @@
|
|||
<span class="sd"> :param drop_p: drop probability for dropout (default 0.5)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">vocabulary_size</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">,</span> <span class="n">embedding_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">hidden_size</span><span class="o">=</span><span class="mi">256</span><span class="p">,</span> <span class="n">repr_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">vocabulary_size</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">,</span> <span class="n">embedding_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">hidden_size</span><span class="o">=</span><span class="mi">256</span><span class="p">,</span> <span class="n">repr_size</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span>
|
||||
<span class="n">kernel_heights</span><span class="o">=</span><span class="p">[</span><span class="mi">3</span><span class="p">,</span> <span class="mi">5</span><span class="p">,</span> <span class="mi">7</span><span class="p">],</span> <span class="n">stride</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">padding</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span> <span class="n">drop_p</span><span class="o">=</span><span class="mf">0.5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">(</span><span class="n">CNNnet</span><span class="p">,</span> <span class="bp">self</span><span class="p">)</span><span class="o">.</span><span class="fm">__init__</span><span class="p">()</span>
|
||||
|
||||
|
|
@ -627,7 +929,7 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">doc_embedder</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Linear</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">kernel_heights</span><span class="p">)</span> <span class="o">*</span> <span class="n">hidden_size</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">dim</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">output</span> <span class="o">=</span> <span class="n">nn</span><span class="o">.</span><span class="n">Linear</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">dim</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">__conv_block</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="nb">input</span><span class="p">,</span> <span class="n">conv_layer</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__conv_block</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="nb">input</span><span class="p">,</span> <span class="n">conv_layer</span><span class="p">):</span>
|
||||
<span class="n">conv_out</span> <span class="o">=</span> <span class="n">conv_layer</span><span class="p">(</span><span class="nb">input</span><span class="p">)</span> <span class="c1"># conv_out.size() = (batch_size, out_channels, dim, 1)</span>
|
||||
<span class="n">activation</span> <span class="o">=</span> <span class="n">F</span><span class="o">.</span><span class="n">relu</span><span class="p">(</span><span class="n">conv_out</span><span class="o">.</span><span class="n">squeeze</span><span class="p">(</span><span class="mi">3</span><span class="p">))</span> <span class="c1"># activation.size() = (batch_size, out_channels, dim1)</span>
|
||||
<span class="n">max_out</span> <span class="o">=</span> <span class="n">F</span><span class="o">.</span><span class="n">max_pool1d</span><span class="p">(</span><span class="n">activation</span><span class="p">,</span> <span class="n">activation</span><span class="o">.</span><span class="n">size</span><span class="p">()[</span><span class="mi">2</span><span class="p">])</span><span class="o">.</span><span class="n">squeeze</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span> <span class="c1"># maxpool_out.size() = (batch_size, out_channels)</span>
|
||||
|
|
@ -635,7 +937,7 @@
|
|||
|
||||
<div class="viewcode-block" id="CNNnet.document_embedding">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.CNNnet.document_embedding">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">document_embedding</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="nb">input</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">document_embedding</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="nb">input</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Embeds documents (i.e., performs the forward pass up to the</span>
|
||||
<span class="sd"> next-to-last layer).</span>
|
||||
|
||||
|
|
@ -660,7 +962,7 @@
|
|||
|
||||
<div class="viewcode-block" id="CNNnet.get_params">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.neural.CNNnet.get_params">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Get hyper-parameters for this estimator</span>
|
||||
|
||||
|
|
@ -670,7 +972,7 @@
|
|||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Return the size of the vocabulary</span>
|
||||
|
||||
|
|
@ -685,31 +987,75 @@
|
|||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=a95f357e85573c9b56d5"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=a95f357e85573c9b56d5"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.19.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,90 +1,391 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en" data-content_root="../../../">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.svmperf — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=92fd9be5" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css?v=19f00094" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.classification.svmperf — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script src="../../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../../_static/documentation_options.js?v=22607128"></script>
|
||||
<script src="../../../_static/doctools.js?v=9a2dae69"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=dc90522c"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/classification/svmperf';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.classification.svmperf</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.classification.svmperf</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.classification.svmperf</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">random</span>
|
||||
<span class="kn">import</span> <span class="nn">shutil</span>
|
||||
<span class="kn">import</span> <span class="nn">subprocess</span>
|
||||
<span class="kn">import</span> <span class="nn">tempfile</span>
|
||||
<span class="kn">from</span> <span class="nn">os</span> <span class="kn">import</span> <span class="n">remove</span><span class="p">,</span> <span class="n">makedirs</span>
|
||||
<span class="kn">from</span> <span class="nn">os.path</span> <span class="kn">import</span> <span class="n">join</span><span class="p">,</span> <span class="n">exists</span>
|
||||
<span class="kn">from</span> <span class="nn">subprocess</span> <span class="kn">import</span> <span class="n">PIPE</span><span class="p">,</span> <span class="n">STDOUT</span>
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.base</span> <span class="kn">import</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">ClassifierMixin</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.datasets</span> <span class="kn">import</span> <span class="n">dump_svmlight_file</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">logging</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">random</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">shutil</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">subprocess</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">tempfile</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">os</span><span class="w"> </span><span class="kn">import</span> <span class="n">remove</span><span class="p">,</span> <span class="n">makedirs</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">os.path</span><span class="w"> </span><span class="kn">import</span> <span class="n">join</span><span class="p">,</span> <span class="n">exists</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">subprocess</span><span class="w"> </span><span class="kn">import</span> <span class="n">PIPE</span><span class="p">,</span> <span class="n">STDOUT</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">ClassifierMixin</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.datasets</span><span class="w"> </span><span class="kn">import</span> <span class="n">dump_svmlight_file</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="SVMperf">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.svmperf.SVMperf">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">SVMperf</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">ClassifierMixin</span><span class="p">):</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">SVMperf</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">ClassifierMixin</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""A wrapper for the `SVM-perf package <https://www.cs.cornell.edu/people/tj/svm_light/svm_perf.html>`__ by Thorsten Joachims.</span>
|
||||
<span class="sd"> When using losses for quantification, the source code has to be patched. See</span>
|
||||
<span class="sd"> the `installation documentation <https://hlt-isti.github.io/QuaPy/build/html/Installation.html#svm-perf-with-quantification-oriented-losses>`__</span>
|
||||
|
|
@ -106,31 +407,20 @@
|
|||
<span class="c1"># losses with their respective codes in svm_perf implementation</span>
|
||||
<span class="n">valid_losses</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'01'</span><span class="p">:</span><span class="mi">0</span><span class="p">,</span> <span class="s1">'f1'</span><span class="p">:</span><span class="mi">1</span><span class="p">,</span> <span class="s1">'kld'</span><span class="p">:</span><span class="mi">12</span><span class="p">,</span> <span class="s1">'nkld'</span><span class="p">:</span><span class="mi">13</span><span class="p">,</span> <span class="s1">'q'</span><span class="p">:</span><span class="mi">22</span><span class="p">,</span> <span class="s1">'qacc'</span><span class="p">:</span><span class="mi">23</span><span class="p">,</span> <span class="s1">'qf1'</span><span class="p">:</span><span class="mi">24</span><span class="p">,</span> <span class="s1">'qgm'</span><span class="p">:</span><span class="mi">25</span><span class="p">,</span> <span class="s1">'mae'</span><span class="p">:</span><span class="mi">26</span><span class="p">,</span> <span class="s1">'mrae'</span><span class="p">:</span><span class="mi">27</span><span class="p">}</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">svmperf_base</span><span class="p">,</span> <span class="n">C</span><span class="o">=</span><span class="mf">0.01</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">loss</span><span class="o">=</span><span class="s1">'01'</span><span class="p">,</span> <span class="n">host_folder</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="n">exists</span><span class="p">(</span><span class="n">svmperf_base</span><span class="p">),</span> <span class="sa">f</span><span class="s1">'path </span><span class="si">{</span><span class="n">svmperf_base</span><span class="si">}</span><span class="s1"> does not seem to point to a valid path'</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">svmperf_base</span><span class="p">,</span> <span class="n">C</span><span class="o">=</span><span class="mf">0.01</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">loss</span><span class="o">=</span><span class="s1">'01'</span><span class="p">,</span> <span class="n">host_folder</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="n">exists</span><span class="p">(</span><span class="n">svmperf_base</span><span class="p">),</span> \
|
||||
<span class="p">(</span><span class="sa">f</span><span class="s1">'path </span><span class="si">{</span><span class="n">svmperf_base</span><span class="si">}</span><span class="s1"> does not seem to point to a valid path;'</span>
|
||||
<span class="sa">f</span><span class="s1">'did you install svm-perf? '</span>
|
||||
<span class="sa">f</span><span class="s1">'see instructions in https://hlt-isti.github.io/QuaPy/manuals/explicit-loss-minimization.html'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">svmperf_base</span> <span class="o">=</span> <span class="n">svmperf_base</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">C</span> <span class="o">=</span> <span class="n">C</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">loss</span> <span class="o">=</span> <span class="n">loss</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">host_folder</span> <span class="o">=</span> <span class="n">host_folder</span>
|
||||
|
||||
<span class="c1"># def set_params(self, **parameters):</span>
|
||||
<span class="c1"># """</span>
|
||||
<span class="c1"># Set the hyper-parameters for svm-perf. Currently, only the `C` and `loss` parameters are supported</span>
|
||||
<span class="c1">#</span>
|
||||
<span class="c1"># :param parameters: a `**kwargs` dictionary `{'C': <float>}`</span>
|
||||
<span class="c1"># """</span>
|
||||
<span class="c1"># assert sorted(list(parameters.keys())) == ['C', 'loss'], \</span>
|
||||
<span class="c1"># 'currently, only the C and loss parameters are supported'</span>
|
||||
<span class="c1"># self.C = parameters.get('C', self.C)</span>
|
||||
<span class="c1"># self.loss = parameters.get('loss', self.loss)</span>
|
||||
<span class="c1">#</span>
|
||||
<span class="c1"># def get_params(self, deep=True):</span>
|
||||
<span class="c1"># return {'C': self.C, 'loss': self.loss}</span>
|
||||
|
||||
<div class="viewcode-block" id="SVMperf.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.svmperf.SVMperf.fit">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Trains the SVM for the multivariate performance loss</span>
|
||||
|
||||
|
|
@ -153,8 +443,7 @@
|
|||
<span class="c1"># this would allow to run parallel instances of predict</span>
|
||||
<span class="n">random_code</span> <span class="o">=</span> <span class="s1">'svmperfprocess'</span><span class="o">+</span><span class="s1">'-'</span><span class="o">.</span><span class="n">join</span><span class="p">(</span><span class="nb">str</span><span class="p">(</span><span class="n">local_random</span><span class="o">.</span><span class="n">randint</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="mi">1000000</span><span class="p">))</span> <span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">5</span><span class="p">))</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">host_folder</span> <span class="ow">is</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="c1"># tmp dir are removed after the fit terminates in multiprocessing...</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">tmpdir</span> <span class="o">=</span> <span class="n">tempfile</span><span class="o">.</span><span class="n">TemporaryDirectory</span><span class="p">(</span><span class="n">suffix</span><span class="o">=</span><span class="n">random_code</span><span class="p">)</span><span class="o">.</span><span class="n">name</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">tmpdir</span> <span class="o">=</span> <span class="n">join</span><span class="p">(</span><span class="n">tempfile</span><span class="o">.</span><span class="n">gettempdir</span><span class="p">(),</span> <span class="n">random_code</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">tmpdir</span> <span class="o">=</span> <span class="n">join</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">host_folder</span><span class="p">,</span> <span class="s1">'.'</span> <span class="o">+</span> <span class="n">random_code</span><span class="p">)</span>
|
||||
<span class="n">makedirs</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">tmpdir</span><span class="p">,</span> <span class="n">exist_ok</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
|
|
@ -166,21 +455,21 @@
|
|||
|
||||
<span class="n">cmd</span> <span class="o">=</span> <span class="s1">' '</span><span class="o">.</span><span class="n">join</span><span class="p">([</span><span class="bp">self</span><span class="o">.</span><span class="n">svmperf_learn</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">c_cmd</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss_cmd</span><span class="p">,</span> <span class="n">traindat</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">model</span><span class="p">])</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">verbose</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="s1">'[Running]'</span><span class="p">,</span> <span class="n">cmd</span><span class="p">)</span>
|
||||
<span class="n">p</span> <span class="o">=</span> <span class="n">subprocess</span><span class="o">.</span><span class="n">run</span><span class="p">(</span><span class="n">cmd</span><span class="o">.</span><span class="n">split</span><span class="p">(),</span> <span class="n">stdout</span><span class="o">=</span><span class="n">PIPE</span><span class="p">,</span> <span class="n">stderr</span><span class="o">=</span><span class="n">STDOUT</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="sa">f</span><span class="s1">'[Running] </span><span class="si">{</span><span class="n">cmd</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">p</span> <span class="o">=</span> <span class="n">subprocess</span><span class="o">.</span><span class="n">run</span><span class="p">(</span><span class="n">cmd</span><span class="o">.</span><span class="n">split</span><span class="p">(),</span> <span class="n">stdout</span><span class="o">=</span><span class="n">PIPE</span><span class="p">,</span> <span class="n">stderr</span><span class="o">=</span><span class="n">PIPE</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="ow">not</span> <span class="n">exists</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model</span><span class="p">):</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="n">p</span><span class="o">.</span><span class="n">stderr</span><span class="o">.</span><span class="n">decode</span><span class="p">(</span><span class="s1">'utf-8'</span><span class="p">))</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">error</span><span class="p">(</span><span class="n">p</span><span class="o">.</span><span class="n">stderr</span><span class="o">.</span><span class="n">decode</span><span class="p">(</span><span class="s1">'utf-8'</span><span class="p">))</span>
|
||||
<span class="n">remove</span><span class="p">(</span><span class="n">traindat</span><span class="p">)</span>
|
||||
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">verbose</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="n">p</span><span class="o">.</span><span class="n">stdout</span><span class="o">.</span><span class="n">decode</span><span class="p">(</span><span class="s1">'utf-8'</span><span class="p">))</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="n">p</span><span class="o">.</span><span class="n">stdout</span><span class="o">.</span><span class="n">decode</span><span class="p">(</span><span class="s1">'utf-8'</span><span class="p">))</span>
|
||||
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="SVMperf.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.svmperf.SVMperf.predict">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Predicts labels for the instances `X`</span>
|
||||
|
||||
|
|
@ -195,7 +484,7 @@
|
|||
|
||||
<div class="viewcode-block" id="SVMperf.decision_function">
|
||||
<a class="viewcode-back" href="../../../quapy.classification.html#quapy.classification.svmperf.SVMperf.decision_function">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">decision_function</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">decision_function</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Evaluate the decision function for the samples in `X`.</span>
|
||||
|
||||
|
|
@ -218,11 +507,11 @@
|
|||
|
||||
<span class="n">cmd</span> <span class="o">=</span> <span class="s1">' '</span><span class="o">.</span><span class="n">join</span><span class="p">([</span><span class="bp">self</span><span class="o">.</span><span class="n">svmperf_classify</span><span class="p">,</span> <span class="n">testdat</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">model</span><span class="p">,</span> <span class="n">predictions_path</span><span class="p">])</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">verbose</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="s1">'[Running]'</span><span class="p">,</span> <span class="n">cmd</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="sa">f</span><span class="s1">'[Running] </span><span class="si">{</span><span class="n">cmd</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">p</span> <span class="o">=</span> <span class="n">subprocess</span><span class="o">.</span><span class="n">run</span><span class="p">(</span><span class="n">cmd</span><span class="o">.</span><span class="n">split</span><span class="p">(),</span> <span class="n">stdout</span><span class="o">=</span><span class="n">PIPE</span><span class="p">,</span> <span class="n">stderr</span><span class="o">=</span><span class="n">STDOUT</span><span class="p">)</span>
|
||||
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">verbose</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="n">p</span><span class="o">.</span><span class="n">stdout</span><span class="o">.</span><span class="n">decode</span><span class="p">(</span><span class="s1">'utf-8'</span><span class="p">))</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="n">p</span><span class="o">.</span><span class="n">stdout</span><span class="o">.</span><span class="n">decode</span><span class="p">(</span><span class="s1">'utf-8'</span><span class="p">))</span>
|
||||
|
||||
<span class="n">scores</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">loadtxt</span><span class="p">(</span><span class="n">predictions_path</span><span class="p">)</span>
|
||||
<span class="n">remove</span><span class="p">(</span><span class="n">testdat</span><span class="p">)</span>
|
||||
|
|
@ -231,38 +520,82 @@
|
|||
<span class="k">return</span> <span class="n">scores</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="fm">__del__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__del__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="nb">hasattr</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="s1">'tmpdir'</span><span class="p">):</span>
|
||||
<span class="n">shutil</span><span class="o">.</span><span class="n">rmtree</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">tmpdir</span><span class="p">,</span> <span class="n">ignore_errors</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,91 +1,392 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en" data-content_root="../../../">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.data.base — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=92fd9be5" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css?v=19f00094" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.data.base — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script src="../../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../../_static/documentation_options.js?v=22607128"></script>
|
||||
<script src="../../../_static/doctools.js?v=9a2dae69"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=dc90522c"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/data/base';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.data.base</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
<h1>Source code for quapy.data.base</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">itertools</span>
|
||||
<span class="kn">from</span> <span class="nn">functools</span> <span class="kn">import</span> <span class="n">cached_property</span>
|
||||
<span class="kn">from</span> <span class="nn">typing</span> <span class="kn">import</span> <span class="n">Iterable</span>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">scipy.sparse</span> <span class="kn">import</span> <span class="n">issparse</span>
|
||||
<span class="kn">from</span> <span class="nn">scipy.sparse</span> <span class="kn">import</span> <span class="n">vstack</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.model_selection</span> <span class="kn">import</span> <span class="n">train_test_split</span><span class="p">,</span> <span class="n">RepeatedStratifiedKFold</span>
|
||||
<span class="kn">from</span> <span class="nn">numpy.random</span> <span class="kn">import</span> <span class="n">RandomState</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.functional</span> <span class="kn">import</span> <span class="n">strprev</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.util</span> <span class="kn">import</span> <span class="n">temp_seed</span>
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.data.base</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.data.base</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">itertools</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">functools</span><span class="w"> </span><span class="kn">import</span> <span class="n">cached_property</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">typing</span><span class="w"> </span><span class="kn">import</span> <span class="n">Iterable</span>
|
||||
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">scipy.sparse</span><span class="w"> </span><span class="kn">import</span> <span class="n">issparse</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">scipy.sparse</span><span class="w"> </span><span class="kn">import</span> <span class="n">vstack</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.model_selection</span><span class="w"> </span><span class="kn">import</span> <span class="n">train_test_split</span><span class="p">,</span> <span class="n">RepeatedStratifiedKFold</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">numpy.random</span><span class="w"> </span><span class="kn">import</span> <span class="n">RandomState</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.functional</span><span class="w"> </span><span class="kn">import</span> <span class="n">strprev</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.util</span><span class="w"> </span><span class="kn">import</span> <span class="n">temp_seed</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy.functional</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">F</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="LabelledCollection">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">LabelledCollection</span><span class="p">:</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">LabelledCollection</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> A LabelledCollection is a set of objects each with a label attached to each of them. </span>
|
||||
<span class="sd"> This class implements several sampling routines and other utilities.</span>
|
||||
|
|
@ -97,7 +398,7 @@
|
|||
<span class="sd"> (i.e., a prevalence of 0)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">issparse</span><span class="p">(</span><span class="n">instances</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">instances</span> <span class="o">=</span> <span class="n">instances</span>
|
||||
<span class="k">elif</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">instances</span><span class="p">,</span> <span class="nb">list</span><span class="p">)</span> <span class="ow">and</span> <span class="nb">len</span><span class="p">(</span><span class="n">instances</span><span class="p">)</span> <span class="o">></span> <span class="mi">0</span> <span class="ow">and</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">instances</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="nb">str</span><span class="p">):</span>
|
||||
|
|
@ -106,21 +407,25 @@
|
|||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">instances</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span><span class="n">instances</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">labels</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span><span class="n">labels</span><span class="p">)</span>
|
||||
<span class="n">n_docs</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="n">classes</span> <span class="ow">is</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="o">.</span><span class="n">sort</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="n">F</span><span class="o">.</span><span class="n">classes_from_labels</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span><span class="n">classes</span><span class="p">))</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="o">.</span><span class="n">sort</span><span class="p">()</span>
|
||||
<span class="k">if</span> <span class="nb">len</span><span class="p">(</span><span class="nb">set</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">)</span><span class="o">.</span><span class="n">difference</span><span class="p">(</span><span class="nb">set</span><span class="p">(</span><span class="n">classes</span><span class="p">)))</span> <span class="o">></span> <span class="mi">0</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="sa">f</span><span class="s1">'labels (</span><span class="si">{</span><span class="nb">set</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">)</span><span class="si">}</span><span class="s1">) contain values not included in classes_ (</span><span class="si">{</span><span class="nb">set</span><span class="p">(</span><span class="n">classes</span><span class="p">)</span><span class="si">}</span><span class="s1">)'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">index</span> <span class="o">=</span> <span class="p">{</span><span class="n">class_</span><span class="p">:</span> <span class="n">np</span><span class="o">.</span><span class="n">arange</span><span class="p">(</span><span class="n">n_docs</span><span class="p">)[</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span> <span class="o">==</span> <span class="n">class_</span><span class="p">]</span> <span class="k">for</span> <span class="n">class_</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">}</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_index</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">index</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="ow">not</span> <span class="nb">hasattr</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="s1">'_index'</span><span class="p">)</span> <span class="ow">or</span> <span class="bp">self</span><span class="o">.</span><span class="n">_index</span> <span class="ow">is</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_index</span> <span class="o">=</span> <span class="p">{</span><span class="n">class_</span><span class="p">:</span> <span class="n">np</span><span class="o">.</span><span class="n">arange</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="p">))[</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span> <span class="o">==</span> <span class="n">class_</span><span class="p">]</span> <span class="k">for</span> <span class="n">class_</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">}</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">_index</span>
|
||||
|
||||
<div class="viewcode-block" id="LabelledCollection.load">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.load">[docs]</a>
|
||||
<span class="nd">@classmethod</span>
|
||||
<span class="k">def</span> <span class="nf">load</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">path</span><span class="p">:</span> <span class="nb">str</span><span class="p">,</span> <span class="n">loader_func</span><span class="p">:</span> <span class="n">callable</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="o">**</span><span class="n">loader_kwargs</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">load</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">path</span><span class="p">:</span> <span class="nb">str</span><span class="p">,</span> <span class="n">loader_func</span><span class="p">:</span> <span class="nb">callable</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="o">**</span><span class="n">loader_kwargs</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Loads a labelled set of data and convert it into a :class:`LabelledCollection` instance. The function in charge</span>
|
||||
<span class="sd"> of reading the instances must be specified. This function can be a custom one, or any of the reading functions</span>
|
||||
|
|
@ -137,7 +442,7 @@
|
|||
<span class="k">return</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="o">*</span><span class="n">loader_func</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="o">**</span><span class="n">loader_kwargs</span><span class="p">),</span> <span class="n">classes</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="fm">__len__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__len__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns the length of this collection (number of labelled instances)</span>
|
||||
|
||||
|
|
@ -147,7 +452,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.prevalence">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.prevalence">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">prevalence</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">prevalence</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns the prevalence, or relative frequency, of the classes in the codeframe.</span>
|
||||
|
||||
|
|
@ -159,7 +464,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.counts">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.counts">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">counts</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">counts</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns the number of instances for each of the classes in the codeframe.</span>
|
||||
|
||||
|
|
@ -170,7 +475,7 @@
|
|||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">n_classes</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">n_classes</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> The number of classes</span>
|
||||
|
||||
|
|
@ -179,7 +484,16 @@
|
|||
<span class="k">return</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">binary</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">n_instances</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> The number of instances</span>
|
||||
|
||||
<span class="sd"> :return: integer</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">)</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">binary</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns True if the number of classes is 2</span>
|
||||
|
||||
|
|
@ -189,12 +503,11 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.sampling_index">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.sampling_index">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">sampling_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="o">*</span><span class="n">prevs</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">sampling_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="o">*</span><span class="n">prevs</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns an index to be used to extract a random sample of desired size and desired prevalence values. If the</span>
|
||||
<span class="sd"> prevalence values are not specified, then returns the index of a uniform sampling.</span>
|
||||
<span class="sd"> For each class, the sampling is drawn with replacement if the requested prevalence is larger than</span>
|
||||
<span class="sd"> the actual prevalence of the class, or without replacement otherwise.</span>
|
||||
<span class="sd"> For each class, the sampling is drawn with replacement.</span>
|
||||
|
||||
<span class="sd"> :param size: integer, the requested size</span>
|
||||
<span class="sd"> :param prevs: the prevalence for each class; the prevalence value for the last class can be lead empty since</span>
|
||||
|
|
@ -209,7 +522,7 @@
|
|||
<span class="k">if</span> <span class="nb">len</span><span class="p">(</span><span class="n">prevs</span><span class="p">)</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">n_classes</span> <span class="o">-</span> <span class="mi">1</span><span class="p">:</span>
|
||||
<span class="n">prevs</span> <span class="o">=</span> <span class="n">prevs</span> <span class="o">+</span> <span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="nb">sum</span><span class="p">(</span><span class="n">prevs</span><span class="p">),)</span>
|
||||
<span class="k">assert</span> <span class="nb">len</span><span class="p">(</span><span class="n">prevs</span><span class="p">)</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">n_classes</span><span class="p">,</span> <span class="s1">'unexpected number of prevalences'</span>
|
||||
<span class="k">assert</span> <span class="nb">sum</span><span class="p">(</span><span class="n">prevs</span><span class="p">)</span> <span class="o">==</span> <span class="mi">1</span><span class="p">,</span> <span class="sa">f</span><span class="s1">'prevalences (</span><span class="si">{</span><span class="n">prevs</span><span class="si">}</span><span class="s1">) wrong range (sum=</span><span class="si">{</span><span class="nb">sum</span><span class="p">(</span><span class="n">prevs</span><span class="p">)</span><span class="si">}</span><span class="s1">)'</span>
|
||||
<span class="k">assert</span> <span class="n">np</span><span class="o">.</span><span class="n">isclose</span><span class="p">(</span><span class="nb">sum</span><span class="p">(</span><span class="n">prevs</span><span class="p">),</span> <span class="mi">1</span><span class="p">),</span> <span class="sa">f</span><span class="s1">'prevalences (</span><span class="si">{</span><span class="n">prevs</span><span class="si">}</span><span class="s1">) wrong range (sum=</span><span class="si">{</span><span class="nb">sum</span><span class="p">(</span><span class="n">prevs</span><span class="p">)</span><span class="si">}</span><span class="s1">)'</span>
|
||||
|
||||
<span class="c1"># Decide how many instances should be taken for each class in order to satisfy the requested prevalence</span>
|
||||
<span class="c1"># accurately, and the number of instances in the sample (exactly). If int(size * prevs[i]) (which is</span>
|
||||
|
|
@ -238,7 +551,7 @@
|
|||
<span class="k">for</span> <span class="n">class_</span><span class="p">,</span> <span class="n">n_requested</span> <span class="ow">in</span> <span class="n">n_requests</span><span class="o">.</span><span class="n">items</span><span class="p">():</span>
|
||||
<span class="n">n_candidates</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">index</span><span class="p">[</span><span class="n">class_</span><span class="p">])</span>
|
||||
<span class="n">index_sample</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">index</span><span class="p">[</span><span class="n">class_</span><span class="p">][</span>
|
||||
<span class="n">np</span><span class="o">.</span><span class="n">random</span><span class="o">.</span><span class="n">choice</span><span class="p">(</span><span class="n">n_candidates</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">n_requested</span><span class="p">,</span> <span class="n">replace</span><span class="o">=</span><span class="p">(</span><span class="n">n_requested</span> <span class="o">></span> <span class="n">n_candidates</span><span class="p">))</span>
|
||||
<span class="n">np</span><span class="o">.</span><span class="n">random</span><span class="o">.</span><span class="n">choice</span><span class="p">(</span><span class="n">n_candidates</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">n_requested</span><span class="p">,</span> <span class="n">replace</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
<span class="p">]</span> <span class="k">if</span> <span class="n">n_requested</span> <span class="o">></span> <span class="mi">0</span> <span class="k">else</span> <span class="p">[]</span>
|
||||
|
||||
<span class="n">indexes_sample</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">index_sample</span><span class="p">)</span>
|
||||
|
|
@ -253,11 +566,10 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.uniform_sampling_index">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.uniform_sampling_index">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">uniform_sampling_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">uniform_sampling_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns an index to be used to extract a uniform sample of desired size. The sampling is drawn</span>
|
||||
<span class="sd"> with replacement if the requested size is greater than the number of instances, or without replacement</span>
|
||||
<span class="sd"> otherwise.</span>
|
||||
<span class="sd"> with replacement.</span>
|
||||
|
||||
<span class="sd"> :param size: integer, the size of the uniform sample</span>
|
||||
<span class="sd"> :param random_state: if specified, guarantees reproducibility of the split.</span>
|
||||
|
|
@ -267,16 +579,15 @@
|
|||
<span class="n">ng</span> <span class="o">=</span> <span class="n">RandomState</span><span class="p">(</span><span class="n">seed</span><span class="o">=</span><span class="n">random_state</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="n">ng</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">random</span>
|
||||
<span class="k">return</span> <span class="n">ng</span><span class="o">.</span><span class="n">choice</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="p">),</span> <span class="n">size</span><span class="p">,</span> <span class="n">replace</span><span class="o">=</span><span class="n">size</span> <span class="o">></span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="p">))</span></div>
|
||||
<span class="k">return</span> <span class="n">ng</span><span class="o">.</span><span class="n">choice</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="p">),</span> <span class="n">size</span><span class="p">,</span> <span class="n">replace</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="LabelledCollection.sampling">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.sampling">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">sampling</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="o">*</span><span class="n">prevs</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">sampling</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="o">*</span><span class="n">prevs</span><span class="p">,</span> <span class="n">shuffle</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Return a random sample (an instance of :class:`LabelledCollection`) of desired size and desired prevalence</span>
|
||||
<span class="sd"> values. For each class, the sampling is drawn without replacement if the requested prevalence is larger than</span>
|
||||
<span class="sd"> the actual prevalence of the class, or with replacement otherwise.</span>
|
||||
<span class="sd"> values. For each class, the sampling is drawn with replacement.</span>
|
||||
|
||||
<span class="sd"> :param size: integer, the requested size</span>
|
||||
<span class="sd"> :param prevs: the prevalence for each class; the prevalence value for the last class can be lead empty since</span>
|
||||
|
|
@ -293,11 +604,10 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.uniform_sampling">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.uniform_sampling">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">uniform_sampling</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">uniform_sampling</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">size</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns a uniform sample (an instance of :class:`LabelledCollection`) of desired size. The sampling is drawn</span>
|
||||
<span class="sd"> with replacement if the requested size is greater than the number of instances, or without replacement</span>
|
||||
<span class="sd"> otherwise.</span>
|
||||
<span class="sd"> with replacement.</span>
|
||||
|
||||
<span class="sd"> :param size: integer, the requested size</span>
|
||||
<span class="sd"> :param random_state: if specified, guarantees reproducibility of the split.</span>
|
||||
|
|
@ -309,7 +619,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.sampling_from_index">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.sampling_from_index">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">sampling_from_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">index</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">sampling_from_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">index</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns an instance of :class:`LabelledCollection` whose elements are sampled from this collection using the</span>
|
||||
<span class="sd"> index.</span>
|
||||
|
|
@ -324,7 +634,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.split_stratified">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.split_stratified">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">split_stratified</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">train_prop</span><span class="o">=</span><span class="mf">0.6</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">split_stratified</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">train_prop</span><span class="o">=</span><span class="mf">0.6</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns two instances of :class:`LabelledCollection` split with stratification from this collection, at desired</span>
|
||||
<span class="sd"> proportion.</span>
|
||||
|
|
@ -336,17 +646,17 @@
|
|||
<span class="sd"> :return: two instances of :class:`LabelledCollection`, the first one with `train_prop` elements, and the</span>
|
||||
<span class="sd"> second one with `1-train_prop` elements</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">tr_docs</span><span class="p">,</span> <span class="n">te_docs</span><span class="p">,</span> <span class="n">tr_labels</span><span class="p">,</span> <span class="n">te_labels</span> <span class="o">=</span> <span class="n">train_test_split</span><span class="p">(</span>
|
||||
<span class="n">tr_X</span><span class="p">,</span> <span class="n">te_X</span><span class="p">,</span> <span class="n">tr_y</span><span class="p">,</span> <span class="n">te_y</span> <span class="o">=</span> <span class="n">train_test_split</span><span class="p">(</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">instances</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">train_size</span><span class="o">=</span><span class="n">train_prop</span><span class="p">,</span> <span class="n">stratify</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="n">random_state</span>
|
||||
<span class="p">)</span>
|
||||
<span class="n">training</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">tr_docs</span><span class="p">,</span> <span class="n">tr_labels</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">test</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">te_docs</span><span class="p">,</span> <span class="n">te_labels</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">training</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">tr_X</span><span class="p">,</span> <span class="n">tr_y</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">test</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">te_X</span><span class="p">,</span> <span class="n">te_y</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">training</span><span class="p">,</span> <span class="n">test</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="LabelledCollection.split_random">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.split_random">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">split_random</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">train_prop</span><span class="o">=</span><span class="mf">0.6</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">split_random</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">train_prop</span><span class="o">=</span><span class="mf">0.6</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns two instances of :class:`LabelledCollection` split randomly from this collection, at desired</span>
|
||||
<span class="sd"> proportion.</span>
|
||||
|
|
@ -373,7 +683,7 @@
|
|||
<span class="k">return</span> <span class="n">training</span><span class="p">,</span> <span class="n">test</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="fm">__add__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">other</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__add__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">other</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns a new :class:`LabelledCollection` as the union of this collection with another collection.</span>
|
||||
<span class="sd"> Both labelled collections must have the same classes.</span>
|
||||
|
|
@ -389,7 +699,7 @@
|
|||
<div class="viewcode-block" id="LabelledCollection.join">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.join">[docs]</a>
|
||||
<span class="nd">@classmethod</span>
|
||||
<span class="k">def</span> <span class="nf">join</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">:</span> <span class="n">Iterable</span><span class="p">[</span><span class="s1">'LabelledCollection'</span><span class="p">]):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">join</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">:</span> <span class="n">Iterable</span><span class="p">[</span><span class="s1">'LabelledCollection'</span><span class="p">]):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns a new :class:`LabelledCollection` as the union of the collections given in input.</span>
|
||||
|
||||
|
|
@ -425,12 +735,23 @@
|
|||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="ne">NotImplementedError</span><span class="p">(</span><span class="s1">'unsupported operation for collection types'</span><span class="p">)</span>
|
||||
<span class="n">labels</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">concatenate</span><span class="p">([</span><span class="n">lc</span><span class="o">.</span><span class="n">labels</span> <span class="k">for</span> <span class="n">lc</span> <span class="ow">in</span> <span class="n">args</span><span class="p">])</span>
|
||||
<span class="n">classes</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">labels</span><span class="p">)</span><span class="o">.</span><span class="n">sort</span><span class="p">()</span>
|
||||
<span class="c1"># union of each collection's own classes_, so a class declared but absent from</span>
|
||||
<span class="c1"># this particular join (e.g. an empty fold) is preserved at zero prevalence</span>
|
||||
<span class="n">classes</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">concatenate</span><span class="p">([</span><span class="n">lc</span><span class="o">.</span><span class="n">classes_</span> <span class="k">for</span> <span class="n">lc</span> <span class="ow">in</span> <span class="n">args</span><span class="p">]))</span>
|
||||
<span class="k">return</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">instances</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="n">classes</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">Xy</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">classes</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Gets an array-like with the classes used in this collection</span>
|
||||
|
||||
<span class="sd"> :return: array-like</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">classes_</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">Xy</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Gets the instances and labels. This is useful when working with `sklearn` estimators, e.g.:</span>
|
||||
|
||||
|
|
@ -441,7 +762,7 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">instances</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">labels</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">Xp</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">Xp</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Gets the instances and the true prevalence. This is useful when implementing evaluation protocols from</span>
|
||||
<span class="sd"> a :class:`LabelledCollection` object.</span>
|
||||
|
|
@ -451,7 +772,7 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">instances</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">prevalence</span><span class="p">()</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">X</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">X</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> An alias to self.instances</span>
|
||||
|
||||
|
|
@ -460,7 +781,7 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">instances</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">y</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">y</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> An alias to self.labels</span>
|
||||
|
||||
|
|
@ -469,7 +790,7 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">labels</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">p</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">p</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> An alias to self.prevalence()</span>
|
||||
|
||||
|
|
@ -480,7 +801,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.stats">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.stats">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">stats</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">show</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">stats</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">show</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns (and eventually prints) a dictionary with some stats of this collection. E.g.,:</span>
|
||||
|
||||
|
|
@ -515,7 +836,7 @@
|
|||
|
||||
<div class="viewcode-block" id="LabelledCollection.kFCV">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.LabelledCollection.kFCV">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">kFCV</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">nfolds</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">nrepeats</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">kFCV</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">nfolds</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">nrepeats</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Generator of stratified folds to be used in k-fold cross validation.</span>
|
||||
|
||||
|
|
@ -529,13 +850,18 @@
|
|||
<span class="n">train</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">sampling_from_index</span><span class="p">(</span><span class="n">train_index</span><span class="p">)</span>
|
||||
<span class="n">test</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">sampling_from_index</span><span class="p">(</span><span class="n">test_index</span><span class="p">)</span>
|
||||
<span class="k">yield</span> <span class="n">train</span><span class="p">,</span> <span class="n">test</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__repr__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="nb">repr</span><span class="o">=</span><span class="sa">f</span><span class="s1">'<</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">n_instances</span><span class="si">}</span><span class="s1"> instances (dtype=</span><span class="si">{</span><span class="nb">type</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">instances</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span><span class="si">}</span><span class="s1">), '</span>
|
||||
<span class="nb">repr</span><span class="o">+=</span><span class="sa">f</span><span class="s1">'n_classes=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">n_classes</span><span class="si">}</span><span class="s1"> </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="si">}</span><span class="s1">, prevalence=</span><span class="si">{</span><span class="n">F</span><span class="o">.</span><span class="n">strprev</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">prevalence</span><span class="p">())</span><span class="si">}</span><span class="s1">>'</span>
|
||||
<span class="k">return</span> <span class="nb">repr</span></div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="Dataset">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.Dataset">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">Dataset</span><span class="p">:</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">Dataset</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Abstraction of training and test :class:`LabelledCollection` objects.</span>
|
||||
|
||||
|
|
@ -545,7 +871,7 @@
|
|||
<span class="sd"> :param name: a string representing the name of the dataset</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">training</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">test</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">vocabulary</span><span class="p">:</span> <span class="nb">dict</span> <span class="o">=</span> <span class="kc">None</span><span class="p">,</span> <span class="n">name</span><span class="o">=</span><span class="s1">''</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">training</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">test</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">vocabulary</span><span class="p">:</span> <span class="nb">dict</span> <span class="o">=</span> <span class="kc">None</span><span class="p">,</span> <span class="n">name</span><span class="o">=</span><span class="s1">''</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="nb">set</span><span class="p">(</span><span class="n">training</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span> <span class="o">==</span> <span class="nb">set</span><span class="p">(</span><span class="n">test</span><span class="o">.</span><span class="n">classes_</span><span class="p">),</span> <span class="s1">'incompatible labels in training and test collections'</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">training</span> <span class="o">=</span> <span class="n">training</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">test</span> <span class="o">=</span> <span class="n">test</span>
|
||||
|
|
@ -555,7 +881,7 @@
|
|||
<div class="viewcode-block" id="Dataset.SplitStratified">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.Dataset.SplitStratified">[docs]</a>
|
||||
<span class="nd">@classmethod</span>
|
||||
<span class="k">def</span> <span class="nf">SplitStratified</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">collection</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">train_size</span><span class="o">=</span><span class="mf">0.6</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">SplitStratified</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">collection</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">train_size</span><span class="o">=</span><span class="mf">0.6</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Generates a :class:`Dataset` from a stratified split of a :class:`LabelledCollection` instance.</span>
|
||||
<span class="sd"> See :meth:`LabelledCollection.split_stratified`</span>
|
||||
|
|
@ -568,7 +894,7 @@
|
|||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">classes_</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">classes_</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> The classes according to which the training collection is labelled</span>
|
||||
|
||||
|
|
@ -577,7 +903,7 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">classes_</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">n_classes</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">n_classes</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> The number of classes according to which the training collection is labelled</span>
|
||||
|
||||
|
|
@ -586,7 +912,7 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">n_classes</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">binary</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">binary</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns True if the training collection is labelled according to two classes</span>
|
||||
|
||||
|
|
@ -597,7 +923,7 @@
|
|||
<div class="viewcode-block" id="Dataset.load">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.Dataset.load">[docs]</a>
|
||||
<span class="nd">@classmethod</span>
|
||||
<span class="k">def</span> <span class="nf">load</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">train_path</span><span class="p">,</span> <span class="n">test_path</span><span class="p">,</span> <span class="n">loader_func</span><span class="p">:</span> <span class="n">callable</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="o">**</span><span class="n">loader_kwargs</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">load</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">train_path</span><span class="p">,</span> <span class="n">test_path</span><span class="p">,</span> <span class="n">loader_func</span><span class="p">:</span> <span class="nb">callable</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="o">**</span><span class="n">loader_kwargs</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Loads a training and a test labelled set of data and convert it into a :class:`Dataset` instance.</span>
|
||||
<span class="sd"> The function in charge of reading the instances must be specified. This function can be a custom one, or any of</span>
|
||||
|
|
@ -619,7 +945,7 @@
|
|||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> If the dataset is textual, and the vocabulary was indicated, returns the size of the vocabulary</span>
|
||||
|
||||
|
|
@ -628,7 +954,7 @@
|
|||
<span class="k">return</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">vocabulary</span><span class="p">)</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">train_test</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">train_test</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Alias to `self.training` and `self.test`</span>
|
||||
|
||||
|
|
@ -639,7 +965,7 @@
|
|||
|
||||
<div class="viewcode-block" id="Dataset.stats">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.Dataset.stats">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">stats</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">show</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">stats</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">show</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns (and eventually prints) a dictionary with some stats of this dataset. E.g.,:</span>
|
||||
|
||||
|
|
@ -666,7 +992,7 @@
|
|||
<div class="viewcode-block" id="Dataset.kFCV">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.Dataset.kFCV">[docs]</a>
|
||||
<span class="nd">@classmethod</span>
|
||||
<span class="k">def</span> <span class="nf">kFCV</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">nfolds</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">nrepeats</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="mi">0</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">kFCV</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">nfolds</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">nrepeats</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="mi">0</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Generator of stratified folds to be used in k-fold cross validation. This function is only a wrapper around</span>
|
||||
<span class="sd"> :meth:`LabelledCollection.kFCV` that returns :class:`Dataset` instances made of training and test folds.</span>
|
||||
|
|
@ -683,7 +1009,7 @@
|
|||
|
||||
<div class="viewcode-block" id="Dataset.reduce">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.base.Dataset.reduce">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">reduce</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">n_train</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">n_test</span><span class="o">=</span><span class="mi">100</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">reduce</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">n_train</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">n_test</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Reduce the number of instances in place for quick experiments. Preserves the prevalence of each set.</span>
|
||||
|
||||
|
|
@ -691,38 +1017,93 @@
|
|||
<span class="sd"> :param n_test: number of test documents to keep (default 100)</span>
|
||||
<span class="sd"> :return: self</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">training</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">sampling</span><span class="p">(</span><span class="n">n_train</span><span class="p">,</span> <span class="o">*</span><span class="bp">self</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">prevalence</span><span class="p">())</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">test</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">sampling</span><span class="p">(</span><span class="n">n_test</span><span class="p">,</span> <span class="o">*</span><span class="bp">self</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">prevalence</span><span class="p">())</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">training</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">sampling</span><span class="p">(</span>
|
||||
<span class="n">n_train</span><span class="p">,</span>
|
||||
<span class="o">*</span><span class="bp">self</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">prevalence</span><span class="p">(),</span>
|
||||
<span class="n">random_state</span> <span class="o">=</span> <span class="n">random_state</span>
|
||||
<span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">test</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">sampling</span><span class="p">(</span>
|
||||
<span class="n">n_test</span><span class="p">,</span>
|
||||
<span class="o">*</span><span class="bp">self</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">prevalence</span><span class="p">(),</span>
|
||||
<span class="n">random_state</span> <span class="o">=</span> <span class="n">random_state</span>
|
||||
<span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__repr__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="sa">f</span><span class="s1">'training=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">training</span><span class="si">}</span><span class="s1">; test=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">test</span><span class="si">}</span><span class="s1">'</span></div>
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,22 +1,20 @@
|
|||
|
||||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en" data-content_root="../../../">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.data.preprocessing — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=92fd9be5" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css?v=19f00094" />
|
||||
<title>quapy.data.preprocessing — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=b86133f3" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css?v=9edc463e" />
|
||||
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script src="../../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../../_static/documentation_options.js?v=22607128"></script>
|
||||
<script src="../../../_static/doctools.js?v=9a2dae69"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=dc90522c"></script>
|
||||
<script src="../../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
|
|
@ -42,7 +40,13 @@
|
|||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../index.html">Quickstart</a></li>
|
||||
</ul>
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../manuals.html">Manuals</a></li>
|
||||
</ul>
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../quapy.html">API</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
|
|
@ -70,21 +74,55 @@
|
|||
<div itemprop="articleBody">
|
||||
|
||||
<h1>Source code for quapy.data.preprocessing</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">scipy.sparse</span> <span class="kn">import</span> <span class="n">spmatrix</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.feature_extraction.text</span> <span class="kn">import</span> <span class="n">TfidfVectorizer</span><span class="p">,</span> <span class="n">CountVectorizer</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.preprocessing</span> <span class="kn">import</span> <span class="n">StandardScaler</span>
|
||||
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">scipy.sparse</span><span class="w"> </span><span class="kn">import</span> <span class="n">spmatrix</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.feature_extraction.text</span><span class="w"> </span><span class="kn">import</span> <span class="n">TfidfVectorizer</span><span class="p">,</span> <span class="n">CountVectorizer</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.preprocessing</span><span class="w"> </span><span class="kn">import</span> <span class="n">StandardScaler</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">tqdm</span><span class="w"> </span><span class="kn">import</span> <span class="n">tqdm</span>
|
||||
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.data.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">Dataset</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.util</span><span class="w"> </span><span class="kn">import</span> <span class="n">map_parallel</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="instance_transformation">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.instance_transformation">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">instance_transformation</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span><span class="n">Dataset</span><span class="p">,</span> <span class="n">transformer</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Transforms a :class:`quapy.data.base.Dataset` applying the `fit_transform` and `transform` functions</span>
|
||||
<span class="sd"> of a (sklearn's) transformer.</span>
|
||||
|
||||
<span class="sd"> :param dataset: a :class:`quapy.data.base.Dataset` where the instances of training and test collections are</span>
|
||||
<span class="sd"> lists of str</span>
|
||||
<span class="sd"> :param transformer: TransformerMixin implementing `fit_transform` and `transform` functions</span>
|
||||
<span class="sd"> :param inplace: whether or not to apply the transformation inplace (True), or to a new copy (False, default)</span>
|
||||
<span class="sd"> :return: a new :class:`quapy.data.base.Dataset` with transformed instances (if inplace=False) or a reference to the</span>
|
||||
<span class="sd"> current Dataset (if inplace=True) where the instances have been transformed</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">training_transformed</span> <span class="o">=</span> <span class="n">transformer</span><span class="o">.</span><span class="n">fit_transform</span><span class="p">(</span><span class="o">*</span><span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">Xy</span><span class="p">)</span>
|
||||
<span class="n">test_transformed</span> <span class="o">=</span> <span class="n">transformer</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">orig_name</span> <span class="o">=</span> <span class="n">dataset</span><span class="o">.</span><span class="n">name</span>
|
||||
|
||||
<span class="k">if</span> <span class="n">inplace</span><span class="p">:</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">training</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">training_transformed</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">test</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">test_transformed</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="nb">hasattr</span><span class="p">(</span><span class="n">transformer</span><span class="p">,</span> <span class="s1">'vocabulary_'</span><span class="p">):</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">vocabulary</span> <span class="o">=</span> <span class="n">transformer</span><span class="o">.</span><span class="n">vocabulary_</span>
|
||||
<span class="k">return</span> <span class="n">dataset</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="n">training</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">training_transformed</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">labels</span><span class="o">.</span><span class="n">copy</span><span class="p">(),</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">test</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">test_transformed</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">labels</span><span class="o">.</span><span class="n">copy</span><span class="p">(),</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">vocab</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="k">if</span> <span class="nb">hasattr</span><span class="p">(</span><span class="n">transformer</span><span class="p">,</span> <span class="s1">'vocabulary_'</span><span class="p">):</span>
|
||||
<span class="n">vocab</span> <span class="o">=</span> <span class="n">transformer</span><span class="o">.</span><span class="n">vocabulary_</span>
|
||||
<span class="k">return</span> <span class="n">Dataset</span><span class="p">(</span><span class="n">training</span><span class="p">,</span> <span class="n">test</span><span class="p">,</span> <span class="n">vocabulary</span><span class="o">=</span><span class="n">vocab</span><span class="p">,</span> <span class="n">name</span><span class="o">=</span><span class="n">orig_name</span><span class="p">)</span></div>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">quapy</span> <span class="k">as</span> <span class="nn">qp</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.data.base</span> <span class="kn">import</span> <span class="n">Dataset</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.util</span> <span class="kn">import</span> <span class="n">map_parallel</span>
|
||||
<span class="kn">from</span> <span class="nn">.base</span> <span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="text2tfidf">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.text2tfidf">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">text2tfidf</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span><span class="n">Dataset</span><span class="p">,</span> <span class="n">min_df</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">sublinear_tf</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">text2tfidf</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span><span class="n">Dataset</span><span class="p">,</span> <span class="n">min_df</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">sublinear_tf</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Transforms a :class:`quapy.data.base.Dataset` of textual instances into a :class:`quapy.data.base.Dataset` of</span>
|
||||
<span class="sd"> tfidf weighted sparse vectors</span>
|
||||
|
|
@ -103,24 +141,13 @@
|
|||
<span class="n">__check_type</span><span class="p">(</span><span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">instances</span><span class="p">,</span> <span class="n">np</span><span class="o">.</span><span class="n">ndarray</span><span class="p">,</span> <span class="nb">str</span><span class="p">)</span>
|
||||
|
||||
<span class="n">vectorizer</span> <span class="o">=</span> <span class="n">TfidfVectorizer</span><span class="p">(</span><span class="n">min_df</span><span class="o">=</span><span class="n">min_df</span><span class="p">,</span> <span class="n">sublinear_tf</span><span class="o">=</span><span class="n">sublinear_tf</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
|
||||
<span class="n">training_documents</span> <span class="o">=</span> <span class="n">vectorizer</span><span class="o">.</span><span class="n">fit_transform</span><span class="p">(</span><span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">instances</span><span class="p">)</span>
|
||||
<span class="n">test_documents</span> <span class="o">=</span> <span class="n">vectorizer</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">instances</span><span class="p">)</span>
|
||||
|
||||
<span class="k">if</span> <span class="n">inplace</span><span class="p">:</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">training</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">training_documents</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">test</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">test_documents</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">vocabulary</span> <span class="o">=</span> <span class="n">vectorizer</span><span class="o">.</span><span class="n">vocabulary_</span>
|
||||
<span class="k">return</span> <span class="n">dataset</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="n">training</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">training_documents</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">labels</span><span class="o">.</span><span class="n">copy</span><span class="p">(),</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">test</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">test_documents</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">labels</span><span class="o">.</span><span class="n">copy</span><span class="p">(),</span> <span class="n">dataset</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">Dataset</span><span class="p">(</span><span class="n">training</span><span class="p">,</span> <span class="n">test</span><span class="p">,</span> <span class="n">vectorizer</span><span class="o">.</span><span class="n">vocabulary_</span><span class="p">)</span></div>
|
||||
<span class="k">return</span> <span class="n">instance_transformation</span><span class="p">(</span><span class="n">dataset</span><span class="p">,</span> <span class="n">vectorizer</span><span class="p">,</span> <span class="n">inplace</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="reduce_columns">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.reduce_columns">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">reduce_columns</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span> <span class="n">Dataset</span><span class="p">,</span> <span class="n">min_df</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">reduce_columns</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span> <span class="n">Dataset</span><span class="p">,</span> <span class="n">min_df</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Reduces the dimensionality of the instances, represented as a `csr_matrix` (or any subtype of</span>
|
||||
<span class="sd"> `scipy.sparse.spmatrix`), of training and test documents by removing the columns of words which are not present</span>
|
||||
|
|
@ -138,7 +165,7 @@
|
|||
<span class="n">__check_type</span><span class="p">(</span><span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">instances</span><span class="p">,</span> <span class="n">spmatrix</span><span class="p">)</span>
|
||||
<span class="k">assert</span> <span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">instances</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">==</span> <span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">instances</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="s1">'unaligned vector spaces'</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">filter_by_occurrences</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">W</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">filter_by_occurrences</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">W</span><span class="p">):</span>
|
||||
<span class="n">column_prevalence</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">((</span><span class="n">X</span> <span class="o">></span> <span class="mi">0</span><span class="p">)</span><span class="o">.</span><span class="n">sum</span><span class="p">(</span><span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">))</span><span class="o">.</span><span class="n">flatten</span><span class="p">()</span>
|
||||
<span class="n">take_columns</span> <span class="o">=</span> <span class="n">column_prevalence</span> <span class="o">>=</span> <span class="n">min_df</span>
|
||||
<span class="n">X</span> <span class="o">=</span> <span class="n">X</span><span class="p">[:,</span> <span class="n">take_columns</span><span class="p">]</span>
|
||||
|
|
@ -159,7 +186,7 @@
|
|||
|
||||
<div class="viewcode-block" id="standardize">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.standardize">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">standardize</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span> <span class="n">Dataset</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">standardize</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span> <span class="n">Dataset</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Standardizes the real-valued columns of a :class:`quapy.data.base.Dataset`.</span>
|
||||
<span class="sd"> Standardization, aka z-scoring, of a variable `X` comes down to subtracting the average and normalizing by the</span>
|
||||
|
|
@ -170,19 +197,24 @@
|
|||
<span class="sd"> :class:`quapy.data.base.Dataset` is to be returned</span>
|
||||
<span class="sd"> :return: an instance of :class:`quapy.data.base.Dataset`</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">s</span> <span class="o">=</span> <span class="n">StandardScaler</span><span class="p">(</span><span class="n">copy</span><span class="o">=</span><span class="ow">not</span> <span class="n">inplace</span><span class="p">)</span>
|
||||
<span class="n">training</span> <span class="o">=</span> <span class="n">s</span><span class="o">.</span><span class="n">fit_transform</span><span class="p">(</span><span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">instances</span><span class="p">)</span>
|
||||
<span class="n">test</span> <span class="o">=</span> <span class="n">s</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">instances</span><span class="p">)</span>
|
||||
<span class="n">s</span> <span class="o">=</span> <span class="n">StandardScaler</span><span class="p">()</span>
|
||||
<span class="n">train</span><span class="p">,</span> <span class="n">test</span> <span class="o">=</span> <span class="n">dataset</span><span class="o">.</span><span class="n">train_test</span>
|
||||
<span class="n">std_train_X</span> <span class="o">=</span> <span class="n">s</span><span class="o">.</span><span class="n">fit_transform</span><span class="p">(</span><span class="n">train</span><span class="o">.</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">std_test_X</span> <span class="o">=</span> <span class="n">s</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">test</span><span class="o">.</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="n">inplace</span><span class="p">:</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">training</span><span class="o">.</span><span class="n">instances</span> <span class="o">=</span> <span class="n">std_train_X</span>
|
||||
<span class="n">dataset</span><span class="o">.</span><span class="n">test</span><span class="o">.</span><span class="n">instances</span> <span class="o">=</span> <span class="n">std_test_X</span>
|
||||
<span class="k">return</span> <span class="n">dataset</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="n">training</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">std_train_X</span><span class="p">,</span> <span class="n">train</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="n">train</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="n">test</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">std_test_X</span><span class="p">,</span> <span class="n">test</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="n">test</span><span class="o">.</span><span class="n">classes_</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">Dataset</span><span class="p">(</span><span class="n">training</span><span class="p">,</span> <span class="n">test</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">vocabulary</span><span class="p">,</span> <span class="n">dataset</span><span class="o">.</span><span class="n">name</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="index">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.index">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">index</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span> <span class="n">Dataset</span><span class="p">,</span> <span class="n">min_df</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">index</span><span class="p">(</span><span class="n">dataset</span><span class="p">:</span> <span class="n">Dataset</span><span class="p">,</span> <span class="n">min_df</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">inplace</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Indexes the tokens of a textual :class:`quapy.data.base.Dataset` of string documents.</span>
|
||||
<span class="sd"> To index a document means to replace each different token by a unique numerical index.</span>
|
||||
|
|
@ -219,7 +251,7 @@
|
|||
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">__check_type</span><span class="p">(</span><span class="n">container</span><span class="p">,</span> <span class="n">container_type</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">element_type</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__check_type</span><span class="p">(</span><span class="n">container</span><span class="p">,</span> <span class="n">container_type</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">element_type</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">container_type</span><span class="p">:</span>
|
||||
<span class="k">assert</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">container</span><span class="p">,</span> <span class="n">container_type</span><span class="p">),</span> \
|
||||
<span class="sa">f</span><span class="s1">'unexpected type of container (expected </span><span class="si">{</span><span class="n">container_type</span><span class="si">}</span><span class="s1">, found </span><span class="si">{</span><span class="nb">type</span><span class="p">(</span><span class="n">container</span><span class="p">)</span><span class="si">}</span><span class="s1">)'</span>
|
||||
|
|
@ -230,7 +262,7 @@
|
|||
|
||||
<div class="viewcode-block" id="IndexTransformer">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.IndexTransformer">[docs]</a>
|
||||
<span class="k">class</span> <span class="nc">IndexTransformer</span><span class="p">:</span>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">IndexTransformer</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> This class implements a sklearn's-style transformer that indexes text as numerical ids for the tokens it</span>
|
||||
<span class="sd"> contains, and that would be generated by sklearn's</span>
|
||||
|
|
@ -240,14 +272,14 @@
|
|||
<span class="sd"> `CountVectorizer <https://scikit-learn.org/stable/modules/generated/sklearn.feature_extraction.text.CountVectorizer.html>`_</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">vect</span> <span class="o">=</span> <span class="n">CountVectorizer</span><span class="p">(</span><span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">unk</span> <span class="o">=</span> <span class="o">-</span><span class="mi">1</span> <span class="c1"># a valid index is assigned after fit</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">pad</span> <span class="o">=</span> <span class="o">-</span><span class="mi">2</span> <span class="c1"># a valid index is assigned after fit</span>
|
||||
|
||||
<div class="viewcode-block" id="IndexTransformer.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.IndexTransformer.fit">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fits the transformer, i.e., decides on the vocabulary, given a list of strings.</span>
|
||||
|
||||
|
|
@ -264,7 +296,7 @@
|
|||
|
||||
<div class="viewcode-block" id="IndexTransformer.transform">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.IndexTransformer.transform">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Transforms the strings in `X` as lists of numerical ids</span>
|
||||
|
||||
|
|
@ -279,13 +311,13 @@
|
|||
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">documents</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_index</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">documents</span><span class="p">):</span>
|
||||
<span class="n">vocab</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">vocabulary_</span><span class="o">.</span><span class="n">copy</span><span class="p">()</span>
|
||||
<span class="k">return</span> <span class="p">[[</span><span class="n">vocab</span><span class="o">.</span><span class="n">get</span><span class="p">(</span><span class="n">word</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">unk</span><span class="p">)</span> <span class="k">for</span> <span class="n">word</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">analyzer</span><span class="p">(</span><span class="n">doc</span><span class="p">)]</span> <span class="k">for</span> <span class="n">doc</span> <span class="ow">in</span> <span class="n">tqdm</span><span class="p">(</span><span class="n">documents</span><span class="p">,</span> <span class="s1">'indexing'</span><span class="p">)]</span>
|
||||
|
||||
<div class="viewcode-block" id="IndexTransformer.fit_transform">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.IndexTransformer.fit_transform">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">fit_transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit_transform</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Fits the transform on `X` and transforms it.</span>
|
||||
|
||||
|
|
@ -298,7 +330,7 @@
|
|||
|
||||
<div class="viewcode-block" id="IndexTransformer.vocabulary_size">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.IndexTransformer.vocabulary_size">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">vocabulary_size</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Gets the length of the vocabulary according to which the document tokens have been indexed</span>
|
||||
|
||||
|
|
@ -309,7 +341,7 @@
|
|||
|
||||
<div class="viewcode-block" id="IndexTransformer.add_word">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.preprocessing.IndexTransformer.add_word">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">add_word</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">word</span><span class="p">,</span> <span class="nb">id</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">nogaps</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">add_word</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">word</span><span class="p">,</span> <span class="nb">id</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">nogaps</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Adds a new token (regardless of whether it has been found in the text or not), with dedicated id.</span>
|
||||
<span class="sd"> Useful to define special tokens for codifying unknown words, or padding tokens.</span>
|
||||
|
|
|
|||
|
|
@ -1,83 +1,385 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en" data-content_root="../../../">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.data.reader — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=92fd9be5" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css?v=19f00094" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.data.reader — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script src="../../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../../_static/documentation_options.js?v=22607128"></script>
|
||||
<script src="../../../_static/doctools.js?v=9a2dae69"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=dc90522c"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/data/reader';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.data.reader</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.data.reader</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.data.reader</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">scipy.sparse</span> <span class="kn">import</span> <span class="n">dok_matrix</span>
|
||||
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">logging</span>
|
||||
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">scipy.sparse</span><span class="w"> </span><span class="kn">import</span> <span class="n">dok_matrix</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">tqdm</span><span class="w"> </span><span class="kn">import</span> <span class="n">tqdm</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="from_text">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.reader.from_text">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">from_text</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">encoding</span><span class="o">=</span><span class="s1">'utf-8'</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">class2int</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">from_text</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">encoding</span><span class="o">=</span><span class="s1">'utf-8'</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">class2int</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Reads a labelled colletion of documents.</span>
|
||||
<span class="sd"> File fomart <0 or 1>\t<document>\n</span>
|
||||
|
|
@ -104,14 +406,14 @@
|
|||
<span class="n">all_sentences</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">sentence</span><span class="p">)</span>
|
||||
<span class="n">all_labels</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">label</span><span class="p">)</span>
|
||||
<span class="k">except</span> <span class="ne">ValueError</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="sa">f</span><span class="s1">'format error in </span><span class="si">{</span><span class="n">line</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">warning</span><span class="p">(</span><span class="sa">f</span><span class="s1">'format error in </span><span class="si">{</span><span class="n">line</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">all_sentences</span><span class="p">,</span> <span class="n">all_labels</span></div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="from_sparse">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.reader.from_sparse">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">from_sparse</span><span class="p">(</span><span class="n">path</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">from_sparse</span><span class="p">(</span><span class="n">path</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Reads a labelled collection of real-valued instances expressed in sparse format</span>
|
||||
<span class="sd"> File format <-1 or 0 or 1>[\s col(int):val(float)]\n</span>
|
||||
|
|
@ -120,7 +422,7 @@
|
|||
<span class="sd"> :return: a `csr_matrix` containing the instances (rows), and a ndarray containing the labels</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">split_col_val</span><span class="p">(</span><span class="n">col_val</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">split_col_val</span><span class="p">(</span><span class="n">col_val</span><span class="p">):</span>
|
||||
<span class="n">col</span><span class="p">,</span> <span class="n">val</span> <span class="o">=</span> <span class="n">col_val</span><span class="o">.</span><span class="n">split</span><span class="p">(</span><span class="s1">':'</span><span class="p">)</span>
|
||||
<span class="n">col</span><span class="p">,</span> <span class="n">val</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">col</span><span class="p">)</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="n">val</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">col</span><span class="p">,</span> <span class="n">val</span>
|
||||
|
|
@ -148,7 +450,7 @@
|
|||
|
||||
<div class="viewcode-block" id="from_csv">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.reader.from_csv">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">from_csv</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">encoding</span><span class="o">=</span><span class="s1">'utf-8'</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">from_csv</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">encoding</span><span class="o">=</span><span class="s1">'utf-8'</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Reads a csv file in which columns are separated by ','.</span>
|
||||
<span class="sd"> File format <label>,<feat1>,<feat2>,...,<featn>\n</span>
|
||||
|
|
@ -171,7 +473,7 @@
|
|||
|
||||
<div class="viewcode-block" id="reindex_labels">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.reader.reindex_labels">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">reindex_labels</span><span class="p">(</span><span class="n">y</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">reindex_labels</span><span class="p">(</span><span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Re-indexes a list of labels as a list of indexes, and returns the classnames corresponding to the indexes.</span>
|
||||
<span class="sd"> E.g.:</span>
|
||||
|
|
@ -194,7 +496,7 @@
|
|||
|
||||
<div class="viewcode-block" id="binarize">
|
||||
<a class="viewcode-back" href="../../../quapy.data.html#quapy.data.reader.binarize">[docs]</a>
|
||||
<span class="k">def</span> <span class="nf">binarize</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">pos_class</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">binarize</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">pos_class</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Binarizes a categorical array-like collection of labels towards the positive class `pos_class`. E.g.,:</span>
|
||||
|
||||
|
|
@ -214,31 +516,75 @@
|
|||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,23 +1,20 @@
|
|||
|
||||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<html class="writer-html5" lang="en" data-content_root="../../">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.evaluation — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/css/theme.css" />
|
||||
<title>quapy.evaluation — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/pygments.css?v=b86133f3" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/css/theme.css?v=9edc463e" />
|
||||
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../" id="documentation_options" src="../../_static/documentation_options.js"></script>
|
||||
<script src="../../_static/jquery.js"></script>
|
||||
<script src="../../_static/underscore.js"></script>
|
||||
<script src="../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../_static/doctools.js"></script>
|
||||
<script src="../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../_static/jquery.js?v=5d32c60e"></script>
|
||||
<script src="../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c"></script>
|
||||
<script src="../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../_static/js/theme.js"></script>
|
||||
<link rel="index" title="Index" href="../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../search.html" />
|
||||
|
|
@ -43,7 +40,13 @@
|
|||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../modules.html">quapy</a></li>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../index.html">Quickstart</a></li>
|
||||
</ul>
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../manuals.html">Manuals</a></li>
|
||||
</ul>
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../quapy.html">API</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
|
|
@ -71,16 +74,18 @@
|
|||
<div itemprop="articleBody">
|
||||
|
||||
<h1>Source code for quapy.evaluation</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">from</span> <span class="nn">typing</span> <span class="kn">import</span> <span class="n">Union</span><span class="p">,</span> <span class="n">Callable</span><span class="p">,</span> <span class="n">Iterable</span>
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
|
||||
<span class="kn">import</span> <span class="nn">quapy</span> <span class="k">as</span> <span class="nn">qp</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.protocol</span> <span class="kn">import</span> <span class="n">AbstractProtocol</span><span class="p">,</span> <span class="n">OnLabelledCollectionProtocol</span><span class="p">,</span> <span class="n">IterateProtocol</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.method.base</span> <span class="kn">import</span> <span class="n">BaseQuantifier</span>
|
||||
<span class="kn">import</span> <span class="nn">pandas</span> <span class="k">as</span> <span class="nn">pd</span>
|
||||
<span></span><span class="kn">from</span><span class="w"> </span><span class="nn">typing</span><span class="w"> </span><span class="kn">import</span> <span class="n">Union</span><span class="p">,</span> <span class="n">Callable</span><span class="p">,</span> <span class="n">Iterable</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">tqdm</span><span class="w"> </span><span class="kn">import</span> <span class="n">tqdm</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.protocol</span><span class="w"> </span><span class="kn">import</span> <span class="n">AbstractProtocol</span><span class="p">,</span> <span class="n">OnLabelledCollectionProtocol</span><span class="p">,</span> <span class="n">IterateProtocol</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseQuantifier</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">pandas</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">pd</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="prediction"><a class="viewcode-back" href="../../quapy.html#quapy.evaluation.prediction">[docs]</a><span class="k">def</span> <span class="nf">prediction</span><span class="p">(</span>
|
||||
<div class="viewcode-block" id="prediction">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.evaluation.prediction">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">prediction</span><span class="p">(</span>
|
||||
<span class="n">model</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span>
|
||||
<span class="n">protocol</span><span class="p">:</span> <span class="n">AbstractProtocol</span><span class="p">,</span>
|
||||
<span class="n">aggr_speedup</span><span class="p">:</span> <span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="nb">bool</span><span class="p">]</span> <span class="o">=</span> <span class="s1">'auto'</span><span class="p">,</span>
|
||||
|
|
@ -118,7 +123,7 @@
|
|||
<span class="c1"># checks whether the prediction can be made more efficiently; this check consists in verifying if the model is</span>
|
||||
<span class="c1"># of type aggregative, if the protocol is based on LabelledCollection, and if the total number of documents to</span>
|
||||
<span class="c1"># classify using the protocol would exceed the number of test documents in the original collection</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.method.aggregative</span> <span class="kn">import</span> <span class="n">AggregativeQuantifier</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.aggregative</span><span class="w"> </span><span class="kn">import</span> <span class="n">AggregativeQuantifier</span>
|
||||
<span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">AggregativeQuantifier</span><span class="p">)</span> <span class="ow">and</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">protocol</span><span class="p">,</span> <span class="n">OnLabelledCollectionProtocol</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">aggr_speedup</span> <span class="o">==</span> <span class="s1">'force'</span><span class="p">:</span>
|
||||
<span class="n">apply_optimization</span> <span class="o">=</span> <span class="kc">True</span>
|
||||
|
|
@ -136,10 +141,11 @@
|
|||
<span class="n">protocol_with_predictions</span> <span class="o">=</span> <span class="n">protocol</span><span class="o">.</span><span class="n">on_preclassified_instances</span><span class="p">(</span><span class="n">pre_classified</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">__prediction_helper</span><span class="p">(</span><span class="n">model</span><span class="o">.</span><span class="n">aggregate</span><span class="p">,</span> <span class="n">protocol_with_predictions</span><span class="p">,</span> <span class="n">verbose</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="n">__prediction_helper</span><span class="p">(</span><span class="n">model</span><span class="o">.</span><span class="n">quantify</span><span class="p">,</span> <span class="n">protocol</span><span class="p">,</span> <span class="n">verbose</span><span class="p">)</span></div>
|
||||
<span class="k">return</span> <span class="n">__prediction_helper</span><span class="p">(</span><span class="n">model</span><span class="o">.</span><span class="n">predict</span><span class="p">,</span> <span class="n">protocol</span><span class="p">,</span> <span class="n">verbose</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">__prediction_helper</span><span class="p">(</span><span class="n">quantification_fn</span><span class="p">,</span> <span class="n">protocol</span><span class="p">:</span> <span class="n">AbstractProtocol</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__prediction_helper</span><span class="p">(</span><span class="n">quantification_fn</span><span class="p">,</span> <span class="n">protocol</span><span class="p">:</span> <span class="n">AbstractProtocol</span><span class="p">,</span> <span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="n">true_prevs</span><span class="p">,</span> <span class="n">estim_prevs</span> <span class="o">=</span> <span class="p">[],</span> <span class="p">[]</span>
|
||||
<span class="k">for</span> <span class="n">sample_instances</span><span class="p">,</span> <span class="n">sample_prev</span> <span class="ow">in</span> <span class="n">tqdm</span><span class="p">(</span><span class="n">protocol</span><span class="p">(),</span> <span class="n">total</span><span class="o">=</span><span class="n">protocol</span><span class="o">.</span><span class="n">total</span><span class="p">(),</span> <span class="n">desc</span><span class="o">=</span><span class="s1">'predicting'</span><span class="p">)</span> <span class="k">if</span> <span class="n">verbose</span> <span class="k">else</span> <span class="n">protocol</span><span class="p">():</span>
|
||||
<span class="n">estim_prevs</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">quantification_fn</span><span class="p">(</span><span class="n">sample_instances</span><span class="p">))</span>
|
||||
|
|
@ -151,7 +157,9 @@
|
|||
<span class="k">return</span> <span class="n">true_prevs</span><span class="p">,</span> <span class="n">estim_prevs</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="evaluation_report"><a class="viewcode-back" href="../../quapy.html#quapy.evaluation.evaluation_report">[docs]</a><span class="k">def</span> <span class="nf">evaluation_report</span><span class="p">(</span><span class="n">model</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span>
|
||||
<div class="viewcode-block" id="evaluation_report">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.evaluation.evaluation_report">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">evaluation_report</span><span class="p">(</span><span class="n">model</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span>
|
||||
<span class="n">protocol</span><span class="p">:</span> <span class="n">AbstractProtocol</span><span class="p">,</span>
|
||||
<span class="n">error_metrics</span><span class="p">:</span> <span class="n">Iterable</span><span class="p">[</span><span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span><span class="n">Callable</span><span class="p">]]</span> <span class="o">=</span> <span class="s1">'mae'</span><span class="p">,</span>
|
||||
<span class="n">aggr_speedup</span><span class="p">:</span> <span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="nb">bool</span><span class="p">]</span> <span class="o">=</span> <span class="s1">'auto'</span><span class="p">,</span>
|
||||
|
|
@ -182,7 +190,8 @@
|
|||
<span class="k">return</span> <span class="n">_prevalence_report</span><span class="p">(</span><span class="n">true_prevs</span><span class="p">,</span> <span class="n">estim_prevs</span><span class="p">,</span> <span class="n">error_metrics</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">_prevalence_report</span><span class="p">(</span><span class="n">true_prevs</span><span class="p">,</span> <span class="n">estim_prevs</span><span class="p">,</span> <span class="n">error_metrics</span><span class="p">:</span> <span class="n">Iterable</span><span class="p">[</span><span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="n">Callable</span><span class="p">]]</span> <span class="o">=</span> <span class="s1">'mae'</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_prevalence_report</span><span class="p">(</span><span class="n">true_prevs</span><span class="p">,</span> <span class="n">estim_prevs</span><span class="p">,</span> <span class="n">error_metrics</span><span class="p">:</span> <span class="n">Iterable</span><span class="p">[</span><span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="n">Callable</span><span class="p">]]</span> <span class="o">=</span> <span class="s1">'mae'</span><span class="p">):</span>
|
||||
|
||||
<span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">error_metrics</span><span class="p">,</span> <span class="nb">str</span><span class="p">):</span>
|
||||
<span class="n">error_metrics</span> <span class="o">=</span> <span class="p">[</span><span class="n">error_metrics</span><span class="p">]</span>
|
||||
|
|
@ -203,7 +212,9 @@
|
|||
<span class="k">return</span> <span class="n">df</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="evaluate"><a class="viewcode-back" href="../../quapy.html#quapy.evaluation.evaluate">[docs]</a><span class="k">def</span> <span class="nf">evaluate</span><span class="p">(</span>
|
||||
<div class="viewcode-block" id="evaluate">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.evaluation.evaluate">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">evaluate</span><span class="p">(</span>
|
||||
<span class="n">model</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span>
|
||||
<span class="n">protocol</span><span class="p">:</span> <span class="n">AbstractProtocol</span><span class="p">,</span>
|
||||
<span class="n">error_metric</span><span class="p">:</span> <span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="n">Callable</span><span class="p">],</span>
|
||||
|
|
@ -235,7 +246,10 @@
|
|||
<span class="k">return</span> <span class="n">error_metric</span><span class="p">(</span><span class="n">true_prevs</span><span class="p">,</span> <span class="n">estim_prevs</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="evaluate_on_samples"><a class="viewcode-back" href="../../quapy.html#quapy.evaluation.evaluate_on_samples">[docs]</a><span class="k">def</span> <span class="nf">evaluate_on_samples</span><span class="p">(</span>
|
||||
|
||||
<div class="viewcode-block" id="evaluate_on_samples">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.evaluation.evaluate_on_samples">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">evaluate_on_samples</span><span class="p">(</span>
|
||||
<span class="n">model</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span>
|
||||
<span class="n">samples</span><span class="p">:</span> <span class="n">Iterable</span><span class="p">[</span><span class="n">qp</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">LabelledCollection</span><span class="p">],</span>
|
||||
<span class="n">error_metric</span><span class="p">:</span> <span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="n">Callable</span><span class="p">],</span>
|
||||
|
|
@ -259,6 +273,7 @@
|
|||
|
||||
|
||||
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -1,91 +1,393 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method._neural — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method._neural — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../../" id="documentation_options" src="../../../_static/documentation_options.js"></script>
|
||||
<script src="../../../_static/jquery.js"></script>
|
||||
<script src="../../../_static/underscore.js"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../../_static/doctools.js"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=a95f357e85573c9b56d5" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=a95f357e85573c9b56d5" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=a95f357e85573c9b56d5"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=a95f357e85573c9b56d5" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=a95f357e85573c9b56d5" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/method/_neural';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.method._neural</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.method._neural</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.method._neural</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">os</span>
|
||||
<span class="kn">from</span> <span class="nn">pathlib</span> <span class="kn">import</span> <span class="n">Path</span>
|
||||
<span class="kn">import</span> <span class="nn">random</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">logging</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">os</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">pathlib</span><span class="w"> </span><span class="kn">import</span> <span class="n">Path</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">random</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">torch</span>
|
||||
<span class="kn">from</span> <span class="nn">torch.nn</span> <span class="kn">import</span> <span class="n">MSELoss</span>
|
||||
<span class="kn">from</span> <span class="nn">torch.nn.functional</span> <span class="kn">import</span> <span class="n">relu</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">torch</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">torch.nn</span><span class="w"> </span><span class="kn">import</span> <span class="n">MSELoss</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">torch.nn.functional</span><span class="w"> </span><span class="kn">import</span> <span class="n">relu</span>
|
||||
|
||||
<span class="kn">from</span> <span class="nn">quapy.protocol</span> <span class="kn">import</span> <span class="n">UPP</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.method.aggregative</span> <span class="kn">import</span> <span class="o">*</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.util</span> <span class="kn">import</span> <span class="n">EarlyStop</span>
|
||||
<span class="kn">from</span> <span class="nn">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.protocol</span><span class="w"> </span><span class="kn">import</span> <span class="n">UPP</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.aggregative</span><span class="w"> </span><span class="kn">import</span> <span class="o">*</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.util</span><span class="w"> </span><span class="kn">import</span> <span class="n">EarlyStop</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">tqdm</span><span class="w"> </span><span class="kn">import</span> <span class="n">tqdm</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer">[docs]</a><span class="k">class</span> <span class="nc">QuaNetTrainer</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="QuaNetTrainer">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">QuaNetTrainer</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Implementation of `QuaNet <https://dl.acm.org/doi/abs/10.1145/3269206.3269287>`_, a neural network for</span>
|
||||
<span class="sd"> quantification. This implementation uses `PyTorch <https://pytorch.org/>`_ and can take advantage of GPU</span>
|
||||
|
|
@ -100,7 +402,7 @@
|
|||
<span class="sd"> >>> # use samples of 100 elements</span>
|
||||
<span class="sd"> >>> qp.environ['SAMPLE_SIZE'] = 100</span>
|
||||
<span class="sd"> >>></span>
|
||||
<span class="sd"> >>> # load the kindle dataset as text, and convert words to numerical indexes</span>
|
||||
<span class="sd"> >>> # load the Kindle dataset as text, and convert words to numerical indexes</span>
|
||||
<span class="sd"> >>> dataset = qp.datasets.fetch_reviews('kindle', pickle=True)</span>
|
||||
<span class="sd"> >>> qp.train.preprocessing.index(dataset, min_df=5, inplace=True)</span>
|
||||
<span class="sd"> >>></span>
|
||||
|
|
@ -110,12 +412,14 @@
|
|||
<span class="sd"> >>></span>
|
||||
<span class="sd"> >>> # train QuaNet (QuaNet is an alias to QuaNetTrainer)</span>
|
||||
<span class="sd"> >>> model = QuaNet(classifier, qp.environ['SAMPLE_SIZE'], device='cuda')</span>
|
||||
<span class="sd"> >>> model.fit(dataset.training)</span>
|
||||
<span class="sd"> >>> estim_prevalence = model.quantify(dataset.test.instances)</span>
|
||||
<span class="sd"> >>> model.fit(*dataset.training.Xy)</span>
|
||||
<span class="sd"> >>> estim_prevalence = model.predict(dataset.test.instances)</span>
|
||||
|
||||
<span class="sd"> :param classifier: an object implementing `fit` (i.e., that can be trained on labelled data),</span>
|
||||
<span class="sd"> `predict_proba` (i.e., that can generate posterior probabilities of unlabelled examples) and</span>
|
||||
<span class="sd"> `transform` (i.e., that can generate embedded representations of the unlabelled instances).</span>
|
||||
<span class="sd"> :param fit_classifier: whether to train the learner (default is True). Set to False if the</span>
|
||||
<span class="sd"> learner has been trained outside the quantifier.</span>
|
||||
<span class="sd"> :param sample_size: integer, the sample size; default is None, meaning that the sample size should be</span>
|
||||
<span class="sd"> taken from qp.environ["SAMPLE_SIZE"]</span>
|
||||
<span class="sd"> :param n_epochs: integer, maximum number of training epochs</span>
|
||||
|
|
@ -135,8 +439,9 @@
|
|||
<span class="sd"> :param device: string, indicate "cpu" or "cuda"</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="n">classifier</span><span class="p">,</span>
|
||||
<span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span>
|
||||
<span class="n">sample_size</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span>
|
||||
<span class="n">n_epochs</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span>
|
||||
<span class="n">tr_iter_per_poch</span><span class="o">=</span><span class="mi">500</span><span class="p">,</span>
|
||||
|
|
@ -159,6 +464,7 @@
|
|||
<span class="sa">f</span><span class="s1">'the classifier </span><span class="si">{</span><span class="n">classifier</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1"> does not seem to be able to produce posterior probabilities '</span> \
|
||||
<span class="sa">f</span><span class="s1">'since it does not implement the method "predict_proba"'</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">fit_classifier</span> <span class="o">=</span> <span class="n">fit_classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">sample_size</span> <span class="o">=</span> <span class="n">qp</span><span class="o">.</span><span class="n">_get_sample_size</span><span class="p">(</span><span class="n">sample_size</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_epochs</span> <span class="o">=</span> <span class="n">n_epochs</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">tr_iter</span> <span class="o">=</span> <span class="n">tr_iter_per_poch</span>
|
||||
|
|
@ -184,20 +490,23 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">__check_params_colision</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">quanet_params</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">get_params</span><span class="p">())</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_classes_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.fit"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.fit">[docs]</a> <span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="QuaNetTrainer.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Trains QuaNet.</span>
|
||||
|
||||
<span class="sd"> :param data: the training data on which to train QuaNet. If `fit_classifier=True`, the data will be split in</span>
|
||||
<span class="sd"> :param X: the training instances on which to train QuaNet. If `fit_classifier=True`, the data will be split in</span>
|
||||
<span class="sd"> 40/40/20 for training the classifier, training QuaNet, and validating QuaNet, respectively. If</span>
|
||||
<span class="sd"> `fit_classifier=False`, the data will be split in 66/34 for training QuaNet and validating it, respectively.</span>
|
||||
<span class="sd"> :param fit_classifier: if True, trains the classifier on a split containing 40% of the data</span>
|
||||
<span class="sd"> :param y: the labels of X</span>
|
||||
<span class="sd"> :return: self</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">data</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_classes_</span> <span class="o">=</span> <span class="n">data</span><span class="o">.</span><span class="n">classes_</span>
|
||||
<span class="n">os</span><span class="o">.</span><span class="n">makedirs</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">checkpointdir</span><span class="p">,</span> <span class="n">exist_ok</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
|
||||
<span class="k">if</span> <span class="n">fit_classifier</span><span class="p">:</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">fit_classifier</span><span class="p">:</span>
|
||||
<span class="n">classifier_data</span><span class="p">,</span> <span class="n">unused_data</span> <span class="o">=</span> <span class="n">data</span><span class="o">.</span><span class="n">split_stratified</span><span class="p">(</span><span class="mf">0.4</span><span class="p">)</span>
|
||||
<span class="n">train_data</span><span class="p">,</span> <span class="n">valid_data</span> <span class="o">=</span> <span class="n">unused_data</span><span class="o">.</span><span class="n">split_stratified</span><span class="p">(</span><span class="mf">0.66</span><span class="p">)</span> <span class="c1"># 0.66 split of 60% makes 40% and 20%</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">classifier_data</span><span class="o">.</span><span class="n">Xy</span><span class="p">)</span>
|
||||
|
|
@ -217,13 +526,13 @@
|
|||
<span class="n">train_data_embed</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">train_data</span><span class="o">.</span><span class="n">instances</span><span class="p">),</span> <span class="n">train_data</span><span class="o">.</span><span class="n">labels</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">_classes_</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">quantifiers</span> <span class="o">=</span> <span class="p">{</span>
|
||||
<span class="s1">'cc'</span><span class="p">:</span> <span class="n">CC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">),</span>
|
||||
<span class="s1">'acc'</span><span class="p">:</span> <span class="n">ACC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="n">valid_data</span><span class="p">),</span>
|
||||
<span class="s1">'pcc'</span><span class="p">:</span> <span class="n">PCC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">),</span>
|
||||
<span class="s1">'pacc'</span><span class="p">:</span> <span class="n">PACC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="n">valid_data</span><span class="p">),</span>
|
||||
<span class="s1">'cc'</span><span class="p">:</span> <span class="n">CC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">valid_data</span><span class="o">.</span><span class="n">Xy</span><span class="p">),</span>
|
||||
<span class="s1">'acc'</span><span class="p">:</span> <span class="n">ACC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">valid_data</span><span class="o">.</span><span class="n">Xy</span><span class="p">),</span>
|
||||
<span class="s1">'pcc'</span><span class="p">:</span> <span class="n">PCC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">valid_data</span><span class="o">.</span><span class="n">Xy</span><span class="p">),</span>
|
||||
<span class="s1">'pacc'</span><span class="p">:</span> <span class="n">PACC</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">valid_data</span><span class="o">.</span><span class="n">Xy</span><span class="p">),</span>
|
||||
<span class="p">}</span>
|
||||
<span class="k">if</span> <span class="n">classifier_data</span> <span class="ow">is</span> <span class="ow">not</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">quantifiers</span><span class="p">[</span><span class="s1">'emq'</span><span class="p">]</span> <span class="o">=</span> <span class="n">EMQ</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">classifier_data</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">quantifiers</span><span class="p">[</span><span class="s1">'emq'</span><span class="p">]</span> <span class="o">=</span> <span class="n">EMQ</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">False</span><span class="p">)</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">valid_data</span><span class="o">.</span><span class="n">Xy</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">status</span> <span class="o">=</span> <span class="p">{</span>
|
||||
<span class="s1">'tr-loss'</span><span class="p">:</span> <span class="o">-</span><span class="mi">1</span><span class="p">,</span>
|
||||
|
|
@ -241,7 +550,7 @@
|
|||
<span class="n">order_by</span><span class="o">=</span><span class="mi">0</span> <span class="k">if</span> <span class="n">data</span><span class="o">.</span><span class="n">binary</span> <span class="k">else</span> <span class="kc">None</span><span class="p">,</span>
|
||||
<span class="o">**</span><span class="bp">self</span><span class="o">.</span><span class="n">quanet_params</span>
|
||||
<span class="p">)</span><span class="o">.</span><span class="n">to</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">device</span><span class="p">)</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">quanet</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">debug</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">quanet</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">optim</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">optim</span><span class="o">.</span><span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">quanet</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">lr</span><span class="p">)</span>
|
||||
<span class="n">early_stop</span> <span class="o">=</span> <span class="n">EarlyStop</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">patience</span><span class="p">,</span> <span class="n">lower_is_better</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
|
|
@ -256,14 +565,16 @@
|
|||
<span class="k">if</span> <span class="n">early_stop</span><span class="o">.</span><span class="n">IMPROVED</span><span class="p">:</span>
|
||||
<span class="n">torch</span><span class="o">.</span><span class="n">save</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">quanet</span><span class="o">.</span><span class="n">state_dict</span><span class="p">(),</span> <span class="n">checkpoint</span><span class="p">)</span>
|
||||
<span class="k">elif</span> <span class="n">early_stop</span><span class="o">.</span><span class="n">STOP</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="sa">f</span><span class="s1">'training ended by patience exhausted; loading best model parameters in </span><span class="si">{</span><span class="n">checkpoint</span><span class="si">}</span><span class="s1"> '</span>
|
||||
<span class="sa">f</span><span class="s1">'for epoch </span><span class="si">{</span><span class="n">early_stop</span><span class="o">.</span><span class="n">best_epoch</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span>
|
||||
<span class="sa">f</span><span class="s1">'training ended by patience exhausted; loading best model parameters in </span><span class="si">{</span><span class="n">checkpoint</span><span class="si">}</span><span class="s1"> '</span>
|
||||
<span class="sa">f</span><span class="s1">'for epoch </span><span class="si">{</span><span class="n">early_stop</span><span class="o">.</span><span class="n">best_epoch</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">quanet</span><span class="o">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">load</span><span class="p">(</span><span class="n">checkpoint</span><span class="p">))</span>
|
||||
<span class="k">break</span>
|
||||
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_get_aggregative_estims</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">posteriors</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_get_aggregative_estims</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">posteriors</span><span class="p">):</span>
|
||||
<span class="n">label_predictions</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">posteriors</span><span class="p">,</span> <span class="n">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
|
||||
<span class="n">prevs_estim</span> <span class="o">=</span> <span class="p">[]</span>
|
||||
<span class="k">for</span> <span class="n">quantifier</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">quantifiers</span><span class="o">.</span><span class="n">values</span><span class="p">():</span>
|
||||
|
|
@ -274,9 +585,11 @@
|
|||
|
||||
<span class="k">return</span> <span class="n">prevs_estim</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.quantify"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.quantify">[docs]</a> <span class="k">def</span> <span class="nf">quantify</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="n">posteriors</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">predict_proba</span><span class="p">(</span><span class="n">instances</span><span class="p">)</span>
|
||||
<span class="n">embeddings</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">instances</span><span class="p">)</span>
|
||||
<div class="viewcode-block" id="QuaNetTrainer.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="n">posteriors</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">predict_proba</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">embeddings</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">transform</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">quant_estims</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_get_aggregative_estims</span><span class="p">(</span><span class="n">posteriors</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">quanet</span><span class="o">.</span><span class="n">eval</span><span class="p">()</span>
|
||||
<span class="k">with</span> <span class="n">torch</span><span class="o">.</span><span class="n">no_grad</span><span class="p">():</span>
|
||||
|
|
@ -286,7 +599,8 @@
|
|||
<span class="n">prevalence</span> <span class="o">=</span> <span class="n">prevalence</span><span class="o">.</span><span class="n">numpy</span><span class="p">()</span><span class="o">.</span><span class="n">flatten</span><span class="p">()</span>
|
||||
<span class="k">return</span> <span class="n">prevalence</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_epoch</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">posteriors</span><span class="p">,</span> <span class="n">iterations</span><span class="p">,</span> <span class="n">epoch</span><span class="p">,</span> <span class="n">early_stop</span><span class="p">,</span> <span class="n">train</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_epoch</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">posteriors</span><span class="p">,</span> <span class="n">iterations</span><span class="p">,</span> <span class="n">epoch</span><span class="p">,</span> <span class="n">early_stop</span><span class="p">,</span> <span class="n">train</span><span class="p">):</span>
|
||||
<span class="n">mse_loss</span> <span class="o">=</span> <span class="n">MSELoss</span><span class="p">()</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">quanet</span><span class="o">.</span><span class="n">train</span><span class="p">(</span><span class="n">mode</span><span class="o">=</span><span class="n">train</span><span class="p">)</span>
|
||||
|
|
@ -336,12 +650,17 @@
|
|||
<span class="sa">f</span><span class="s1">'val-mseloss=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">status</span><span class="p">[</span><span class="s2">"va-loss"</span><span class="p">]</span><span class="si">:</span><span class="s1">.5f</span><span class="si">}</span><span class="s1"> val-maeloss=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">status</span><span class="p">[</span><span class="s2">"va-mae"</span><span class="p">]</span><span class="si">:</span><span class="s1">.5f</span><span class="si">}</span><span class="s1"> '</span>
|
||||
<span class="sa">f</span><span class="s1">'patience=</span><span class="si">{</span><span class="n">early_stop</span><span class="o">.</span><span class="n">patience</span><span class="si">}</span><span class="s1">/</span><span class="si">{</span><span class="n">early_stop</span><span class="o">.</span><span class="n">PATIENCE_LIMIT</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.get_params"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.get_params">[docs]</a> <span class="k">def</span> <span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">deep</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="QuaNetTrainer.get_params">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.get_params">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">deep</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="n">classifier_params</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">get_params</span><span class="p">()</span>
|
||||
<span class="n">classifier_params</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'classifier__'</span><span class="o">+</span><span class="n">k</span><span class="p">:</span><span class="n">v</span> <span class="k">for</span> <span class="n">k</span><span class="p">,</span><span class="n">v</span> <span class="ow">in</span> <span class="n">classifier_params</span><span class="o">.</span><span class="n">items</span><span class="p">()}</span>
|
||||
<span class="k">return</span> <span class="p">{</span><span class="o">**</span><span class="n">classifier_params</span><span class="p">,</span> <span class="o">**</span><span class="bp">self</span><span class="o">.</span><span class="n">quanet_params</span><span class="p">}</span></div>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.set_params"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.set_params">[docs]</a> <span class="k">def</span> <span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">parameters</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.set_params">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.set_params">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">parameters</span><span class="p">):</span>
|
||||
<span class="n">learner_params</span> <span class="o">=</span> <span class="p">{}</span>
|
||||
<span class="k">for</span> <span class="n">key</span><span class="p">,</span> <span class="n">val</span> <span class="ow">in</span> <span class="n">parameters</span><span class="o">.</span><span class="n">items</span><span class="p">():</span>
|
||||
<span class="k">if</span> <span class="n">key</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">quanet_params</span><span class="p">:</span>
|
||||
|
|
@ -352,7 +671,8 @@
|
|||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="s1">'unknown parameter '</span><span class="p">,</span> <span class="n">key</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="n">set_params</span><span class="p">(</span><span class="o">**</span><span class="n">learner_params</span><span class="p">)</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">__check_params_colision</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">quanet_params</span><span class="p">,</span> <span class="n">learner_params</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__check_params_colision</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">quanet_params</span><span class="p">,</span> <span class="n">learner_params</span><span class="p">):</span>
|
||||
<span class="n">quanet_keys</span> <span class="o">=</span> <span class="nb">set</span><span class="p">(</span><span class="n">quanet_params</span><span class="o">.</span><span class="n">keys</span><span class="p">())</span>
|
||||
<span class="n">learner_keys</span> <span class="o">=</span> <span class="nb">set</span><span class="p">(</span><span class="n">learner_params</span><span class="o">.</span><span class="n">keys</span><span class="p">())</span>
|
||||
<span class="n">intersection</span> <span class="o">=</span> <span class="n">quanet_keys</span><span class="o">.</span><span class="n">intersection</span><span class="p">(</span><span class="n">learner_keys</span><span class="p">)</span>
|
||||
|
|
@ -360,25 +680,34 @@
|
|||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="sa">f</span><span class="s1">'the use of parameters </span><span class="si">{</span><span class="n">intersection</span><span class="si">}</span><span class="s1"> is ambiguous sine those can refer to '</span>
|
||||
<span class="sa">f</span><span class="s1">'the parameters of QuaNet or the learner </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">classifier</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.clean_checkpoint"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.clean_checkpoint">[docs]</a> <span class="k">def</span> <span class="nf">clean_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="QuaNetTrainer.clean_checkpoint">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.clean_checkpoint">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">clean_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Removes the checkpoint</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">os</span><span class="o">.</span><span class="n">remove</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">checkpoint</span><span class="p">)</span></div>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.clean_checkpoint_dir"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.clean_checkpoint_dir">[docs]</a> <span class="k">def</span> <span class="nf">clean_checkpoint_dir</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetTrainer.clean_checkpoint_dir">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetTrainer.clean_checkpoint_dir">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">clean_checkpoint_dir</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Removes anything contained in the checkpoint directory</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="kn">import</span> <span class="nn">shutil</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">shutil</span>
|
||||
<span class="n">shutil</span><span class="o">.</span><span class="n">rmtree</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">checkpointdir</span><span class="p">,</span> <span class="n">ignore_errors</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">classes_</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">classes_</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">_classes_</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="mae_loss"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.mae_loss">[docs]</a><span class="k">def</span> <span class="nf">mae_loss</span><span class="p">(</span><span class="n">output</span><span class="p">,</span> <span class="n">target</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="mae_loss">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.mae_loss">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">mae_loss</span><span class="p">(</span><span class="n">output</span><span class="p">,</span> <span class="n">target</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Torch-like wrapper for the Mean Absolute Error</span>
|
||||
|
||||
|
|
@ -389,7 +718,10 @@
|
|||
<span class="k">return</span> <span class="n">torch</span><span class="o">.</span><span class="n">mean</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">abs</span><span class="p">(</span><span class="n">output</span> <span class="o">-</span> <span class="n">target</span><span class="p">))</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="QuaNetModule"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetModule">[docs]</a><span class="k">class</span> <span class="nc">QuaNetModule</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Module</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetModule">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetModule">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">QuaNetModule</span><span class="p">(</span><span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Module</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Implements the `QuaNet <https://dl.acm.org/doi/abs/10.1145/3269206.3269287>`_ forward pass.</span>
|
||||
<span class="sd"> See :class:`QuaNetTrainer` for training QuaNet.</span>
|
||||
|
|
@ -406,7 +738,7 @@
|
|||
<span class="sd"> :param order_by: integer, class for which the document embeddings are to be sorted</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="n">doc_embedding_size</span><span class="p">,</span>
|
||||
<span class="n">n_classes</span><span class="p">,</span>
|
||||
<span class="n">stats_size</span><span class="p">,</span>
|
||||
|
|
@ -441,10 +773,10 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">output</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Linear</span><span class="p">(</span><span class="n">prev_size</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">)</span>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">device</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">device</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="n">torch</span><span class="o">.</span><span class="n">device</span><span class="p">(</span><span class="s1">'cuda'</span><span class="p">)</span> <span class="k">if</span> <span class="nb">next</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">parameters</span><span class="p">())</span><span class="o">.</span><span class="n">is_cuda</span> <span class="k">else</span> <span class="n">torch</span><span class="o">.</span><span class="n">device</span><span class="p">(</span><span class="s1">'cpu'</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_init_hidden</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_init_hidden</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="n">directions</span> <span class="o">=</span> <span class="mi">2</span> <span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">bidirectional</span> <span class="k">else</span> <span class="mi">1</span>
|
||||
<span class="n">var_hidden</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">nlayers</span> <span class="o">*</span> <span class="n">directions</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">hidden_size</span><span class="p">)</span>
|
||||
<span class="n">var_cell</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">nlayers</span> <span class="o">*</span> <span class="n">directions</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">hidden_size</span><span class="p">)</span>
|
||||
|
|
@ -452,7 +784,9 @@
|
|||
<span class="n">var_hidden</span><span class="p">,</span> <span class="n">var_cell</span> <span class="o">=</span> <span class="n">var_hidden</span><span class="o">.</span><span class="n">cuda</span><span class="p">(),</span> <span class="n">var_cell</span><span class="o">.</span><span class="n">cuda</span><span class="p">()</span>
|
||||
<span class="k">return</span> <span class="n">var_hidden</span><span class="p">,</span> <span class="n">var_cell</span>
|
||||
|
||||
<div class="viewcode-block" id="QuaNetModule.forward"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetModule.forward">[docs]</a> <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">doc_embeddings</span><span class="p">,</span> <span class="n">doc_posteriors</span><span class="p">,</span> <span class="n">statistics</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="QuaNetModule.forward">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._neural.QuaNetModule.forward">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">doc_embeddings</span><span class="p">,</span> <span class="n">doc_posteriors</span><span class="p">,</span> <span class="n">statistics</span><span class="p">):</span>
|
||||
<span class="n">device</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">device</span>
|
||||
<span class="n">doc_embeddings</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">as_tensor</span><span class="p">(</span><span class="n">doc_embeddings</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="o">.</span><span class="n">float</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>
|
||||
<span class="n">doc_posteriors</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">as_tensor</span><span class="p">(</span><span class="n">doc_posteriors</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="o">.</span><span class="n">float</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>
|
||||
|
|
@ -482,7 +816,9 @@
|
|||
<span class="n">logits</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">output</span><span class="p">(</span><span class="n">abstracted</span><span class="p">)</span><span class="o">.</span><span class="n">view</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">)</span>
|
||||
<span class="n">prevalence</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">logits</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">)</span>
|
||||
|
||||
<span class="k">return</span> <span class="n">prevalence</span></div></div>
|
||||
<span class="k">return</span> <span class="n">prevalence</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -490,31 +826,75 @@
|
|||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=a95f357e85573c9b56d5"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=a95f357e85573c9b56d5"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.19.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,87 +1,388 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method._threshold_optim — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method._threshold_optim — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../../" id="documentation_options" src="../../../_static/documentation_options.js"></script>
|
||||
<script src="../../../_static/jquery.js"></script>
|
||||
<script src="../../../_static/underscore.js"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../../_static/doctools.js"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/method/_threshold_optim';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.method._threshold_optim</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.method._threshold_optim</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.method._threshold_optim</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">from</span> <span class="nn">abc</span> <span class="kn">import</span> <span class="n">abstractmethod</span>
|
||||
<span></span><span class="kn">from</span><span class="w"> </span><span class="nn">abc</span><span class="w"> </span><span class="kn">import</span> <span class="n">abstractmethod</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.base</span> <span class="kn">import</span> <span class="n">BaseEstimator</span>
|
||||
<span class="kn">import</span> <span class="nn">quapy</span> <span class="k">as</span> <span class="nn">qp</span>
|
||||
<span class="kn">import</span> <span class="nn">quapy.functional</span> <span class="k">as</span> <span class="nn">F</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.data</span> <span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.method.aggregative</span> <span class="kn">import</span> <span class="n">BinaryAggregativeQuantifier</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseEstimator</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy.functional</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">F</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.data</span><span class="w"> </span><span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.aggregative</span><span class="w"> </span><span class="kn">import</span> <span class="n">BinaryAggregativeQuantifier</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization">[docs]</a><span class="k">class</span> <span class="nc">ThresholdOptimization</span><span class="p">(</span><span class="n">BinaryAggregativeQuantifier</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="ThresholdOptimization">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">ThresholdOptimization</span><span class="p">(</span><span class="n">BinaryAggregativeQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Abstract class of Threshold Optimization variants for :class:`ACC` as proposed by</span>
|
||||
<span class="sd"> `Forman 2006 <https://dl.acm.org/doi/abs/10.1145/1150402.1150423>`_ and</span>
|
||||
|
|
@ -91,22 +392,29 @@
|
|||
<span class="sd"> that would allow for more true positives and many more false positives, on the grounds this</span>
|
||||
<span class="sd"> would deliver larger denominators.</span>
|
||||
|
||||
<span class="sd"> :param classifier: a sklearn's Estimator that generates a classifier</span>
|
||||
<span class="sd"> :param val_split: indicates the proportion of data to be used as a stratified held-out validation set in which the</span>
|
||||
<span class="sd"> misclassification rates are to be estimated.</span>
|
||||
<span class="sd"> This parameter can be indicated as a real value (between 0 and 1), representing a proportion of</span>
|
||||
<span class="sd"> validation data, or as an integer, indicating that the misclassification rates should be estimated via</span>
|
||||
<span class="sd"> `k`-fold cross validation (this integer stands for the number of folds `k`, defaults 5), or as a</span>
|
||||
<span class="sd"> :class:`quapy.data.base.LabelledCollection` (the split itself).</span>
|
||||
<span class="sd"> :param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be</span>
|
||||
<span class="sd"> the one indicated in `qp.environ['DEFAULT_CLS']`</span>
|
||||
|
||||
<span class="sd"> :param fit_classifier: whether to train the learner (default is True). Set to False if the</span>
|
||||
<span class="sd"> learner has been trained outside the quantifier.</span>
|
||||
|
||||
<span class="sd"> :param val_split: specifies the data used for generating classifier predictions. This specification</span>
|
||||
<span class="sd"> can be made as float in (0, 1) indicating the proportion of stratified held-out validation set to</span>
|
||||
<span class="sd"> be extracted from the training set; or as an integer (default 5), indicating that the predictions</span>
|
||||
<span class="sd"> are to be generated in a `k`-fold cross-validation manner (with this integer indicating the value</span>
|
||||
<span class="sd"> for `k`); or as a tuple (X,y) defining the specific set of data to use for validation.</span>
|
||||
|
||||
<span class="sd"> :param n_jobs: number of parallel workers</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classifier</span> <span class="o">=</span> <span class="n">classifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">val_split</span> <span class="o">=</span> <span class="n">val_split</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span> <span class="o">=</span> <span class="n">qp</span><span class="o">.</span><span class="n">_get_njobs</span><span class="p">(</span><span class="n">n_jobs</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization.condition"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.condition">[docs]</a> <span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span> <span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<div class="viewcode-block" id="ThresholdOptimization.condition">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.condition">[docs]</a>
|
||||
<span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Implements the criterion according to which the threshold should be selected.</span>
|
||||
<span class="sd"> This function should return the (float) score to be minimized.</span>
|
||||
|
|
@ -117,7 +425,10 @@
|
|||
<span class="sd"> """</span>
|
||||
<span class="o">...</span></div>
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization.discard"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.discard">[docs]</a> <span class="k">def</span> <span class="nf">discard</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">bool</span><span class="p">:</span>
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization.discard">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.discard">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">discard</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">bool</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Indicates whether a combination of tpr and fpr should be discarded</span>
|
||||
|
||||
|
|
@ -128,7 +439,8 @@
|
|||
<span class="k">return</span> <span class="p">(</span><span class="n">tpr</span> <span class="o">-</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">==</span> <span class="mi">0</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">_eval_candidate_thresholds</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_eval_candidate_thresholds</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Seeks for the best `tpr` and `fpr` according to the score obtained at different</span>
|
||||
<span class="sd"> decision thresholds. The scoring function is implemented in function `_condition`.</span>
|
||||
|
|
@ -163,7 +475,9 @@
|
|||
|
||||
<span class="k">return</span> <span class="n">candidates</span>
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization.aggregate_with_threshold"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.aggregate_with_threshold">[docs]</a> <span class="k">def</span> <span class="nf">aggregate_with_threshold</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">,</span> <span class="n">tprs</span><span class="p">,</span> <span class="n">fprs</span><span class="p">,</span> <span class="n">thresholds</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="ThresholdOptimization.aggregate_with_threshold">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.aggregate_with_threshold">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">aggregate_with_threshold</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">,</span> <span class="n">tprs</span><span class="p">,</span> <span class="n">fprs</span><span class="p">,</span> <span class="n">thresholds</span><span class="p">):</span>
|
||||
<span class="c1"># This function performs the adjusted count for given tpr, fpr, and threshold.</span>
|
||||
<span class="c1"># Note that, due to broadcasting, tprs, fprs, and thresholds could be arrays of length > 1</span>
|
||||
<span class="n">prevs_estims</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">mean</span><span class="p">(</span><span class="n">classif_predictions</span><span class="p">[:,</span> <span class="kc">None</span><span class="p">]</span> <span class="o">>=</span> <span class="n">thresholds</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
|
||||
|
|
@ -171,35 +485,45 @@
|
|||
<span class="n">prevs_estims</span> <span class="o">=</span> <span class="n">F</span><span class="o">.</span><span class="n">as_binary_prevalence</span><span class="p">(</span><span class="n">prevs_estims</span><span class="p">,</span> <span class="n">clip_if_necessary</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">prevs_estims</span><span class="o">.</span><span class="n">squeeze</span><span class="p">()</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_compute_table</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">y_</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_compute_table</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">y_</span><span class="p">):</span>
|
||||
<span class="n">TP</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">logical_and</span><span class="p">(</span><span class="n">y</span> <span class="o">==</span> <span class="n">y_</span><span class="p">,</span> <span class="n">y</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">pos_label</span><span class="p">)</span><span class="o">.</span><span class="n">sum</span><span class="p">()</span>
|
||||
<span class="n">FP</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">logical_and</span><span class="p">(</span><span class="n">y</span> <span class="o">!=</span> <span class="n">y_</span><span class="p">,</span> <span class="n">y</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">neg_label</span><span class="p">)</span><span class="o">.</span><span class="n">sum</span><span class="p">()</span>
|
||||
<span class="n">FN</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">logical_and</span><span class="p">(</span><span class="n">y</span> <span class="o">!=</span> <span class="n">y_</span><span class="p">,</span> <span class="n">y</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">pos_label</span><span class="p">)</span><span class="o">.</span><span class="n">sum</span><span class="p">()</span>
|
||||
<span class="n">TN</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">logical_and</span><span class="p">(</span><span class="n">y</span> <span class="o">==</span> <span class="n">y_</span><span class="p">,</span> <span class="n">y</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">neg_label</span><span class="p">)</span><span class="o">.</span><span class="n">sum</span><span class="p">()</span>
|
||||
<span class="k">return</span> <span class="n">TP</span><span class="p">,</span> <span class="n">FP</span><span class="p">,</span> <span class="n">FN</span><span class="p">,</span> <span class="n">TN</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_compute_tpr</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">TP</span><span class="p">,</span> <span class="n">FP</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">TP</span> <span class="o">+</span> <span class="n">FP</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_compute_tpr</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">TP</span><span class="p">,</span> <span class="n">FN</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">TP</span> <span class="o">+</span> <span class="n">FN</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="mi">1</span>
|
||||
<span class="k">return</span> <span class="n">TP</span> <span class="o">/</span> <span class="p">(</span><span class="n">TP</span> <span class="o">+</span> <span class="n">FP</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">TP</span> <span class="o">/</span> <span class="p">(</span><span class="n">TP</span> <span class="o">+</span> <span class="n">FN</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_compute_fpr</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">FP</span><span class="p">,</span> <span class="n">TN</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_compute_fpr</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">FP</span><span class="p">,</span> <span class="n">TN</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">FP</span> <span class="o">+</span> <span class="n">TN</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="mi">0</span>
|
||||
<span class="k">return</span> <span class="n">FP</span> <span class="o">/</span> <span class="p">(</span><span class="n">FP</span> <span class="o">+</span> <span class="n">TN</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization.aggregation_fit"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.aggregation_fit">[docs]</a> <span class="k">def</span> <span class="nf">aggregation_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">):</span>
|
||||
<span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">classif_predictions</span><span class="o">.</span><span class="n">Xy</span>
|
||||
<div class="viewcode-block" id="ThresholdOptimization.aggregation_fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.aggregation_fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">aggregation_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">,</span> <span class="n">labels</span><span class="p">):</span>
|
||||
<span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">classif_predictions</span><span class="p">,</span> <span class="n">labels</span>
|
||||
<span class="c1"># the standard behavior is to keep the best threshold only</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">tpr</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">fpr</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">threshold</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_eval_candidate_thresholds</span><span class="p">(</span><span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization.aggregate"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.aggregate">[docs]</a> <span class="k">def</span> <span class="nf">aggregate</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">:</span> <span class="n">np</span><span class="o">.</span><span class="n">ndarray</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="ThresholdOptimization.aggregate">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.ThresholdOptimization.aggregate">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">aggregate</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">:</span> <span class="n">np</span><span class="o">.</span><span class="n">ndarray</span><span class="p">):</span>
|
||||
<span class="c1"># the standard behavior is to compute the adjusted count using the best threshold found</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">aggregate_with_threshold</span><span class="p">(</span><span class="n">classif_predictions</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">tpr</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">fpr</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">threshold</span><span class="p">)</span></div></div>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">aggregate_with_threshold</span><span class="p">(</span><span class="n">classif_predictions</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">tpr</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">fpr</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">threshold</span><span class="p">)</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="T50"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.T50">[docs]</a><span class="k">class</span> <span class="nc">T50</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="T50">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.T50">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">T50</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Threshold Optimization variant for :class:`ACC` as proposed by</span>
|
||||
<span class="sd"> `Forman 2006 <https://dl.acm.org/doi/abs/10.1145/1150402.1150423>`_ and</span>
|
||||
|
|
@ -207,23 +531,34 @@
|
|||
<span class="sd"> for the threshold that makes `tpr` closest to 0.5.</span>
|
||||
<span class="sd"> The goal is to bring improved stability to the denominator of the adjustment.</span>
|
||||
|
||||
<span class="sd"> :param classifier: a sklearn's Estimator that generates a classifier</span>
|
||||
<span class="sd"> :param val_split: indicates the proportion of data to be used as a stratified held-out validation set in which the</span>
|
||||
<span class="sd"> misclassification rates are to be estimated.</span>
|
||||
<span class="sd"> This parameter can be indicated as a real value (between 0 and 1), representing a proportion of</span>
|
||||
<span class="sd"> validation data, or as an integer, indicating that the misclassification rates should be estimated via</span>
|
||||
<span class="sd"> `k`-fold cross validation (this integer stands for the number of folds `k`, defaults 5), or as a</span>
|
||||
<span class="sd"> :class:`quapy.data.base.LabelledCollection` (the split itself).</span>
|
||||
<span class="sd"> :param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be</span>
|
||||
<span class="sd"> the one indicated in `qp.environ['DEFAULT_CLS']`</span>
|
||||
|
||||
<span class="sd"> :param fit_classifier: whether to train the learner (default is True). Set to False if the</span>
|
||||
<span class="sd"> learner has been trained outside the quantifier.</span>
|
||||
|
||||
<span class="sd"> :param val_split: specifies the data used for generating classifier predictions. This specification</span>
|
||||
<span class="sd"> can be made as float in (0, 1) indicating the proportion of stratified held-out validation set to</span>
|
||||
<span class="sd"> be extracted from the training set; or as an integer (default 5), indicating that the predictions</span>
|
||||
<span class="sd"> are to be generated in a `k`-fold cross-validation manner (with this integer indicating the value</span>
|
||||
<span class="sd"> for `k`); or as a tuple (X,y) defining the specific set of data to use for validation.</span>
|
||||
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="T50.condition"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.T50.condition">[docs]</a> <span class="k">def</span> <span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="nb">abs</span><span class="p">(</span><span class="n">tpr</span> <span class="o">-</span> <span class="mf">0.5</span><span class="p">)</span></div></div>
|
||||
<div class="viewcode-block" id="T50.condition">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.T50.condition">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="nb">abs</span><span class="p">(</span><span class="n">tpr</span> <span class="o">-</span> <span class="mf">0.5</span><span class="p">)</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="MAX"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MAX">[docs]</a><span class="k">class</span> <span class="nc">MAX</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="MAX">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MAX">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">MAX</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Threshold Optimization variant for :class:`ACC` as proposed by</span>
|
||||
<span class="sd"> `Forman 2006 <https://dl.acm.org/doi/abs/10.1145/1150402.1150423>`_ and</span>
|
||||
|
|
@ -231,24 +566,33 @@
|
|||
<span class="sd"> for the threshold that maximizes `tpr-fpr`.</span>
|
||||
<span class="sd"> The goal is to bring improved stability to the denominator of the adjustment.</span>
|
||||
|
||||
<span class="sd"> :param classifier: a sklearn's Estimator that generates a classifier</span>
|
||||
<span class="sd"> :param val_split: indicates the proportion of data to be used as a stratified held-out validation set in which the</span>
|
||||
<span class="sd"> misclassification rates are to be estimated.</span>
|
||||
<span class="sd"> This parameter can be indicated as a real value (between 0 and 1), representing a proportion of</span>
|
||||
<span class="sd"> validation data, or as an integer, indicating that the misclassification rates should be estimated via</span>
|
||||
<span class="sd"> `k`-fold cross validation (this integer stands for the number of folds `k`, defaults 5), or as a</span>
|
||||
<span class="sd"> :class:`quapy.data.base.LabelledCollection` (the split itself).</span>
|
||||
<span class="sd"> :param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be</span>
|
||||
<span class="sd"> the one indicated in `qp.environ['DEFAULT_CLS']`</span>
|
||||
<span class="sd"> :param fit_classifier: whether to train the learner (default is True). Set to False if the</span>
|
||||
<span class="sd"> learner has been trained outside the quantifier.</span>
|
||||
<span class="sd"> :param val_split: specifies the data used for generating classifier predictions. This specification</span>
|
||||
<span class="sd"> can be made as float in (0, 1) indicating the proportion of stratified held-out validation set to</span>
|
||||
<span class="sd"> be extracted from the training set; or as an integer (default 5), indicating that the predictions</span>
|
||||
<span class="sd"> are to be generated in a `k`-fold cross-validation manner (with this integer indicating the value</span>
|
||||
<span class="sd"> for `k`); or as a tuple (X,y) defining the specific set of data to use for validation.</span>
|
||||
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="MAX.condition"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MAX.condition">[docs]</a> <span class="k">def</span> <span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<div class="viewcode-block" id="MAX.condition">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MAX.condition">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="c1"># MAX strives to maximize (tpr - fpr), which is equivalent to minimize (fpr - tpr)</span>
|
||||
<span class="k">return</span> <span class="p">(</span><span class="n">fpr</span> <span class="o">-</span> <span class="n">tpr</span><span class="p">)</span></div></div>
|
||||
<span class="k">return</span> <span class="p">(</span><span class="n">fpr</span> <span class="o">-</span> <span class="n">tpr</span><span class="p">)</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="X"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.X">[docs]</a><span class="k">class</span> <span class="nc">X</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="X">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.X">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">X</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Threshold Optimization variant for :class:`ACC` as proposed by</span>
|
||||
<span class="sd"> `Forman 2006 <https://dl.acm.org/doi/abs/10.1145/1150402.1150423>`_ and</span>
|
||||
|
|
@ -256,23 +600,32 @@
|
|||
<span class="sd"> for the threshold that yields `tpr=1-fpr`.</span>
|
||||
<span class="sd"> The goal is to bring improved stability to the denominator of the adjustment.</span>
|
||||
|
||||
<span class="sd"> :param classifier: a sklearn's Estimator that generates a classifier</span>
|
||||
<span class="sd"> :param val_split: indicates the proportion of data to be used as a stratified held-out validation set in which the</span>
|
||||
<span class="sd"> misclassification rates are to be estimated.</span>
|
||||
<span class="sd"> This parameter can be indicated as a real value (between 0 and 1), representing a proportion of</span>
|
||||
<span class="sd"> validation data, or as an integer, indicating that the misclassification rates should be estimated via</span>
|
||||
<span class="sd"> `k`-fold cross validation (this integer stands for the number of folds `k`, defaults 5), or as a</span>
|
||||
<span class="sd"> :class:`quapy.data.base.LabelledCollection` (the split itself).</span>
|
||||
<span class="sd"> :param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be</span>
|
||||
<span class="sd"> the one indicated in `qp.environ['DEFAULT_CLS']`</span>
|
||||
<span class="sd"> :param fit_classifier: whether to train the learner (default is True). Set to False if the</span>
|
||||
<span class="sd"> learner has been trained outside the quantifier.</span>
|
||||
<span class="sd"> :param val_split: specifies the data used for generating classifier predictions. This specification</span>
|
||||
<span class="sd"> can be made as float in (0, 1) indicating the proportion of stratified held-out validation set to</span>
|
||||
<span class="sd"> be extracted from the training set; or as an integer (default 5), indicating that the predictions</span>
|
||||
<span class="sd"> are to be generated in a `k`-fold cross-validation manner (with this integer indicating the value</span>
|
||||
<span class="sd"> for `k`); or as a tuple (X,y) defining the specific set of data to use for validation.</span>
|
||||
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="X.condition"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.X.condition">[docs]</a> <span class="k">def</span> <span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="nb">abs</span><span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="p">(</span><span class="n">tpr</span> <span class="o">+</span> <span class="n">fpr</span><span class="p">))</span></div></div>
|
||||
<div class="viewcode-block" id="X.condition">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.X.condition">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="nb">abs</span><span class="p">(</span><span class="mi">1</span> <span class="o">-</span> <span class="p">(</span><span class="n">tpr</span> <span class="o">+</span> <span class="n">fpr</span><span class="p">))</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="MS"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS">[docs]</a><span class="k">class</span> <span class="nc">MS</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="MS">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">MS</span><span class="p">(</span><span class="n">ThresholdOptimization</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Median Sweep. Threshold Optimization variant for :class:`ACC` as proposed by</span>
|
||||
<span class="sd"> `Forman 2006 <https://dl.acm.org/doi/abs/10.1145/1150402.1150423>`_ and</span>
|
||||
|
|
@ -280,22 +633,30 @@
|
|||
<span class="sd"> class prevalence estimates for all decision thresholds and returns the median of them all.</span>
|
||||
<span class="sd"> The goal is to bring improved stability to the denominator of the adjustment.</span>
|
||||
|
||||
<span class="sd"> :param classifier: a sklearn's Estimator that generates a classifier</span>
|
||||
<span class="sd"> :param val_split: indicates the proportion of data to be used as a stratified held-out validation set in which the</span>
|
||||
<span class="sd"> misclassification rates are to be estimated.</span>
|
||||
<span class="sd"> This parameter can be indicated as a real value (between 0 and 1), representing a proportion of</span>
|
||||
<span class="sd"> validation data, or as an integer, indicating that the misclassification rates should be estimated via</span>
|
||||
<span class="sd"> `k`-fold cross validation (this integer stands for the number of folds `k`, defaults 5), or as a</span>
|
||||
<span class="sd"> :class:`quapy.data.base.LabelledCollection` (the split itself).</span>
|
||||
<span class="sd"> :param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be</span>
|
||||
<span class="sd"> the one indicated in `qp.environ['DEFAULT_CLS']`</span>
|
||||
<span class="sd"> :param fit_classifier: whether to train the learner (default is True). Set to False if the</span>
|
||||
<span class="sd"> learner has been trained outside the quantifier.</span>
|
||||
<span class="sd"> :param val_split: specifies the data used for generating classifier predictions. This specification</span>
|
||||
<span class="sd"> can be made as float in (0, 1) indicating the proportion of stratified held-out validation set to</span>
|
||||
<span class="sd"> be extracted from the training set; or as an integer (default 5), indicating that the predictions</span>
|
||||
<span class="sd"> are to be generated in a `k`-fold cross-validation manner (with this integer indicating the value</span>
|
||||
<span class="sd"> for `k`); or as a tuple (X,y) defining the specific set of data to use for validation.</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="MS.condition"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS.condition">[docs]</a> <span class="k">def</span> <span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="MS.condition">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS.condition">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">condition</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">float</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="mi">1</span></div>
|
||||
|
||||
<div class="viewcode-block" id="MS.aggregation_fit"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS.aggregation_fit">[docs]</a> <span class="k">def</span> <span class="nf">aggregation_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">):</span>
|
||||
<span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">classif_predictions</span><span class="o">.</span><span class="n">Xy</span>
|
||||
|
||||
<div class="viewcode-block" id="MS.aggregation_fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS.aggregation_fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">aggregation_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">,</span> <span class="n">labels</span><span class="p">):</span>
|
||||
<span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">classif_predictions</span><span class="p">,</span> <span class="n">labels</span>
|
||||
<span class="c1"># keeps all candidates</span>
|
||||
<span class="n">tprs_fprs_thresholds</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_eval_candidate_thresholds</span><span class="p">(</span><span class="n">decision_scores</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">tprs</span> <span class="o">=</span> <span class="n">tprs_fprs_thresholds</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">]</span>
|
||||
|
|
@ -303,14 +664,21 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">thresholds</span> <span class="o">=</span> <span class="n">tprs_fprs_thresholds</span><span class="p">[:,</span> <span class="mi">2</span><span class="p">]</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<div class="viewcode-block" id="MS.aggregate"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS.aggregate">[docs]</a> <span class="k">def</span> <span class="nf">aggregate</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">:</span> <span class="n">np</span><span class="o">.</span><span class="n">ndarray</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="MS.aggregate">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS.aggregate">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">aggregate</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classif_predictions</span><span class="p">:</span> <span class="n">np</span><span class="o">.</span><span class="n">ndarray</span><span class="p">):</span>
|
||||
<span class="n">prevalences</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">aggregate_with_threshold</span><span class="p">(</span><span class="n">classif_predictions</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">tprs</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">fprs</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">thresholds</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="n">prevalences</span><span class="o">.</span><span class="n">ndim</span><span class="o">==</span><span class="mi">2</span><span class="p">:</span>
|
||||
<span class="n">prevalences</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">median</span><span class="p">(</span><span class="n">prevalences</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">prevalences</span></div></div>
|
||||
<span class="k">return</span> <span class="n">prevalences</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="MS2"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS2">[docs]</a><span class="k">class</span> <span class="nc">MS2</span><span class="p">(</span><span class="n">MS</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="MS2">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS2">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">MS2</span><span class="p">(</span><span class="n">MS</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Median Sweep 2. Threshold Optimization variant for :class:`ACC` as proposed by</span>
|
||||
<span class="sd"> `Forman 2006 <https://dl.acm.org/doi/abs/10.1145/1150402.1150423>`_ and</span>
|
||||
|
|
@ -319,46 +687,97 @@
|
|||
<span class="sd"> which `tpr-fpr>0.25`</span>
|
||||
<span class="sd"> The goal is to bring improved stability to the denominator of the adjustment.</span>
|
||||
|
||||
<span class="sd"> :param classifier: a sklearn's Estimator that generates a classifier</span>
|
||||
<span class="sd"> :param val_split: indicates the proportion of data to be used as a stratified held-out validation set in which the</span>
|
||||
<span class="sd"> misclassification rates are to be estimated.</span>
|
||||
<span class="sd"> This parameter can be indicated as a real value (between 0 and 1), representing a proportion of</span>
|
||||
<span class="sd"> validation data, or as an integer, indicating that the misclassification rates should be estimated via</span>
|
||||
<span class="sd"> `k`-fold cross validation (this integer stands for the number of folds `k`, defaults 5), or as a</span>
|
||||
<span class="sd"> :class:`quapy.data.base.LabelledCollection` (the split itself).</span>
|
||||
<span class="sd"> :param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be</span>
|
||||
<span class="sd"> the one indicated in `qp.environ['DEFAULT_CLS']`</span>
|
||||
<span class="sd"> :param fit_classifier: whether to train the learner (default is True). Set to False if the</span>
|
||||
<span class="sd"> learner has been trained outside the quantifier.</span>
|
||||
<span class="sd"> :param val_split: specifies the data used for generating classifier predictions. This specification</span>
|
||||
<span class="sd"> can be made as float in (0, 1) indicating the proportion of stratified held-out validation set to</span>
|
||||
<span class="sd"> be extracted from the training set; or as an integer (default 5), indicating that the predictions</span>
|
||||
<span class="sd"> are to be generated in a `k`-fold cross-validation manner (with this integer indicating the value</span>
|
||||
<span class="sd"> for `k`); or as a tuple (X,y) defining the specific set of data to use for validation.</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="MS2.discard"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS2.discard">[docs]</a> <span class="k">def</span> <span class="nf">discard</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">bool</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="p">(</span><span class="n">tpr</span><span class="o">-</span><span class="n">fpr</span><span class="p">)</span> <span class="o"><=</span> <span class="mf">0.25</span></div></div>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">classifier</span><span class="p">:</span> <span class="n">BaseEstimator</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">val_split</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
|
||||
<span class="nb">super</span><span class="p">()</span><span class="o">.</span><span class="fm">__init__</span><span class="p">(</span><span class="n">classifier</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="p">,</span> <span class="n">val_split</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="MS2.discard">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method._threshold_optim.MS2.discard">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">discard</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">tpr</span><span class="p">,</span> <span class="n">fpr</span><span class="p">)</span> <span class="o">-></span> <span class="nb">bool</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="p">(</span><span class="n">tpr</span><span class="o">-</span><span class="n">fpr</span><span class="p">)</span> <span class="o"><=</span> <span class="mf">0.25</span></div>
|
||||
</div>
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,133 +1,464 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method.base — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method.base — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../../" id="documentation_options" src="../../../_static/documentation_options.js"></script>
|
||||
<script src="../../../_static/jquery.js"></script>
|
||||
<script src="../../../_static/underscore.js"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../../_static/doctools.js"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/method/base';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.method.base</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.method.base</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.method.base</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">from</span> <span class="nn">abc</span> <span class="kn">import</span> <span class="n">ABCMeta</span><span class="p">,</span> <span class="n">abstractmethod</span>
|
||||
<span class="kn">from</span> <span class="nn">copy</span> <span class="kn">import</span> <span class="n">deepcopy</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">warnings</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">abc</span><span class="w"> </span><span class="kn">import</span> <span class="n">ABCMeta</span><span class="p">,</span> <span class="n">abstractmethod</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">copy</span><span class="w"> </span><span class="kn">import</span> <span class="n">deepcopy</span>
|
||||
|
||||
<span class="kn">from</span> <span class="nn">joblib</span> <span class="kn">import</span> <span class="n">Parallel</span><span class="p">,</span> <span class="n">delayed</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn.base</span> <span class="kn">import</span> <span class="n">BaseEstimator</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">joblib</span><span class="w"> </span><span class="kn">import</span> <span class="n">Parallel</span><span class="p">,</span> <span class="n">delayed</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseEstimator</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">quapy</span> <span class="k">as</span> <span class="nn">qp</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.data</span> <span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.data</span><span class="w"> </span><span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
|
||||
|
||||
<span class="c1"># Base Quantifier abstract class</span>
|
||||
<span class="c1"># ------------------------------------</span>
|
||||
<div class="viewcode-block" id="BaseQuantifier"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BaseQuantifier">[docs]</a><span class="k">class</span> <span class="nc">BaseQuantifier</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="BaseQuantifier">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BaseQuantifier">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">BaseQuantifier</span><span class="p">(</span><span class="n">BaseEstimator</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Abstract Quantifier. A quantifier is defined as an object of a class that implements the method :meth:`fit` on</span>
|
||||
<span class="sd"> :class:`quapy.data.base.LabelledCollection`, the method :meth:`quantify`, and the :meth:`set_params` and</span>
|
||||
<span class="sd"> a pair X, y, the method :meth:`predict`, and the :meth:`set_params` and</span>
|
||||
<span class="sd"> :meth:`get_params` for model selection (see :meth:`quapy.model_selection.GridSearchQ`)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<div class="viewcode-block" id="BaseQuantifier.fit"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BaseQuantifier.fit">[docs]</a> <span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="BaseQuantifier.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BaseQuantifier.fit">[docs]</a>
|
||||
<span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Trains a quantifier.</span>
|
||||
<span class="sd"> Generates a quantifier.</span>
|
||||
|
||||
<span class="sd"> :param data: a :class:`quapy.data.base.LabelledCollection` consisting of the training data</span>
|
||||
<span class="sd"> :param X: array-like, the training instances</span>
|
||||
<span class="sd"> :param y: array-like, the labels</span>
|
||||
<span class="sd"> :return: self</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="o">...</span></div>
|
||||
|
||||
<div class="viewcode-block" id="BaseQuantifier.quantify"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BaseQuantifier.quantify">[docs]</a> <span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span> <span class="nf">quantify</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="BaseQuantifier.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BaseQuantifier.predict">[docs]</a>
|
||||
<span class="nd">@abstractmethod</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Generate class prevalence estimates for the sample's instances</span>
|
||||
|
||||
<span class="sd"> :param instances: array-like</span>
|
||||
<span class="sd"> :param X: array-like, the test instances</span>
|
||||
<span class="sd"> :return: `np.ndarray` of shape `(n_classes,)` with class prevalence estimates.</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="o">...</span></div></div>
|
||||
<span class="o">...</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="BinaryQuantifier"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BinaryQuantifier">[docs]</a><span class="k">class</span> <span class="nc">BinaryQuantifier</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="BaseQuantifier.quantify">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BaseQuantifier.quantify">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">quantify</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Alias to :meth:`predict`, for old compatibility</span>
|
||||
|
||||
<span class="sd"> :param X: array-like</span>
|
||||
<span class="sd"> :return: `np.ndarray` of shape `(n_classes,)` with class prevalence estimates.</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">predict</span><span class="p">(</span><span class="n">X</span><span class="p">)</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="BinaryQuantifier">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.BinaryQuantifier">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">BinaryQuantifier</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Abstract class of binary quantifiers, i.e., quantifiers estimating class prevalence values for only two classes</span>
|
||||
<span class="sd"> (typically, to be interpreted as one class and its complement).</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_check_binary</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">quantifier_name</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="n">data</span><span class="o">.</span><span class="n">binary</span><span class="p">,</span> <span class="sa">f</span><span class="s1">'</span><span class="si">{</span><span class="n">quantifier_name</span><span class="si">}</span><span class="s1"> works only on problems of binary classification. '</span> \
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_check_binary</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">quantifier_name</span><span class="p">):</span>
|
||||
<span class="n">n_classes</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="nb">set</span><span class="p">(</span><span class="n">y</span><span class="p">))</span>
|
||||
<span class="k">assert</span> <span class="n">n_classes</span><span class="o">==</span><span class="mi">2</span><span class="p">,</span> <span class="sa">f</span><span class="s1">'</span><span class="si">{</span><span class="n">quantifier_name</span><span class="si">}</span><span class="s1"> works only on problems of binary classification. '</span> \
|
||||
<span class="sa">f</span><span class="s1">'Use the class OneVsAll to enable </span><span class="si">{</span><span class="n">quantifier_name</span><span class="si">}</span><span class="s1"> work on single-label data.'</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="OneVsAll"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAll">[docs]</a><span class="k">class</span> <span class="nc">OneVsAll</span><span class="p">:</span>
|
||||
|
||||
<div class="viewcode-block" id="OneVsAll">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAll">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">OneVsAll</span><span class="p">:</span>
|
||||
<span class="k">pass</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="newOneVsAll"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.newOneVsAll">[docs]</a><span class="k">def</span> <span class="nf">newOneVsAll</span><span class="p">(</span><span class="n">binary_quantifier</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="newOneVsAll">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.newOneVsAll">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">newOneVsAll</span><span class="p">(</span><span class="n">binary_quantifier</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">binary_quantifier</span><span class="p">,</span> <span class="n">BaseQuantifier</span><span class="p">),</span> \
|
||||
<span class="sa">f</span><span class="s1">'</span><span class="si">{</span><span class="n">binary_quantifier</span><span class="si">}</span><span class="s1"> does not seem to be a Quantifier'</span>
|
||||
<span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">binary_quantifier</span><span class="p">,</span> <span class="n">qp</span><span class="o">.</span><span class="n">method</span><span class="o">.</span><span class="n">aggregative</span><span class="o">.</span><span class="n">AggregativeQuantifier</span><span class="p">):</span>
|
||||
|
|
@ -136,77 +467,131 @@
|
|||
<span class="k">return</span> <span class="n">OneVsAllGeneric</span><span class="p">(</span><span class="n">binary_quantifier</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="OneVsAllGeneric"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAllGeneric">[docs]</a><span class="k">class</span> <span class="nc">OneVsAllGeneric</span><span class="p">(</span><span class="n">OneVsAll</span><span class="p">,</span> <span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="OneVsAllGeneric">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAllGeneric">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">OneVsAllGeneric</span><span class="p">(</span><span class="n">OneVsAll</span><span class="p">,</span> <span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Allows any binary quantifier to perform quantification on single-label datasets. The method maintains one binary</span>
|
||||
<span class="sd"> quantifier for each class, and then l1-normalizes the outputs so that the class prevelence values sum up to 1.</span>
|
||||
<span class="sd"> quantifier for each class, and then l1-normalizes the outputs so that the class prevalence values sum up to 1.</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">binary_quantifier</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">binary_quantifier</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">binary_quantifier</span><span class="p">,</span> <span class="n">BaseQuantifier</span><span class="p">),</span> \
|
||||
<span class="sa">f</span><span class="s1">'</span><span class="si">{</span><span class="n">binary_quantifier</span><span class="si">}</span><span class="s1"> does not seem to be a Quantifier'</span>
|
||||
<span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">binary_quantifier</span><span class="p">,</span> <span class="n">qp</span><span class="o">.</span><span class="n">method</span><span class="o">.</span><span class="n">aggregative</span><span class="o">.</span><span class="n">AggregativeQuantifier</span><span class="p">):</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="s1">'[warning] the quantifier seems to be an instance of qp.method.aggregative.AggregativeQuantifier; '</span>
|
||||
<span class="sa">f</span><span class="s1">'you might prefer instantiating </span><span class="si">{</span><span class="n">qp</span><span class="o">.</span><span class="n">method</span><span class="o">.</span><span class="n">aggregative</span><span class="o">.</span><span class="n">OneVsAllAggregative</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">warnings</span><span class="o">.</span><span class="n">warn</span><span class="p">(</span><span class="s1">'the quantifier seems to be an instance of qp.method.aggregative.AggregativeQuantifier; '</span>
|
||||
<span class="sa">f</span><span class="s1">'you might prefer instantiating </span><span class="si">{</span><span class="n">qp</span><span class="o">.</span><span class="n">method</span><span class="o">.</span><span class="n">aggregative</span><span class="o">.</span><span class="n">OneVsAllAggregative</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">binary_quantifier</span> <span class="o">=</span> <span class="n">binary_quantifier</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span> <span class="o">=</span> <span class="n">qp</span><span class="o">.</span><span class="n">_get_njobs</span><span class="p">(</span><span class="n">n_jobs</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="OneVsAllGeneric.fit"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAllGeneric.fit">[docs]</a> <span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">fit_classifier</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="ow">not</span> <span class="n">data</span><span class="o">.</span><span class="n">binary</span><span class="p">,</span> <span class="sa">f</span><span class="s1">'</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1"> expect non-binary data'</span>
|
||||
<span class="k">assert</span> <span class="n">fit_classifier</span> <span class="o">==</span> <span class="kc">True</span><span class="p">,</span> <span class="s1">'fit_classifier must be True'</span>
|
||||
<div class="viewcode-block" id="OneVsAllGeneric.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAllGeneric.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes</span> <span class="o">=</span> <span class="nb">sorted</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">y</span><span class="p">))</span>
|
||||
<span class="k">assert</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classes</span><span class="p">)</span><span class="o">!=</span><span class="mi">2</span><span class="p">,</span> <span class="sa">f</span><span class="s1">'</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1"> expect non-binary data'</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">dict_binary_quantifiers</span> <span class="o">=</span> <span class="p">{</span><span class="n">c</span><span class="p">:</span> <span class="n">deepcopy</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">binary_quantifier</span><span class="p">)</span> <span class="k">for</span> <span class="n">c</span> <span class="ow">in</span> <span class="n">data</span><span class="o">.</span><span class="n">classes_</span><span class="p">}</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_parallel</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_delayed_binary_fit</span><span class="p">,</span> <span class="n">data</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">dict_binary_quantifiers</span> <span class="o">=</span> <span class="p">{</span><span class="n">c</span><span class="p">:</span> <span class="n">deepcopy</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">binary_quantifier</span><span class="p">)</span> <span class="k">for</span> <span class="n">c</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">classes</span><span class="p">}</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_parallel</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_delayed_binary_fit</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_parallel</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">func</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_parallel</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">func</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span>
|
||||
<span class="n">Parallel</span><span class="p">(</span><span class="n">n_jobs</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span><span class="p">,</span> <span class="n">backend</span><span class="o">=</span><span class="s1">'threading'</span><span class="p">)(</span>
|
||||
<span class="n">delayed</span><span class="p">(</span><span class="n">func</span><span class="p">)(</span><span class="n">c</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span> <span class="k">for</span> <span class="n">c</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">classes_</span>
|
||||
<span class="n">delayed</span><span class="p">(</span><span class="n">func</span><span class="p">)(</span><span class="n">c</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span> <span class="k">for</span> <span class="n">c</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">classes</span>
|
||||
<span class="p">)</span>
|
||||
<span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="OneVsAllGeneric.quantify"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAllGeneric.quantify">[docs]</a> <span class="k">def</span> <span class="nf">quantify</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
<span class="n">prevalences</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_parallel</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_delayed_binary_predict</span><span class="p">,</span> <span class="n">instances</span><span class="p">)</span>
|
||||
<div class="viewcode-block" id="OneVsAllGeneric.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.base.OneVsAllGeneric.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="n">prevalences</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_parallel</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_delayed_binary_predict</span><span class="p">,</span> <span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">qp</span><span class="o">.</span><span class="n">functional</span><span class="o">.</span><span class="n">normalize_prevalence</span><span class="p">(</span><span class="n">prevalences</span><span class="p">)</span></div>
|
||||
|
||||
<span class="nd">@property</span>
|
||||
<span class="k">def</span> <span class="nf">classes_</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="nb">sorted</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">dict_binary_quantifiers</span><span class="o">.</span><span class="n">keys</span><span class="p">())</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_delayed_binary_predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">c</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">dict_binary_quantifiers</span><span class="p">[</span><span class="n">c</span><span class="p">]</span><span class="o">.</span><span class="n">quantify</span><span class="p">(</span><span class="n">X</span><span class="p">)[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="c1"># @property</span>
|
||||
<span class="c1"># def classes_(self):</span>
|
||||
<span class="c1"># return sorted(self.dict_binary_quantifiers.keys())</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_delayed_binary_predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">c</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">dict_binary_quantifiers</span><span class="p">[</span><span class="n">c</span><span class="p">]</span><span class="o">.</span><span class="n">predict</span><span class="p">(</span><span class="n">X</span><span class="p">)[</span><span class="mi">1</span><span class="p">]</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_delayed_binary_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">c</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="n">bindata</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">y</span> <span class="o">==</span> <span class="n">c</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="p">[</span><span class="kc">False</span><span class="p">,</span> <span class="kc">True</span><span class="p">])</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">dict_binary_quantifiers</span><span class="p">[</span><span class="n">c</span><span class="p">]</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">bindata</span><span class="o">.</span><span class="n">Xy</span><span class="p">)</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_delayed_binary_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">c</span><span class="p">,</span> <span class="n">data</span><span class="p">):</span>
|
||||
<span class="n">bindata</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">data</span><span class="o">.</span><span class="n">instances</span><span class="p">,</span> <span class="n">data</span><span class="o">.</span><span class="n">labels</span> <span class="o">==</span> <span class="n">c</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="p">[</span><span class="kc">False</span><span class="p">,</span> <span class="kc">True</span><span class="p">])</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">dict_binary_quantifiers</span><span class="p">[</span><span class="n">c</span><span class="p">]</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">bindata</span><span class="p">)</span></div>
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,86 +1,397 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method.non_aggregative — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.method.non_aggregative — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../../" id="documentation_options" src="../../../_static/documentation_options.js"></script>
|
||||
<script src="../../../_static/jquery.js"></script>
|
||||
<script src="../../../_static/underscore.js"></script>
|
||||
<script src="../../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../../_static/doctools.js"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/method/non_aggregative';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../../_static/searchtools.js"></script>
|
||||
<script src="../../../_static/language_data.js"></script>
|
||||
<script src="../../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.method.non_aggregative</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.method.non_aggregative</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.method.non_aggregative</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">from</span> <span class="nn">typing</span> <span class="kn">import</span> <span class="n">Union</span><span class="p">,</span> <span class="n">Callable</span>
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span></span><span class="kn">from</span><span class="w"> </span><span class="nn">itertools</span><span class="w"> </span><span class="kn">import</span> <span class="n">product</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">tqdm</span><span class="w"> </span><span class="kn">import</span> <span class="n">tqdm</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">typing</span><span class="w"> </span><span class="kn">import</span> <span class="n">Union</span><span class="p">,</span> <span class="n">Callable</span><span class="p">,</span> <span class="n">Counter</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.feature_extraction.text</span><span class="w"> </span><span class="kn">import</span> <span class="n">CountVectorizer</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.utils</span><span class="w"> </span><span class="kn">import</span> <span class="n">resample</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn.preprocessing</span><span class="w"> </span><span class="kn">import</span> <span class="n">normalize</span>
|
||||
|
||||
<span class="kn">from</span> <span class="nn">quapy.functional</span> <span class="kn">import</span> <span class="n">get_divergence</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.data</span> <span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.method.base</span> <span class="kn">import</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">BinaryQuantifier</span>
|
||||
<span class="kn">import</span> <span class="nn">quapy.functional</span> <span class="k">as</span> <span class="nn">F</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.confidence</span><span class="w"> </span><span class="kn">import</span> <span class="n">WithConfidenceABC</span><span class="p">,</span> <span class="n">ConfidenceRegionABC</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.functional</span><span class="w"> </span><span class="kn">import</span> <span class="n">get_divergence</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">BinaryQuantifier</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method._helper</span><span class="w"> </span><span class="kn">import</span> <span class="n">_labels_to_indices</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method._energy</span><span class="w"> </span><span class="kn">import</span> <span class="n">_EnergyDistanceCore</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy.functional</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">F</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">scipy.optimize</span><span class="w"> </span><span class="kn">import</span> <span class="n">lsq_linear</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">scipy</span><span class="w"> </span><span class="kn">import</span> <span class="n">sparse</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="MaximumLikelihoodPrevalenceEstimation"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.MaximumLikelihoodPrevalenceEstimation">[docs]</a><span class="k">class</span> <span class="nc">MaximumLikelihoodPrevalenceEstimation</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="MaximumLikelihoodPrevalenceEstimation">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.MaximumLikelihoodPrevalenceEstimation">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">MaximumLikelihoodPrevalenceEstimation</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> The `Maximum Likelihood Prevalence Estimation` (MLPE) method is a lazy method that assumes there is no prior</span>
|
||||
<span class="sd"> probability shift between training and test instances (put it other way, that the i.i.d. assumpion holds).</span>
|
||||
|
|
@ -89,30 +400,41 @@
|
|||
<span class="sd"> any quantification method should beat.</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_classes_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
|
||||
<div class="viewcode-block" id="MaximumLikelihoodPrevalenceEstimation.fit"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.MaximumLikelihoodPrevalenceEstimation.fit">[docs]</a> <span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="MaximumLikelihoodPrevalenceEstimation.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.MaximumLikelihoodPrevalenceEstimation.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Computes the training prevalence and stores it.</span>
|
||||
|
||||
<span class="sd"> :param data: the training sample</span>
|
||||
<span class="sd"> :param X: array-like of shape `(n_samples, n_features)`, the training instances</span>
|
||||
<span class="sd"> :param y: array-like of shape `(n_samples,)`, the labels</span>
|
||||
<span class="sd"> :return: self</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">estimated_prevalence</span> <span class="o">=</span> <span class="n">data</span><span class="o">.</span><span class="n">prevalence</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_classes_</span> <span class="o">=</span> <span class="n">F</span><span class="o">.</span><span class="n">classes_from_labels</span><span class="p">(</span><span class="n">labels</span><span class="o">=</span><span class="n">y</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">estimated_prevalence</span> <span class="o">=</span> <span class="n">F</span><span class="o">.</span><span class="n">prevalence_from_labels</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">_classes_</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<div class="viewcode-block" id="MaximumLikelihoodPrevalenceEstimation.quantify"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.MaximumLikelihoodPrevalenceEstimation.quantify">[docs]</a> <span class="k">def</span> <span class="nf">quantify</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="MaximumLikelihoodPrevalenceEstimation.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.MaximumLikelihoodPrevalenceEstimation.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Ignores the input instances and returns, as the class prevalence estimantes, the training prevalence.</span>
|
||||
|
||||
<span class="sd"> :param instances: array-like (ignored)</span>
|
||||
<span class="sd"> :param X: array-like (ignored)</span>
|
||||
<span class="sd"> :return: the class prevalence seen during training</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">estimated_prevalence</span></div></div>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">estimated_prevalence</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="DMx"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx">[docs]</a><span class="k">class</span> <span class="nc">DMx</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="DMx">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">DMx</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Generic Distribution Matching quantifier for binary or multiclass quantification based on the space of covariates.</span>
|
||||
<span class="sd"> This implementation takes the number of bins, the divergence, and the possibility to work on CDF as hyperparameters.</span>
|
||||
|
|
@ -122,18 +444,23 @@
|
|||
<span class="sd"> or a callable function taking two ndarrays of the same dimension as input (default "HD", meaning Hellinger</span>
|
||||
<span class="sd"> Distance)</span>
|
||||
<span class="sd"> :param cdf: whether to use CDF instead of PDF (default False)</span>
|
||||
<span class="sd"> :param search: string indicating the search strategy used to estimate the prevalence values.</span>
|
||||
<span class="sd"> Valid options are `optim_minimize` (default, works for binary and multiclass problems),</span>
|
||||
<span class="sd"> `linear_search` (binary only), and `ternary_search` (binary only)</span>
|
||||
<span class="sd"> :param n_jobs: number of parallel workers (default None)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">nbins</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">divergence</span><span class="p">:</span> <span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="n">Callable</span><span class="p">]</span><span class="o">=</span><span class="s1">'HD'</span><span class="p">,</span> <span class="n">cdf</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">search</span><span class="o">=</span><span class="s1">'optim_minimize'</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">nbins</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">divergence</span><span class="p">:</span> <span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="n">Callable</span><span class="p">]</span><span class="o">=</span><span class="s1">'HD'</span><span class="p">,</span> <span class="n">cdf</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">search</span><span class="o">=</span><span class="s1">'optim_minimize'</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">nbins</span> <span class="o">=</span> <span class="n">nbins</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">divergence</span> <span class="o">=</span> <span class="n">divergence</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">cdf</span> <span class="o">=</span> <span class="n">cdf</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">search</span> <span class="o">=</span> <span class="n">search</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span> <span class="o">=</span> <span class="n">n_jobs</span>
|
||||
|
||||
<div class="viewcode-block" id="DMx.HDx"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx.HDx">[docs]</a> <span class="nd">@classmethod</span>
|
||||
<span class="k">def</span> <span class="nf">HDx</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="DMx.HDx">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx.HDx">[docs]</a>
|
||||
<span class="nd">@classmethod</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">HDx</span><span class="p">(</span><span class="bp">cls</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> `Hellinger Distance x <https://www.sciencedirect.com/science/article/pii/S0020025512004069>`_ (HDx).</span>
|
||||
<span class="sd"> HDx is a method for training binary quantifiers, that models quantification as the problem of</span>
|
||||
|
|
@ -149,15 +476,15 @@
|
|||
<span class="sd"> :return: an instance of this class setup to mimick the performance of the HDx as originally proposed by</span>
|
||||
<span class="sd"> González-Castro, Alaiz-Rodríguez, Alegre (2013)</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.method.meta</span> <span class="kn">import</span> <span class="n">MedianEstimator</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.meta</span><span class="w"> </span><span class="kn">import</span> <span class="n">MedianEstimator</span>
|
||||
|
||||
<span class="n">dmx</span> <span class="o">=</span> <span class="n">DMx</span><span class="p">(</span><span class="n">divergence</span><span class="o">=</span><span class="s1">'HD'</span><span class="p">,</span> <span class="n">cdf</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">search</span><span class="o">=</span><span class="s1">'linear_search'</span><span class="p">)</span>
|
||||
<span class="n">nbins</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'nbins'</span><span class="p">:</span> <span class="n">np</span><span class="o">.</span><span class="n">linspace</span><span class="p">(</span><span class="mi">10</span><span class="p">,</span> <span class="mi">110</span><span class="p">,</span> <span class="mi">11</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="nb">int</span><span class="p">)}</span>
|
||||
<span class="n">hdx</span> <span class="o">=</span> <span class="n">MedianEstimator</span><span class="p">(</span><span class="n">base_quantifier</span><span class="o">=</span><span class="n">dmx</span><span class="p">,</span> <span class="n">param_grid</span><span class="o">=</span><span class="n">nbins</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="n">n_jobs</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">hdx</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">__get_distributions</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__get_distributions</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="n">histograms</span> <span class="o">=</span> <span class="p">[]</span>
|
||||
<span class="k">for</span> <span class="n">feat_idx</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">nfeats</span><span class="p">):</span>
|
||||
<span class="n">feature</span> <span class="o">=</span> <span class="n">X</span><span class="p">[:,</span> <span class="n">feat_idx</span><span class="p">]</span>
|
||||
|
|
@ -172,7 +499,9 @@
|
|||
|
||||
<span class="k">return</span> <span class="n">distributions</span>
|
||||
|
||||
<div class="viewcode-block" id="DMx.fit"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx.fit">[docs]</a> <span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="DMx.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Generates the validation distributions out of the training data (covariates).</span>
|
||||
<span class="sd"> The validation distributions have shape `(n, nfeats, nbins)`, with `n` the number of classes, `nfeats`</span>
|
||||
|
|
@ -181,46 +510,300 @@
|
|||
<span class="sd"> training data labelled with class `i`; while `dij = di[j]` is the discrete distribution for feature j in</span>
|
||||
<span class="sd"> training data labelled with class `i`, and `dij[k]` is the fraction of instances with a value in the `k`-th bin.</span>
|
||||
|
||||
<span class="sd"> :param data: the training set</span>
|
||||
<span class="sd"> :param X: array-like of shape `(n_samples, n_features)`, the training instances</span>
|
||||
<span class="sd"> :param y: array-like of shape `(n_samples,)`, the labels</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">X</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">data</span><span class="o">.</span><span class="n">Xy</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">nfeats</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">feat_ranges</span> <span class="o">=</span> <span class="n">_get_features_range</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">classes</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>
|
||||
<span class="n">y</span> <span class="o">=</span> <span class="n">_labels_to_indices</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">classes</span><span class="p">)</span>
|
||||
<span class="n">n_classes</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="n">classes</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">validation_distribution</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span>
|
||||
<span class="p">[</span><span class="bp">self</span><span class="o">.</span><span class="n">__get_distributions</span><span class="p">(</span><span class="n">X</span><span class="p">[</span><span class="n">y</span><span class="o">==</span><span class="n">cat</span><span class="p">])</span> <span class="k">for</span> <span class="n">cat</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">data</span><span class="o">.</span><span class="n">n_classes</span><span class="p">)]</span>
|
||||
<span class="p">[</span><span class="bp">self</span><span class="o">.</span><span class="n">__get_distributions</span><span class="p">(</span><span class="n">X</span><span class="p">[</span><span class="n">y</span><span class="o">==</span><span class="n">cat</span><span class="p">])</span> <span class="k">for</span> <span class="n">cat</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">n_classes</span><span class="p">)]</span>
|
||||
<span class="p">)</span>
|
||||
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<div class="viewcode-block" id="DMx.quantify"><a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx.quantify">[docs]</a> <span class="k">def</span> <span class="nf">quantify</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="DMx.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.DMx.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Searches for the mixture model parameter (the sought prevalence values) that yields a validation distribution</span>
|
||||
<span class="sd"> (the mixture) that best matches the test distribution, in terms of the divergence measure of choice.</span>
|
||||
<span class="sd"> The matching is computed as the average dissimilarity (in terms of the dissimilarity measure of choice)</span>
|
||||
<span class="sd"> between all feature-specific discrete distributions.</span>
|
||||
|
||||
<span class="sd"> :param instances: instances in the sample</span>
|
||||
<span class="sd"> :param X: instances in the sample</span>
|
||||
<span class="sd"> :return: a vector of class prevalence estimates</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">assert</span> <span class="n">instances</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">nfeats</span><span class="p">,</span> <span class="sa">f</span><span class="s1">'wrong shape; expected </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">nfeats</span><span class="si">}</span><span class="s1">, found </span><span class="si">{</span><span class="n">instances</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="si">}</span><span class="s1">'</span>
|
||||
<span class="k">assert</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">nfeats</span><span class="p">,</span> <span class="sa">f</span><span class="s1">'wrong shape; expected </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">nfeats</span><span class="si">}</span><span class="s1">, found </span><span class="si">{</span><span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="si">}</span><span class="s1">'</span>
|
||||
|
||||
<span class="n">test_distribution</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">__get_distributions</span><span class="p">(</span><span class="n">instances</span><span class="p">)</span>
|
||||
<span class="n">test_distribution</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">__get_distributions</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">divergence</span> <span class="o">=</span> <span class="n">get_divergence</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">divergence</span><span class="p">)</span>
|
||||
<span class="n">n_classes</span><span class="p">,</span> <span class="n">n_feats</span><span class="p">,</span> <span class="n">nbins</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">validation_distribution</span><span class="o">.</span><span class="n">shape</span>
|
||||
<span class="k">def</span> <span class="nf">loss</span><span class="p">(</span><span class="n">prev</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">loss</span><span class="p">(</span><span class="n">prev</span><span class="p">):</span>
|
||||
<span class="n">prev</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">expand_dims</span><span class="p">(</span><span class="n">prev</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
|
||||
<span class="n">mixture_distribution</span> <span class="o">=</span> <span class="p">(</span><span class="n">prev</span> <span class="o">@</span> <span class="bp">self</span><span class="o">.</span><span class="n">validation_distribution</span><span class="o">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">n_classes</span><span class="p">,</span><span class="o">-</span><span class="mi">1</span><span class="p">))</span><span class="o">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">n_feats</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">)</span>
|
||||
<span class="n">divs</span> <span class="o">=</span> <span class="p">[</span><span class="n">divergence</span><span class="p">(</span><span class="n">test_distribution</span><span class="p">[</span><span class="n">feat</span><span class="p">],</span> <span class="n">mixture_distribution</span><span class="p">[</span><span class="n">feat</span><span class="p">])</span> <span class="k">for</span> <span class="n">feat</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">n_feats</span><span class="p">)]</span>
|
||||
<span class="k">return</span> <span class="n">np</span><span class="o">.</span><span class="n">mean</span><span class="p">(</span><span class="n">divs</span><span class="p">)</span>
|
||||
|
||||
<span class="k">return</span> <span class="n">F</span><span class="o">.</span><span class="n">argmin_prevalence</span><span class="p">(</span><span class="n">loss</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">,</span> <span class="n">method</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">search</span><span class="p">)</span></div></div>
|
||||
<span class="k">return</span> <span class="n">F</span><span class="o">.</span><span class="n">argmin_prevalence</span><span class="p">(</span><span class="n">loss</span><span class="p">,</span> <span class="n">n_classes</span><span class="p">,</span> <span class="n">method</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">search</span><span class="p">)</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">_get_features_range</span><span class="p">(</span><span class="n">X</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="EDx">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.EDx">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">EDx</span><span class="p">(</span><span class="n">_EnergyDistanceCore</span><span class="p">,</span> <span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Energy Distance x (EDx), a covariate-space distribution-matching</span>
|
||||
<span class="sd"> quantifier based on energy distance.</span>
|
||||
|
||||
<span class="sd"> EDx is the classifier-free counterpart of :class:`quapy.method.aggregative.EDy`.</span>
|
||||
<span class="sd"> Instead of representing each class through posterior-probability vectors, it</span>
|
||||
<span class="sd"> represents each class by the cloud of raw feature vectors observed in the</span>
|
||||
<span class="sd"> training set and estimates the test prevalence vector by solving the same</span>
|
||||
<span class="sd"> energy-distance quadratic program directly in feature space.</span>
|
||||
|
||||
<span class="sd"> This implementation works for binary and multiclass single-label</span>
|
||||
<span class="sd"> quantification and relies on the optional ``quadprog`` dependency. The</span>
|
||||
<span class="sd"> current QuaPy adaptation shares its numerical core with EDy and keeps</span>
|
||||
<span class="sd"> credit to the original implementation available in</span>
|
||||
<span class="sd"> `quantificationlib <https://github.com/AICGijon/quantificationlib>`_.</span>
|
||||
|
||||
<span class="sd"> The formulation follows the same references as EDy, namely:</span>
|
||||
|
||||
<span class="sd"> * Alberto Castaño, Laura Morán-Fernández, Jaime Alonso,</span>
|
||||
<span class="sd"> Verónica Bolón-Canedo, Amparo Alonso-Betanzos, and Juan José del Coz.</span>
|
||||
<span class="sd"> *An analysis of quantification methods based on matching distributions*.</span>
|
||||
<span class="sd"> * Hideko Kawakubo, Marthinus Christoffel du Plessis, and Masashi Sugiyama</span>
|
||||
<span class="sd"> (2016). *Computationally efficient class-prior estimation under class</span>
|
||||
<span class="sd"> balance change using energy distance*. IEICE Transactions on Information</span>
|
||||
<span class="sd"> and Systems, 99(1):176-186.</span>
|
||||
|
||||
<span class="sd"> :param distance: distance used to compare feature vectors. Valid string</span>
|
||||
<span class="sd"> aliases are ``'manhattan'`` (default) and ``'euclidean'``; a custom</span>
|
||||
<span class="sd"> callable compatible with pairwise-distance signatures can also be used</span>
|
||||
<span class="sd"> :param n_jobs: number of parallel workers (default ``None``, meaning the</span>
|
||||
<span class="sd"> value is taken from the environment)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">distance</span><span class="p">:</span> <span class="n">Union</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="n">Callable</span><span class="p">]</span> <span class="o">=</span> <span class="s1">'manhattan'</span><span class="p">,</span> <span class="n">n_jobs</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">distance</span> <span class="o">=</span> <span class="n">distance</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span> <span class="o">=</span> <span class="n">qp</span><span class="o">.</span><span class="n">_get_njobs</span><span class="p">(</span><span class="n">n_jobs</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_features_in_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">train_distrib_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">train_n_cls_i_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">K_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">G_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">C_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">b_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">a_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
|
||||
<div class="viewcode-block" id="EDx.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.EDx.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Fit class-conditional feature-space distributions from training data."""</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_check_ed_init_parameters</span><span class="p">()</span>
|
||||
<span class="n">labels</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">labels</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">n_features_in_</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="n">train_distrib</span> <span class="o">=</span> <span class="p">[</span><span class="n">X</span><span class="p">[</span><span class="n">labels</span> <span class="o">==</span> <span class="n">class_</span><span class="p">]</span> <span class="k">for</span> <span class="n">class_</span> <span class="ow">in</span> <span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">]</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">_fit_energy_model</span><span class="p">(</span><span class="n">train_distrib</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="EDx.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.EDx.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Estimate class prevalences for a test sample of raw instances."""</span>
|
||||
<span class="k">assert</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">==</span> <span class="bp">self</span><span class="o">.</span><span class="n">n_features_in_</span><span class="p">,</span> <span class="p">(</span>
|
||||
<span class="sa">f</span><span class="s1">'wrong shape; expected </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">n_features_in_</span><span class="si">}</span><span class="s1">, found </span><span class="si">{</span><span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span><span class="si">}</span><span class="s1">'</span>
|
||||
<span class="p">)</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">_predict_energy</span><span class="p">(</span><span class="n">X</span><span class="p">)</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="ReadMe">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.ReadMe">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">ReadMe</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">WithConfidenceABC</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> ReadMe is a non-aggregative quantification system proposed by</span>
|
||||
<span class="sd"> `Daniel Hopkins and Gary King, 2007. A method of automated nonparametric content analysis for</span>
|
||||
<span class="sd"> social science. American Journal of Political Science, 54(1):229–247.</span>
|
||||
<span class="sd"> <https://onlinelibrary.wiley.com/doi/abs/10.1111/j.1540-5907.2009.00428.x>`_.</span>
|
||||
<span class="sd"> The idea is to estimate `Q(Y=i)` directly from:</span>
|
||||
|
||||
<span class="sd"> :math:`Q(X)=\\sum_{i=1} Q(X|Y=i) Q(Y=i)`</span>
|
||||
|
||||
<span class="sd"> via least-squares regression, i.e., without incurring the cost of computing posterior probabilities.</span>
|
||||
<span class="sd"> However, this poses a very difficult representation in which the vector `Q(X)` and the matrix `Q(X|Y=i)`</span>
|
||||
<span class="sd"> can be of very high dimensions. In order to render the problem tracktable, ReadMe performs bagging in</span>
|
||||
<span class="sd"> the feature space. ReadMe also combines bagging with bootstrap in order to derive confidence intervals</span>
|
||||
<span class="sd"> around point estimations.</span>
|
||||
|
||||
<span class="sd"> We use the same default parameters as in the official</span>
|
||||
<span class="sd"> `R implementation <https://github.com/iqss-research/ReadMeV1/blob/master/R/prototype.R>`_.</span>
|
||||
|
||||
<span class="sd"> :param prob_model: str ('naive', or 'full'), selects the modality in which the probabilities `Q(X)` and</span>
|
||||
<span class="sd"> `Q(X|Y)` are to be modelled. Options include "full", which corresponds to the original formulation of</span>
|
||||
<span class="sd"> ReadMe, in which X is constrained to be a binary matrix (e.g., of term presence/absence) and in which</span>
|
||||
<span class="sd"> `Q(X)` and `Q(X|Y)` are modelled, respectively, as matrices of `(2^K, 1)` and `(2^K, n)` values, where</span>
|
||||
<span class="sd"> `K` is the number of columns in the data matrix (i.e., `bagging_range`), and `n` is the number of classes.</span>
|
||||
<span class="sd"> Of course, this approach is computationally prohibited for large `K`, so the authors advised against computing it</span>
|
||||
<span class="sd"> for matrices with `K>25` (although we recommend even smaller values of `K`). A much faster model is "naive", which</span>
|
||||
<span class="sd"> considers the `Q(X)` and `Q(X|Y)` be multinomial distributions under the `bag-of-words` perspective. In this</span>
|
||||
<span class="sd"> case, `bagging_range` can be set to much larger values. Default is "full" (i.e., original ReadMe behavior).</span>
|
||||
<span class="sd"> :param bootstrap_trials: int, number of bootstrap trials (default 300)</span>
|
||||
<span class="sd"> :param bagging_trials: int, number of bagging trials (default 300)</span>
|
||||
<span class="sd"> :param bagging_range: int, number of features to keep for each bagging trial (default 15)</span>
|
||||
<span class="sd"> :param confidence_level: float, a value in (0,1) reflecting the desired confidence level (default 0.95)</span>
|
||||
<span class="sd"> :param region: str in 'intervals', 'ellipse', 'ellipse-clr'; indicates the preferred method for</span>
|
||||
<span class="sd"> defining the confidence region (see :class:`WithConfidenceABC`)</span>
|
||||
<span class="sd"> :param bonferroni: bool (default False), whether to apply Bonferroni correction when</span>
|
||||
<span class="sd"> `region='intervals'`. This parameter has no effect for ellipse-based regions.</span>
|
||||
<span class="sd"> :param random_state: int or None, allows replicability (default None)</span>
|
||||
<span class="sd"> :param verbose: bool, whether to display information during the process (default False)</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="n">MAX_FEATURES_FOR_EMPIRICAL_ESTIMATION</span> <span class="o">=</span> <span class="mi">25</span>
|
||||
<span class="n">PROBABILISTIC_MODELS</span> <span class="o">=</span> <span class="p">[</span><span class="s2">"naive"</span><span class="p">,</span> <span class="s2">"full"</span><span class="p">]</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="n">prob_model</span><span class="o">=</span><span class="s2">"full"</span><span class="p">,</span>
|
||||
<span class="n">bootstrap_trials</span><span class="o">=</span><span class="mi">300</span><span class="p">,</span>
|
||||
<span class="n">bagging_trials</span><span class="o">=</span><span class="mi">300</span><span class="p">,</span>
|
||||
<span class="n">bagging_range</span><span class="o">=</span><span class="mi">15</span><span class="p">,</span>
|
||||
<span class="n">confidence_level</span><span class="o">=</span><span class="mf">0.95</span><span class="p">,</span>
|
||||
<span class="n">region</span><span class="o">=</span><span class="s1">'intervals'</span><span class="p">,</span>
|
||||
<span class="n">bonferroni</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span>
|
||||
<span class="n">random_state</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span>
|
||||
<span class="n">verbose</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">assert</span> <span class="n">prob_model</span> <span class="ow">in</span> <span class="n">ReadMe</span><span class="o">.</span><span class="n">PROBABILISTIC_MODELS</span><span class="p">,</span> \
|
||||
<span class="sa">f</span><span class="s1">'unknown </span><span class="si">{</span><span class="n">prob_model</span><span class="si">=}</span><span class="s1">, valid ones are </span><span class="si">{</span><span class="n">ReadMe</span><span class="o">.</span><span class="n">PROBABILISTIC_MODELS</span><span class="si">=}</span><span class="s1">'</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">prob_model</span> <span class="o">=</span> <span class="n">prob_model</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">bootstrap_trials</span> <span class="o">=</span> <span class="n">bootstrap_trials</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">bagging_trials</span> <span class="o">=</span> <span class="n">bagging_trials</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">bagging_range</span> <span class="o">=</span> <span class="n">bagging_range</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">confidence_level</span> <span class="o">=</span> <span class="n">confidence_level</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">region</span> <span class="o">=</span> <span class="n">region</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">bonferroni</span> <span class="o">=</span> <span class="n">bonferroni</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">random_state</span> <span class="o">=</span> <span class="n">random_state</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span>
|
||||
|
||||
<div class="viewcode-block" id="ReadMe.fit">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.ReadMe.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_check_matrix</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">rng</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">random</span><span class="o">.</span><span class="n">default_rng</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">random_state</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">classes_</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">unique</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>
|
||||
|
||||
<span class="n">Xsize</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
|
||||
|
||||
<span class="c1"># Bootstrap loop</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">Xboots</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">yboots</span> <span class="o">=</span> <span class="p">[],</span> <span class="p">[]</span>
|
||||
<span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">bootstrap_trials</span><span class="p">):</span>
|
||||
<span class="n">idx</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">rng</span><span class="o">.</span><span class="n">choice</span><span class="p">(</span><span class="n">Xsize</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">Xsize</span><span class="p">,</span> <span class="n">replace</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">Xboots</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">X</span><span class="p">[</span><span class="n">idx</span><span class="p">])</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">yboots</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">y</span><span class="p">[</span><span class="n">idx</span><span class="p">])</span>
|
||||
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="ReadMe.predict_conf">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.ReadMe.predict_conf">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict_conf</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">confidence_level</span><span class="o">=</span><span class="kc">None</span><span class="p">)</span> <span class="o">-></span> <span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">ndarray</span><span class="p">,</span> <span class="n">ConfidenceRegionABC</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_check_matrix</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="n">confidence_level</span> <span class="ow">is</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="n">confidence_level</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">confidence_level</span>
|
||||
|
||||
<span class="n">n_features</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="n">boots_prevalences</span> <span class="o">=</span> <span class="p">[]</span>
|
||||
<span class="k">for</span> <span class="n">Xboots</span><span class="p">,</span> <span class="n">yboots</span> <span class="ow">in</span> <span class="n">tqdm</span><span class="p">(</span>
|
||||
<span class="nb">zip</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">Xboots</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">yboots</span><span class="p">),</span>
|
||||
<span class="n">desc</span><span class="o">=</span><span class="s1">'bootstrap predictions'</span><span class="p">,</span> <span class="n">total</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">bootstrap_trials</span><span class="p">,</span> <span class="n">disable</span><span class="o">=</span><span class="ow">not</span> <span class="bp">self</span><span class="o">.</span><span class="n">verbose</span>
|
||||
<span class="p">):</span>
|
||||
<span class="n">bagging_estimates</span> <span class="o">=</span> <span class="p">[]</span>
|
||||
<span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">bagging_trials</span><span class="p">):</span>
|
||||
<span class="n">feat_idx</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">rng</span><span class="o">.</span><span class="n">choice</span><span class="p">(</span><span class="n">n_features</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">bagging_range</span><span class="p">,</span> <span class="n">replace</span><span class="o">=</span><span class="kc">False</span><span class="p">)</span>
|
||||
<span class="n">Xboots_bagging</span> <span class="o">=</span> <span class="n">Xboots</span><span class="p">[:,</span> <span class="n">feat_idx</span><span class="p">]</span>
|
||||
<span class="n">X_boots_bagging</span> <span class="o">=</span> <span class="n">X</span><span class="p">[:,</span> <span class="n">feat_idx</span><span class="p">]</span>
|
||||
<span class="n">bagging_prev</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_quantify_iteration</span><span class="p">(</span><span class="n">Xboots_bagging</span><span class="p">,</span> <span class="n">yboots</span><span class="p">,</span> <span class="n">X_boots_bagging</span><span class="p">)</span>
|
||||
<span class="n">bagging_estimates</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">bagging_prev</span><span class="p">)</span>
|
||||
|
||||
<span class="n">boots_prevalences</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">mean</span><span class="p">(</span><span class="n">bagging_estimates</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">))</span>
|
||||
|
||||
<span class="n">conf</span> <span class="o">=</span> <span class="n">WithConfidenceABC</span><span class="o">.</span><span class="n">construct_region</span><span class="p">(</span><span class="n">boots_prevalences</span><span class="p">,</span> <span class="n">confidence_level</span><span class="p">,</span> <span class="n">method</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">region</span><span class="p">,</span> <span class="n">bonferroni</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">bonferroni</span><span class="p">)</span>
|
||||
<span class="n">prev_estim</span> <span class="o">=</span> <span class="n">conf</span><span class="o">.</span><span class="n">point_estimate</span><span class="p">()</span>
|
||||
|
||||
<span class="k">return</span> <span class="n">prev_estim</span><span class="p">,</span> <span class="n">conf</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="ReadMe.predict">
|
||||
<a class="viewcode-back" href="../../../quapy.method.html#quapy.method.non_aggregative.ReadMe.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="n">prev_estim</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">predict_conf</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">prev_estim</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_quantify_iteration</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">Xtr</span><span class="p">,</span> <span class="n">ytr</span><span class="p">,</span> <span class="n">Xte</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Single ReadMe estimate."""</span>
|
||||
<span class="n">PX_given_Y</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">([</span><span class="bp">self</span><span class="o">.</span><span class="n">_compute_P</span><span class="p">(</span><span class="n">Xtr</span><span class="p">[</span><span class="n">ytr</span> <span class="o">==</span> <span class="n">c</span><span class="p">])</span> <span class="k">for</span> <span class="n">i</span><span class="p">,</span><span class="n">c</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">classes_</span><span class="p">)])</span>
|
||||
<span class="n">PX</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_compute_P</span><span class="p">(</span><span class="n">Xte</span><span class="p">)</span>
|
||||
|
||||
<span class="n">res</span> <span class="o">=</span> <span class="n">lsq_linear</span><span class="p">(</span><span class="n">A</span><span class="o">=</span><span class="n">PX_given_Y</span><span class="o">.</span><span class="n">T</span><span class="p">,</span> <span class="n">b</span><span class="o">=</span><span class="n">PX</span><span class="p">,</span> <span class="n">bounds</span><span class="o">=</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="mi">1</span><span class="p">))</span>
|
||||
<span class="n">pY</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">maximum</span><span class="p">(</span><span class="n">res</span><span class="o">.</span><span class="n">x</span><span class="p">,</span> <span class="mi">0</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">pY</span> <span class="o">/</span> <span class="n">pY</span><span class="o">.</span><span class="n">sum</span><span class="p">()</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_check_matrix</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""the "full" model requires estimating empirical distributions; due to the high computational cost,</span>
|
||||
<span class="sd"> this function is only made available for binary matrices"""</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">prob_model</span> <span class="o">==</span> <span class="s1">'full'</span> <span class="ow">and</span> <span class="ow">not</span> <span class="bp">self</span><span class="o">.</span><span class="n">_is_binary_matrix</span><span class="p">(</span><span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="s1">'the empirical distribution can only be computed efficiently on binary matrices'</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_is_binary_matrix</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="n">data</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">data</span> <span class="k">if</span> <span class="n">sparse</span><span class="o">.</span><span class="n">issparse</span><span class="p">(</span><span class="n">X</span><span class="p">)</span> <span class="k">else</span> <span class="n">X</span>
|
||||
<span class="k">return</span> <span class="n">np</span><span class="o">.</span><span class="n">all</span><span class="p">((</span><span class="n">data</span> <span class="o">==</span> <span class="mi">0</span><span class="p">)</span> <span class="o">|</span> <span class="p">(</span><span class="n">data</span> <span class="o">==</span> <span class="mi">1</span><span class="p">))</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_compute_P</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">prob_model</span> <span class="o">==</span> <span class="s1">'naive'</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">_multinomial_distribution</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">elif</span> <span class="bp">self</span><span class="o">.</span><span class="n">prob_model</span> <span class="o">==</span> <span class="s1">'full'</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">_empirical_distribution</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="sa">f</span><span class="s1">'unknown </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">prob_model</span><span class="si">}</span><span class="s1">; valid ones are </span><span class="si">{</span><span class="n">ReadMe</span><span class="o">.</span><span class="n">PROBABILISTIC_MODELS</span><span class="si">=}</span><span class="s1">'</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_empirical_distribution</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
|
||||
<span class="k">if</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span> <span class="o">></span> <span class="bp">self</span><span class="o">.</span><span class="n">MAX_FEATURES_FOR_EMPIRICAL_ESTIMATION</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="sa">f</span><span class="s1">'the empirical distribution can only be computed efficiently for dimensions '</span>
|
||||
<span class="sa">f</span><span class="s1">'less or equal than </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">MAX_FEATURES_FOR_EMPIRICAL_ESTIMATION</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
|
||||
<span class="c1"># we first convert every binary row (e.g., 0 0 1 0 1) into the equivalent number (e.g., 5);</span>
|
||||
<span class="c1"># this will speed up subsequent comparisons a lot</span>
|
||||
<span class="n">K</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="n">binary_powers</span> <span class="o">=</span> <span class="mi">1</span> <span class="o"><<</span> <span class="n">np</span><span class="o">.</span><span class="n">arange</span><span class="p">(</span><span class="n">K</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">)</span> <span class="c1"># (2^K, ..., 32, 16, 8, 4, 2, 1)</span>
|
||||
<span class="n">X_as_binary_numbers</span> <span class="o">=</span> <span class="n">X</span> <span class="o">@</span> <span class="n">binary_powers</span> <span class="c1"># e.g., [0 0 1 0 1] @ [16, 8, 4, 2, 1] = 5</span>
|
||||
|
||||
<span class="c1"># count occurrences and compute probs</span>
|
||||
<span class="n">counts</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">bincount</span><span class="p">(</span><span class="n">X_as_binary_numbers</span><span class="p">,</span> <span class="n">minlength</span><span class="o">=</span><span class="mi">2</span> <span class="o">**</span> <span class="n">K</span><span class="p">)</span><span class="o">.</span><span class="n">astype</span><span class="p">(</span><span class="nb">float</span><span class="p">)</span>
|
||||
<span class="n">probs</span> <span class="o">=</span> <span class="n">counts</span> <span class="o">/</span> <span class="n">counts</span><span class="o">.</span><span class="n">sum</span><span class="p">()</span>
|
||||
|
||||
<span class="k">return</span> <span class="n">probs</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_multinomial_distribution</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="n">PX</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span><span class="n">X</span><span class="o">.</span><span class="n">sum</span><span class="p">(</span><span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">))</span>
|
||||
<span class="n">PX</span> <span class="o">=</span> <span class="n">normalize</span><span class="p">(</span><span class="n">PX</span><span class="p">,</span> <span class="n">norm</span><span class="o">=</span><span class="s1">'l1'</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">PX</span><span class="o">.</span><span class="n">ravel</span><span class="p">()</span></div>
|
||||
|
||||
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_get_features_range</span><span class="p">(</span><span class="n">X</span><span class="p">):</span>
|
||||
<span class="n">feat_ranges</span> <span class="o">=</span> <span class="p">[]</span>
|
||||
<span class="n">ncols</span> <span class="o">=</span> <span class="n">X</span><span class="o">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="k">for</span> <span class="n">col_idx</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">ncols</span><span class="p">):</span>
|
||||
|
|
@ -233,34 +816,82 @@
|
|||
<span class="c1"># aliases</span>
|
||||
<span class="c1">#---------------------------------------------------------------</span>
|
||||
|
||||
|
||||
<span class="n">HDx</span> <span class="o">=</span> <span class="n">DMx</span><span class="o">.</span><span class="n">HDx</span>
|
||||
<span class="n">DistributionMatchingX</span> <span class="o">=</span> <span class="n">DMx</span>
|
||||
<span class="n">EnergyDistanceX</span> <span class="o">=</span> <span class="n">EDx</span>
|
||||
<span class="n">HellingerDistanceX</span> <span class="o">=</span> <span class="n">HDx</span>
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,122 +1,437 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.model_selection — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.model_selection — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../" id="documentation_options" src="../../_static/documentation_options.js"></script>
|
||||
<script src="../../_static/jquery.js"></script>
|
||||
<script src="../../_static/underscore.js"></script>
|
||||
<script src="../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../_static/doctools.js"></script>
|
||||
<script src="../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/model_selection';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../_static/searchtools.js"></script>
|
||||
<script src="../../_static/language_data.js"></script>
|
||||
<script src="../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.model_selection</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.model_selection</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.model_selection</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">itertools</span>
|
||||
<span class="kn">import</span> <span class="nn">signal</span>
|
||||
<span class="kn">from</span> <span class="nn">copy</span> <span class="kn">import</span> <span class="n">deepcopy</span>
|
||||
<span class="kn">from</span> <span class="nn">enum</span> <span class="kn">import</span> <span class="n">Enum</span>
|
||||
<span class="kn">from</span> <span class="nn">typing</span> <span class="kn">import</span> <span class="n">Union</span><span class="p">,</span> <span class="n">Callable</span>
|
||||
<span class="kn">from</span> <span class="nn">functools</span> <span class="kn">import</span> <span class="n">wraps</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">itertools</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">logging</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">signal</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">copy</span><span class="w"> </span><span class="kn">import</span> <span class="n">deepcopy</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">enum</span><span class="w"> </span><span class="kn">import</span> <span class="n">Enum</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">typing</span><span class="w"> </span><span class="kn">import</span> <span class="n">Union</span><span class="p">,</span> <span class="n">Callable</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">functools</span><span class="w"> </span><span class="kn">import</span> <span class="n">wraps</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">sklearn</span> <span class="kn">import</span> <span class="n">clone</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">sklearn</span><span class="w"> </span><span class="kn">import</span> <span class="n">clone</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">quapy</span> <span class="k">as</span> <span class="nn">qp</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy</span> <span class="kn">import</span> <span class="n">evaluation</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.protocol</span> <span class="kn">import</span> <span class="n">AbstractProtocol</span><span class="p">,</span> <span class="n">OnLabelledCollectionProtocol</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.data.base</span> <span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.method.aggregative</span> <span class="kn">import</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">AggregativeQuantifier</span>
|
||||
<span class="kn">from</span> <span class="nn">quapy.util</span> <span class="kn">import</span> <span class="n">timeout</span>
|
||||
<span class="kn">from</span> <span class="nn">time</span> <span class="kn">import</span> <span class="n">time</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="kn">import</span> <span class="n">evaluation</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.protocol</span><span class="w"> </span><span class="kn">import</span> <span class="n">AbstractProtocol</span><span class="p">,</span> <span class="n">OnLabelledCollectionProtocol</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.data.base</span><span class="w"> </span><span class="kn">import</span> <span class="n">LabelledCollection</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.method.aggregative</span><span class="w"> </span><span class="kn">import</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">AggregativeQuantifier</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">quapy.util</span><span class="w"> </span><span class="kn">import</span> <span class="n">timeout</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">time</span><span class="w"> </span><span class="kn">import</span> <span class="n">time</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="Status"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.Status">[docs]</a><span class="k">class</span> <span class="nc">Status</span><span class="p">(</span><span class="n">Enum</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="Status">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.Status">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">Status</span><span class="p">(</span><span class="n">Enum</span><span class="p">):</span>
|
||||
<span class="n">SUCCESS</span> <span class="o">=</span> <span class="mi">1</span>
|
||||
<span class="n">TIMEOUT</span> <span class="o">=</span> <span class="mi">2</span>
|
||||
<span class="n">INVALID</span> <span class="o">=</span> <span class="mi">3</span>
|
||||
<span class="n">ERROR</span> <span class="o">=</span> <span class="mi">4</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="ConfigStatus"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.ConfigStatus">[docs]</a><span class="k">class</span> <span class="nc">ConfigStatus</span><span class="p">:</span>
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">params</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">msg</span><span class="o">=</span><span class="s1">''</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="ConfigStatus">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.ConfigStatus">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">ConfigStatus</span><span class="p">:</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">params</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">msg</span><span class="o">=</span><span class="s1">''</span><span class="p">):</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">params</span> <span class="o">=</span> <span class="n">params</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">status</span> <span class="o">=</span> <span class="n">status</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">msg</span> <span class="o">=</span> <span class="n">msg</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__str__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__str__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="sa">f</span><span class="s1">':params:</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">params</span><span class="si">}</span><span class="s1"> :status:</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">status</span><span class="si">}</span><span class="s1"> '</span> <span class="o">+</span> <span class="bp">self</span><span class="o">.</span><span class="n">msg</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__repr__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__repr__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="nb">str</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="ConfigStatus.success"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.ConfigStatus.success">[docs]</a> <span class="k">def</span> <span class="nf">success</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="ConfigStatus.success">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.ConfigStatus.success">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">success</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">status</span> <span class="o">==</span> <span class="n">Status</span><span class="o">.</span><span class="n">SUCCESS</span></div>
|
||||
|
||||
<div class="viewcode-block" id="ConfigStatus.failed"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.ConfigStatus.failed">[docs]</a> <span class="k">def</span> <span class="nf">failed</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">status</span> <span class="o">!=</span> <span class="n">Status</span><span class="o">.</span><span class="n">SUCCESS</span></div></div>
|
||||
|
||||
<div class="viewcode-block" id="ConfigStatus.failed">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.ConfigStatus.failed">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">failed</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">status</span> <span class="o">!=</span> <span class="n">Status</span><span class="o">.</span><span class="n">SUCCESS</span></div>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ">[docs]</a><span class="k">class</span> <span class="nc">GridSearchQ</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">GridSearchQ</span><span class="p">(</span><span class="n">BaseQuantifier</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Grid Search optimization targeting a quantification-oriented metric.</span>
|
||||
|
||||
<span class="sd"> Optimizes the hyperparameters of a quantification method, based on an evaluation method and on an evaluation</span>
|
||||
|
|
@ -139,7 +454,7 @@
|
|||
<span class="sd"> :param verbose: set to True to get information through the stdout</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span>
|
||||
<span class="n">model</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span>
|
||||
<span class="n">param_grid</span><span class="p">:</span> <span class="nb">dict</span><span class="p">,</span>
|
||||
<span class="n">protocol</span><span class="p">:</span> <span class="n">AbstractProtocol</span><span class="p">,</span>
|
||||
|
|
@ -158,14 +473,14 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span> <span class="o">=</span> <span class="n">qp</span><span class="o">.</span><span class="n">_get_njobs</span><span class="p">(</span><span class="n">n_jobs</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">raise_errors</span> <span class="o">=</span> <span class="n">raise_errors</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">verbose</span> <span class="o">=</span> <span class="n">verbose</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">__check_error</span><span class="p">(</span><span class="n">error</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">__check_error_measure</span><span class="p">(</span><span class="n">error</span><span class="p">)</span>
|
||||
<span class="k">assert</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">protocol</span><span class="p">,</span> <span class="n">AbstractProtocol</span><span class="p">),</span> <span class="s1">'unknown protocol'</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_sout</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">msg</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_sout</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">msg</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">verbose</span><span class="p">:</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="sa">f</span><span class="s1">'[</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">:</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">model</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">]: </span><span class="si">{</span><span class="n">msg</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="n">logging</span><span class="o">.</span><span class="n">getLogger</span><span class="p">(</span><span class="vm">__name__</span><span class="p">)</span><span class="o">.</span><span class="n">info</span><span class="p">(</span><span class="sa">f</span><span class="s1">'[</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">:</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">model</span><span class="o">.</span><span class="vm">__class__</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1">]: </span><span class="si">{</span><span class="n">msg</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">__check_error</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">error</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">__check_error_measure</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">error</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">error</span> <span class="ow">in</span> <span class="n">qp</span><span class="o">.</span><span class="n">error</span><span class="o">.</span><span class="n">QUANTIFICATION_ERROR</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">error</span> <span class="o">=</span> <span class="n">error</span>
|
||||
<span class="k">elif</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">error</span><span class="p">,</span> <span class="nb">str</span><span class="p">):</span>
|
||||
|
|
@ -176,26 +491,27 @@
|
|||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="sa">f</span><span class="s1">'unexpected error type; must either be a callable function or a str representing</span><span class="se">\n</span><span class="s1">'</span>
|
||||
<span class="sa">f</span><span class="s1">'the name of an error function in </span><span class="si">{</span><span class="n">qp</span><span class="o">.</span><span class="n">error</span><span class="o">.</span><span class="n">QUANTIFICATION_ERROR_NAMES</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_prepare_classifier</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">cls_params</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_prepare_classifier</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">cls_params</span><span class="p">):</span>
|
||||
<span class="n">model</span> <span class="o">=</span> <span class="n">deepcopy</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">job</span><span class="p">(</span><span class="n">cls_params</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">job</span><span class="p">(</span><span class="n">cls_params</span><span class="p">):</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">set_params</span><span class="p">(</span><span class="o">**</span><span class="n">cls_params</span><span class="p">)</span>
|
||||
<span class="n">predictions</span> <span class="o">=</span> <span class="n">model</span><span class="o">.</span><span class="n">classifier_fit_predict</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_training</span><span class="p">)</span>
|
||||
<span class="n">predictions</span> <span class="o">=</span> <span class="n">model</span><span class="o">.</span><span class="n">classifier_fit_predict</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_training_X</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">_training_y</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">predictions</span>
|
||||
|
||||
<span class="n">predictions</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">took</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_error_handler</span><span class="p">(</span><span class="n">job</span><span class="p">,</span> <span class="n">cls_params</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_sout</span><span class="p">(</span><span class="sa">f</span><span class="s1">'[classifier fit] hyperparams=</span><span class="si">{</span><span class="n">cls_params</span><span class="si">}</span><span class="s1"> [took </span><span class="si">{</span><span class="n">took</span><span class="si">:</span><span class="s1">.3f</span><span class="si">}</span><span class="s1">s]'</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">model</span><span class="p">,</span> <span class="n">predictions</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">took</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_prepare_aggregation</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">args</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_prepare_aggregation</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">args</span><span class="p">):</span>
|
||||
<span class="n">model</span><span class="p">,</span> <span class="n">predictions</span><span class="p">,</span> <span class="n">cls_took</span><span class="p">,</span> <span class="n">cls_params</span><span class="p">,</span> <span class="n">q_params</span> <span class="o">=</span> <span class="n">args</span>
|
||||
<span class="n">model</span> <span class="o">=</span> <span class="n">deepcopy</span><span class="p">(</span><span class="n">model</span><span class="p">)</span>
|
||||
<span class="n">params</span> <span class="o">=</span> <span class="p">{</span><span class="o">**</span><span class="n">cls_params</span><span class="p">,</span> <span class="o">**</span><span class="n">q_params</span><span class="p">}</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">job</span><span class="p">(</span><span class="n">q_params</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">job</span><span class="p">(</span><span class="n">q_params</span><span class="p">):</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">set_params</span><span class="p">(</span><span class="o">**</span><span class="n">q_params</span><span class="p">)</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">aggregation_fit</span><span class="p">(</span><span class="n">predictions</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">_training</span><span class="p">)</span>
|
||||
<span class="n">P</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">predictions</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">aggregation_fit</span><span class="p">(</span><span class="n">P</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
||||
<span class="n">score</span> <span class="o">=</span> <span class="n">evaluation</span><span class="o">.</span><span class="n">evaluate</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">protocol</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">protocol</span><span class="p">,</span> <span class="n">error_metric</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">error</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">score</span>
|
||||
|
||||
|
|
@ -203,12 +519,12 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">_print_status</span><span class="p">(</span><span class="n">params</span><span class="p">,</span> <span class="n">score</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">aggr_took</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">model</span><span class="p">,</span> <span class="n">params</span><span class="p">,</span> <span class="n">score</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="p">(</span><span class="n">cls_took</span><span class="o">+</span><span class="n">aggr_took</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_prepare_nonaggr_model</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">params</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_prepare_nonaggr_model</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">params</span><span class="p">):</span>
|
||||
<span class="n">model</span> <span class="o">=</span> <span class="n">deepcopy</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model</span><span class="p">)</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">job</span><span class="p">(</span><span class="n">params</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">job</span><span class="p">(</span><span class="n">params</span><span class="p">):</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">set_params</span><span class="p">(</span><span class="o">**</span><span class="n">params</span><span class="p">)</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_training</span><span class="p">)</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">_training_X</span><span class="p">,</span> <span class="bp">self</span><span class="o">.</span><span class="n">_training_y</span><span class="p">)</span>
|
||||
<span class="n">score</span> <span class="o">=</span> <span class="n">evaluation</span><span class="o">.</span><span class="n">evaluate</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">protocol</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">protocol</span><span class="p">,</span> <span class="n">error_metric</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">error</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">score</span>
|
||||
|
||||
|
|
@ -216,7 +532,7 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">_print_status</span><span class="p">(</span><span class="n">params</span><span class="p">,</span> <span class="n">score</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">took</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">model</span><span class="p">,</span> <span class="n">params</span><span class="p">,</span> <span class="n">score</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">took</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_break_down_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_break_down_fit</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Decides whether to break down the fit phase in two (classifier-fit followed by aggregation-fit).</span>
|
||||
<span class="sd"> In order to do so, some conditions should be met: a) the quantifier is of type aggregative,</span>
|
||||
|
|
@ -231,17 +547,19 @@
|
|||
<span class="k">return</span> <span class="kc">False</span>
|
||||
<span class="k">return</span> <span class="kc">True</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_compute_scores_aggregative</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">training</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_compute_scores_aggregative</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="c1"># break down the set of hyperparameters into two: classifier-specific, quantifier-specific</span>
|
||||
<span class="n">cls_configs</span><span class="p">,</span> <span class="n">q_configs</span> <span class="o">=</span> <span class="n">group_params</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">param_grid</span><span class="p">)</span>
|
||||
|
||||
<span class="c1"># train all classifiers and get the predictions</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_training</span> <span class="o">=</span> <span class="n">training</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_training_X</span> <span class="o">=</span> <span class="n">X</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_training_y</span> <span class="o">=</span> <span class="n">y</span>
|
||||
<span class="n">cls_outs</span> <span class="o">=</span> <span class="n">qp</span><span class="o">.</span><span class="n">util</span><span class="o">.</span><span class="n">parallel</span><span class="p">(</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_prepare_classifier</span><span class="p">,</span>
|
||||
<span class="n">cls_configs</span><span class="p">,</span>
|
||||
<span class="n">seed</span><span class="o">=</span><span class="n">qp</span><span class="o">.</span><span class="n">environ</span><span class="o">.</span><span class="n">get</span><span class="p">(</span><span class="s1">'_R_SEED'</span><span class="p">,</span> <span class="kc">None</span><span class="p">),</span>
|
||||
<span class="n">n_jobs</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span>
|
||||
<span class="n">n_jobs</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span><span class="p">,</span>
|
||||
<span class="n">asarray</span><span class="o">=</span><span class="kc">False</span>
|
||||
<span class="p">)</span>
|
||||
|
||||
<span class="c1"># filter out classifier configurations that yielded any error</span>
|
||||
|
|
@ -266,9 +584,10 @@
|
|||
|
||||
<span class="k">return</span> <span class="n">aggr_outs</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_compute_scores_nonaggregative</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">training</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_compute_scores_nonaggregative</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="n">configs</span> <span class="o">=</span> <span class="n">expand_grid</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">param_grid</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_training</span> <span class="o">=</span> <span class="n">training</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_training_X</span> <span class="o">=</span> <span class="n">X</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_training_y</span> <span class="o">=</span> <span class="n">y</span>
|
||||
<span class="n">scores</span> <span class="o">=</span> <span class="n">qp</span><span class="o">.</span><span class="n">util</span><span class="o">.</span><span class="n">parallel</span><span class="p">(</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_prepare_nonaggr_model</span><span class="p">,</span>
|
||||
<span class="n">configs</span><span class="p">,</span>
|
||||
|
|
@ -277,17 +596,20 @@
|
|||
<span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">scores</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_print_status</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">params</span><span class="p">,</span> <span class="n">score</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">took</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_print_status</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">params</span><span class="p">,</span> <span class="n">score</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">took</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">status</span><span class="o">.</span><span class="n">success</span><span class="p">():</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_sout</span><span class="p">(</span><span class="sa">f</span><span class="s1">'hyperparams=[</span><span class="si">{</span><span class="n">params</span><span class="si">}</span><span class="s1">]</span><span class="se">\t</span><span class="s1"> got </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">error</span><span class="o">.</span><span class="vm">__name__</span><span class="si">}</span><span class="s1"> = </span><span class="si">{</span><span class="n">score</span><span class="si">:</span><span class="s1">.5f</span><span class="si">}</span><span class="s1"> [took </span><span class="si">{</span><span class="n">took</span><span class="si">:</span><span class="s1">.3f</span><span class="si">}</span><span class="s1">s]'</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_sout</span><span class="p">(</span><span class="sa">f</span><span class="s1">'error=</span><span class="si">{</span><span class="n">status</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.fit"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.fit">[docs]</a> <span class="k">def</span> <span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">training</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="GridSearchQ.fit">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.fit">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">fit</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">""" Learning routine. Fits methods with all combinations of hyperparameters and selects the one minimizing</span>
|
||||
<span class="sd"> the error metric.</span>
|
||||
|
||||
<span class="sd"> :param training: the training set on which to optimize the hyperparameters</span>
|
||||
<span class="sd"> :param X: array-like, training covariates</span>
|
||||
<span class="sd"> :param y: array-like, labels of training data</span>
|
||||
<span class="sd"> :return: self</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
|
|
@ -303,9 +625,9 @@
|
|||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_sout</span><span class="p">(</span><span class="sa">f</span><span class="s1">'starting model selection with n_jobs=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">n_jobs</span><span class="si">}</span><span class="s1">'</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">_break_down_fit</span><span class="p">():</span>
|
||||
<span class="n">results</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_compute_scores_aggregative</span><span class="p">(</span><span class="n">training</span><span class="p">)</span>
|
||||
<span class="n">results</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_compute_scores_aggregative</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="n">results</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_compute_scores_nonaggregative</span><span class="p">(</span><span class="n">training</span><span class="p">)</span>
|
||||
<span class="n">results</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">_compute_scores_nonaggregative</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">param_scores_</span> <span class="o">=</span> <span class="p">{}</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">best_score_</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
|
|
@ -320,13 +642,13 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">param_scores_</span><span class="p">[</span><span class="nb">str</span><span class="p">(</span><span class="n">params</span><span class="p">)]</span> <span class="o">=</span> <span class="n">status</span><span class="o">.</span><span class="n">status</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">error_collector</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">status</span><span class="p">)</span>
|
||||
|
||||
<span class="n">tend</span> <span class="o">=</span> <span class="n">time</span><span class="p">()</span><span class="o">-</span><span class="n">tinit</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">fit_time_</span> <span class="o">=</span> <span class="n">time</span><span class="p">()</span><span class="o">-</span><span class="n">tinit</span>
|
||||
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">best_score_</span> <span class="ow">is</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="s1">'no combination of hyperparameters seemed to work'</span><span class="p">)</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_sout</span><span class="p">(</span><span class="sa">f</span><span class="s1">'optimization finished: best params </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">best_params_</span><span class="si">}</span><span class="s1"> (score=</span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">best_score_</span><span class="si">:</span><span class="s1">.5f</span><span class="si">}</span><span class="s1">) '</span>
|
||||
<span class="sa">f</span><span class="s1">'[took </span><span class="si">{</span><span class="n">tend</span><span class="si">:</span><span class="s1">.4f</span><span class="si">}</span><span class="s1">s]'</span><span class="p">)</span>
|
||||
<span class="sa">f</span><span class="s1">'[took </span><span class="si">{</span><span class="bp">self</span><span class="o">.</span><span class="n">fit_time_</span><span class="si">:</span><span class="s1">.4f</span><span class="si">}</span><span class="s1">s]'</span><span class="p">)</span>
|
||||
|
||||
<span class="n">no_errors</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">error_collector</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="n">no_errors</span><span class="o">></span><span class="mi">0</span><span class="p">:</span>
|
||||
|
|
@ -338,7 +660,10 @@
|
|||
<span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">protocol</span><span class="p">,</span> <span class="n">OnLabelledCollectionProtocol</span><span class="p">):</span>
|
||||
<span class="n">tinit</span> <span class="o">=</span> <span class="n">time</span><span class="p">()</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">_sout</span><span class="p">(</span><span class="sa">f</span><span class="s1">'refitting on the whole development set'</span><span class="p">)</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">best_model_</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">training</span> <span class="o">+</span> <span class="bp">self</span><span class="o">.</span><span class="n">protocol</span><span class="o">.</span><span class="n">get_labelled_collection</span><span class="p">())</span>
|
||||
<span class="n">validation_collection</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">protocol</span><span class="o">.</span><span class="n">get_labelled_collection</span><span class="p">()</span>
|
||||
<span class="n">training_collection</span> <span class="o">=</span> <span class="n">LabelledCollection</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">classes</span><span class="o">=</span><span class="n">validation_collection</span><span class="o">.</span><span class="n">classes</span><span class="p">)</span>
|
||||
<span class="n">devel_collection</span> <span class="o">=</span> <span class="n">training_collection</span> <span class="o">+</span> <span class="n">validation_collection</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">best_model_</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">devel_collection</span><span class="o">.</span><span class="n">Xy</span><span class="p">)</span>
|
||||
<span class="n">tend</span> <span class="o">=</span> <span class="n">time</span><span class="p">()</span> <span class="o">-</span> <span class="n">tinit</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">refit_time_</span> <span class="o">=</span> <span class="n">tend</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
|
|
@ -347,24 +672,33 @@
|
|||
|
||||
<span class="k">return</span> <span class="bp">self</span></div>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.quantify"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.quantify">[docs]</a> <span class="k">def</span> <span class="nf">quantify</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">instances</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.predict">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">predict</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">X</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Estimate class prevalence values using the best model found after calling the :meth:`fit` method.</span>
|
||||
|
||||
<span class="sd"> :param instances: sample contanining the instances</span>
|
||||
<span class="sd"> :param X: sample contanining the instances</span>
|
||||
<span class="sd"> :return: a ndarray of shape `(n_classes)` with class prevalence estimates as according to the best model found</span>
|
||||
<span class="sd"> by the model selection process.</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">assert</span> <span class="nb">hasattr</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="s1">'best_model_'</span><span class="p">),</span> <span class="s1">'quantify called before fit'</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">best_model</span><span class="p">()</span><span class="o">.</span><span class="n">quantify</span><span class="p">(</span><span class="n">instances</span><span class="p">)</span></div>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">best_model</span><span class="p">()</span><span class="o">.</span><span class="n">predict</span><span class="p">(</span><span class="n">X</span><span class="p">)</span></div>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.set_params"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.set_params">[docs]</a> <span class="k">def</span> <span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">parameters</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.set_params">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.set_params">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">set_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="o">**</span><span class="n">parameters</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Sets the hyper-parameters to explore.</span>
|
||||
|
||||
<span class="sd"> :param parameters: a dictionary with keys the parameter names and values the list of values to explore</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">param_grid</span> <span class="o">=</span> <span class="n">parameters</span></div>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.get_params"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.get_params">[docs]</a> <span class="k">def</span> <span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">deep</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.get_params">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.get_params">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_params</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">deep</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""Returns the dictionary of hyper-parameters to explore (`param_grid`)</span>
|
||||
|
||||
<span class="sd"> :param deep: Unused</span>
|
||||
|
|
@ -372,7 +706,10 @@
|
|||
<span class="sd"> """</span>
|
||||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">param_grid</span></div>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.best_model"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.best_model">[docs]</a> <span class="k">def</span> <span class="nf">best_model</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="GridSearchQ.best_model">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.GridSearchQ.best_model">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">best_model</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Returns the best model found after calling the :meth:`fit` method, i.e., the one trained on the combination</span>
|
||||
<span class="sd"> of hyper-parameters that minimized the error function.</span>
|
||||
|
|
@ -383,7 +720,8 @@
|
|||
<span class="k">return</span> <span class="bp">self</span><span class="o">.</span><span class="n">best_model_</span>
|
||||
<span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="s1">'best_model called before fit'</span><span class="p">)</span></div>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_error_handler</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">func</span><span class="p">,</span> <span class="n">params</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_error_handler</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">func</span><span class="p">,</span> <span class="n">params</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Endorses one job with two returned values: the status, and the time of execution</span>
|
||||
|
||||
|
|
@ -396,11 +734,11 @@
|
|||
|
||||
<span class="n">output</span> <span class="o">=</span> <span class="kc">None</span>
|
||||
|
||||
<span class="k">def</span> <span class="nf">_handle</span><span class="p">(</span><span class="n">status</span><span class="p">,</span> <span class="n">exception</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_handle</span><span class="p">(</span><span class="n">status</span><span class="p">,</span> <span class="n">exception</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">raise_errors</span><span class="p">:</span>
|
||||
<span class="k">raise</span> <span class="n">exception</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="n">ConfigStatus</span><span class="p">(</span><span class="n">params</span><span class="p">,</span> <span class="n">status</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">ConfigStatus</span><span class="p">(</span><span class="n">params</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">msg</span><span class="o">=</span><span class="nb">str</span><span class="p">(</span><span class="n">exception</span><span class="p">))</span>
|
||||
|
||||
<span class="k">try</span><span class="p">:</span>
|
||||
<span class="k">with</span> <span class="n">timeout</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">timeout</span><span class="p">):</span>
|
||||
|
|
@ -421,7 +759,10 @@
|
|||
<span class="k">return</span> <span class="n">output</span><span class="p">,</span> <span class="n">status</span><span class="p">,</span> <span class="n">took</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="cross_val_predict"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.cross_val_predict">[docs]</a><span class="k">def</span> <span class="nf">cross_val_predict</span><span class="p">(</span><span class="n">quantifier</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">nfolds</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="mi">0</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="cross_val_predict">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.cross_val_predict">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">cross_val_predict</span><span class="p">(</span><span class="n">quantifier</span><span class="p">:</span> <span class="n">BaseQuantifier</span><span class="p">,</span> <span class="n">data</span><span class="p">:</span> <span class="n">LabelledCollection</span><span class="p">,</span> <span class="n">nfolds</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="mi">0</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Akin to `scikit-learn's cross_val_predict <https://scikit-learn.org/stable/modules/generated/sklearn.model_selection.cross_val_predict.html>`_</span>
|
||||
<span class="sd"> but for quantification.</span>
|
||||
|
|
@ -436,15 +777,18 @@
|
|||
<span class="n">total_prev</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">zeros</span><span class="p">(</span><span class="n">shape</span><span class="o">=</span><span class="n">data</span><span class="o">.</span><span class="n">n_classes</span><span class="p">)</span>
|
||||
|
||||
<span class="k">for</span> <span class="n">train</span><span class="p">,</span> <span class="n">test</span> <span class="ow">in</span> <span class="n">data</span><span class="o">.</span><span class="n">kFCV</span><span class="p">(</span><span class="n">nfolds</span><span class="o">=</span><span class="n">nfolds</span><span class="p">,</span> <span class="n">random_state</span><span class="o">=</span><span class="n">random_state</span><span class="p">):</span>
|
||||
<span class="n">quantifier</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">train</span><span class="p">)</span>
|
||||
<span class="n">fold_prev</span> <span class="o">=</span> <span class="n">quantifier</span><span class="o">.</span><span class="n">quantify</span><span class="p">(</span><span class="n">test</span><span class="o">.</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">quantifier</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="o">*</span><span class="n">train</span><span class="o">.</span><span class="n">Xy</span><span class="p">)</span>
|
||||
<span class="n">fold_prev</span> <span class="o">=</span> <span class="n">quantifier</span><span class="o">.</span><span class="n">predict</span><span class="p">(</span><span class="n">test</span><span class="o">.</span><span class="n">X</span><span class="p">)</span>
|
||||
<span class="n">rel_size</span> <span class="o">=</span> <span class="mf">1.</span> <span class="o">*</span> <span class="nb">len</span><span class="p">(</span><span class="n">test</span><span class="p">)</span> <span class="o">/</span> <span class="nb">len</span><span class="p">(</span><span class="n">data</span><span class="p">)</span>
|
||||
<span class="n">total_prev</span> <span class="o">+=</span> <span class="n">fold_prev</span><span class="o">*</span><span class="n">rel_size</span>
|
||||
|
||||
<span class="k">return</span> <span class="n">total_prev</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="expand_grid"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.expand_grid">[docs]</a><span class="k">def</span> <span class="nf">expand_grid</span><span class="p">(</span><span class="n">param_grid</span><span class="p">:</span> <span class="nb">dict</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="expand_grid">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.expand_grid">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">expand_grid</span><span class="p">(</span><span class="n">param_grid</span><span class="p">:</span> <span class="nb">dict</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Expands a param_grid dictionary as a list of configurations.</span>
|
||||
<span class="sd"> Example:</span>
|
||||
|
|
@ -463,7 +807,10 @@
|
|||
<span class="k">return</span> <span class="n">configs</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="group_params"><a class="viewcode-back" href="../../quapy.html#quapy.model_selection.group_params">[docs]</a><span class="k">def</span> <span class="nf">group_params</span><span class="p">(</span><span class="n">param_grid</span><span class="p">:</span> <span class="nb">dict</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="group_params">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.model_selection.group_params">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">group_params</span><span class="p">(</span><span class="n">param_grid</span><span class="p">:</span> <span class="nb">dict</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Partitions a param_grid dictionary as two lists of configurations, one for the classifier-specific</span>
|
||||
<span class="sd"> hyper-parameters, and another for que quantifier-specific hyper-parameters</span>
|
||||
|
|
@ -484,33 +831,78 @@
|
|||
|
||||
<span class="k">return</span> <span class="n">classifier_configs</span><span class="p">,</span> <span class="n">quantifier_configs</span></div>
|
||||
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,93 +1,395 @@
|
|||
|
||||
<!DOCTYPE html>
|
||||
<html class="writer-html5" lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.util — QuaPy: A Python-based open-source framework for quantification 0.1.8 documentation</title>
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/pygments.css" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/css/theme.css" />
|
||||
|
||||
|
||||
<html lang="en" data-content_root="../../" >
|
||||
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>quapy.util — QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation</title>
|
||||
|
||||
<!--[if lt IE 9]>
|
||||
<script src="../../_static/js/html5shiv.min.js"></script>
|
||||
<![endif]-->
|
||||
|
||||
<script data-url_root="../../" id="documentation_options" src="../../_static/documentation_options.js"></script>
|
||||
<script src="../../_static/jquery.js"></script>
|
||||
<script src="../../_static/underscore.js"></script>
|
||||
<script src="../../_static/_sphinx_javascript_frameworks_compat.js"></script>
|
||||
<script src="../../_static/doctools.js"></script>
|
||||
<script src="../../_static/sphinx_highlight.js"></script>
|
||||
<script src="../../_static/js/theme.js"></script>
|
||||
|
||||
<script data-cfasync="false">
|
||||
document.documentElement.dataset.mode = localStorage.getItem("mode") || "";
|
||||
document.documentElement.dataset.theme = localStorage.getItem("theme") || "";
|
||||
</script>
|
||||
<!--
|
||||
this give us a css class that will be invisible only if js is disabled
|
||||
-->
|
||||
<noscript>
|
||||
<style>
|
||||
.pst-js-only { display: none !important; }
|
||||
|
||||
</style>
|
||||
</noscript>
|
||||
|
||||
<!-- Loaded before other Sphinx assets -->
|
||||
<link href="../../_static/styles/theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
<link href="../../_static/styles/pydata-sphinx-theme.css?digest=90905a2f556bf617f1a9" rel="stylesheet" />
|
||||
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/pygments.css?v=8f2a1f02" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/sphinx-design.min.css?v=95c83b7e" />
|
||||
<link rel="stylesheet" type="text/css" href="../../_static/custom.css?v=9f2a2228" />
|
||||
|
||||
<!-- So that users can add custom icons -->
|
||||
<script defer src="../../_static/scripts/fontawesome.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<!-- Pre-loaded scripts that we'll load fully later -->
|
||||
<link rel="preload" as="script" href="../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9" />
|
||||
<link rel="preload" as="script" href="../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9" />
|
||||
|
||||
<script src="../../_static/documentation_options.js?v=37f418d5"></script>
|
||||
<script src="../../_static/doctools.js?v=fd6eb6e6"></script>
|
||||
<script src="../../_static/sphinx_highlight.js?v=6ffebe34"></script>
|
||||
<script src="../../_static/design-tabs.js?v=f930bc37"></script>
|
||||
<script>DOCUMENTATION_OPTIONS.pagename = '_modules/quapy/util';</script>
|
||||
<script>DOCUMENTATION_OPTIONS.search_as_you_type = false;</script>
|
||||
<link rel="index" title="Index" href="../../genindex.html" />
|
||||
<link rel="search" title="Search" href="../../search.html" />
|
||||
</head>
|
||||
<link rel="search" title="Search" href="../../search.html" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1"/>
|
||||
<meta name="docsearch:language" content="en"/>
|
||||
<meta name="docsearch:version" content="" />
|
||||
|
||||
|
||||
<script src="../../_static/searchtools.js"></script>
|
||||
<script src="../../_static/language_data.js"></script>
|
||||
<script src="../../searchindex.js"></script>
|
||||
|
||||
</head>
|
||||
<body data-default-mode="">
|
||||
|
||||
|
||||
<div id="pst-skip-link" class="skip-link d-print-none"><a href="#main-content">Skip to main content</a></div>
|
||||
|
||||
<body class="wy-body-for-nav">
|
||||
<div class="wy-grid-for-nav">
|
||||
<nav data-toggle="wy-nav-shift" class="wy-nav-side">
|
||||
<div class="wy-side-scroll">
|
||||
<div class="wy-side-nav-search" >
|
||||
|
||||
<div id="pst-scroll-pixel-helper"></div>
|
||||
|
||||
<button type="button" class="btn rounded-pill" id="pst-back-to-top">
|
||||
<i class="fa-solid fa-arrow-up"></i>Back to top</button>
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-search-dialog">
|
||||
|
||||
<form class="bd-search d-flex align-items-center"
|
||||
action="../../search.html"
|
||||
method="get">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<input type="search"
|
||||
class="form-control"
|
||||
name="q"
|
||||
placeholder="Search the docs ..."
|
||||
aria-label="Search the docs ..."
|
||||
autocomplete="off"
|
||||
autocorrect="off"
|
||||
autocapitalize="off"
|
||||
spellcheck="false"/>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd>K</kbd></span>
|
||||
</form>
|
||||
</dialog>
|
||||
|
||||
|
||||
|
||||
<a href="../../index.html" class="icon icon-home">
|
||||
QuaPy: A Python-based open-source framework for quantification
|
||||
</a>
|
||||
<div role="search">
|
||||
<form id="rtd-search-form" class="wy-form" action="../../search.html" method="get">
|
||||
<input type="text" name="q" placeholder="Search docs" aria-label="Search docs" />
|
||||
<input type="hidden" name="check_keywords" value="yes" />
|
||||
<input type="hidden" name="area" value="default" />
|
||||
</form>
|
||||
<div class="pst-async-banner-revealer d-none">
|
||||
<aside id="bd-header-version-warning" class="d-none d-print-none" aria-label="Version warning"></aside>
|
||||
</div>
|
||||
</div><div class="wy-menu wy-menu-vertical" data-spy="affix" role="navigation" aria-label="Navigation menu">
|
||||
<ul>
|
||||
<li class="toctree-l1"><a class="reference internal" href="../../modules.html">quapy</a></li>
|
||||
</ul>
|
||||
|
||||
</div>
|
||||
</div>
|
||||
</nav>
|
||||
|
||||
<header id="pst-header" class="bd-header navbar navbar-expand-lg bd-navbar d-print-none">
|
||||
<div class="bd-header__inner bd-page-width">
|
||||
<button class="pst-navbar-icon sidebar-toggle primary-toggle" aria-label="Site navigation">
|
||||
<span class="fa-solid fa-bars"></span>
|
||||
</button>
|
||||
|
||||
|
||||
<div class="col-lg-3 navbar-header-items__start">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<section data-toggle="wy-nav-shift" class="wy-nav-content-wrap"><nav class="wy-nav-top" aria-label="Mobile navigation menu" >
|
||||
<i data-toggle="wy-nav-top" class="fa fa-bars"></i>
|
||||
<a href="../../index.html">QuaPy: A Python-based open-source framework for quantification</a>
|
||||
</nav>
|
||||
|
||||
|
||||
|
||||
|
||||
<a class="navbar-brand logo" href="../../index.html">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<img src="../../_static/quapy_logo.png" class="logo__image only-light" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
<img src="../../_static/quapy_logo_dark.png" class="logo__image only-dark pst-js-only" alt="QuaPy: A Python-based open-source framework for quantification 0.2.1 documentation - Home"/>
|
||||
|
||||
|
||||
</a></div>
|
||||
|
||||
</div>
|
||||
|
||||
<div class="col-lg-9 navbar-header-items">
|
||||
|
||||
<div class="me-auto navbar-header-items__center">
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<div class="wy-nav-content">
|
||||
<div class="rst-content">
|
||||
<div role="navigation" aria-label="Page navigation">
|
||||
<ul class="wy-breadcrumbs">
|
||||
<li><a href="../../index.html" class="icon icon-home" aria-label="Home"></a></li>
|
||||
<li class="breadcrumb-item"><a href="../index.html">Module code</a></li>
|
||||
<li class="breadcrumb-item active">quapy.util</li>
|
||||
<li class="wy-breadcrumbs-aside">
|
||||
</li>
|
||||
</ul>
|
||||
<hr/>
|
||||
</nav></div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-header-items__end">
|
||||
|
||||
<div class="navbar-item navbar-persistent--container">
|
||||
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="navbar-persistent--mobile">
|
||||
|
||||
<button class="btn search-button-field search-button__button pst-js-only" title="Search" aria-label="Search" data-bs-placement="bottom" data-bs-toggle="tooltip">
|
||||
<i class="fa-solid fa-magnifying-glass"></i>
|
||||
<span class="search-button__default-text">Search</span>
|
||||
<span class="search-button__kbd-shortcut"><kbd class="kbd-shortcut__modifier">Ctrl</kbd>+<kbd class="kbd-shortcut__modifier">K</kbd></span>
|
||||
</button>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<div role="main" class="document" itemscope="itemscope" itemtype="http://schema.org/Article">
|
||||
<div itemprop="articleBody">
|
||||
|
||||
|
||||
</header>
|
||||
|
||||
|
||||
<div class="bd-container">
|
||||
<div class="bd-container__inner bd-page-width">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<dialog id="pst-primary-sidebar-modal"></dialog>
|
||||
<div id="pst-primary-sidebar" class="bd-sidebar-primary bd-sidebar hide-on-wide">
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items sidebar-primary__section">
|
||||
|
||||
|
||||
<div class="sidebar-header-items__center">
|
||||
|
||||
|
||||
|
||||
<div class="navbar-item">
|
||||
<nav>
|
||||
<ul class="bd-navbar-elements navbar-nav">
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../index.html">
|
||||
Home
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../manuals.html">
|
||||
Manuals
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
<li class="nav-item ">
|
||||
<a class="nav-link nav-internal" href="../../quapy.html">
|
||||
API
|
||||
</a>
|
||||
</li>
|
||||
|
||||
</ul>
|
||||
</nav></div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="sidebar-header-items__end">
|
||||
|
||||
<div class="navbar-item">
|
||||
|
||||
<div class="theme-switch-container dropdown pst-js-only" data-bs-toggle="tooltip" data-bs-placement="bottom" title="Color mode">
|
||||
<button class="btn btn-sm nav-link pst-navbar-icon theme-switch-button dropdown-toggle" aria-label="Color mode" data-bs-toggle="dropdown">
|
||||
<i class="theme-switch fa-solid fa-sun fa-lg fa-fw" data-mode="light" title="Light"></i>
|
||||
<i class="theme-switch fa-solid fa-moon fa-lg fa-fw" data-mode="dark" title="Dark"></i>
|
||||
<i class="theme-switch fa-solid fa-circle-half-stroke fa-lg fa-fw" data-mode="auto" title="System Settings"></i>
|
||||
</button>
|
||||
<ul class="dropdown-menu dropdown-menu-end">
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="auto"><i class="fa-solid fa-circle-half-stroke fa-lg fa-fw me-1"></i>System Settings</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="light"><i class="fa-solid fa-sun fa-lg fa-fw me-1"></i>Light</button></li>
|
||||
<li><button class="dropdown-item d-flex align-items-center theme-change-button" data-mode="dark"><i class="fa-solid fa-moon fa-lg fa-fw me-1"></i>Dark</button></li>
|
||||
</ul>
|
||||
</div></div>
|
||||
|
||||
<div class="navbar-item"><ul class="navbar-icon-links"
|
||||
aria-label="Icon Links">
|
||||
<li class="nav-item">
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<a href="https://github.com/HLT-ISTI/QuaPy" title="GitHub" class="nav-link pst-navbar-icon" rel="noopener" target="_blank" data-bs-toggle="tooltip" data-bs-placement="bottom"><i class="fa-brands fa-github fa-lg" aria-hidden="true"></i><span class="visually-hidden">GitHub</span></a>
|
||||
</li>
|
||||
</ul></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
<div class="sidebar-primary-items__end sidebar-primary__section">
|
||||
<div class="sidebar-primary-item">
|
||||
<div id="ethical-ad-placement"
|
||||
class="flat"
|
||||
data-ea-publisher="readthedocs"
|
||||
data-ea-type="readthedocs-sidebar"
|
||||
data-ea-manual="true">
|
||||
</div></div>
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
|
||||
<main id="main-content" class="bd-main" role="main">
|
||||
|
||||
|
||||
<div class="bd-content">
|
||||
<div class="bd-article-container">
|
||||
|
||||
<div class="bd-header-article d-print-none">
|
||||
<div class="header-article-items header-article__inner">
|
||||
|
||||
<div class="header-article-items__start">
|
||||
|
||||
<div class="header-article-item">
|
||||
|
||||
<nav aria-label="Breadcrumb" class="d-print-none">
|
||||
<ul class="bd-breadcrumbs">
|
||||
|
||||
<li class="breadcrumb-item breadcrumb-home">
|
||||
<a href="../../index.html" class="nav-link" aria-label="Home">
|
||||
<i class="fa-solid fa-home"></i>
|
||||
</a>
|
||||
</li>
|
||||
|
||||
<li class="breadcrumb-item"><a href="../index.html" class="nav-link">Module code</a></li>
|
||||
|
||||
<li class="breadcrumb-item active" aria-current="page"><span class="ellipsis">quapy.util</span></li>
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
</div>
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
<div id="searchbox"></div>
|
||||
<article class="bd-article">
|
||||
|
||||
<h1>Source code for quapy.util</h1><div class="highlight"><pre>
|
||||
<span></span><span class="kn">import</span> <span class="nn">contextlib</span>
|
||||
<span class="kn">import</span> <span class="nn">itertools</span>
|
||||
<span class="kn">import</span> <span class="nn">multiprocessing</span>
|
||||
<span class="kn">import</span> <span class="nn">os</span>
|
||||
<span class="kn">import</span> <span class="nn">pickle</span>
|
||||
<span class="kn">import</span> <span class="nn">urllib</span>
|
||||
<span class="kn">from</span> <span class="nn">pathlib</span> <span class="kn">import</span> <span class="n">Path</span>
|
||||
<span class="kn">from</span> <span class="nn">contextlib</span> <span class="kn">import</span> <span class="n">ExitStack</span>
|
||||
<span class="kn">import</span> <span class="nn">quapy</span> <span class="k">as</span> <span class="nn">qp</span>
|
||||
<span></span><span class="kn">import</span><span class="w"> </span><span class="nn">contextlib</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">itertools</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">multiprocessing</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">os</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">pickle</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">urllib</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">pathlib</span><span class="w"> </span><span class="kn">import</span> <span class="n">Path</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">contextlib</span><span class="w"> </span><span class="kn">import</span> <span class="n">ExitStack</span>
|
||||
|
||||
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="nn">np</span>
|
||||
<span class="kn">from</span> <span class="nn">joblib</span> <span class="kn">import</span> <span class="n">Parallel</span><span class="p">,</span> <span class="n">delayed</span>
|
||||
<span class="kn">from</span> <span class="nn">time</span> <span class="kn">import</span> <span class="n">time</span>
|
||||
<span class="kn">import</span> <span class="nn">signal</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">pandas</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">pd</span>
|
||||
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">quapy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">qp</span>
|
||||
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">numpy</span><span class="w"> </span><span class="k">as</span><span class="w"> </span><span class="nn">np</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">joblib</span><span class="w"> </span><span class="kn">import</span> <span class="n">Parallel</span><span class="p">,</span> <span class="n">delayed</span>
|
||||
<span class="kn">from</span><span class="w"> </span><span class="nn">time</span><span class="w"> </span><span class="kn">import</span> <span class="n">time</span>
|
||||
<span class="kn">import</span><span class="w"> </span><span class="nn">signal</span>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">_get_parallel_slices</span><span class="p">(</span><span class="n">n_tasks</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_get_parallel_slices</span><span class="p">(</span><span class="n">n_tasks</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">n_jobs</span> <span class="o">==</span> <span class="o">-</span><span class="mi">1</span><span class="p">:</span>
|
||||
<span class="n">n_jobs</span> <span class="o">=</span> <span class="n">multiprocessing</span><span class="o">.</span><span class="n">cpu_count</span><span class="p">()</span>
|
||||
<span class="n">batch</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">n_tasks</span> <span class="o">/</span> <span class="n">n_jobs</span><span class="p">)</span>
|
||||
|
|
@ -95,7 +397,9 @@
|
|||
<span class="k">return</span> <span class="p">[</span><span class="nb">slice</span><span class="p">(</span><span class="n">job</span> <span class="o">*</span> <span class="n">batch</span><span class="p">,</span> <span class="p">(</span><span class="n">job</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span> <span class="o">*</span> <span class="n">batch</span> <span class="o">+</span> <span class="p">(</span><span class="n">remainder</span> <span class="k">if</span> <span class="n">job</span> <span class="o">==</span> <span class="n">n_jobs</span> <span class="o">-</span> <span class="mi">1</span> <span class="k">else</span> <span class="mi">0</span><span class="p">))</span> <span class="k">for</span> <span class="n">job</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">n_jobs</span><span class="p">)]</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="map_parallel"><a class="viewcode-back" href="../../quapy.html#quapy.util.map_parallel">[docs]</a><span class="k">def</span> <span class="nf">map_parallel</span><span class="p">(</span><span class="n">func</span><span class="p">,</span> <span class="n">args</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">):</span>
|
||||
<div class="viewcode-block" id="map_parallel">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.map_parallel">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">map_parallel</span><span class="p">(</span><span class="n">func</span><span class="p">,</span> <span class="n">args</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Applies func to n_jobs slices of args. E.g., if args is an array of 99 items and n_jobs=2, then</span>
|
||||
<span class="sd"> func is applied in two parallel processes to args[0:50] and to args[50:99]. func is a function</span>
|
||||
|
|
@ -113,7 +417,10 @@
|
|||
<span class="k">return</span> <span class="nb">list</span><span class="p">(</span><span class="n">itertools</span><span class="o">.</span><span class="n">chain</span><span class="o">.</span><span class="n">from_iterable</span><span class="p">(</span><span class="n">results</span><span class="p">))</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="parallel"><a class="viewcode-back" href="../../quapy.html#quapy.util.parallel">[docs]</a><span class="k">def</span> <span class="nf">parallel</span><span class="p">(</span><span class="n">func</span><span class="p">,</span> <span class="n">args</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">,</span> <span class="n">seed</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">asarray</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">backend</span><span class="o">=</span><span class="s1">'loky'</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="parallel">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.parallel">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">parallel</span><span class="p">(</span><span class="n">func</span><span class="p">,</span> <span class="n">args</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">,</span> <span class="n">seed</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">asarray</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">backend</span><span class="o">=</span><span class="s1">'loky'</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> A wrapper of multiprocessing:</span>
|
||||
|
||||
|
|
@ -129,8 +436,9 @@
|
|||
<span class="sd"> :param seed: the numeric seed</span>
|
||||
<span class="sd"> :param asarray: set to True to return a np.ndarray instead of a list</span>
|
||||
<span class="sd"> :param backend: indicates the backend used for handling parallel works</span>
|
||||
<span class="sd"> :param open_args: if True, then the delayed function is called on *args_i, instead of on args_i</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">def</span> <span class="nf">func_dec</span><span class="p">(</span><span class="n">environ</span><span class="p">,</span> <span class="n">seed</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">func_dec</span><span class="p">(</span><span class="n">environ</span><span class="p">,</span> <span class="n">seed</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
<span class="n">qp</span><span class="o">.</span><span class="n">environ</span> <span class="o">=</span> <span class="n">environ</span><span class="o">.</span><span class="n">copy</span><span class="p">()</span>
|
||||
<span class="n">qp</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">'N_JOBS'</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>
|
||||
<span class="c1">#set a context with a temporal seed to ensure results are reproducibles in parallel</span>
|
||||
|
|
@ -147,8 +455,48 @@
|
|||
<span class="k">return</span> <span class="n">out</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="temp_seed"><a class="viewcode-back" href="../../quapy.html#quapy.util.temp_seed">[docs]</a><span class="nd">@contextlib</span><span class="o">.</span><span class="n">contextmanager</span>
|
||||
<span class="k">def</span> <span class="nf">temp_seed</span><span class="p">(</span><span class="n">random_state</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="parallel_unpack">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.parallel_unpack">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">parallel_unpack</span><span class="p">(</span><span class="n">func</span><span class="p">,</span> <span class="n">args</span><span class="p">,</span> <span class="n">n_jobs</span><span class="p">,</span> <span class="n">seed</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">asarray</span><span class="o">=</span><span class="kc">True</span><span class="p">,</span> <span class="n">backend</span><span class="o">=</span><span class="s1">'loky'</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> A wrapper of multiprocessing:</span>
|
||||
|
||||
<span class="sd"> >>> Parallel(n_jobs=n_jobs)(</span>
|
||||
<span class="sd"> >>> delayed(func)(*args_i) for args_i in args</span>
|
||||
<span class="sd"> >>> )</span>
|
||||
|
||||
<span class="sd"> that takes the `quapy.environ` variable as input silently.</span>
|
||||
<span class="sd"> Seeds the child processes to ensure reproducibility when n_jobs>1.</span>
|
||||
|
||||
<span class="sd"> :param func: callable</span>
|
||||
<span class="sd"> :param args: args of func</span>
|
||||
<span class="sd"> :param seed: the numeric seed</span>
|
||||
<span class="sd"> :param asarray: set to True to return a np.ndarray instead of a list</span>
|
||||
<span class="sd"> :param backend: indicates the backend used for handling parallel works</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">func_dec</span><span class="p">(</span><span class="n">environ</span><span class="p">,</span> <span class="n">seed</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
<span class="n">qp</span><span class="o">.</span><span class="n">environ</span> <span class="o">=</span> <span class="n">environ</span><span class="o">.</span><span class="n">copy</span><span class="p">()</span>
|
||||
<span class="n">qp</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">'N_JOBS'</span><span class="p">]</span> <span class="o">=</span> <span class="mi">1</span>
|
||||
<span class="c1"># set a context with a temporal seed to ensure results are reproducibles in parallel</span>
|
||||
<span class="k">with</span> <span class="n">ExitStack</span><span class="p">()</span> <span class="k">as</span> <span class="n">stack</span><span class="p">:</span>
|
||||
<span class="k">if</span> <span class="n">seed</span> <span class="ow">is</span> <span class="ow">not</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="n">stack</span><span class="o">.</span><span class="n">enter_context</span><span class="p">(</span><span class="n">qp</span><span class="o">.</span><span class="n">util</span><span class="o">.</span><span class="n">temp_seed</span><span class="p">(</span><span class="n">seed</span><span class="p">))</span>
|
||||
<span class="k">return</span> <span class="n">func</span><span class="p">(</span><span class="o">*</span><span class="n">args</span><span class="p">)</span>
|
||||
|
||||
<span class="n">out</span> <span class="o">=</span> <span class="n">Parallel</span><span class="p">(</span><span class="n">n_jobs</span><span class="o">=</span><span class="n">n_jobs</span><span class="p">,</span> <span class="n">backend</span><span class="o">=</span><span class="n">backend</span><span class="p">)(</span>
|
||||
<span class="n">delayed</span><span class="p">(</span><span class="n">func_dec</span><span class="p">)(</span><span class="n">qp</span><span class="o">.</span><span class="n">environ</span><span class="p">,</span> <span class="kc">None</span> <span class="k">if</span> <span class="n">seed</span> <span class="ow">is</span> <span class="kc">None</span> <span class="k">else</span> <span class="n">seed</span> <span class="o">+</span> <span class="n">i</span><span class="p">,</span> <span class="o">*</span><span class="n">args_i</span><span class="p">)</span> <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">args_i</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">args</span><span class="p">)</span>
|
||||
<span class="p">)</span>
|
||||
<span class="k">if</span> <span class="n">asarray</span><span class="p">:</span>
|
||||
<span class="n">out</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">out</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="temp_seed">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.temp_seed">[docs]</a>
|
||||
<span class="nd">@contextlib</span><span class="o">.</span><span class="n">contextmanager</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">temp_seed</span><span class="p">(</span><span class="n">random_state</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Can be used in a "with" context to set a temporal seed without modifying the outer numpy's current state. E.g.:</span>
|
||||
|
||||
|
|
@ -169,14 +517,17 @@
|
|||
<span class="n">np</span><span class="o">.</span><span class="n">random</span><span class="o">.</span><span class="n">set_state</span><span class="p">(</span><span class="n">state</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="download_file"><a class="viewcode-back" href="../../quapy.html#quapy.util.download_file">[docs]</a><span class="k">def</span> <span class="nf">download_file</span><span class="p">(</span><span class="n">url</span><span class="p">,</span> <span class="n">archive_filename</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="download_file">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.download_file">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">download_file</span><span class="p">(</span><span class="n">url</span><span class="p">,</span> <span class="n">archive_filename</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Downloads a file from a url</span>
|
||||
|
||||
<span class="sd"> :param url: the url</span>
|
||||
<span class="sd"> :param archive_filename: destination filename</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">def</span> <span class="nf">progress</span><span class="p">(</span><span class="n">blocknum</span><span class="p">,</span> <span class="n">bs</span><span class="p">,</span> <span class="n">size</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">progress</span><span class="p">(</span><span class="n">blocknum</span><span class="p">,</span> <span class="n">bs</span><span class="p">,</span> <span class="n">size</span><span class="p">):</span>
|
||||
<span class="n">total_sz_mb</span> <span class="o">=</span> <span class="s1">'</span><span class="si">%.2f</span><span class="s1"> MB'</span> <span class="o">%</span> <span class="p">(</span><span class="n">size</span> <span class="o">/</span> <span class="mf">1e6</span><span class="p">)</span>
|
||||
<span class="n">current_sz_mb</span> <span class="o">=</span> <span class="s1">'</span><span class="si">%.2f</span><span class="s1"> MB'</span> <span class="o">%</span> <span class="p">((</span><span class="n">blocknum</span> <span class="o">*</span> <span class="n">bs</span><span class="p">)</span> <span class="o">/</span> <span class="mf">1e6</span><span class="p">)</span>
|
||||
<span class="nb">print</span><span class="p">(</span><span class="s1">'</span><span class="se">\r</span><span class="s1">downloaded </span><span class="si">%s</span><span class="s1"> / </span><span class="si">%s</span><span class="s1">'</span> <span class="o">%</span> <span class="p">(</span><span class="n">current_sz_mb</span><span class="p">,</span> <span class="n">total_sz_mb</span><span class="p">),</span> <span class="n">end</span><span class="o">=</span><span class="s1">''</span><span class="p">)</span>
|
||||
|
|
@ -185,9 +536,12 @@
|
|||
<span class="nb">print</span><span class="p">(</span><span class="s2">""</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="download_file_if_not_exists"><a class="viewcode-back" href="../../quapy.html#quapy.util.download_file_if_not_exists">[docs]</a><span class="k">def</span> <span class="nf">download_file_if_not_exists</span><span class="p">(</span><span class="n">url</span><span class="p">,</span> <span class="n">archive_filename</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="download_file_if_not_exists">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.download_file_if_not_exists">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">download_file_if_not_exists</span><span class="p">(</span><span class="n">url</span><span class="p">,</span> <span class="n">archive_filename</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Dowloads a function (using :meth:`download_file`) if the file does not exist.</span>
|
||||
<span class="sd"> Downloads a file (using :meth:`download_file`) if the file does not exist.</span>
|
||||
|
||||
<span class="sd"> :param url: the url</span>
|
||||
<span class="sd"> :param archive_filename: destination filename</span>
|
||||
|
|
@ -198,7 +552,10 @@
|
|||
<span class="n">download_file</span><span class="p">(</span><span class="n">url</span><span class="p">,</span> <span class="n">archive_filename</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="create_if_not_exist"><a class="viewcode-back" href="../../quapy.html#quapy.util.create_if_not_exist">[docs]</a><span class="k">def</span> <span class="nf">create_if_not_exist</span><span class="p">(</span><span class="n">path</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="create_if_not_exist">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.create_if_not_exist">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">create_if_not_exist</span><span class="p">(</span><span class="n">path</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> An alias to `os.makedirs(path, exist_ok=True)` that also returns the path. This is useful in cases like, e.g.:</span>
|
||||
|
||||
|
|
@ -211,7 +568,10 @@
|
|||
<span class="k">return</span> <span class="n">path</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="get_quapy_home"><a class="viewcode-back" href="../../quapy.html#quapy.util.get_quapy_home">[docs]</a><span class="k">def</span> <span class="nf">get_quapy_home</span><span class="p">():</span>
|
||||
|
||||
<div class="viewcode-block" id="get_quapy_home">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.get_quapy_home">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">get_quapy_home</span><span class="p">():</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Gets the home directory of QuaPy, i.e., the directory where QuaPy saves permanent data, such as dowloaded datasets.</span>
|
||||
<span class="sd"> This directory is `~/quapy_data`</span>
|
||||
|
|
@ -223,7 +583,10 @@
|
|||
<span class="k">return</span> <span class="n">home</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="create_parent_dir"><a class="viewcode-back" href="../../quapy.html#quapy.util.create_parent_dir">[docs]</a><span class="k">def</span> <span class="nf">create_parent_dir</span><span class="p">(</span><span class="n">path</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="create_parent_dir">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.create_parent_dir">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">create_parent_dir</span><span class="p">(</span><span class="n">path</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Creates the parent dir (if any) of a given path, if not exists. E.g., for `./path/to/file.txt`, the path `./path/to`</span>
|
||||
<span class="sd"> is created.</span>
|
||||
|
|
@ -235,7 +598,10 @@
|
|||
<span class="n">os</span><span class="o">.</span><span class="n">makedirs</span><span class="p">(</span><span class="n">parentdir</span><span class="p">,</span> <span class="n">exist_ok</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="save_text_file"><a class="viewcode-back" href="../../quapy.html#quapy.util.save_text_file">[docs]</a><span class="k">def</span> <span class="nf">save_text_file</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">text</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="save_text_file">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.save_text_file">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">save_text_file</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">text</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Saves a text file to disk, given its full path, and creates the parent directory if missing.</span>
|
||||
|
||||
|
|
@ -243,11 +609,14 @@
|
|||
<span class="sd"> :param text: text to save.</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="n">create_parent_dir</span><span class="p">(</span><span class="n">path</span><span class="p">)</span>
|
||||
<span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="n">text</span><span class="p">,</span> <span class="s1">'wt'</span><span class="p">)</span> <span class="k">as</span> <span class="n">fout</span><span class="p">:</span>
|
||||
<span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="s1">'wt'</span><span class="p">)</span> <span class="k">as</span> <span class="n">fout</span><span class="p">:</span>
|
||||
<span class="n">fout</span><span class="o">.</span><span class="n">write</span><span class="p">(</span><span class="n">text</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="pickled_resource"><a class="viewcode-back" href="../../quapy.html#quapy.util.pickled_resource">[docs]</a><span class="k">def</span> <span class="nf">pickled_resource</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">:</span><span class="nb">str</span><span class="p">,</span> <span class="n">generation_func</span><span class="p">:</span><span class="n">callable</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="pickled_resource">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.pickled_resource">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">pickled_resource</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">:</span><span class="nb">str</span><span class="p">,</span> <span class="n">generation_func</span><span class="p">:</span><span class="nb">callable</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Allows for fast reuse of resources that are generated only once by calling generation_func(\\*args). The next times</span>
|
||||
<span class="sd"> this function is invoked, it loads the pickled resource. Example:</span>
|
||||
|
|
@ -266,15 +635,18 @@
|
|||
<span class="k">return</span> <span class="n">generation_func</span><span class="p">(</span><span class="o">*</span><span class="n">args</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="k">if</span> <span class="n">os</span><span class="o">.</span><span class="n">path</span><span class="o">.</span><span class="n">exists</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">):</span>
|
||||
<span class="k">return</span> <span class="n">pickle</span><span class="o">.</span><span class="n">load</span><span class="p">(</span><span class="nb">open</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">,</span> <span class="s1">'rb'</span><span class="p">))</span>
|
||||
<span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">,</span> <span class="s1">'rb'</span><span class="p">)</span> <span class="k">as</span> <span class="n">fin</span><span class="p">:</span>
|
||||
<span class="n">instance</span> <span class="o">=</span> <span class="n">pickle</span><span class="o">.</span><span class="n">load</span><span class="p">(</span><span class="n">fin</span><span class="p">)</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="n">instance</span> <span class="o">=</span> <span class="n">generation_func</span><span class="p">(</span><span class="o">*</span><span class="n">args</span><span class="p">)</span>
|
||||
<span class="n">os</span><span class="o">.</span><span class="n">makedirs</span><span class="p">(</span><span class="nb">str</span><span class="p">(</span><span class="n">Path</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">)</span><span class="o">.</span><span class="n">parent</span><span class="p">),</span> <span class="n">exist_ok</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
|
||||
<span class="n">pickle</span><span class="o">.</span><span class="n">dump</span><span class="p">(</span><span class="n">instance</span><span class="p">,</span> <span class="nb">open</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">,</span> <span class="s1">'wb'</span><span class="p">),</span> <span class="n">pickle</span><span class="o">.</span><span class="n">HIGHEST_PROTOCOL</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">instance</span></div>
|
||||
<span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="n">pickle_path</span><span class="p">,</span> <span class="s1">'wb'</span><span class="p">)</span> <span class="k">as</span> <span class="n">foo</span><span class="p">:</span>
|
||||
<span class="n">pickle</span><span class="o">.</span><span class="n">dump</span><span class="p">(</span><span class="n">instance</span><span class="p">,</span> <span class="n">foo</span><span class="p">,</span> <span class="n">pickle</span><span class="o">.</span><span class="n">HIGHEST_PROTOCOL</span><span class="p">)</span>
|
||||
<span class="k">return</span> <span class="n">instance</span></div>
|
||||
|
||||
|
||||
<span class="k">def</span> <span class="nf">_check_sample_size</span><span class="p">(</span><span class="n">sample_size</span><span class="p">):</span>
|
||||
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">_check_sample_size</span><span class="p">(</span><span class="n">sample_size</span><span class="p">):</span>
|
||||
<span class="k">if</span> <span class="n">sample_size</span> <span class="ow">is</span> <span class="kc">None</span><span class="p">:</span>
|
||||
<span class="k">assert</span> <span class="n">qp</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">'SAMPLE_SIZE'</span><span class="p">]</span> <span class="ow">is</span> <span class="ow">not</span> <span class="kc">None</span><span class="p">,</span> \
|
||||
<span class="s1">'error: sample_size set to None, and cannot be resolved from the environment'</span>
|
||||
|
|
@ -284,7 +656,34 @@
|
|||
<span class="k">return</span> <span class="n">sample_size</span>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="EarlyStop"><a class="viewcode-back" href="../../quapy.html#quapy.util.EarlyStop">[docs]</a><span class="k">class</span> <span class="nc">EarlyStop</span><span class="p">:</span>
|
||||
<div class="viewcode-block" id="load_report">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.load_report">[docs]</a>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">load_report</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">as_dict</span><span class="o">=</span><span class="kc">False</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">str2prev_arr</span><span class="p">(</span><span class="n">strprev</span><span class="p">):</span>
|
||||
<span class="n">within</span> <span class="o">=</span> <span class="n">strprev</span><span class="o">.</span><span class="n">strip</span><span class="p">(</span><span class="s1">'[]'</span><span class="p">)</span><span class="o">.</span><span class="n">split</span><span class="p">()</span>
|
||||
<span class="n">float_list</span> <span class="o">=</span> <span class="p">[</span><span class="nb">float</span><span class="p">(</span><span class="n">p</span><span class="p">)</span> <span class="k">for</span> <span class="n">p</span> <span class="ow">in</span> <span class="n">within</span><span class="p">]</span>
|
||||
<span class="n">float_list</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">]</span> <span class="o">=</span> <span class="mf">1.</span> <span class="o">-</span> <span class="nb">sum</span><span class="p">(</span><span class="n">float_list</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">])</span>
|
||||
<span class="k">return</span> <span class="n">np</span><span class="o">.</span><span class="n">asarray</span><span class="p">(</span><span class="n">float_list</span><span class="p">)</span>
|
||||
|
||||
<span class="n">df</span> <span class="o">=</span> <span class="n">pd</span><span class="o">.</span><span class="n">read_csv</span><span class="p">(</span><span class="n">path</span><span class="p">,</span> <span class="n">index_col</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
|
||||
<span class="n">df</span><span class="p">[</span><span class="s1">'true-prev'</span><span class="p">]</span> <span class="o">=</span> <span class="n">df</span><span class="p">[</span><span class="s1">'true-prev'</span><span class="p">]</span><span class="o">.</span><span class="n">apply</span><span class="p">(</span><span class="n">str2prev_arr</span><span class="p">)</span>
|
||||
<span class="n">df</span><span class="p">[</span><span class="s1">'estim-prev'</span><span class="p">]</span> <span class="o">=</span> <span class="n">df</span><span class="p">[</span><span class="s1">'estim-prev'</span><span class="p">]</span><span class="o">.</span><span class="n">apply</span><span class="p">(</span><span class="n">str2prev_arr</span><span class="p">)</span>
|
||||
<span class="k">if</span> <span class="n">as_dict</span><span class="p">:</span>
|
||||
<span class="n">d</span> <span class="o">=</span> <span class="p">{}</span>
|
||||
<span class="k">for</span> <span class="n">col</span> <span class="ow">in</span> <span class="n">df</span><span class="o">.</span><span class="n">columns</span><span class="o">.</span><span class="n">values</span><span class="p">:</span>
|
||||
<span class="n">vals</span> <span class="o">=</span> <span class="n">df</span><span class="p">[</span><span class="n">col</span><span class="p">]</span><span class="o">.</span><span class="n">values</span>
|
||||
<span class="k">if</span> <span class="n">col</span> <span class="ow">in</span> <span class="p">[</span><span class="s1">'true-prev'</span><span class="p">,</span> <span class="s1">'estim-prev'</span><span class="p">]:</span>
|
||||
<span class="n">vals</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">vstack</span><span class="p">(</span><span class="n">vals</span><span class="p">)</span>
|
||||
<span class="n">d</span><span class="p">[</span><span class="n">col</span><span class="p">]</span> <span class="o">=</span> <span class="n">vals</span>
|
||||
<span class="k">return</span> <span class="n">d</span>
|
||||
<span class="k">else</span><span class="p">:</span>
|
||||
<span class="k">return</span> <span class="n">df</span></div>
|
||||
|
||||
|
||||
|
||||
<div class="viewcode-block" id="EarlyStop">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.EarlyStop">[docs]</a>
|
||||
<span class="k">class</span><span class="w"> </span><span class="nc">EarlyStop</span><span class="p">:</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> A class implementing the early-stopping condition typically used for training neural networks.</span>
|
||||
|
||||
|
|
@ -309,7 +708,7 @@
|
|||
<span class="sd"> :ivar IMPROVED: flag (boolean) indicating whether there was an improvement in the last call</span>
|
||||
<span class="sd"> """</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">patience</span><span class="p">,</span> <span class="n">lower_is_better</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">patience</span><span class="p">,</span> <span class="n">lower_is_better</span><span class="o">=</span><span class="kc">True</span><span class="p">):</span>
|
||||
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">PATIENCE_LIMIT</span> <span class="o">=</span> <span class="n">patience</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">better</span> <span class="o">=</span> <span class="k">lambda</span> <span class="n">a</span><span class="p">,</span><span class="n">b</span><span class="p">:</span> <span class="n">a</span><span class="o"><</span><span class="n">b</span> <span class="k">if</span> <span class="n">lower_is_better</span> <span class="k">else</span> <span class="n">a</span><span class="o">></span><span class="n">b</span>
|
||||
|
|
@ -319,7 +718,7 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">STOP</span> <span class="o">=</span> <span class="kc">False</span>
|
||||
<span class="bp">self</span><span class="o">.</span><span class="n">IMPROVED</span> <span class="o">=</span> <span class="kc">False</span>
|
||||
|
||||
<span class="k">def</span> <span class="fm">__call__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">watch_score</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="fm">__call__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">watch_score</span><span class="p">,</span> <span class="n">epoch</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Commits the new score found in epoch `epoch`. If the score improves over the best score found so far, then</span>
|
||||
<span class="sd"> the patiente counter gets reset. If otherwise, the patience counter is decreased, and in case it reachs 0,</span>
|
||||
|
|
@ -339,8 +738,11 @@
|
|||
<span class="bp">self</span><span class="o">.</span><span class="n">STOP</span> <span class="o">=</span> <span class="kc">True</span></div>
|
||||
|
||||
|
||||
<div class="viewcode-block" id="timeout"><a class="viewcode-back" href="../../quapy.html#quapy.util.timeout">[docs]</a><span class="nd">@contextlib</span><span class="o">.</span><span class="n">contextmanager</span>
|
||||
<span class="k">def</span> <span class="nf">timeout</span><span class="p">(</span><span class="n">seconds</span><span class="p">):</span>
|
||||
|
||||
<div class="viewcode-block" id="timeout">
|
||||
<a class="viewcode-back" href="../../quapy.html#quapy.util.timeout">[docs]</a>
|
||||
<span class="nd">@contextlib</span><span class="o">.</span><span class="n">contextmanager</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">timeout</span><span class="p">(</span><span class="n">seconds</span><span class="p">):</span>
|
||||
<span class="w"> </span><span class="sd">"""</span>
|
||||
<span class="sd"> Opens a context that will launch an exception if not closed after a given number of seconds</span>
|
||||
|
||||
|
|
@ -359,7 +761,7 @@
|
|||
<span class="sd"> :param seconds: number of seconds, set to <=0 to ignore the timer</span>
|
||||
<span class="sd"> """</span>
|
||||
<span class="k">if</span> <span class="n">seconds</span> <span class="o">></span> <span class="mi">0</span><span class="p">:</span>
|
||||
<span class="k">def</span> <span class="nf">handler</span><span class="p">(</span><span class="n">signum</span><span class="p">,</span> <span class="n">frame</span><span class="p">):</span>
|
||||
<span class="k">def</span><span class="w"> </span><span class="nf">handler</span><span class="p">(</span><span class="n">signum</span><span class="p">,</span> <span class="n">frame</span><span class="p">):</span>
|
||||
<span class="k">raise</span> <span class="ne">TimeoutError</span><span class="p">()</span>
|
||||
|
||||
<span class="n">signal</span><span class="o">.</span><span class="n">signal</span><span class="p">(</span><span class="n">signal</span><span class="o">.</span><span class="n">SIGALRM</span><span class="p">,</span> <span class="n">handler</span><span class="p">)</span>
|
||||
|
|
@ -370,33 +772,78 @@
|
|||
<span class="k">if</span> <span class="n">seconds</span> <span class="o">></span> <span class="mi">0</span><span class="p">:</span>
|
||||
<span class="n">signal</span><span class="o">.</span><span class="n">alarm</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span></div>
|
||||
|
||||
|
||||
</pre></div>
|
||||
|
||||
</div>
|
||||
</article>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<footer class="prev-next-footer d-print-none">
|
||||
|
||||
<div class="prev-next-area">
|
||||
</div>
|
||||
</footer>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
|
||||
</div>
|
||||
<footer>
|
||||
|
||||
<hr/>
|
||||
|
||||
<div role="contentinfo">
|
||||
<p>© Copyright 2024, Alejandro Moreo.</p>
|
||||
<footer class="bd-footer-content">
|
||||
|
||||
</footer>
|
||||
|
||||
</main>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Scripts loaded after <body> so the DOM is not blocked -->
|
||||
<script defer src="../../_static/scripts/bootstrap.js?digest=90905a2f556bf617f1a9"></script>
|
||||
<script defer src="../../_static/scripts/pydata-sphinx-theme.js?digest=90905a2f556bf617f1a9"></script>
|
||||
|
||||
Built with <a href="https://www.sphinx-doc.org/">Sphinx</a> using a
|
||||
<a href="https://github.com/readthedocs/sphinx_rtd_theme">theme</a>
|
||||
provided by <a href="https://readthedocs.org">Read the Docs</a>.
|
||||
|
||||
<footer class="bd-footer">
|
||||
<div class="bd-footer__inner bd-page-width">
|
||||
|
||||
<div class="footer-items__start">
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</footer>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</div>
|
||||
<script>
|
||||
jQuery(function () {
|
||||
SphinxRtdTheme.Navigation.enable(true);
|
||||
});
|
||||
</script>
|
||||
<p class="copyright">
|
||||
|
||||
© Copyright 2024, Alejandro Moreo.
|
||||
<br/>
|
||||
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div class="footer-item">
|
||||
|
||||
</body>
|
||||
<p class="sphinx-version">
|
||||
Created using <a href="https://www.sphinx-doc.org/">Sphinx</a> 9.0.4.
|
||||
<br/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
|
||||
<div class="footer-items__end">
|
||||
|
||||
<div class="footer-item">
|
||||
<p class="theme-version">
|
||||
<!-- # L10n: Setting the PST URL as an argument as this does not need to be localized -->
|
||||
Built with the <a href="https://pydata-sphinx-theme.readthedocs.io/en/stable/index.html">PyData Sphinx Theme</a> 0.20.0.
|
||||
</p></div>
|
||||
|
||||
</div>
|
||||
|
||||
</div>
|
||||
|
||||
</footer>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -1,20 +1,9 @@
|
|||
/*
|
||||
* _sphinx_javascript_frameworks_compat.js
|
||||
* ~~~~~~~~~~
|
||||
*
|
||||
* Compatability shim for jQuery and underscores.js.
|
||||
*
|
||||
* WILL BE REMOVED IN Sphinx 6.0
|
||||
* xref RemovedInSphinx60Warning
|
||||
/* Compatability shim for jQuery and underscores.js.
|
||||
*
|
||||
* Copyright Sphinx contributors
|
||||
* Released under the two clause BSD licence
|
||||
*/
|
||||
|
||||
/**
|
||||
* select a different prefix for underscore
|
||||
*/
|
||||
$u = _.noConflict();
|
||||
|
||||
|
||||
/**
|
||||
* small helper function to urldecode strings
|
||||
*
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
.clearfix{*zoom:1}.clearfix:after,.clearfix:before{display:table;content:""}.clearfix:after{clear:both}@font-face{font-family:FontAwesome;font-style:normal;font-weight:400;src:url(fonts/fontawesome-webfont.eot?674f50d287a8c48dc19ba404d20fe713?#iefix) format("embedded-opentype"),url(fonts/fontawesome-webfont.woff2?af7ae505a9eed503f8b8e6982036873e) format("woff2"),url(fonts/fontawesome-webfont.woff?fee66e712a8a08eef5805a46892932ad) format("woff"),url(fonts/fontawesome-webfont.ttf?b06871f281fee6b241d60582ae9369b9) format("truetype"),url(fonts/fontawesome-webfont.svg?912ec66d7572ff821749319396470bde#FontAwesome) format("svg")}.fa:before{font-family:FontAwesome;font-style:normal;font-weight:400;line-height:1}.fa:before,a .fa{text-decoration:inherit}.fa:before,a .fa,li .fa{display:inline-block}li .fa-large:before{width:1.875em}ul.fas{list-style-type:none;margin-left:2em;text-indent:-.8em}ul.fas li .fa{width:.8em}ul.fas li .fa-large:before{vertical-align:baseline}.fa-book:before,.icon-book:before{content:"\f02d"}.fa-caret-down:before,.icon-caret-down:before{content:"\f0d7"}.fa-caret-up:before,.icon-caret-up:before{content:"\f0d8"}.fa-caret-left:before,.icon-caret-left:before{content:"\f0d9"}.fa-caret-right:before,.icon-caret-right:before{content:"\f0da"}.rst-versions{position:fixed;bottom:0;left:0;width:300px;color:#fcfcfc;background:#1f1d1d;font-family:Lato,proxima-nova,Helvetica Neue,Arial,sans-serif;z-index:400}.rst-versions a{color:#2980b9;text-decoration:none}.rst-versions .rst-badge-small{display:none}.rst-versions .rst-current-version{padding:12px;background-color:#272525;display:block;text-align:right;font-size:90%;cursor:pointer;color:#27ae60}.rst-versions .rst-current-version:after{clear:both;content:"";display:block}.rst-versions .rst-current-version .fa{color:#fcfcfc}.rst-versions .rst-current-version .fa-book,.rst-versions .rst-current-version .icon-book{float:left}.rst-versions .rst-current-version.rst-out-of-date{background-color:#e74c3c;color:#fff}.rst-versions .rst-current-version.rst-active-old-version{background-color:#f1c40f;color:#000}.rst-versions.shift-up{height:auto;max-height:100%;overflow-y:scroll}.rst-versions.shift-up .rst-other-versions{display:block}.rst-versions .rst-other-versions{font-size:90%;padding:12px;color:grey;display:none}.rst-versions .rst-other-versions hr{display:block;height:1px;border:0;margin:20px 0;padding:0;border-top:1px solid #413d3d}.rst-versions .rst-other-versions dd{display:inline-block;margin:0}.rst-versions .rst-other-versions dd a{display:inline-block;padding:6px;color:#fcfcfc}.rst-versions.rst-badge{width:auto;bottom:20px;right:20px;left:auto;border:none;max-width:300px;max-height:90%}.rst-versions.rst-badge .fa-book,.rst-versions.rst-badge .icon-book{float:none;line-height:30px}.rst-versions.rst-badge.shift-up .rst-current-version{text-align:right}.rst-versions.rst-badge.shift-up .rst-current-version .fa-book,.rst-versions.rst-badge.shift-up .rst-current-version .icon-book{float:left}.rst-versions.rst-badge>.rst-current-version{width:auto;height:30px;line-height:30px;padding:0 6px;display:block;text-align:center}@media screen and (max-width:768px){.rst-versions{width:85%;display:none}.rst-versions.shift{display:block}}
|
||||
.clearfix{*zoom:1}.clearfix:after,.clearfix:before{display:table;content:""}.clearfix:after{clear:both}@font-face{font-family:FontAwesome;font-style:normal;font-weight:400;src:url(fonts/fontawesome-webfont.eot?674f50d287a8c48dc19ba404d20fe713?#iefix) format("embedded-opentype"),url(fonts/fontawesome-webfont.woff2?af7ae505a9eed503f8b8e6982036873e) format("woff2"),url(fonts/fontawesome-webfont.woff?fee66e712a8a08eef5805a46892932ad) format("woff"),url(fonts/fontawesome-webfont.ttf?b06871f281fee6b241d60582ae9369b9) format("truetype"),url(fonts/fontawesome-webfont.svg?912ec66d7572ff821749319396470bde#FontAwesome) format("svg")}.fa:before{font-family:FontAwesome;font-style:normal;font-weight:400;line-height:1}.fa:before,a .fa{text-decoration:inherit}.fa:before,a .fa,li .fa{display:inline-block}li .fa-large:before{width:1.875em}ul.fas{list-style-type:none;margin-left:2em;text-indent:-.8em}ul.fas li .fa{width:.8em}ul.fas li .fa-large:before{vertical-align:baseline}.fa-book:before,.icon-book:before{content:"\f02d"}.fa-caret-down:before,.icon-caret-down:before{content:"\f0d7"}.fa-caret-up:before,.icon-caret-up:before{content:"\f0d8"}.fa-caret-left:before,.icon-caret-left:before{content:"\f0d9"}.fa-caret-right:before,.icon-caret-right:before{content:"\f0da"}.rst-versions{position:fixed;bottom:0;left:0;width:300px;color:#fcfcfc;background:#1f1d1d;font-family:Lato,proxima-nova,Helvetica Neue,Arial,sans-serif;z-index:400}.rst-versions a{color:#2980b9;text-decoration:none}.rst-versions .rst-badge-small{display:none}.rst-versions .rst-current-version{padding:12px;background-color:#272525;display:block;text-align:right;font-size:90%;cursor:pointer;color:#27ae60}.rst-versions .rst-current-version:after{clear:both;content:"";display:block}.rst-versions .rst-current-version .fa{color:#fcfcfc}.rst-versions .rst-current-version .fa-book,.rst-versions .rst-current-version .icon-book{float:left}.rst-versions .rst-current-version.rst-out-of-date{background-color:#e74c3c;color:#fff}.rst-versions .rst-current-version.rst-active-old-version{background-color:#f1c40f;color:#000}.rst-versions.shift-up{height:auto;max-height:100%;overflow-y:scroll}.rst-versions.shift-up .rst-other-versions{display:block}.rst-versions .rst-other-versions{font-size:90%;padding:12px;color:grey;display:none}.rst-versions .rst-other-versions hr{display:block;height:1px;border:0;margin:20px 0;padding:0;border-top:1px solid #413d3d}.rst-versions .rst-other-versions dd{display:inline-block;margin:0}.rst-versions .rst-other-versions dd a{display:inline-block;padding:6px;color:#fcfcfc}.rst-versions .rst-other-versions .rtd-current-item{font-weight:700}.rst-versions.rst-badge{width:auto;bottom:20px;right:20px;left:auto;border:none;max-width:300px;max-height:90%}.rst-versions.rst-badge .fa-book,.rst-versions.rst-badge .icon-book{float:none;line-height:30px}.rst-versions.rst-badge.shift-up .rst-current-version{text-align:right}.rst-versions.rst-badge.shift-up .rst-current-version .fa-book,.rst-versions.rst-badge.shift-up .rst-current-version .icon-book{float:left}.rst-versions.rst-badge>.rst-current-version{width:auto;height:30px;line-height:30px;padding:0 6px;display:block;text-align:center}@media screen and (max-width:768px){.rst-versions{width:85%;display:none}.rst-versions.shift{display:block}}#flyout-search-form{padding:6px}
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
/* Highlighting utilities for Sphinx HTML documentation. */
|
||||
"use strict";
|
||||
|
||||
const SPHINX_HIGHLIGHT_ENABLED = true
|
||||
const SPHINX_HIGHLIGHT_ENABLED = true;
|
||||
|
||||
/**
|
||||
* highlight a given string on a node by wrapping it in
|
||||
|
|
@ -13,9 +13,9 @@ const _highlight = (node, addItems, text, className) => {
|
|||
const parent = node.parentNode;
|
||||
const pos = val.toLowerCase().indexOf(text);
|
||||
if (
|
||||
pos >= 0 &&
|
||||
!parent.classList.contains(className) &&
|
||||
!parent.classList.contains("nohighlight")
|
||||
pos >= 0
|
||||
&& !parent.classList.contains(className)
|
||||
&& !parent.classList.contains("nohighlight")
|
||||
) {
|
||||
let span;
|
||||
|
||||
|
|
@ -29,19 +29,18 @@ const _highlight = (node, addItems, text, className) => {
|
|||
}
|
||||
|
||||
span.appendChild(document.createTextNode(val.substr(pos, text.length)));
|
||||
parent.insertBefore(
|
||||
span,
|
||||
parent.insertBefore(
|
||||
document.createTextNode(val.substr(pos + text.length)),
|
||||
node.nextSibling
|
||||
)
|
||||
);
|
||||
const rest = document.createTextNode(val.substr(pos + text.length));
|
||||
parent.insertBefore(span, parent.insertBefore(rest, node.nextSibling));
|
||||
node.nodeValue = val.substr(0, pos);
|
||||
/* There may be more occurrences of search term in this node. So call this
|
||||
* function recursively on the remaining fragment.
|
||||
*/
|
||||
_highlight(rest, addItems, text, className);
|
||||
|
||||
if (isInSVG) {
|
||||
const rect = document.createElementNS(
|
||||
"http://www.w3.org/2000/svg",
|
||||
"rect"
|
||||
"rect",
|
||||
);
|
||||
const bbox = parent.getBBox();
|
||||
rect.x.baseVal.value = bbox.x;
|
||||
|
|
@ -60,7 +59,7 @@ const _highlightText = (thisNode, text, className) => {
|
|||
let addItems = [];
|
||||
_highlight(thisNode, addItems, text, className);
|
||||
addItems.forEach((obj) =>
|
||||
obj.parent.insertAdjacentElement("beforebegin", obj.target)
|
||||
obj.parent.insertAdjacentElement("beforebegin", obj.target),
|
||||
);
|
||||
};
|
||||
|
||||
|
|
@ -68,25 +67,31 @@ const _highlightText = (thisNode, text, className) => {
|
|||
* Small JavaScript module for the documentation.
|
||||
*/
|
||||
const SphinxHighlight = {
|
||||
|
||||
/**
|
||||
* highlight the search words provided in localstorage in the text
|
||||
*/
|
||||
highlightSearchWords: () => {
|
||||
if (!SPHINX_HIGHLIGHT_ENABLED) return; // bail if no highlight
|
||||
if (!SPHINX_HIGHLIGHT_ENABLED) return; // bail if no highlight
|
||||
|
||||
// get and clear terms from localstorage
|
||||
const url = new URL(window.location);
|
||||
const highlight =
|
||||
localStorage.getItem("sphinx_highlight_terms")
|
||||
|| url.searchParams.get("highlight")
|
||||
|| "";
|
||||
localStorage.removeItem("sphinx_highlight_terms")
|
||||
url.searchParams.delete("highlight");
|
||||
window.history.replaceState({}, "", url);
|
||||
localStorage.getItem("sphinx_highlight_terms")
|
||||
|| url.searchParams.get("highlight")
|
||||
|| "";
|
||||
localStorage.removeItem("sphinx_highlight_terms");
|
||||
// Update history only if '?highlight' is present; otherwise it
|
||||
// clears text fragments (not set in window.location by the browser)
|
||||
if (url.searchParams.has("highlight")) {
|
||||
url.searchParams.delete("highlight");
|
||||
window.history.replaceState({}, "", url);
|
||||
}
|
||||
|
||||
// get individual terms from highlight string
|
||||
const terms = highlight.toLowerCase().split(/\s+/).filter(x => x);
|
||||
const terms = highlight
|
||||
.toLowerCase()
|
||||
.split(/\s+/)
|
||||
.filter((x) => x);
|
||||
if (terms.length === 0) return; // nothing to do
|
||||
|
||||
// There should never be more than one element matching "div.body"
|
||||
|
|
@ -102,11 +107,11 @@ const SphinxHighlight = {
|
|||
document
|
||||
.createRange()
|
||||
.createContextualFragment(
|
||||
'<p class="highlight-link">' +
|
||||
'<a href="javascript:SphinxHighlight.hideSearchWords()">' +
|
||||
_("Hide Search Matches") +
|
||||
"</a></p>"
|
||||
)
|
||||
'<p class="highlight-link">'
|
||||
+ '<a href="javascript:SphinxHighlight.hideSearchWords()">'
|
||||
+ _("Hide Search Matches")
|
||||
+ "</a></p>",
|
||||
),
|
||||
);
|
||||
},
|
||||
|
||||
|
|
@ -120,7 +125,7 @@ const SphinxHighlight = {
|
|||
document
|
||||
.querySelectorAll("span.highlighted")
|
||||
.forEach((el) => el.classList.remove("highlighted"));
|
||||
localStorage.removeItem("sphinx_highlight_terms")
|
||||
localStorage.removeItem("sphinx_highlight_terms");
|
||||
},
|
||||
|
||||
initEscapeListener: () => {
|
||||
|
|
@ -129,10 +134,15 @@ const SphinxHighlight = {
|
|||
|
||||
document.addEventListener("keydown", (event) => {
|
||||
// bail for input elements
|
||||
if (BLACKLISTED_KEY_CONTROL_ELEMENTS.has(document.activeElement.tagName)) return;
|
||||
if (BLACKLISTED_KEY_CONTROL_ELEMENTS.has(document.activeElement.tagName))
|
||||
return;
|
||||
// bail with special keys
|
||||
if (event.shiftKey || event.altKey || event.ctrlKey || event.metaKey) return;
|
||||
if (DOCUMENTATION_OPTIONS.ENABLE_SEARCH_SHORTCUTS && (event.key === "Escape")) {
|
||||
if (event.shiftKey || event.altKey || event.ctrlKey || event.metaKey)
|
||||
return;
|
||||
if (
|
||||
DOCUMENTATION_OPTIONS.ENABLE_SEARCH_SHORTCUTS
|
||||
&& event.key === "Escape"
|
||||
) {
|
||||
SphinxHighlight.hideSearchWords();
|
||||
event.preventDefault();
|
||||
}
|
||||
|
|
@ -140,5 +150,10 @@ const SphinxHighlight = {
|
|||
},
|
||||
};
|
||||
|
||||
_ready(SphinxHighlight.highlightSearchWords);
|
||||
_ready(SphinxHighlight.initEscapeListener);
|
||||
_ready(() => {
|
||||
/* Do not call highlightSearchWords() when we are on the search page.
|
||||
* It will highlight words from the *previous* search query.
|
||||
*/
|
||||
if (typeof Search === "undefined") SphinxHighlight.highlightSearchWords();
|
||||
SphinxHighlight.initEscapeListener();
|
||||
});
|
||||
|
|
|
|||
|
|
@ -0,0 +1,53 @@
|
|||
.navbar-brand.logo .title.logo__title {
|
||||
display: none;
|
||||
}
|
||||
|
||||
.navbar-brand img {
|
||||
max-height: 2.2rem;
|
||||
width: auto;
|
||||
}
|
||||
|
||||
.navbar-brand.logo .title.logo__title {
|
||||
font-size: 1.05rem;
|
||||
line-height: 1.15;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
@media (max-width: 1200px) {
|
||||
.bd-header .navbar-header-items {
|
||||
min-width: 0;
|
||||
}
|
||||
|
||||
.bd-header .navbar-header-items__center {
|
||||
overflow-x: auto;
|
||||
}
|
||||
|
||||
.bd-header .bd-navbar-elements {
|
||||
flex-wrap: nowrap;
|
||||
}
|
||||
}
|
||||
|
||||
.hero-copy {
|
||||
font-size: 1.15rem;
|
||||
line-height: 1.7;
|
||||
max-width: 56rem;
|
||||
margin: 0 0 1.5rem 0;
|
||||
}
|
||||
|
||||
.landing-grid {
|
||||
margin: 1.2rem 0 2rem 0;
|
||||
}
|
||||
|
||||
.landing-card {
|
||||
border-radius: 1rem;
|
||||
border: 1px solid var(--pst-color-border, #d0d7de);
|
||||
box-shadow: 0 10px 24px rgba(15, 23, 42, 0.08);
|
||||
}
|
||||
|
||||
.landing-card .sd-card-title {
|
||||
font-size: 1.1rem;
|
||||
}
|
||||
|
||||
[data-theme="dark"] .landing-card {
|
||||
box-shadow: 0 10px 24px rgba(0, 0, 0, 0.22);
|
||||
}
|
||||
|
After Width: | Height: | Size: 18 KiB |
|
After Width: | Height: | Size: 18 KiB |
|
|
@ -44,12 +44,15 @@ extensions = [
|
|||
'sphinx.ext.napoleon',
|
||||
'sphinx.ext.intersphinx',
|
||||
'myst_parser',
|
||||
'sphinx_design',
|
||||
]
|
||||
|
||||
autosectionlabel_prefix_document = True
|
||||
|
||||
source_suffix = ['.rst', '.md']
|
||||
|
||||
myst_enable_extensions = ['colon_fence']
|
||||
|
||||
templates_path = ['_templates']
|
||||
|
||||
# List of patterns, relative to source directory, that match files and
|
||||
|
|
@ -61,10 +64,26 @@ exclude_patterns = ['_build', 'Thumbs.db', '.DS_Store']
|
|||
# -- Options for HTML output -------------------------------------------------
|
||||
# https://www.sphinx-doc.org/en/master/usage/configuration.html#options-for-html-output
|
||||
|
||||
html_theme = 'sphinx_rtd_theme'
|
||||
#html_theme = 'sphinx_rtd_theme'
|
||||
html_theme = 'pydata_sphinx_theme'
|
||||
# html_theme = 'furo'
|
||||
# need to be installed: pip install furo (not working...)
|
||||
# html_static_path = ['_static']
|
||||
html_static_path = ['_static']
|
||||
html_css_files = ['custom.css']
|
||||
html_theme_options = {
|
||||
'logo': {
|
||||
'image_light': '_static/quapy_logo.png',
|
||||
'image_dark': '_static/quapy_logo_dark.png',
|
||||
},
|
||||
'icon_links': [
|
||||
{
|
||||
'name': 'GitHub',
|
||||
'url': 'https://github.com/HLT-ISTI/QuaPy',
|
||||
'icon': 'fa-brands fa-github',
|
||||
'type': 'fontawesome',
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
# intersphinx configuration
|
||||
intersphinx_mapping = {
|
||||
|
|
|
|||
|
|
@ -1,16 +1,68 @@
|
|||
```{toctree}
|
||||
:hidden:
|
||||
|
||||
self
|
||||
Home <self>
|
||||
manuals
|
||||
API <quapy>
|
||||
```
|
||||
|
||||
# Quickstart
|
||||
# QuaPy
|
||||
|
||||
QuaPy is an open source framework for quantification (a.k.a. supervised prevalence estimation, or learning to quantify) written in Python.
|
||||
```{div} hero-copy
|
||||
QuaPy is an open-source Python framework for quantification, also known as
|
||||
supervised prevalence estimation or learning to quantify. It is designed with
|
||||
research and experimental analysis in mind, and combines datasets, protocols,
|
||||
evaluation measures, visualization tools, and a broad collection of
|
||||
quantification methods in a single workflow.
|
||||
```
|
||||
|
||||
QuaPy is based on the concept of "data sample", and provides implementations of the most important aspects of the quantification workflow, such as (baseline and advanced) quantification methods, quantification-oriented model selection mechanisms, evaluation measures, and evaluations protocols used for evaluating quantification methods. QuaPy also makes available commonly used datasets, and offers visualization tools for facilitating the analysis and interpretation of the experimental results.
|
||||
`````{grid} 1 1 2 2
|
||||
:gutter: 3
|
||||
:class-container: landing-grid
|
||||
|
||||
QuaPy is hosted on GitHub at [https://github.com/HLT-ISTI/QuaPy](https://github.com/HLT-ISTI/QuaPy).
|
||||
````{grid-item-card} Quickstart
|
||||
:class-card: landing-card
|
||||
Install QuaPy and run your first quantifier in a few lines of code.
|
||||
+++
|
||||
```{button-link} #installation
|
||||
:color: primary
|
||||
Get Started
|
||||
```
|
||||
````
|
||||
|
||||
````{grid-item-card} Manuals
|
||||
:class-card: landing-card
|
||||
Hands-on guides with methodological context, literature pointers, and reproducible workflows.
|
||||
+++
|
||||
```{button-ref} manuals
|
||||
:ref-type: doc
|
||||
:color: primary
|
||||
Open Manuals
|
||||
```
|
||||
````
|
||||
|
||||
````{grid-item-card} API
|
||||
:class-card: landing-card
|
||||
Browse the full reference for `quapy`, including methods, datasets, utilities, and research-oriented extensions.
|
||||
+++
|
||||
```{button-ref} quapy
|
||||
:ref-type: doc
|
||||
:color: primary
|
||||
Browse API
|
||||
```
|
||||
````
|
||||
|
||||
````{grid-item-card} GitHub
|
||||
:class-card: landing-card
|
||||
Explore the source code, open issues, and current development branch activity.
|
||||
+++
|
||||
```{button-link} https://github.com/HLT-ISTI/QuaPy
|
||||
:color: primary
|
||||
Open GitHub
|
||||
```
|
||||
````
|
||||
|
||||
`````
|
||||
|
||||
## Installation
|
||||
|
||||
|
|
@ -18,16 +70,37 @@ QuaPy is hosted on GitHub at [https://github.com/HLT-ISTI/QuaPy](https://github.
|
|||
pip install quapy
|
||||
```
|
||||
|
||||
## Usage
|
||||
## Why QuaPy
|
||||
|
||||
The following script fetches a dataset of tweets, trains, applies, and evaluates a quantifier based on the *Adjusted Classify & Count* quantification method, using, as the evaluation measure, the *Mean Absolute Error* (MAE) between the predicted and the true class prevalence values of the test set:
|
||||
QuaPy is built around the concept of a data sample and supports the main tasks
|
||||
in the quantification workflow: training quantifiers, generating evaluation
|
||||
samples, measuring quantification error, selecting models under distribution
|
||||
shift, and visualizing experimental behaviour. The framework is especially
|
||||
suited for research settings, where one often needs not only implementations,
|
||||
but also methodological context, literature links, and reproducible evaluation
|
||||
procedures.
|
||||
|
||||
Some of the main features are:
|
||||
|
||||
* Implementation of many popular quantification methods, including Classify & Count and its variants,
|
||||
Expectation Maximization, HDy, QuaNet, quantification ensembles, and Bayesian extensions.
|
||||
* Evaluation protocols for generating test samples under prior probability shift.
|
||||
* A broad set of quantification-oriented evaluation metrics.
|
||||
* Ready-to-use textual, numeric, and benchmark competition datasets.
|
||||
* Method documentation that points back to the relevant literature and original papers.
|
||||
* Native support for binary and single-label multiclass quantification.
|
||||
* Visualization tools for analysing predictions, drift, confidence regions, and ternary prevalences.
|
||||
|
||||
## First Example
|
||||
|
||||
The following script fetches a binary dataset, trains an Adjusted Classify & Count quantifier,
|
||||
and evaluates the resulting prevalence prediction with Mean Absolute Error.
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
|
||||
training, test = qp.datasets.fetch_UCIBinaryDataset("yeast").train_test
|
||||
|
||||
# create an "Adjusted Classify & Count" quantifier
|
||||
model = qp.method.aggregative.ACC()
|
||||
Xtr, ytr = training.Xy
|
||||
model.fit(Xtr, ytr)
|
||||
|
|
@ -39,43 +112,14 @@ error = qp.error.mae(true_prevalence, estim_prevalence)
|
|||
print(f'Mean Absolute Error (MAE)={error:.3f}')
|
||||
```
|
||||
|
||||
Quantification is useful in scenarios characterized by prior probability shift. In other words, we would be little interested in estimating the class prevalence values of the test set if we could assume the IID assumption to hold, as this prevalence would be roughly equivalent to the class prevalence of the training set. For this reason, any quantification model should be tested across many samples, even ones characterized by class prevalence values different or very different from those found in the training set. QuaPy implements sampling procedures and evaluation protocols that automate this workflow. See the [](./manuals) for detailed examples.
|
||||
|
||||
## Manuals
|
||||
|
||||
The following manuals illustrate several aspects of QuaPy through examples:
|
||||
|
||||
```{toctree}
|
||||
:maxdepth: 3
|
||||
|
||||
manuals
|
||||
```
|
||||
|
||||
```{toctree}
|
||||
:hidden:
|
||||
|
||||
API <quapy>
|
||||
```
|
||||
|
||||
## Features
|
||||
|
||||
* Implementation of many popular quantification methods (Classify-&-Count and its variants, Expectation Maximization,
|
||||
quantification methods based on structured output learning, HDy, QuaNet, quantification ensembles, among others).
|
||||
* Versatile functionality for performing evaluation based on sampling generation protocols (e.g., APP, NPP, etc.).
|
||||
* Implementation of most commonly used evaluation metrics (e.g., AE, RAE, NAE, NRAE, SE, KLD, NKLD, etc.).
|
||||
* Datasets frequently used in quantification (textual and numeric), including:
|
||||
* 32 UCI Machine Learning datasets.
|
||||
* 11 Twitter quantification-by-sentiment datasets.
|
||||
* 3 product reviews quantification-by-sentiment datasets.
|
||||
* 4 tasks from LeQua 2022 competition and 4 tasks from LeQua 2024 competition
|
||||
* IFCB for Plancton quantification
|
||||
* Native support for binary and single-label multiclass quantification scenarios.
|
||||
* Model selection functionality that minimizes quantification-oriented loss functions.
|
||||
* Visualization tools for analysing the experimental results.
|
||||
Quantification is especially useful when the class prevalence of the test data
|
||||
may differ from that of the training data. QuaPy implements protocols that make
|
||||
it easy to evaluate methods across many such shifts. See the [](./manuals) for
|
||||
worked examples.
|
||||
|
||||
## Citing QuaPy
|
||||
|
||||
If you find QuaPy useful (and we hope you will), please consider citing the original paper in your research.
|
||||
If you find QuaPy useful, please consider citing the original paper.
|
||||
|
||||
```bibtex
|
||||
@inproceedings{moreo2021quapy,
|
||||
|
|
@ -89,7 +133,8 @@ If you find QuaPy useful (and we hope you will), please consider citing the orig
|
|||
|
||||
## Contributing
|
||||
|
||||
In case you want to contribute improvements to quapy, please generate pull request to the "devel" branch.
|
||||
If you want to contribute improvements to QuaPy, please open a pull request
|
||||
against the `devel` branch.
|
||||
|
||||
## Acknowledgments
|
||||
|
||||
|
|
@ -98,6 +143,6 @@ In case you want to contribute improvements to quapy, please generate pull reque
|
|||
:alt: SoBigData++
|
||||
```
|
||||
|
||||
This work has been supported by the QuaDaSh project
|
||||
_"Finanziato dall’Unione europea---Next Generation EU,
|
||||
This work has been supported by the QuaDaSh project
|
||||
_"Finanziato dall'Unione europea---Next Generation EU,
|
||||
Missione 4 Componente 2 CUP B53D23026250001"_.
|
||||
|
|
|
|||
|
|
@ -7,7 +7,6 @@ Manuals
|
|||
|
||||
manuals/datasets
|
||||
manuals/evaluation
|
||||
manuals/explicit-loss-minimization
|
||||
manuals/methods
|
||||
manuals/model-selection
|
||||
manuals/plotting
|
||||
|
|
|
|||
|
|
@ -412,11 +412,82 @@ ECML-PKDD 2024, Vilnius, Lithuania.
|
|||
```
|
||||
|
||||
|
||||
## Image Embedding Datasets
|
||||
|
||||
QuaPy also provides a collection of image datasets in the form of pre-generated
|
||||
embeddings.
|
||||
These
|
||||
embeddings were generated using [this extraction script](https://github.com/pglez82/visiondatasets_quapy)
|
||||
and are hosted in [Zenodo](https://zenodo.org/records/21131944).
|
||||
|
||||
An example of current public interface is:
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
|
||||
data = qp.datasets.fetch_image_embeddings(
|
||||
dataset_name='cifar10',
|
||||
embedding='features',
|
||||
heldout_only=True,
|
||||
)
|
||||
train, test = data.train_test
|
||||
```
|
||||
|
||||
The available datasets are in `qp.datasets.IMAGE_DATASETS`, and include 6 datasets:
|
||||
|
||||
* `cifar10`, `cifar100`, and `cifar100coarse`:
|
||||
[Alex Krizhevsky and Geoffrey Hinton. Learning multiple layers of features from tiny images. Technical report, University of Toronto, 2009.](https://cave.cs.toronto.edu/kriz/learning-features-2009-TR.pdf)
|
||||
* `mnist`:
|
||||
[Yann LeCun, Corinna Cortes, and Christopher J. C. Burges. The MNIST database of handwritten digits. 1998.](http://yann.lecun.com/exdb/mnist/)
|
||||
* `fashionmnist`:
|
||||
[Han Xiao, Kashif Rasul, and Roland Vollgraf. Fashion-MNIST: a novel image dataset for benchmarking machine learning algorithms. arXiv preprint arXiv:1708.07747, 2017.](https://arxiv.org/abs/1708.07747)
|
||||
* `svhn`:
|
||||
[Yuval Netzer, Tao Wang, Adam Coates, Alessandro Bissacco, Baolin Wu, Andrew Y. Ng, et al. Reading digits in natural images with unsupervised feature learning. NIPS Workshop, 2011.](https://static.googleusercontent.com/media/research.google.com/es//pubs/archive/37648.pdf)
|
||||
|
||||
|
||||
The available embedding types are in `qp.datasets.IMAGE_EMBEDDINGS`, and include:
|
||||
|
||||
* `features` are the penultimate-layer representations
|
||||
* `logits` are the pre-activation outputs of the neural model
|
||||
* `predictions` are the post-softmax posterior probabilities
|
||||
|
||||
The datasets correspond to frozen neural representations extracted from models
|
||||
trained on image classification tasks. QuaPy downloads them automatically on
|
||||
first use and stores them locally for fast reuse.
|
||||
|
||||
Each dataset is internally organised into three splits: `train`, `val`, and
|
||||
`test`. The `train` split was used to train the neural model that produced the
|
||||
embeddings, while `val` and `test` were not seen during neural training.
|
||||
For this reason, the default setting indicates `heldout_only=True`, meaning
|
||||
that the returned dataset will take the validation partition as the training
|
||||
set, and the test partition as the test set.
|
||||
This is often the most convenient choice for quantification experiments, since
|
||||
it avoids training quantifiers on examples that were already used to train the
|
||||
embedding model.
|
||||
|
||||
If instead you want to use all the available non-test data, you can set `heldout_only=False`,
|
||||
in which case, the returned training set is the union of the original neural
|
||||
training split and the validation split.
|
||||
|
||||
Some statistics are shown in the following table:
|
||||
|
||||
| Dataset | backbone | classes | neural network train size | validation size | test size | feature dim | logit dim | prediction dim | type |
|
||||
|---|---|:---:|:---:|:---:|:---:|:---:|:---:|:---:|---|
|
||||
| cifar100 | resnet18 | 100 | 35000 | 15000 | 10000 | 512 | 100 | 100 | dense |
|
||||
| cifar10 | resnet18 | 10 | 35000 | 15000 | 10000 | 512 | 10 | 10 | dense |
|
||||
| cifar100coarse | resnet18 | 20 | 35000 | 15000 | 10000 | 512 | 20 | 20 | dense |
|
||||
| mnist | basiccnn | 10 |42000 | 18000 | 10000 | 128 | 10 | 10 | dense |
|
||||
| fashionmnist | basiccnn | 10 | 42000 | 18000 | 10000 | 128 | 10 | 10 | dense |
|
||||
| svhn | resnet18 | 10 | 51280 | 21977 | 26032 | 512 | 10 | 10 | dense |
|
||||
|
||||
|
||||
|
||||
|
||||
## IFCB Plankton dataset
|
||||
|
||||
IFCB is a dataset of plankton species in water samples hosted in `Zenodo <https://zenodo.org/records/10036244>`_.
|
||||
This dataset is based on the data available publicly at `WHOI-Plankton repo <https://github.com/hsosik/WHOI-Plankton>`_
|
||||
and in the scripts for the processing are available at `P. González's repo <https://github.com/pglez82/IFCB_Zenodo>`_.
|
||||
IFCB is a dataset of plankton species in water samples hosted in [Zenodo](https://zenodo.org/records/10036244).
|
||||
This dataset is based on the data available publicly at [WHOI-Plankton repo](https://github.com/hsosik/WHOI-Plankton)
|
||||
and the scripts for the processing are available at [P. González's repo](https://github.com/pglez82/IFCB_Zenodo).
|
||||
|
||||
This dataset comes with precomputed features for testing quantification algorithms.
|
||||
|
||||
|
|
|
|||
|
|
@ -65,6 +65,116 @@ error_function = qp.error.from_name('mse')
|
|||
error = error_function(true_prev, estim_prev)
|
||||
```
|
||||
|
||||
The main quantification measures currently available in `qp.error` are the
|
||||
following. As a rule of thumb, names starting with `m` indicate the mean value
|
||||
across many sample pairs, while the corresponding unprefixed function returns
|
||||
the sample-wise quantity. Let `p` denote the true prevalence vector,
|
||||
`\hat{p}` the predicted prevalence vector, `\mathcal{Y}` the set of classes,
|
||||
and `p^{tr}` the training prevalence vector.
|
||||
|
||||
### Prevalence-vector measures
|
||||
|
||||
Absolute error and its mean version:
|
||||
|
||||
```{math}
|
||||
AE(p,\hat{p}) = \frac{1}{|\mathcal{Y}|}\sum_{y \in \mathcal{Y}} |\hat{p}(y)-p(y)|
|
||||
```
|
||||
|
||||
Implemented as `ae` and `mae`.
|
||||
|
||||
Normalized absolute error and its mean version:
|
||||
|
||||
```{math}
|
||||
NAE(p,\hat{p}) = \frac{AE(p,\hat{p})}{z_{AE}},\qquad
|
||||
z_{AE}=\frac{2(1-\min_{y \in \mathcal{Y}} p(y))}{|\mathcal{Y}|}
|
||||
```
|
||||
|
||||
Implemented as `nae` and `mnae`.
|
||||
|
||||
Squared error and its mean version:
|
||||
|
||||
```{math}
|
||||
SE(p,\hat{p}) = \frac{1}{|\mathcal{Y}|}\sum_{y \in \mathcal{Y}} (\hat{p}(y)-p(y))^2
|
||||
```
|
||||
|
||||
Implemented as `se` and `mse`.
|
||||
|
||||
Relative absolute error and its mean version:
|
||||
|
||||
```{math}
|
||||
RAE(p,\hat{p}) = \frac{1}{|\mathcal{Y}|}\sum_{y \in \mathcal{Y}}\frac{|\hat{p}(y)-p(y)|}{p(y)}
|
||||
```
|
||||
|
||||
Implemented as `rae` and `mrae`.
|
||||
|
||||
Normalized relative absolute error and its mean version:
|
||||
|
||||
```{math}
|
||||
NRAE(p,\hat{p}) = \frac{RAE(p,\hat{p})}{z_{RAE}},\qquad
|
||||
z_{RAE}=\frac{|\mathcal{Y}|-1+\frac{1-\min_{y \in \mathcal{Y}} p(y)}{\min_{y \in \mathcal{Y}} p(y)}}{|\mathcal{Y}|}
|
||||
```
|
||||
|
||||
Implemented as `nrae` and `mnrae`.
|
||||
|
||||
Kullback-Leibler divergence and its mean version:
|
||||
|
||||
```{math}
|
||||
KLD(p,\hat{p}) = \sum_{y \in \mathcal{Y}} p(y)\log\frac{p(y)}{\hat{p}(y)}
|
||||
```
|
||||
|
||||
Implemented as `kld` and `mkld`.
|
||||
|
||||
Normalized Kullback-Leibler divergence and its mean version:
|
||||
|
||||
```{math}
|
||||
NKLD(p,\hat{p}) = 2\frac{e^{KLD(p,\hat{p})}}{e^{KLD(p,\hat{p})}+1}-1
|
||||
```
|
||||
|
||||
Implemented as `nkld` and `mnkld`.
|
||||
|
||||
Squared ratio error and its mean version:
|
||||
|
||||
```{math}
|
||||
SRE(p,\hat{p},p^{tr}) = \frac{1}{|\mathcal{Y}|}\sum_{i \in \mathcal{Y}} (w_i-\hat{w}_i)^2,\qquad
|
||||
w_i=\frac{p_i}{p^{tr}_i}
|
||||
```
|
||||
|
||||
Implemented as `sre` and `msre`.
|
||||
|
||||
The Aitchison Quantification Error (AQE) and its mean version (MAQE) are implemented as `aqe` and `maqe` using the
|
||||
Aitchison Distance (available in `qp.functional.AitchisonDistance`, here denoted `d_A`):
|
||||
|
||||
```{math}
|
||||
d_A(p,\hat{p}) = \|\mathrm{clr}(p)-\mathrm{clr}(\hat{p})\|_2
|
||||
```
|
||||
|
||||
### Additional measures
|
||||
|
||||
Match distance computes the cumulative-distribution discrepancy under the
|
||||
assumption that moving mass from class `i` to class `i+1` has unit cost:
|
||||
|
||||
```{math}
|
||||
MD(p,\hat{p}) = \sum_{i=1}^{|\mathcal{Y}|-1} \left|\sum_{j=1}^{i} p_j - \sum_{j=1}^{i} \hat{p}_j\right|
|
||||
```
|
||||
|
||||
Implemented as `md`. Its normalized variant `nmd` rescales this quantity by
|
||||
`1/(|\mathcal{Y}|-1)`.
|
||||
|
||||
For binary quantification, QuaPy also provides the signed bias of the positive
|
||||
class and its mean value:
|
||||
|
||||
```{math}
|
||||
bias(p,\hat{p}) = \hat{p}_1 - p_1
|
||||
```
|
||||
|
||||
Implemented as `bias_binary` and `mean_bias_binary`.
|
||||
|
||||
### Classification measures
|
||||
|
||||
The same module also exposes two classification-oriented error measures, which
|
||||
can occasionally be useful for diagnostics: `acce` (accuracy error, i.e.,
|
||||
`1-accuracy`) and `f1e` (macro-`F_1` error, i.e., `1-F_1^M`).
|
||||
|
||||
## Evaluation Protocols
|
||||
|
||||
An _evaluation protocol_ is an evaluation procedure that uses
|
||||
|
|
|
|||
|
|
@ -1,26 +0,0 @@
|
|||
# Explicit Loss Minimization
|
||||
|
||||
QuaPy makes available several Explicit Loss Minimization (ELM) methods, including
|
||||
SVM(Q), SVM(KLD), SVM(NKLD), SVM(AE), or SVM(RAE).
|
||||
These methods require to first download the
|
||||
[svmperf](http://www.cs.cornell.edu/people/tj/svm_light/svm_perf.html)
|
||||
package, apply the patch
|
||||
[svm-perf-quantification-ext.patch](https://github.com/HLT-ISTI/QuaPy/blob/master/svm-perf-quantification-ext.patch), and compile the sources.
|
||||
The script [prepare_svmperf.sh](https://github.com/HLT-ISTI/QuaPy/blob/master/prepare_svmperf.sh) does all the job. Simply run:
|
||||
|
||||
```
|
||||
./prepare_svmperf.sh
|
||||
```
|
||||
|
||||
The resulting directory `svm_perf_quantification/` contains the
|
||||
patched version of _svmperf_ with quantification-oriented losses.
|
||||
|
||||
The [svm-perf-quantification-ext.patch](https://github.com/HLT-ISTI/QuaPy/blob/master/prepare_svmperf.sh) is an extension of the patch made available by
|
||||
[Esuli et al. 2015](https://dl.acm.org/doi/abs/10.1145/2700406?casa_token=8D2fHsGCVn0AAAAA:ZfThYOvrzWxMGfZYlQW_y8Cagg-o_l6X_PcF09mdETQ4Tu7jK98mxFbGSXp9ZSO14JkUIYuDGFG0)
|
||||
that allows SVMperf to optimize for
|
||||
the _Q_ measure as proposed by [Barranquero et al. 2015](https://www.sciencedirect.com/science/article/abs/pii/S003132031400291X)
|
||||
and for the _KLD_ and _NKLD_ measures as proposed by [Esuli et al. 2015](https://dl.acm.org/doi/abs/10.1145/2700406?casa_token=8D2fHsGCVn0AAAAA:ZfThYOvrzWxMGfZYlQW_y8Cagg-o_l6X_PcF09mdETQ4Tu7jK98mxFbGSXp9ZSO14JkUIYuDGFG0).
|
||||
This patch extends the above one by also allowing SVMperf to optimize for
|
||||
_AE_ and _RAE_.
|
||||
See the [](./methods) manual for more details and code examples.
|
||||
|
||||
|
|
@ -2,11 +2,17 @@
|
|||
|
||||
Quantification methods can be categorized as belonging to
|
||||
`aggregative`, `non-aggregative`, and `meta-learning` groups.
|
||||
Most methods included in QuaPy at the moment are of type `aggregative`
|
||||
(though we plan to add many more methods in the near future), i.e.,
|
||||
are methods characterized by the fact that
|
||||
quantification is performed as an aggregation function of the individual
|
||||
products of classification.
|
||||
`Aggregative` quantifiers rely on a surrogate classifier as an intermediate
|
||||
step, and devise different aggregation functions over the classifier outputs.
|
||||
By contrast, `non-aggregative` methods perform quantification without requiring
|
||||
an underlying classifier.
|
||||
`Meta-learning` refers to quantification methods that are constructed over simpler
|
||||
quantification methods, and implement high-level orchestration functions.
|
||||
|
||||
Beyond these three traditional categories of methods, we here present an additional,
|
||||
orthogonal one: `Bayesian quantifiers`, i.e., quantification methods that do not simply
|
||||
return point-estimates of class prevalence, but are also able to provide a measure of
|
||||
uncertaintly around them.
|
||||
|
||||
Any quantifier in QuaPy shoud extend the class `BaseQuantifier`,
|
||||
and implement some abstract methods:
|
||||
|
|
@ -112,11 +118,15 @@ in evaluation.
|
|||
QuaPy implements the four CC variants, i.e.:
|
||||
|
||||
* _CC_ (Classify & Count), the simplest aggregative quantifier; one that
|
||||
simply relies on the label predictions of a classifier to deliver class estimates.
|
||||
* _ACC_ (Adjusted Classify & Count), the adjusted variant of CC.
|
||||
classifies all instances and computes the prevalence of the predicted labels.
|
||||
This baseline is discussed, among others, in [Forman (2008)](https://link.springer.com/article/10.1007/s10618-008-0097-y).
|
||||
* _ACC_ (Adjusted Classify & Count), the adjusted variant of CC, originally
|
||||
proposed in [Forman (2008)](https://link.springer.com/article/10.1007/s10618-008-0097-y).
|
||||
* _PCC_ (Probabilistic Classify & Count), the probabilistic variant of CC that
|
||||
relies on the soft estimations (or posterior probabilities) returned by a (probabilistic) classifier.
|
||||
* _PACC_ (Probabilistic Adjusted Classify & Count), the adjusted variant of PCC.
|
||||
relies on the posterior probabilities returned by a probabilistic classifier,
|
||||
introduced in [Bella et al. (2010)](https://ieeexplore.ieee.org/abstract/document/5694031).
|
||||
* _PACC_ (Probabilistic Adjusted Classify & Count), the adjusted variant of PCC,
|
||||
also introduced in [Bella et al. (2010)](https://ieeexplore.ieee.org/abstract/document/5694031).
|
||||
|
||||
The following code serves as a complete example using CC equipped
|
||||
with a SVM as the classifier:
|
||||
|
|
@ -136,7 +146,7 @@ svm = LinearSVC()
|
|||
# (an alias is available in qp.method.aggregative.ClassifyAndCount)
|
||||
model = qp.method.aggregative.CC(svm)
|
||||
model.fit(Xtr, ytr)
|
||||
estim_prevalence = model.predict(test.instances)
|
||||
estim_prevalence = model.predict(test.X)
|
||||
```
|
||||
|
||||
The same code could be used to instantiate an ACC, by simply replacing
|
||||
|
|
@ -184,6 +194,10 @@ will be raised otherwise.
|
|||
Lastly, everything we said about ACC and PCC
|
||||
applies to PACC as well.
|
||||
|
||||
A Bayesian counterpart of the ACC family is also available; see the
|
||||
{ref}`Bayesian Quantification Methods section <manuals/methods:Bayesian Quantification Methods>`
|
||||
for `BayesianCC`.
|
||||
|
||||
_New in v0.1.9_: quantifiers ACC and PACC now have three additional arguments: `method`, `solver` and `norm`:
|
||||
|
||||
* Argument `method` specifies how to solve, for `p`, the linear system `q = Mp` (where `q` is the unadjusted counts for the
|
||||
|
|
@ -213,24 +227,23 @@ Options are:
|
|||
* `"condsoftmax"` applies softmax normalization only if the prevalence vector lies outside of the probability simplex.
|
||||
|
||||
|
||||
#### BayesianCC
|
||||
### Threshold Optimization methods
|
||||
|
||||
The `BayesianCC` is a variant of ACC introduced in
|
||||
[Ziegler, A. and Czyż, P. "Bayesian quantification with black-box estimators", arXiv (2023)](https://arxiv.org/abs/2302.09159),
|
||||
which models the probabilities `q = Mp` using latent random variables with weak Bayesian priors, rather than
|
||||
plug-in probability estimates. In particular, it uses Markov Chain Monte Carlo sampling to find the values of
|
||||
`p` compatible with the observed quantities.
|
||||
The `aggregate` method returns the posterior mean and the `get_prevalence_samples` method can be used to find
|
||||
uncertainty around `p` estimates (conditional on the observed data and the trained classifier)
|
||||
and is suitable for problems in which the `q = Mp` matrix is nearly non-invertible.
|
||||
QuaPy implements Forman's threshold optimization methods;
|
||||
see, e.g., [(Forman 2006)](https://dl.acm.org/doi/abs/10.1145/1150402.1150423)
|
||||
and [(Forman 2008)](https://link.springer.com/article/10.1007/s10618-008-0097-y).
|
||||
These include: `T50`, `MAX`, `X`, Median Sweep (`MS`), and its variant `MS2`.
|
||||
|
||||
Note that this quantification method requires `val_split` to be a `float` and installation of additional dependencies (`$ pip install quapy[bayes]`) needed to run Markov chain Monte Carlo sampling. Markov Chain Monte Carlo is is slower than matrix inversion methods, but is guaranteed to sample proper probability vectors, so no clipping strategies are required.
|
||||
An example presenting how to run the method and use posterior samples is available in `examples/bayesian_quantification.py`.
|
||||
These methods are binary-only and implement different heuristics for
|
||||
improving the stability of the denominator of the ACC adjustment (`tpr-fpr`).
|
||||
The methods are called "threshold" since said heuristics have to do
|
||||
with different choices of the underlying classifier's threshold.
|
||||
|
||||
### Expectation Maximization (EMQ)
|
||||
|
||||
The Expectation Maximization Quantifier (EMQ), also known as
|
||||
the SLD, is available at `qp.method.aggregative.EMQ` or via the
|
||||
### Expectation Maximization (EMQ) / Maximum Likelihood for Label Shift (MLLS)
|
||||
|
||||
The Expectation Maximization Quantifier (EMQ) (also known as
|
||||
SLD after the name of the proponets, or Maximum Likelihood for Label Shift, MLLS) , is available at `qp.method.aggregative.EMQ` or via the
|
||||
alias `qp.method.aggregative.ExpectationMaximizationQuantifier`.
|
||||
The method is described in:
|
||||
|
||||
|
|
@ -274,8 +287,57 @@ or Temperature Scaling (`ts`); default is `None` (no calibration).
|
|||
You can use the class method `EMQ_BCTS` to effortlessly instantiate EMQ with the best performing
|
||||
heuristics found by [Alexandari et al. (2020)](http://proceedings.mlr.press/v119/alexandari20a.html). See the API documentation for further details.
|
||||
|
||||
For a Bayesian label-shift counterpart based on the same general family of ideas,
|
||||
see the {ref}`Bayesian Quantification Methods section <manuals/methods:Bayesian Quantification Methods>`
|
||||
for `BayesianMAPLS`.
|
||||
|
||||
### Hellinger Distance y (HDy)
|
||||
### Regularized Learning under Label Shift (RLLS)
|
||||
|
||||
`RLLS` is available at `qp.method.aggregative.RLLS` and ports the regularized
|
||||
importance-weight estimation procedure of
|
||||
[Azizzadenesheli, K., Liu, A., Yang, F., and Anandkumar, A. (2019). Regularized
|
||||
Learning for Domain Adaptation under Label Shifts.
|
||||
ICLR 2019](https://arxiv.org/abs/1903.09734) to QuaPy's aggregative interface.
|
||||
The method estimates the label-shift importance weights `w = q(y)/p(y)` from
|
||||
the classifier's validation posteriors (or, in `mode='hard'`, its argmax
|
||||
predictions) and the corresponding source labels, regularizing the estimation
|
||||
by an amount controlled by `alpha` (scaled by a finite-sample confidence term
|
||||
governed by `delta`). The resulting weights are then used to rescale the
|
||||
training prevalence into the target prevalence estimate.
|
||||
|
||||
Like ACC and PACC, RLLS requires validation predictions and therefore expects
|
||||
`val_split` to be set (as an integer for k-fold cross-validation, a float for
|
||||
a held-out split, or an explicit `(X, y)` tuple) whenever `fit_classifier=True`.
|
||||
This method relies on the optional `cvxpy` dependency, which must be
|
||||
installed separately (`$ pip install cvxpy`).
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
from quapy.method.aggregative import RLLS
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
train, test = qp.datasets.fetch_UCIBinaryDataset('haberman').train_test
|
||||
|
||||
model = RLLS(LogisticRegression(max_iter=2000), val_split=5)
|
||||
model.fit(*train.Xy)
|
||||
estim_prevalence = model.predict(test.X)
|
||||
```
|
||||
|
||||
### Distribution Matching
|
||||
|
||||
Distribution Matching (DM) methods search for the mixture parameter (the sought class prevalence values)
|
||||
yielding the mixture between the class-wise representations that best matches the test distribution.
|
||||
Different criteria for deciding how this matching is assessed, and different ways for modelling the
|
||||
distributions give rise to different instantiations of DM methods.
|
||||
|
||||
The following methods are here discussed because they rely on a surrogate classifier for representing
|
||||
the distributions, albeit different non-aggregative variants of them do often exist. Aside from this,
|
||||
the formulation of DM methods is flexible enough as to accomodate methods that were proposed under a different
|
||||
framework; examples include ACC and PACC.
|
||||
|
||||
See the frameworks by [Firat](https://arxiv.org/abs/1606.00868), [Bunse](https://dl.gi.de/items/5a61f30f-6c84-4165-bd92-9098bd9e91aa), [Garg et al.](https://dl.acm.org/doi/10.5555/3495724.3496001), or [Dussap](https://theses.hal.science/tel-04931123), for more details.
|
||||
|
||||
#### Hellinger Distance y (HDy)
|
||||
|
||||
Implementation of the method based on the Hellinger Distance y (HDy) proposed by
|
||||
[González-Castro, V., Alaiz-Rodríguez, R., and Alegre, E. (2013). Class distribution
|
||||
|
|
@ -310,31 +372,94 @@ model.fit(*dataset.training.Xy)
|
|||
estim_prevalence = model.predict(dataset.test.X)
|
||||
```
|
||||
|
||||
QuaPy also provides an implementation of the generalized
|
||||
"Distribution Matching" approaches for multiclass, inspired by the framework
|
||||
of [Firat (2016)](https://arxiv.org/abs/1606.00868). One can instantiate
|
||||
a variant of HDy for multiclass quantification as follows:
|
||||
#### Generalized Distribution Matching y (DMy)
|
||||
|
||||
QuaPy also provides a generalized posterior-space distribution-matching
|
||||
quantifier for binary or multiclass problems, implemented as
|
||||
`qp.method.aggregative.DMy`. This class follows the generic distribution
|
||||
matching view discussed by [Firat (2016)](https://arxiv.org/abs/1606.00868):
|
||||
it represents class-conditional posterior distributions by histograms and then
|
||||
searches for the prevalence vector whose mixture best matches the test
|
||||
distribution.
|
||||
|
||||
`DMy` is intentionally flexible and exposes three main design choices: the
|
||||
number of histogram bins (`nbins`), the divergence to minimize (`divergence`,
|
||||
e.g., `'HD'` or `'topsoe'`), and whether to match PDFs or CDFs (`cdf`). The
|
||||
optimization routine can also be selected through `search`; the default
|
||||
`'optim_minimize'` works for multiclass problems, while `'linear_search'` and
|
||||
`'ternary_search'` are binary-only. A multiclass HDy-like instance can be
|
||||
obtained as:
|
||||
|
||||
```python
|
||||
mutliclassHDy = qp.method.aggregative.DMy(classifier=LogisticRegression(), divergence='HD', cdf=False)
|
||||
```
|
||||
multiclass_hdy = qp.method.aggregative.DMy(
|
||||
classifier=LogisticRegression(),
|
||||
divergence='HD',
|
||||
cdf=False,
|
||||
)
|
||||
```
|
||||
|
||||
QuaPy also provides an implementation of the "DyS"
|
||||
framework proposed by [Maletzke et al (2020)](https://ojs.aaai.org/index.php/AAAI/article/view/4376)
|
||||
and the "SMM" method proposed by [Hassan et al (2019)](https://ieeexplore.ieee.org/document/9260028)
|
||||
(thanks to _Pablo González_ for the contributions!)
|
||||
#### DyS
|
||||
|
||||
### Threshold Optimization methods
|
||||
QuaPy implements the binary `DyS` framework proposed by
|
||||
[Maletzke et al. (2020)](https://ojs.aaai.org/index.php/AAAI/article/view/4376)
|
||||
as `qp.method.aggregative.DyS`. Conceptually, `DyS` can be seen as a
|
||||
generalization of HDy in which the prevalence is found by ternary search over a
|
||||
distribution-matching objective. In QuaPy, the user can select the number of
|
||||
histogram bins (`n_bins`), the divergence (`divergence`), and the optimization
|
||||
tolerance (`tol`).
|
||||
|
||||
QuaPy implements Forman's threshold optimization methods;
|
||||
see, e.g., [(Forman 2006)](https://dl.acm.org/doi/abs/10.1145/1150402.1150423)
|
||||
and [(Forman 2008)](https://link.springer.com/article/10.1007/s10618-008-0097-y).
|
||||
These include: `T50`, `MAX`, `X`, Median Sweep (`MS`), and its variant `MS2`.
|
||||
#### Energy Distance y (EDy)
|
||||
|
||||
QuaPy also adapts `EDy` from [quantificationlib](https://github.com/AICGijon/quantificationlib),
|
||||
which is available as `qp.method.aggregative.EDy`.
|
||||
|
||||
This
|
||||
method replaces histogram matching with an energy-distance formulation defined
|
||||
directly on posterior-probability vectors and solves the resulting optimization
|
||||
problem by quadratic programming. The method is proposed in
|
||||
[Castaño et al.'s (2024)](https://ieeexplore.ieee.org/document/9791435/) paper.
|
||||
|
||||
In QuaPy, `EDy` works for binary and
|
||||
multiclass problems and lets the user choose the pairwise distance through the
|
||||
`distance` parameter (`'manhattan'`, `'euclidean'`, or a custom callable).
|
||||
Because the optimization relies on `quadprog`, this method requires the
|
||||
optional dependency `pip install quadprog`.
|
||||
|
||||
#### SMM
|
||||
|
||||
QuaPy also includes the binary `SMM` method of
|
||||
[Hassan et al. (2019)](https://ieeexplore.ieee.org/document/9260028),
|
||||
available as `qp.method.aggregative.SMM`. This is a very lightweight
|
||||
distribution-matching variant in which the posterior representation is reduced
|
||||
to class-wise means rather than full histograms, making it conceptually close
|
||||
to PACC.
|
||||
|
||||
|
||||
#### Kernel Density Estimation methods (KDEy)
|
||||
|
||||
QuaPy provides implementations for the three variants
|
||||
of KDE-based methods proposed in
|
||||
_[Moreo, A., González, P. and del Coz, J.J..
|
||||
Kernel Density Estimation for Multiclass Quantification.
|
||||
Machine Learning. Vol 114 (92), 2025](https://link.springer.com/article/10.1007/s10994-024-06726-5)_
|
||||
(a [preprint](https://arxiv.org/abs/2401.00490) is available online).
|
||||
The variants differ in the divergence metric to be minimized:
|
||||
|
||||
- KDEy-HD: minimizes the (squared) Hellinger Distance and solves the problem via a Monte Carlo approach
|
||||
- KDEy-CS: minimizes the Cauchy-Schwarz divergence and solves the problem via a closed-form solution
|
||||
- KDEy-ML: minimizes the Kullback-Leibler divergence and solves the problem via maximum-likelihood
|
||||
|
||||
These methods are specifically devised for multiclass problems (although they can tackle
|
||||
binary problems too).
|
||||
|
||||
All KDE-based methods depend on the hyperparameter `bandwidth` of the kernel. Typical values
|
||||
that can be explored in model selection range in [0.01, 0.25]. Previous experiments reveal the methods' performance
|
||||
varies smoothly at small variations of this hyperparameter.
|
||||
|
||||
A Bayesian counterpart is available as well; see the
|
||||
{ref}`Bayesian Quantification Methods section <manuals/methods:Bayesian Quantification Methods>`
|
||||
for `BayesianKDEy`.
|
||||
|
||||
These methods are binary-only and implement different heuristics for
|
||||
improving the stability of the denominator of the ACC adjustment (`tpr-fpr`).
|
||||
The methods are called "threshold" since said heuristics have to do
|
||||
with different choices of the underlying classifier's threshold.
|
||||
|
||||
### Explicit Loss Minimization
|
||||
|
||||
|
|
@ -363,84 +488,210 @@ the last two methods (SVM(AE) and SVM(RAE)) have been implemented in
|
|||
QuaPy in order to make available ELM variants for what nowadays
|
||||
are considered the most well-behaved evaluation metrics in quantification.
|
||||
|
||||
In order to make these models work, you would need to run the script
|
||||
`prepare_svmperf.sh` (distributed along with QuaPy) that
|
||||
downloads `SVMperf`' source code, applies a patch that
|
||||
implements the quantification oriented losses, and compiles the
|
||||
sources.
|
||||
#### Installing the SVMperf backend
|
||||
|
||||
If you want to add any custom loss, you would need to modify
|
||||
the source code of `SVMperf` in order to implement it, and
|
||||
assign a valid loss code to it. Then you must re-compile
|
||||
the whole thing and instantiate the quantifier in QuaPy
|
||||
as follows:
|
||||
These methods rely on Joachim's [SVMperf](https://www.cs.cornell.edu/people/tj/svm_light/svm_perf.html),
|
||||
patched with quantification-oriented losses. QuaPy provides the script
|
||||
[`prepare_svmperf.sh`](https://github.com/HLT-ISTI/QuaPy/blob/master/prepare_svmperf.sh),
|
||||
which downloads the original sources, applies the patch, and compiles the
|
||||
resulting binary. In practice, this amounts to running:
|
||||
|
||||
```python
|
||||
# you can either set the path to your custom svm_perf_quantification implementation
|
||||
# in the environment variable, or as an argument to the constructor of ELM
|
||||
qp.environ['SVMPERF_HOME'] = './path/to/svm_perf_quantification'
|
||||
|
||||
# assign an alias to your custom loss and the id you have assigned to it
|
||||
svmperf = qp.classification.svmperf.SVMperf
|
||||
svmperf.valid_losses['mycustomloss'] = 28
|
||||
|
||||
# instantiate the ELM method indicating the loss
|
||||
model = qp.method.aggregative.ELM(loss='mycustomloss')
|
||||
```sh
|
||||
./prepare_svmperf.sh
|
||||
```
|
||||
|
||||
All ELM are binary quantifiers since they rely on `SVMperf`, that
|
||||
currently supports only binary classification.
|
||||
ELM variants (any binary quantifier in general) can be extended
|
||||
to operate in single-label scenarios trivially by adopting a
|
||||
"one-vs-all" strategy (as, e.g., in
|
||||
[_Gao, W. and Sebastiani, F. (2016). From classification to quantification in tweet sentiment
|
||||
analysis. Social Network Analysis and Mining, 6(19):1–22_](https://link.springer.com/article/10.1007/s13278-016-0327-z)).
|
||||
In QuaPy this is possible by using the `OneVsAll` class.
|
||||
This creates a directory `svm_perf_quantification/`. Once this is available,
|
||||
you can point QuaPy to it with:
|
||||
|
||||
There are two ways for instantiating this class, `OneVsAllGeneric` that works for
|
||||
any quantifier, and `OneVsAllAggregative` that is optimized for aggregative quantifiers.
|
||||
In general, you can simply use the `newOneVsAll` function and QuaPy will choose
|
||||
the more convenient of the two.
|
||||
```python
|
||||
qp.environ['SVMPERF_HOME'] = './svm_perf_quantification'
|
||||
```
|
||||
|
||||
The patch extends the one originally released for
|
||||
[Esuli and Sebastiani (2015)](https://dl.acm.org/doi/abs/10.1145/2700406)
|
||||
and also covers the `Q`, `AE`, and `RAE` losses used by QuaPy's ELM wrappers.
|
||||
|
||||
All ELM methods are binary because `SVMperf` itself is binary. They can still
|
||||
be wrapped in a one-vs-all scheme for single-label multiclass problems, though
|
||||
this strategy is generally considered inappropriate under prior probability
|
||||
shift. See the examples on
|
||||
[explicit loss minimization](https://github.com/HLT-ISTI/QuaPy/blob/devel/examples/17.explicit_loss_minimization.py)
|
||||
and on
|
||||
[one versus all quantification](https://github.com/HLT-ISTI/QuaPy/blob/devel/examples/10.one_vs_all.py)
|
||||
for minimal working code.
|
||||
|
||||
## Non-Aggregative Methods
|
||||
|
||||
Non-aggregative methods are quantifiers that do not follow the two-step
|
||||
(classify, then aggregate) pattern described above for aggregative methods.
|
||||
These methods are implemented in the `qp.method.non_aggregative` module and
|
||||
extend `BaseQuantifier` directly, implementing `fit` and `predict` on their own terms.
|
||||
|
||||
### Maximum Likelihood Prevalence Estimation (MLPE)
|
||||
|
||||
`MaximumLikelihoodPrevalenceEstimation` (MLPE) is a lazy baseline quantifier
|
||||
that assumes the IID assumption holds, i.e., that there is no prior probability
|
||||
shift between the training and the test distributions. Its `fit` method simply
|
||||
computes and stores the training prevalence, and its `predict` method returns
|
||||
that same training prevalence for any test sample, irrespective of the sample
|
||||
itself. MLPE is considered a lower-bound quantifier: any quantification method
|
||||
worth using should outperform it.
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
from quapy.method.aggregative import SVMQ
|
||||
from quapy.method.non_aggregative import MaximumLikelihoodPrevalenceEstimation
|
||||
|
||||
# load a single-label dataset (this one contains 3 classes)
|
||||
dataset = qp.datasets.fetch_twitter('hcr', pickle=True)
|
||||
dataset = qp.datasets.fetch_UCIBinaryDataset('haberman')
|
||||
train, test = dataset.train_test
|
||||
|
||||
# let qp know where svmperf is
|
||||
qp.environ['SVMPERF_HOME'] = '../svm_perf_quantification'
|
||||
|
||||
model = newOneVsAll(SVMQ(), n_jobs=-1) # run them on parallel
|
||||
model.fit(dataset.training)
|
||||
estim_prevalence = model.predict(dataset.test.instances)
|
||||
model = MaximumLikelihoodPrevalenceEstimation()
|
||||
model.fit(*train.Xy)
|
||||
estim_prevalence = model.predict(test.X) # always equals train.prevalence()
|
||||
```
|
||||
|
||||
Check the examples on [explicit loss minimization](https://github.com/HLT-ISTI/QuaPy/blob/devel/examples/17.explicit_loss_minimization.py)
|
||||
and on [one versus all quantification](https://github.com/HLT-ISTI/QuaPy/blob/devel/examples/10.one_vs_all.py) for more details.
|
||||
**Note** that the _one versus all_ approach is considered inappropriate under prior probability shift, though.
|
||||
### Distribution Matching x (DMx) and Hellinger Distance x (HDx)
|
||||
|
||||
### Kernel Density Estimation methods (KDEy)
|
||||
`DMx` is the covariate-space counterpart of the `DMy` distribution-matching
|
||||
quantifier described in {ref}`the Hellinger Distance y (HDy) section <manuals/methods:Hellinger Distance y (HDy)>`:
|
||||
instead of matching distributions built from the classifier's predictions, `DMx` matches
|
||||
distributions built directly from the (discretized) feature space, and thus
|
||||
requires no classifier at all. For each class, `DMx` builds one histogram per
|
||||
feature from the training instances of that class; at prediction time, it
|
||||
searches for the mixture of these class-conditional histograms that best
|
||||
matches the (also histogram-based) representation of the test sample, in
|
||||
terms of a chosen divergence.
|
||||
|
||||
QuaPy provides implementations for the three variants
|
||||
of KDE-based methods proposed in
|
||||
_[Moreo, A., González, P. and del Coz, J.J..
|
||||
Kernel Density Estimation for Multiclass Quantification.
|
||||
Machine Learning. Vol 114 (92), 2025](https://link.springer.com/article/10.1007/s10994-024-06726-5)_
|
||||
(a [preprint](https://arxiv.org/abs/2401.00490) is available online).
|
||||
The variants differ in the divergence metric to be minimized:
|
||||
`DMx` accepts the following hyperparameters in its constructor:
|
||||
|
||||
- KDEy-HD: minimizes the (squared) Hellinger Distance and solves the problem via a Monte Carlo approach
|
||||
- KDEy-CS: minimizes the Cauchy-Schwarz divergence and solves the problem via a closed-form solution
|
||||
- KDEy-ML: minimizes the Kullback-Leibler divergence and solves the problem via maximum-likelihood
|
||||
* `nbins`: the number of bins used to discretize each feature (default 8)
|
||||
* `divergence`: a string ("HD" for Hellinger Distance, or "topsoe") or a
|
||||
callable taking two histograms and returning a divergence value (default "HD")
|
||||
* `cdf`: whether to match cumulative distributions (CDFs) instead of the
|
||||
histograms (PDFs) themselves (default False)
|
||||
* `search`: the strategy used for finding the optimal prevalence; valid
|
||||
options are `optim_minimize` (default, works for binary and multiclass
|
||||
problems), `linear_search`, and `ternary_search` (these last two are
|
||||
binary-only)
|
||||
* `n_jobs`: number of parallel workers (default None)
|
||||
|
||||
These methods are specifically devised for multiclass problems (although they can tackle
|
||||
binary problems too).
|
||||
`DMx` also offers the class method `DMx.HDx` (aliased as
|
||||
`qp.method.non_aggregative.HDx`, and also as `HellingerDistanceX`) that
|
||||
reproduces the original Hellinger Distance x (HDx) method proposed by
|
||||
[González-Castro, Alaiz-Rodríguez, and Alegre (2013)](https://www.sciencedirect.com/science/article/pii/S0020025512004069),
|
||||
the same paper that introduced HDy. HDx is a binary-only method that computes
|
||||
the matching for `nbins` ranging over `[10, 20, ..., 110]` (via a
|
||||
`MedianEstimator`, taking the median of the resulting estimates) and searches
|
||||
for the best prevalence via a linear search stepping by 0.01, rather than via
|
||||
the `optim_minimize` search used by `DMx` by default.
|
||||
|
||||
All KDE-based methods depend on the hyperparameter `bandwidth` of the kernel. Typical values
|
||||
that can be explored in model selection range in [0.01, 0.25]. Previous experiments reveal the methods' performance
|
||||
varies smoothly at small variations of this hyperparameter.
|
||||
The following code, adapted from the example comparing HDy and HDx
|
||||
(`examples/11.comparing_HDy_HDx.py`), shows the two methods side-by-side:
|
||||
|
||||
```python
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
import quapy as qp
|
||||
from quapy.method.aggregative import HDy
|
||||
from quapy.method.non_aggregative import DMx
|
||||
|
||||
train, test = qp.datasets.fetch_UCIBinaryDataset('haberman').train_test
|
||||
Xtr, ytr = train.Xy
|
||||
|
||||
hdy = HDy(LogisticRegression()).fit(Xtr, ytr)
|
||||
estim_prevalence_hdy = hdy.predict(test.X)
|
||||
|
||||
hdx = DMx.HDx(n_jobs=-1).fit(Xtr, ytr)
|
||||
estim_prevalence_hdx = hdx.predict(test.X)
|
||||
```
|
||||
|
||||
Note that, unlike HDy, HDx requires no classifier whatsoever, since it
|
||||
operates directly on the covariates.
|
||||
|
||||
### Energy Distance x (EDx)
|
||||
|
||||
QuaPy also provides `qp.method.non_aggregative.EDx`, which is the
|
||||
feature-space counterpart of `EDy`: it keeps the same energy-distance
|
||||
formulation and quadratic-program solver, but applies them directly to the raw
|
||||
instances instead of first projecting them onto posterior probabilities through
|
||||
a classifier. In this sense, `EDx` is to `EDy` what `DMx` is to `DMy`.
|
||||
|
||||
`EDx` works for binary and multiclass problems, accepts the same `distance`
|
||||
options as `EDy` (`'manhattan'`, `'euclidean'`, or a custom callable), and
|
||||
requires the optional dependency `pip install quadprog`.
|
||||
|
||||
### ReadMe
|
||||
|
||||
`ReadMe` is a non-aggregative quantification method proposed by
|
||||
[Hopkins, D. and King, G. (2007). A method of automated nonparametric content
|
||||
analysis for social science. American Journal of Political Science,
|
||||
54(1):229-247.](https://onlinelibrary.wiley.com/doi/abs/10.1111/j.1540-5907.2009.00428.x)
|
||||
The method estimates `Q(Y=i)` directly from `Q(X) = sum_i Q(X|Y=i) Q(Y=i)` by
|
||||
solving a (constrained) least-squares regression, thus avoiding the cost of
|
||||
estimating posterior probabilities `Q(Y=i|X)` altogether.
|
||||
|
||||
Since `Q(X)` and `Q(X|Y=i)` can be of very high dimension for realistic
|
||||
feature spaces, ReadMe renders the problem tractable by performing bagging in
|
||||
the feature space: many small random subsets of features (of size
|
||||
`bagging_range`) are drawn, the least-squares problem is solved on each
|
||||
subset, and the resulting estimates are averaged. ReadMe additionally
|
||||
combines this bagging procedure with bootstrap resampling of the training
|
||||
instances in order to derive confidence regions around the point estimate;
|
||||
accordingly, `ReadMe` implements the `WithConfidenceABC` interface (see the
|
||||
{ref}`confidence regions section <confidence-regions-for-class-prevalence-estimation>`),
|
||||
and exposes a `predict_conf` method in addition to `predict`.
|
||||
|
||||
`ReadMe` accepts the following hyperparameters:
|
||||
|
||||
* `prob_model`: either `"full"` (default), the original Hopkins and King
|
||||
formulation, in which `Q(X)` and `Q(X|Y)` are modelled empirically and thus
|
||||
require the feature matrix `X` to be binary (e.g., term presence/absence);
|
||||
or `"naive"`, a much faster approximation that models `Q(X)` and `Q(X|Y)` as
|
||||
multinomial (bag-of-words) distributions, and that supports much larger
|
||||
values of `bagging_range`
|
||||
* `bootstrap_trials`: number of bootstrap resamplings of the training data
|
||||
used for deriving the confidence region (default 300)
|
||||
* `bagging_trials`: number of bagging trials, i.e., random feature subsets,
|
||||
averaged for each point estimate (default 300)
|
||||
* `bagging_range`: number of features kept in each bagging trial (default 15);
|
||||
note that, when `prob_model="full"`, this value should typically be kept
|
||||
small (the authors advise against values above 25) since the empirical
|
||||
distribution requires enumerating `2^bagging_range` possible feature
|
||||
configurations
|
||||
* `confidence_level`: the confidence level for the confidence region
|
||||
(default 0.95)
|
||||
* `region`: the type of confidence region to construct, one of `"intervals"`
|
||||
(default), `"ellipse"`, `"ellipse-clr"`, or `"ellipse-ilr"` (see the
|
||||
{ref}`confidence regions section <confidence-regions-for-class-prevalence-estimation>`
|
||||
for details)
|
||||
* `bonferroni`: whether to apply Bonferroni correction when `region="intervals"`
|
||||
(default `False`); this parameter has no effect for ellipse-based regions
|
||||
* `random_state`: an int for replicability, or `None` (default)
|
||||
* `verbose`: whether to display progress information (default False)
|
||||
|
||||
The following minimal example, adapted from
|
||||
`examples/18.ReadMe_for_text_analysis.py`, shows ReadMe applied to a binary
|
||||
bag-of-words text quantification problem:
|
||||
|
||||
```python
|
||||
from sklearn.feature_extraction.text import CountVectorizer
|
||||
from sklearn.pipeline import Pipeline
|
||||
import quapy as qp
|
||||
from quapy.method.non_aggregative import ReadMe
|
||||
|
||||
reviews = qp.datasets.fetch_reviews('imdb').reduce(n_train=1000, random_state=0)
|
||||
|
||||
# ReadMe's "full" model requires a binary feature matrix
|
||||
encode_0_1 = Pipeline([('0_1_terms', CountVectorizer(min_df=5, binary=True))])
|
||||
train, test = qp.data.preprocessing.instance_transformation(reviews, encode_0_1, inplace=True).train_test
|
||||
|
||||
model = ReadMe(prob_model='full', bootstrap_trials=100, bagging_trials=100, bagging_range=20, random_state=0)
|
||||
model.fit(*train.Xy) # lazy: only bootstrap resampling happens here
|
||||
|
||||
estim_prevalence, conf_region = model.predict_conf(test.X)
|
||||
```
|
||||
|
||||
Note that `ReadMe` is computationally expensive: its cost scales with the
|
||||
product of `bootstrap_trials` and `bagging_trials`, each of which requires
|
||||
solving a least-squares problem.
|
||||
|
||||
|
||||
## Composable Methods
|
||||
|
|
@ -600,25 +851,192 @@ model.fit(*dataset.training.Xy)
|
|||
estim_prevalence = model.predict(dataset.test.X)
|
||||
```
|
||||
|
||||
## Confidence Regions for Class Prevalence Estimation
|
||||
(confidence-regions-for-class-prevalence-estimation)=
|
||||
## Quantifiers with Uncertainty Quantification
|
||||
|
||||
_(New in v0.2.0!)_ Some quantification methods go beyond providing a single point estimate of class prevalence values and also produce confidence regions, which characterize the uncertainty around the point estimate. In QuaPy, two such methods are currently implemented:
|
||||
|
||||
* Aggregative Bootstrap: The Aggregative Bootstrap method extends any aggregative quantifier by generating confidence regions for class prevalence estimates through bootstrapping. Key features of this method include:
|
||||
|
||||
* Optimized Computation: The bootstrap is applied to pre-classified instances, significantly speeding up training and inference.
|
||||
During training, bootstrap repetitions are performed only after training the classifier once. These repetitions are used to train multiple aggregation functions.
|
||||
During inference, bootstrap is applied over pre-classified test instances.
|
||||
* General Applicability: Aggregative Bootstrap can be applied to any aggregative quantifier.
|
||||
For further information, check the [example](https://github.com/HLT-ISTI/QuaPy/tree/master/examples/16.confidence_regions.py) provided.
|
||||
|
||||
* BayesianCC: is a Bayesian variant of the Adjusted Classify & Count (ACC) quantifier; see more details in the [example](https://github.com/HLT-ISTI/QuaPy/tree/master/examples/14.bayesian_quantification.py) provided.
|
||||
_(New in v0.2.0!)_ Some quantification methods go beyond providing a single point estimate of class prevalence values and also produce confidence regions, which characterize the uncertainty around the point estimate. In QuaPy, two such families are currently implemented: bootstrap methods and Bayesian methods.
|
||||
|
||||
Confidence regions are constructed around a point estimate, which is typically computed as the mean value of a set of samples.
|
||||
The confidence region can be instantiated in three ways:
|
||||
* Confidence intervals: are standard confidence intervals generated for each class independently (_method="intervals"_).
|
||||
|
||||
The confidence region can be instantiated in four ways:
|
||||
* Confidence intervals: are standard confidence intervals generated for each class independently (_method="intervals"_). Since confidence intervals are independently derived for each class, Bonferroni correction can be applied.
|
||||
* Confidence ellipse in the simplex: an ellipse constructed around the mean point; the ellipse lies on the simplex and takes
|
||||
into account possible inter-class dependencies in the data (_method="ellipse"_).
|
||||
into account possible inter-class dependencies in the data (_method="ellipse"_).
|
||||
* Confidence ellipse in the Centered-Log Ratio (CLR) space: the underlying assumption of the ellipse is that the components are
|
||||
normally distributed. However, we know elements from the simplex have an inner structure. A better approach is to first
|
||||
transform the components into an unconstrained space (the CLR), and then construct the ellipse in such space (_method="ellipse-clr"_).
|
||||
transform the components into an unconstrained space (the CLR), and then construct the ellipse in such space (_method="ellipse-clr"_).
|
||||
* Confidence ellipse in the Isometric-Log Ratio (ILR) space: analogous to the CLR-based ellipse, but built in the ILR space
|
||||
(_method="ellipse-ilr"_).
|
||||
|
||||
### Aggregative Bootstrap
|
||||
|
||||
The Aggregative Bootstrap method extends any aggregative quantifier by generating confidence regions for class prevalence estimates through bootstrapping. The method is described in the paper [Moreo, A., Salvati, N.
|
||||
An Efficient Method for Deriving Confidence Intervals in Aggregative Quantification.
|
||||
Learning to Quantify: Methods and Applications (LQ 2025), co-located at ECML-PKDD 2025.
|
||||
pp 12-33, Porto (Portugal)](https://lq-2025.github.io/proceedings/CompleteVolume.pdf).
|
||||
|
||||
This implementation is optimized for aggregative quantifiers. The bootstrap is applied to pre-classified instances, significantly speeding up training and inference.
|
||||
During training, bootstrap repetitions are performed only after training the classifier once. These repetitions are used to train multiple aggregation functions.
|
||||
During inference, bootstrap is applied over pre-classified test instances.
|
||||
|
||||
Aggregative Bootstrap can be applied to any aggregative quantifier. For further information, check the [example](https://github.com/HLT-ISTI/QuaPy/tree/master/examples/16.confidence_regions.py) provided. A minimal working example is:
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
from quapy.method.aggregative import PACC
|
||||
from quapy.method.confidence import AggregativeBootstrap
|
||||
|
||||
train, test = qp.datasets.fetch_UCIMulticlassDataset('molecular').train_test
|
||||
|
||||
model = AggregativeBootstrap(
|
||||
PACC(),
|
||||
n_test_samples=200,
|
||||
confidence_level=0.95,
|
||||
region='ellipse-clr', # choose among: intervals, ellipse, ellipse-clr, ellipse-ilr
|
||||
random_state=0,
|
||||
)
|
||||
model.fit(*train.Xy)
|
||||
point_estimate, conf_region = model.predict_conf(test.X)
|
||||
```
|
||||
|
||||
Here `region` makes the type of uncertainty region explicit. In practice, `intervals` is often the simplest default, while `ellipse`, `ellipse-clr`, and `ellipse-ilr` provide coupled regions over the simplex. If `region='intervals'`, you can additionally set `bonferroni=True` to apply Bonferroni correction; this flag has no effect for ellipse-based regions.
|
||||
|
||||
Beyond aggregative quantifiers, Bootstrap sampling can be applied to any type of quantification method, although
|
||||
the speedup procedure described above is not applied.
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
### Bayesian Quantification Methods
|
||||
|
||||
QuaPy also provides a number of Bayesian quantifiers. While these methods are
|
||||
usually related to an existing point-estimation family (e.g., ACC, EMQ/MLLS,
|
||||
HDy, or KDEy), they differ enough in their goals and outputs to deserve a
|
||||
separate presentation. In particular, Bayesian quantifiers typically return a
|
||||
posterior mean rather than a single optimization result, expose posterior
|
||||
samples or confidence regions, and often require additional probabilistic
|
||||
inference dependencies.
|
||||
|
||||
The optional dependencies needed for these methods can be installed with:
|
||||
|
||||
```sh
|
||||
pip install quapy[bayes]
|
||||
```
|
||||
|
||||
#### BayesianCC (a Bayesian implementation of ACC)
|
||||
|
||||
The `BayesianCC` is a variant of ACC introduced in
|
||||
[Ziegler, A. and Czyż, P. "Bayesian quantification with black-box estimators", arXiv (2023)](https://arxiv.org/abs/2302.09159),
|
||||
which models the probabilities `q = Mp` using latent random variables with weak Bayesian priors, rather than
|
||||
plug-in probability estimates. In particular, it uses Markov Chain Monte Carlo sampling to find the values of
|
||||
`p` compatible with the observed quantities.
|
||||
The `aggregate` method returns the posterior mean and the `get_prevalence_samples` method can be used to find
|
||||
uncertainty around `p` estimates (conditional on the observed data and the trained classifier)
|
||||
and is suitable for problems in which the `q = Mp` matrix is nearly non-invertible.
|
||||
|
||||
Note that this quantification method requires `val_split` to be a `float` and installation of additional dependencies (`$ pip install quapy[bayes]`) needed to run Markov chain Monte Carlo sampling. Markov Chain Monte Carlo is is slower than matrix inversion methods, but is guaranteed to sample proper probability vectors, so no clipping strategies are required.
|
||||
An example presenting how to run the method and use posterior samples is available in `examples/bayesian_quantification.py`.
|
||||
|
||||
#### BayesianMAPLS (a Bayesian implementation of EMQ/MLLS)
|
||||
|
||||
`BayesianMAPLS` is a Bayesian variant of EMQ/MLLS proposed by
|
||||
Ye, C. et al. (2024). Label shift estimation for class-imbalance problem: A
|
||||
Bayesian approach. Proceedings of the IEEE/CVF Winter Conference on
|
||||
Applications of Computer Vision (WACV 2024). QuaPy's implementation is
|
||||
adapted from the [authors' reference code](https://github.com/ChangkunYe/MAPLS/blob/main/label_shift/mapls.py).
|
||||
Rather than returning a single point estimate for the class prevalence, it
|
||||
places a Dirichlet prior over the sought prevalence vector (in an
|
||||
unconstrained, Isometric-Log-Ratio-transformed space) and samples from the
|
||||
resulting posterior via Markov Chain Monte Carlo (using `numpyro`/`jax`),
|
||||
conditioned on a preliminary MAP estimate obtained via the underlying `mapls`
|
||||
routine. Like `BayesianCC`, its `aggregate` method returns the posterior mean,
|
||||
while `predict_conf` additionally returns a confidence region (`intervals`,
|
||||
`ellipse`, `ellipse-clr`, or `ellipse-ilr`) built from the posterior samples.
|
||||
For interval regions, all Bayesian methods also accept `bonferroni=True` to
|
||||
apply Bonferroni correction; this flag has no effect for ellipse-based regions.
|
||||
|
||||
This method requires installation of additional dependencies
|
||||
(`$ pip install quapy[bayes]`) needed to run MCMC sampling; parameters
|
||||
`num_warmup` and `num_samples` control the length of the chain, and `prior`
|
||||
allows choosing between a uniform Dirichlet prior (default) or one of the
|
||||
data-dependent priors ("map"/"map2") proposed in the original paper.
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
from quapy.method._bayesian import BayesianMAPLS
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
train, test = qp.datasets.fetch_UCIBinaryDataset('haberman').train_test
|
||||
|
||||
model = BayesianMAPLS(LogisticRegression())
|
||||
model.fit(*train.Xy)
|
||||
estim_prevalence, conf_region = model.predict_conf(test.X)
|
||||
```
|
||||
|
||||
#### PQ: Precise Quantifier (a Bayesian implementation of HDy)
|
||||
|
||||
`PQ` (Precise Quantifier), available at `qp.method.confidence.PQ`, is a
|
||||
Bayesian distribution-matching variant of `HDy` proposed in
|
||||
[Igiraneza, A.B., Fraser, C., and Hinch, R. (2025). Estimating prevalence
|
||||
with precision and accuracy.](https://arxiv.org/abs/2507.06061)
|
||||
Rather than matching a single test histogram against a mixture of two
|
||||
class-conditional histograms via a divergence measure (as `HDy` does), `PQ`
|
||||
places the histogram-matching problem in a Bayesian setting and samples the
|
||||
full posterior distribution over the (binary) prevalence value via Markov
|
||||
Chain Monte Carlo (using `stan`). Its `aggregate` method returns the
|
||||
posterior mean, while `predict_conf` additionally returns a confidence
|
||||
region built from the posterior samples (`intervals`, `ellipse`, or
|
||||
`ellipse-clr`).
|
||||
|
||||
`PQ` accepts `nbins` (the number of histogram bins, quantile-based by default,
|
||||
or uniform if `fixed_bins=True`), and the usual MCMC controls `num_warmup`,
|
||||
`num_samples`, and `stan_seed`. This method relies on the optional `stan`
|
||||
dependency, installed via `$ pip install quapy[bayes]`.
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
from quapy.method.confidence import PQ
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
train, test = qp.datasets.fetch_UCIBinaryDataset('haberman').train_test
|
||||
|
||||
model = PQ(LogisticRegression())
|
||||
model.fit(*train.Xy)
|
||||
estim_prevalence, conf_region = model.predict_conf(test.X)
|
||||
```
|
||||
|
||||
#### BayesianKDEy (a Bayesian implementation of KDEyML)
|
||||
|
||||
`BayesianKDEy`, available at `qp.method._bayesian.BayesianKDEy`, is a Bayesian
|
||||
version of KDEy proposed by [Moreo et al. 2026](https://arxiv.org/abs/2607.04977).
|
||||
Instead of solving for the single prevalence vector that
|
||||
minimizes a divergence between the test distribution and a KDE-based mixture
|
||||
model (as the KDEy variants above do), `BayesianKDEy` places a Dirichlet
|
||||
prior over the prevalence vector and samples its posterior via Markov Chain
|
||||
Monte Carlo (using `numpyro`/`jax`), conditioned on the same KDE mixture
|
||||
components. Its `aggregate` method returns the posterior mean, while
|
||||
`predict_conf` additionally returns a confidence region built from the
|
||||
posterior samples.
|
||||
|
||||
In addition to the `kernel` and `bandwidth` hyperparameters (with the same
|
||||
`gaussian`/`aitchison`/`ilr` kernel choice, and `shrinkage` regularization
|
||||
for the latter two, available in `KDEyML`), `BayesianKDEy` exposes the usual
|
||||
MCMC controls: `num_warmup`, `num_samples`, `mcmc_seed`, a `temperature` for
|
||||
posterior calibration, and `prior` for choosing the Dirichlet prior
|
||||
(`'uniform'` by default, or a custom scalar/array). This method relies on
|
||||
the optional MCMC dependencies, installed via
|
||||
`$ pip install quapy[bayes]`.
|
||||
|
||||
```python
|
||||
import quapy as qp
|
||||
from quapy.method._bayesian import BayesianKDEy
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
train, test = qp.datasets.fetch_UCIBinaryDataset('haberman').train_test
|
||||
|
||||
model = BayesianKDEy(LogisticRegression(), bandwidth=0.1)
|
||||
model.fit(*train.Xy)
|
||||
estim_prevalence, conf_region = model.predict_conf(test.X)
|
||||
```
|
||||
|
||||
|
|
|
|||
|
After Width: | Height: | Size: 100 KiB |
|
After Width: | Height: | Size: 632 KiB |
|
After Width: | Height: | Size: 262 KiB |
|
After Width: | Height: | Size: 90 KiB |
|
After Width: | Height: | Size: 132 KiB |
|
|
@ -251,3 +251,41 @@ In those cases, however, it is likely that the variances of each
|
|||
method get higher, to the detriment of the visualization.
|
||||
We recommend to set _show_std=False_ in those cases
|
||||
in order to hide the color bands.
|
||||
|
||||
## Simplex Visualisation
|
||||
|
||||
For three-class problems, prevalence vectors lie on the 2-dimensional probability simplex.
|
||||
The function `qp.plot.plot_simplex` provides a lightweight ternary plot that can combine
|
||||
scatter layers, shaded regions, and density overlays.
|
||||
|
||||
A simplex plot is specified through optional layers. Point layers are dictionaries with
|
||||
fields `points`, `label`, and `style`; region layers are dictionaries with fields `fn`,
|
||||
`label`, `color`, and `alpha`.
|
||||
|
||||
```python
|
||||
import numpy as np
|
||||
import quapy as qp
|
||||
|
||||
true_prev = np.array([0.20, 0.35, 0.45])
|
||||
train_prev = np.array([0.50, 0.30, 0.20])
|
||||
posterior_cloud = np.random.default_rng(0).dirichlet(alpha=40 * true_prev, size=200)
|
||||
|
||||
qp.plot.plot_simplex(
|
||||
point_layers=[
|
||||
{'points': posterior_cloud, 'label': 'posterior cloud', 'style': {'s': 12, 'alpha': 0.25, 'color': 'steelblue'}},
|
||||
{'points': true_prev, 'label': 'true prevalence', 'style': {'s': 90, 'color': 'black'}},
|
||||
{'points': train_prev, 'label': 'training prevalence', 'style': {'s': 90, 'color': 'darkorange'}},
|
||||
],
|
||||
density_function=lambda p: np.exp(-40 * np.sum((p - true_prev) ** 2, axis=1)),
|
||||
class_names=['class A', 'class B', 'class C'],
|
||||
savepath='./plots/simplex.png',
|
||||
)
|
||||
```
|
||||
|
||||
See the dedicated
|
||||
[example](https://github.com/HLT-ISTI/QuaPy/blob/master/examples/19.visualizing_simplex.py)
|
||||
for a slightly richer illustration. The current example combines a posterior
|
||||
cloud, the true/training/predicted prevalences, a smooth density surface, and
|
||||
a region induced by Bonferroni-corrected 95% confidence intervals.
|
||||
|
||||

|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ QuaPy provides implementations of most popular sample generation protocols
|
|||
used in literature. This is the subject of the following sections.
|
||||
|
||||
|
||||
## Artificial-Prevalence Protocol
|
||||
## APP: Artificial-Prevalence Protocol
|
||||
|
||||
The "artificial-sampling protocol" (APP) proposed by
|
||||
[Forman (2005)](https://link.springer.com/chapter/10.1007/11564096_55)
|
||||
|
|
@ -43,21 +43,21 @@ desired prevalence values covering the full spectrum.
|
|||
|
||||
In APP, the user specifies the number
|
||||
of (equally distant) points to be generated from the interval [0,1];
|
||||
in QuaPy this is achieved by setting _n_prevpoints_.
|
||||
For example, if _n_prevpoints=11_ then, for each class, the prevalence values
|
||||
in QuaPy this is achieved by setting _n_prevalences_.
|
||||
For example, if _n_prevalences=11_ then, for each class, the prevalence values
|
||||
[0., 0.1, 0.2, ..., 1.] will be used. This means that, for two classes,
|
||||
the number of different prevalence values will be 11 (since, once the prevalence
|
||||
of one class is determined, the other one is constrained). For 3 classes,
|
||||
the number of valid combinations can be obtained as 11 + 10 + ... + 1 = 66.
|
||||
In general, the number of valid combinations that will be produced for a given
|
||||
value of n_prevpoints can be consulted by invoking
|
||||
value of _n_prevalences_ can be consulted by invoking
|
||||
_num_prevalence_combinations_, e.g.:
|
||||
|
||||
```python
|
||||
import quapy.functional as F
|
||||
n_prevpoints = 21
|
||||
n_prevalences = 21
|
||||
n_classes = 4
|
||||
n = F.num_prevalence_combinations(n_prevpoints, n_classes, n_repeats=1)
|
||||
n = F.num_prevalence_combinations(n_prevalences, n_classes, n_repeats=1)
|
||||
```
|
||||
|
||||
in this example, _n=1771_. Note the last argument, _n_repeats_, that
|
||||
|
|
@ -74,13 +74,13 @@ _get_nprevpoints_approximation_, e.g.:
|
|||
|
||||
```python
|
||||
budget = 5000
|
||||
n_prevpoints = F.get_nprevpoints_approximation(budget, n_classes, n_repeats=1)
|
||||
n = F.num_prevalence_combinations(n_prevpoints, n_classes, n_repeats=1)
|
||||
print(f'by setting n_prevpoints={n_prevpoints} the number of evaluations for {n_classes} classes will be {n}')
|
||||
n_prevalences = F.get_nprevpoints_approximation(budget, n_classes, n_repeats=1)
|
||||
n = F.num_prevalence_combinations(n_prevalences, n_classes, n_repeats=1)
|
||||
print(f'by setting n_prevalences={n_prevalences} the number of evaluations for {n_classes} classes will be {n}')
|
||||
```
|
||||
this will produce the following output:
|
||||
```
|
||||
by setting n_prevpoints=30 the number of evaluations for 4 classes will be 4960
|
||||
by setting n_prevalences=30 the number of evaluations for 4 classes will be 4960
|
||||
```
|
||||
|
||||
The following code shows an example of usage of APP for model selection
|
||||
|
|
@ -129,7 +129,16 @@ in such cases QuaPy takes the value of _qp.environ['SAMPLE_SIZE']_.
|
|||
This protocol is useful for testing a quantifier under conditions of
|
||||
_prior probability shift_.
|
||||
|
||||
## Sampling from the unit-simplex, the Uniform-Prevalence Protocol (UPP)
|
||||
The following ternary plot, generated by [example 21](https://github.com/HLT-ISTI/QuaPy/blob/master/examples/21.visualizing_protocols.py),
|
||||
shows the prevalence values covered by a grid-based APP in a three-class problem (`academic-success`):
|
||||
|
||||

|
||||
|
||||
Each point corresponds to one sampled prevalence vector. As expected, the
|
||||
points lie on a regular grid over the simplex, ensuring systematic coverage of
|
||||
the prevalence space.
|
||||
|
||||
## UPP: Sampling from the unit-simplex, the Uniform-Prevalence Protocol
|
||||
|
||||
Generating all possible combinations from a grid of prevalence values (APP) in
|
||||
multiclass is cumbersome, and when the number of classes increases it rapidly
|
||||
|
|
@ -148,7 +157,7 @@ for sampling from the unit-simplex as many vectors of prevalence values as indic
|
|||
in the _repeats_ parameter. UPP can be instantiated as:
|
||||
|
||||
```python
|
||||
protocol = qp.in_protocol.UPP(test, repeats=100)
|
||||
protocol = qp.protocol.UPP(test, repeats=100)
|
||||
```
|
||||
|
||||
This is the most convenient protocol for datasets
|
||||
|
|
@ -157,8 +166,22 @@ containing many classes; see, e.g.,
|
|||
and is useful for testing a quantifier under conditions of
|
||||
_prior probability shift_.
|
||||
|
||||
The next plot shows one such protocol, labelled in example 21 as
|
||||
_APP(Kraemer)_, to emphasize that it plays the role of an artificial-prevalence
|
||||
protocol without relying on a fixed grid:
|
||||
|
||||
## Natural-Prevalence Protocol
|
||||

|
||||
|
||||
Unlike grid-based APP, UPP does not force prevalence vectors to lie on a
|
||||
regular lattice. Instead, it spreads samples over the simplex in a
|
||||
statistically uniform way, making it attractive when the number of classes is
|
||||
large and exhaustive grids become impractical.
|
||||
|
||||
*Note* that UPP is actually a different (modern) implementation of the Artificial Prevalence Protocol,
|
||||
and is here given a different name simply to allow both implementations coexist in QuaPy. "UPP" is not a
|
||||
proper accademic name, and practitioners should rather refer to it as APP.
|
||||
|
||||
## NPP: Natural-Prevalence Protocol
|
||||
|
||||
The "natural-prevalence protocol" (NPP) comes down to generating samples drawn
|
||||
uniformly at random from the original labelled collection. This protocol has
|
||||
|
|
@ -168,9 +191,47 @@ All other things being equal, this protocol can be used just like APP or UPP,
|
|||
and is instantiated via:
|
||||
|
||||
```python
|
||||
protocol = qp.in_protocol.NPP(test, repeats=100)
|
||||
protocol = qp.protocol.NPP(test, repeats=100)
|
||||
```
|
||||
|
||||
The prevalence coverage of NPP is much more concentrated, since the samples are
|
||||
obtained by plain random subsampling from the test set and therefore remain
|
||||
close to its natural prevalence:
|
||||
|
||||

|
||||
|
||||
This makes NPP useful when one wants to evaluate quantifiers under mild or
|
||||
realistic drift conditions, but much less suitable than APP or UPP for stress
|
||||
testing performance across the full simplex.
|
||||
|
||||
## Dirichlet Protocol
|
||||
|
||||
QuaPy also implements a :class:`DirichletProtocol`, which samples prevalence
|
||||
vectors from a Dirichlet distribution before drawing the corresponding sample
|
||||
from the labelled collection:
|
||||
|
||||
```python
|
||||
protocol = qp.protocol.DirichletProtocol(test, alpha=0.2, repeats=100)
|
||||
```
|
||||
|
||||
The parameter `alpha` controls how concentrated the protocol is. Small values
|
||||
of `alpha` favour sparse prevalence vectors near the corners of the simplex,
|
||||
while larger values generate more balanced mixtures. When all entries of
|
||||
`alpha` are equal to 1, the protocol becomes uniformly distributed over the
|
||||
simplex, similarly in spirit to UPP.
|
||||
|
||||
The following plot shows the effect of a sparse prior with `alpha=0.2`:
|
||||
|
||||

|
||||
|
||||
Compared to UPP, the mass is clearly pulled towards the vertices and edges,
|
||||
thus producing more extreme label-shift scenarios.
|
||||
|
||||
The parameter `alpha` in the Dirichlet distribution is typically defined as an
|
||||
array of shape `(n_classes)`. When the user specifies a single value, QuaPy
|
||||
broadcasts this value for all classes. Conversely, a different value can be
|
||||
specified for each class.
|
||||
|
||||
## Other protocols
|
||||
|
||||
Other protocols exist in QuaPy and will be added to the `qp.protocol.py` module.
|
||||
|
|
@ -60,6 +60,14 @@ quapy.method.composable module
|
|||
:undoc-members:
|
||||
:show-inheritance:
|
||||
|
||||
quapy.method.confidence module
|
||||
------------------------------
|
||||
|
||||
.. automodule:: quapy.method.confidence
|
||||
:members:
|
||||
:undoc-members:
|
||||
:show-inheritance:
|
||||
|
||||
Module contents
|
||||
---------------
|
||||
|
||||
|
|
|
|||
|
|
@ -36,7 +36,7 @@ with qp.util.temp_seed(0):
|
|||
true_prev = shifted_test.prevalence()
|
||||
|
||||
# by calling "quantify_conf", we obtain the point estimate and the confidence intervals around it
|
||||
pred_prev, conf_intervals = pacc.quantify_conf(shifted_test.X)
|
||||
pred_prev, conf_intervals = pacc.predict_conf(shifted_test.X)
|
||||
|
||||
# conf_intervals is an instance of ConfidenceRegionABC, which provides some useful utilities like:
|
||||
# - coverage: a function which computes the fraction of true values that belong to the confidence region
|
||||
|
|
@ -75,7 +75,8 @@ There are different ways for constructing confidence regions implemented in QuaP
|
|||
convenient for taking into account the inner structure of the probability simplex)
|
||||
use: AggregativeBootstrap(PACC(), confidence_level=0.95, method='ellipse-clr')
|
||||
|
||||
Other methods that return confidence regions in QuaPy include the BayesianCC method.
|
||||
Other methods that return confidence regions in QuaPy include the Bayesian methods (BayesianCC,
|
||||
BayesianMAPLS, PQ, and BayesianKDEy).
|
||||
"""
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,60 @@
|
|||
from sklearn.feature_extraction.text import CountVectorizer
|
||||
from sklearn.feature_selection import SelectKBest, chi2
|
||||
|
||||
import quapy as qp
|
||||
from quapy.method.non_aggregative import ReadMe
|
||||
import quapy.functional as F
|
||||
from sklearn.pipeline import Pipeline
|
||||
|
||||
"""
|
||||
This example showcases how to use the non-aggregative method ReadMe proposed by Hopkins and King.
|
||||
This method is for text analysis, so let us first instantiate a dataset for sentiment quantification (we
|
||||
use IMDb for this example). The method is quite computationally expensive, so we will restrict the training
|
||||
set to 1000 documents only.
|
||||
"""
|
||||
reviews = qp.datasets.fetch_reviews('imdb').reduce(n_train=1000, random_state=0)
|
||||
|
||||
"""
|
||||
We need to convert text to bag-of-words representations. Actually, ReadMe requires the representations to be
|
||||
binary (i.e., storing a 1 whenever a document contains certain word, or 0 otherwise), so we will not use
|
||||
TFIDF weighting. We will also retain the top 1000 most important features according to chi2.
|
||||
"""
|
||||
encode_0_1 = Pipeline([
|
||||
('0_1_terms', CountVectorizer(min_df=5, binary=True)),
|
||||
('feat_sel', SelectKBest(chi2, k=1000))
|
||||
])
|
||||
train, test = qp.data.preprocessing.instance_transformation(reviews, encode_0_1, inplace=True).train_test
|
||||
|
||||
"""
|
||||
We now instantiate ReadMe, with the prob_model='full' (default behaviour, implementing the Hopkins and King original
|
||||
idea). This method consists of estimating Q(Y) by solving:
|
||||
|
||||
Q(X) = \sum_i Q(X|Y=i) Q(Y=i)
|
||||
|
||||
without resorting to estimating the posteriors Q(Y=i|X), by solving a linear least-squares problem.
|
||||
However, since Q(X) and Q(X|Y=i) are matrices of shape (2^K, 1) and (2^K, n), with K the number of features
|
||||
and n the number of classes, their calculation becomes intractable. ReadMe instead performs bagging (i.e., it
|
||||
samples small sets of features and averages the results) thus reducing K to a few terms. In our example we
|
||||
set K (bagging_range) to 20, and the number of bagging_trials to 100.
|
||||
|
||||
ReadMe also computes confidence intervals via bootstrap. We set the number of bootstrap trials to 100.
|
||||
"""
|
||||
readme = ReadMe(prob_model='full', bootstrap_trials=100, bagging_trials=100, bagging_range=20, random_state=0, verbose=True)
|
||||
readme.fit(*train.Xy) # <- there is actually nothing happening here (only bootstrap resampling); the method is "lazy"
|
||||
# and postpones most of the calculations to the test phase.
|
||||
|
||||
# since the method is slow, we will only test 3 cases with different imbalances
|
||||
few_negatives = [0.25, 0.75]
|
||||
balanced = [0.5, 0.5]
|
||||
few_positives = [0.75, 0.25]
|
||||
|
||||
for test_prev in [few_negatives, balanced, few_positives]:
|
||||
sample = reviews.test.sampling(500, *test_prev, random_state=0) # draw sets of 500 documents with desired prevs
|
||||
prev_estim, conf = readme.predict_conf(sample.X)
|
||||
err = qp.error.mae(sample.prevalence(), prev_estim)
|
||||
print(f'true-prevalence={F.strprev(sample.prevalence())},\n'
|
||||
f'predicted-prevalence={F.strprev(prev_estim)}, with confidence intervals {conf},\n'
|
||||
f'MAE={err:.4f}')
|
||||
|
||||
|
||||
|
||||
|
|
@ -0,0 +1,68 @@
|
|||
import numpy as np
|
||||
import quapy as qp
|
||||
from quapy.method.confidence import ConfidenceIntervals
|
||||
|
||||
|
||||
"""
|
||||
A minimal example showing how to visualise ternary prevalences on the simplex.
|
||||
The plot combines a cloud of posterior triplets, a few reference prevalences, a
|
||||
confidence ellipse induced by the cloud, and a smooth density centred around the
|
||||
true prevalence.
|
||||
"""
|
||||
|
||||
rng = np.random.default_rng(0)
|
||||
true_prev = np.array([0.20, 0.35, 0.45])
|
||||
train_prev = np.array([0.50, 0.30, 0.20])
|
||||
pred_prev = np.array([0.18, 0.39, 0.43])
|
||||
posterior_cloud = rng.dirichlet(alpha=45 * true_prev, size=250)
|
||||
|
||||
point_layers = [
|
||||
{
|
||||
'points': posterior_cloud,
|
||||
'label': 'posterior cloud',
|
||||
'style': {'s': 12, 'alpha': 0.25, 'color': 'steelblue', 'edgecolors': 'none'},
|
||||
},
|
||||
{
|
||||
'points': true_prev,
|
||||
'label': 'true prevalence',
|
||||
'style': {'s': 70, 'color': 'black'},
|
||||
},
|
||||
{
|
||||
'points': pred_prev,
|
||||
'label': 'predicted prevalence',
|
||||
'style': {'s': 70, 'color': 'crimson'},
|
||||
},
|
||||
{
|
||||
'points': train_prev,
|
||||
'label': 'training prevalence',
|
||||
'style': {'s': 70, 'color': 'darkorange'},
|
||||
},
|
||||
]
|
||||
|
||||
confidence_region = ConfidenceIntervals(posterior_cloud, confidence_level=0.95, bonferroni_correction=True)
|
||||
|
||||
region_layers = [
|
||||
{
|
||||
'fn': lambda p: float(p in confidence_region),
|
||||
'label': '95% confidence intervals',
|
||||
'color': 'seagreen',
|
||||
'alpha': 0.15,
|
||||
}
|
||||
]
|
||||
|
||||
density = lambda p: np.exp(-45 * np.sum((p - true_prev) ** 2, axis=1))
|
||||
|
||||
qp.plot.plot_simplex(
|
||||
point_layers=point_layers,
|
||||
region_layers=region_layers,
|
||||
density_function=density,
|
||||
density_color='royalblue',
|
||||
class_names=['class A', 'class B', 'class C'],
|
||||
title='Ternary prevalence visualisation',
|
||||
legend_ncol=3,
|
||||
figsize=(7.2, 5.8),
|
||||
class_name_fontsize=9,
|
||||
title_fontsize=10,
|
||||
legend_fontsize=8,
|
||||
savepath='./plots/simplex_visualization.png',
|
||||
)
|
||||
|
|
@ -0,0 +1,64 @@
|
|||
import quapy as qp
|
||||
from quapy.data.datasets import fetch_image_embeddings
|
||||
from quapy.method.aggregative import EMQ, RLLS
|
||||
from quapy.classification.calibration import TemperatureScalingFromLogits
|
||||
from quapy.protocol import UPP
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
|
||||
# This example illustrates how to run experiments with image datasets, in this case with CIFAR10
|
||||
# The datasets available in quapy do not consist of raw image files, but are instead pre-generated
|
||||
# embeddings (see the manuals for further information).
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
# Let us begin with a typical case in which the embeddings come from the penultimate layer of a neural
|
||||
# model (in this case, a resnet18). We get these representations by specifying embedding='features'
|
||||
|
||||
print('fetching cifar10 embeddings')
|
||||
train, test = fetch_image_embeddings(dataset_name='cifar10', embedding='features').train_test
|
||||
|
||||
print('training:', train)
|
||||
print('test:', test)
|
||||
|
||||
Xtr, ytr = train.Xy
|
||||
|
||||
# let us train an Expectation Maximazion Quantifier (EMQ), aka Maximum Likelihood for Label Shift (MLLS)
|
||||
# using a logistic regressor as the underlying classifier, with Bias Corrected Temperature Scaling (BCTS)
|
||||
|
||||
bcts_emq = EMQ(classifier=LogisticRegression(), calib='bcts', val_split=5)
|
||||
|
||||
print(f'fitting quantifier {bcts_emq}')
|
||||
bcts_emq.fit(Xtr, ytr)
|
||||
|
||||
# we generate many samples exhibiting prior probability shift with the artificial prevalence protocol
|
||||
# (we use the multiclass variant UPP instead of the grid-based APP)
|
||||
|
||||
qp.environ["SAMPLE_SIZE"] = 500 # when the sample size is common to all experiments, it is conveniet to set it once and for all
|
||||
artificial_prev_prot = UPP(test, repeats=200)
|
||||
print('generating 200 test bags of 500 instances each')
|
||||
|
||||
bctsemq_report = qp.evaluation.evaluation_report(bcts_emq, protocol=artificial_prev_prot, error_metrics=['mae', 'mrae'])
|
||||
print(bctsemq_report.mean(numeric_only=True))
|
||||
|
||||
# we could instead use the pre-generated logits of the resnet18
|
||||
print('fetching cifar10 logits')
|
||||
train, test = fetch_image_embeddings(dataset_name='cifar10', embedding='logits').train_test
|
||||
Xtr, ytr = train.Xy
|
||||
|
||||
# in this case, the representations are already classification-related outputs;
|
||||
# we can convert them into (hopefully well-) calibrated outputs via TemperatureScaling
|
||||
|
||||
print('generating posterior probabilities out of logits via temperature scaling')
|
||||
emq = EMQ(classifier=TemperatureScalingFromLogits(bias_corrected=True))
|
||||
emq.fit(Xtr, ytr)
|
||||
|
||||
print('generating 200 test bags of 500 instances each')
|
||||
artificial_prev_prot = UPP(test, repeats=200)
|
||||
emq_report = qp.evaluation.evaluation_report(emq, protocol=artificial_prev_prot, error_metrics=['mae', 'mrae'])
|
||||
print(emq_report.mean(numeric_only=True))
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -0,0 +1,51 @@
|
|||
import numpy as np
|
||||
import quapy as qp
|
||||
from quapy.data.datasets import fetch_UCIMulticlassDataset
|
||||
from quapy.protocol import APP, NPP, UPP, DirichletProtocol
|
||||
|
||||
|
||||
"""
|
||||
Ternary plots showcasing different sampling protocols.
|
||||
"""
|
||||
|
||||
rng = np.random.default_rng(0)
|
||||
|
||||
train, test = fetch_UCIMulticlassDataset(dataset_name='academic-success').train_test
|
||||
|
||||
train_prev = {
|
||||
'points': train.prevalence(),
|
||||
'label': 'training prevalence',
|
||||
'style': {'s': 70, 'color': 'darkorange'},
|
||||
}
|
||||
|
||||
def protocols():
|
||||
yield 'app-grid', 'Artificial Prevalence Protocol (grid)', APP(test, n_prevalences=21, repeats=1, sample_size=100)
|
||||
yield 'app-kraemer', 'Artificial Prevalence Protocol (Kraemer)', UPP(test, repeats=5000, sample_size=500)
|
||||
yield 'npp', 'Natural Prevalence Protocol', NPP(test, repeats=1000, sample_size=100)
|
||||
yield 'dirichlet', 'Dirichlet(alpha=0.2)', DirichletProtocol(test, alpha=0.2, repeats=5000, sample_size=100)
|
||||
|
||||
for file_name, prot_name, protocol in protocols():
|
||||
app_points = {
|
||||
'points': [prev for _, prev in protocol()],
|
||||
'label': prot_name,
|
||||
'style': {'s': 15, 'alpha': 0.5, 'color': 'steelblue', 'edgecolors': 'none'},
|
||||
}
|
||||
|
||||
point_layers = [
|
||||
app_points,
|
||||
train_prev,
|
||||
]
|
||||
|
||||
dispersion = 0.1
|
||||
|
||||
qp.plot.plot_simplex(
|
||||
point_layers=point_layers,
|
||||
class_names=['class A', 'class B', 'class C'],
|
||||
#title='Ternary prevalence visualisation',
|
||||
legend_ncol=3,
|
||||
figsize=(7.2, 5.8),
|
||||
class_name_fontsize=9,
|
||||
title_fontsize=10,
|
||||
legend_fontsize=8,
|
||||
savepath=f'./plots/{file_name}.png',
|
||||
)
|
||||
|
|
@ -0,0 +1,254 @@
|
|||
from scipy.sparse import csc_matrix, csr_matrix
|
||||
from sklearn.base import BaseEstimator, TransformerMixin
|
||||
from sklearn.feature_extraction.text import TfidfTransformer, TfidfVectorizer, CountVectorizer
|
||||
import numpy as np
|
||||
from joblib import Parallel, delayed
|
||||
import sklearn
|
||||
import math
|
||||
from scipy.stats import t
|
||||
|
||||
|
||||
class ContTable:
|
||||
def __init__(self, tp=0, tn=0, fp=0, fn=0):
|
||||
self.tp=tp
|
||||
self.tn=tn
|
||||
self.fp=fp
|
||||
self.fn=fn
|
||||
|
||||
def get_d(self): return self.tp + self.tn + self.fp + self.fn
|
||||
|
||||
def get_c(self): return self.tp + self.fn
|
||||
|
||||
def get_not_c(self): return self.tn + self.fp
|
||||
|
||||
def get_f(self): return self.tp + self.fp
|
||||
|
||||
def get_not_f(self): return self.tn + self.fn
|
||||
|
||||
def p_c(self): return (1.0*self.get_c())/self.get_d()
|
||||
|
||||
def p_not_c(self): return 1.0-self.p_c()
|
||||
|
||||
def p_f(self): return (1.0*self.get_f())/self.get_d()
|
||||
|
||||
def p_not_f(self): return 1.0-self.p_f()
|
||||
|
||||
def p_tp(self): return (1.0*self.tp) / self.get_d()
|
||||
|
||||
def p_tn(self): return (1.0*self.tn) / self.get_d()
|
||||
|
||||
def p_fp(self): return (1.0*self.fp) / self.get_d()
|
||||
|
||||
def p_fn(self): return (1.0*self.fn) / self.get_d()
|
||||
|
||||
def tpr(self):
|
||||
c = 1.0*self.get_c()
|
||||
return self.tp / c if c > 0.0 else 0.0
|
||||
|
||||
def fpr(self):
|
||||
_c = 1.0*self.get_not_c()
|
||||
return self.fp / _c if _c > 0.0 else 0.0
|
||||
|
||||
|
||||
def __ig_factor(p_tc, p_t, p_c):
|
||||
den = p_t * p_c
|
||||
if den != 0.0 and p_tc != 0:
|
||||
return p_tc * math.log(p_tc / den, 2)
|
||||
else:
|
||||
return 0.0
|
||||
|
||||
|
||||
def information_gain(cell):
|
||||
return __ig_factor(cell.p_tp(), cell.p_f(), cell.p_c()) + \
|
||||
__ig_factor(cell.p_fp(), cell.p_f(), cell.p_not_c()) +\
|
||||
__ig_factor(cell.p_fn(), cell.p_not_f(), cell.p_c()) + \
|
||||
__ig_factor(cell.p_tn(), cell.p_not_f(), cell.p_not_c())
|
||||
|
||||
|
||||
def squared_information_gain(cell):
|
||||
return information_gain(cell)**2
|
||||
|
||||
|
||||
def posneg_information_gain(cell):
|
||||
ig = information_gain(cell)
|
||||
if cell.tpr() < cell.fpr():
|
||||
return -ig
|
||||
else:
|
||||
return ig
|
||||
|
||||
|
||||
def pos_information_gain(cell):
|
||||
if cell.tpr() < cell.fpr():
|
||||
return 0
|
||||
else:
|
||||
return information_gain(cell)
|
||||
|
||||
def pointwise_mutual_information(cell):
|
||||
return __ig_factor(cell.p_tp(), cell.p_f(), cell.p_c())
|
||||
|
||||
|
||||
def gss(cell):
|
||||
return cell.p_tp()*cell.p_tn() - cell.p_fp()*cell.p_fn()
|
||||
|
||||
|
||||
def chi_square(cell):
|
||||
den = cell.p_f() * cell.p_not_f() * cell.p_c() * cell.p_not_c()
|
||||
if den==0.0: return 0.0
|
||||
num = gss(cell)**2
|
||||
return num / den
|
||||
|
||||
|
||||
def conf_interval(xt, n):
|
||||
if n>30:
|
||||
z2 = 3.84145882069 # norm.ppf(0.5+0.95/2.0)**2
|
||||
else:
|
||||
z2 = t.ppf(0.5 + 0.95 / 2.0, df=max(n-1,1)) ** 2
|
||||
p = (xt + 0.5 * z2) / (n + z2)
|
||||
amplitude = 0.5 * z2 * math.sqrt((p * (1.0 - p)) / (n + z2))
|
||||
return p, amplitude
|
||||
|
||||
|
||||
def strength(minPosRelFreq, minPos, maxNeg):
|
||||
if minPos > maxNeg:
|
||||
return math.log(2.0 * minPosRelFreq, 2.0)
|
||||
else:
|
||||
return 0.0
|
||||
|
||||
|
||||
#set cancel_features=True to allow some features to be weighted as 0 (as in the original article)
|
||||
#however, for some extremely imbalanced dataset caused all documents to be 0
|
||||
def conf_weight(cell, cancel_features=False):
|
||||
c = cell.get_c()
|
||||
not_c = cell.get_not_c()
|
||||
tp = cell.tp
|
||||
fp = cell.fp
|
||||
|
||||
pos_p, pos_amp = conf_interval(tp, c)
|
||||
neg_p, neg_amp = conf_interval(fp, not_c)
|
||||
|
||||
min_pos = pos_p-pos_amp
|
||||
max_neg = neg_p+neg_amp
|
||||
den = (min_pos + max_neg)
|
||||
minpos_relfreq = min_pos / (den if den != 0 else 1)
|
||||
|
||||
str_tplus = strength(minpos_relfreq, min_pos, max_neg);
|
||||
|
||||
if str_tplus == 0 and not cancel_features:
|
||||
return 1e-20
|
||||
|
||||
return str_tplus
|
||||
|
||||
|
||||
def get_tsr_matrix(cell_matrix, tsr_score_funtion):
|
||||
nC = len(cell_matrix)
|
||||
nF = len(cell_matrix[0])
|
||||
tsr_matrix = [[tsr_score_funtion(cell_matrix[c,f]) for f in range(nF)] for c in range(nC)]
|
||||
return np.array(tsr_matrix)
|
||||
|
||||
|
||||
def feature_label_contingency_table(positive_document_indexes, feature_document_indexes, nD):
|
||||
tp_ = len(positive_document_indexes & feature_document_indexes)
|
||||
fp_ = len(feature_document_indexes - positive_document_indexes)
|
||||
fn_ = len(positive_document_indexes - feature_document_indexes)
|
||||
tn_ = nD - (tp_ + fp_ + fn_)
|
||||
return ContTable(tp=tp_, tn=tn_, fp=fp_, fn=fn_)
|
||||
|
||||
|
||||
def category_tables(feature_sets, category_sets, c, nD, nF):
|
||||
return [feature_label_contingency_table(category_sets[c], feature_sets[f], nD) for f in range(nF)]
|
||||
|
||||
|
||||
def get_supervised_matrix(coocurrence_matrix, label_matrix, n_jobs=-1):
|
||||
"""
|
||||
Computes the nC x nF supervised matrix M where Mcf is the 4-cell contingency table for feature f and class c.
|
||||
Efficiency O(nF x nC x log(S)) where S is the sparse factor
|
||||
"""
|
||||
|
||||
nD, nF = coocurrence_matrix.shape
|
||||
nD2, nC = label_matrix.shape
|
||||
|
||||
if nD != nD2:
|
||||
raise ValueError('Number of rows in coocurrence matrix shape %s and label matrix shape %s is not consistent' %
|
||||
(coocurrence_matrix.shape,label_matrix.shape))
|
||||
|
||||
def nonzero_set(matrix, col):
|
||||
return set(matrix[:, col].nonzero()[0])
|
||||
|
||||
if isinstance(coocurrence_matrix, csr_matrix):
|
||||
coocurrence_matrix = csc_matrix(coocurrence_matrix)
|
||||
feature_sets = [nonzero_set(coocurrence_matrix, f) for f in range(nF)]
|
||||
category_sets = [nonzero_set(label_matrix, c) for c in range(nC)]
|
||||
cell_matrix = Parallel(n_jobs=n_jobs, backend="threading")(
|
||||
delayed(category_tables)(feature_sets, category_sets, c, nD, nF) for c in range(nC)
|
||||
)
|
||||
return np.array(cell_matrix)
|
||||
|
||||
|
||||
class TSRweighting(BaseEstimator,TransformerMixin):
|
||||
"""
|
||||
Supervised Term Weighting function based on any Term Selection Reduction (TSR) function (e.g., information gain,
|
||||
chi-square, etc.) or, more generally, on any function that could be computed on the 4-cell contingency table for
|
||||
each category-feature pair.
|
||||
The supervised_4cell_matrix is a `(n_classes, n_words)` matrix containing the 4-cell contingency tables
|
||||
for each class-word pair, and can be pre-computed (e.g., during the feature selection phase) and passed as an
|
||||
argument.
|
||||
When `n_classes>1`, i.e., in multiclass scenarios, a global_policy is used in order to determine a
|
||||
single feature-score which informs about its relevance. Accepted policies include "max" (takes the max score
|
||||
across categories), "ave" and "wave" (take the average, or weighted average, across all categories -- weights
|
||||
correspond to the class prevalence), and "sum" (which sums all category scores).
|
||||
"""
|
||||
|
||||
def __init__(self, tsr_function, global_policy='max', supervised_4cell_matrix=None, sublinear_tf=True, norm='l2', min_df=3, n_jobs=-1):
|
||||
if global_policy not in ['max', 'ave', 'wave', 'sum']: raise ValueError('Global policy should be in {"max", "ave", "wave", "sum"}')
|
||||
self.tsr_function = tsr_function
|
||||
self.global_policy = global_policy
|
||||
self.supervised_4cell_matrix = supervised_4cell_matrix
|
||||
self.sublinear_tf = sublinear_tf
|
||||
self.norm = norm
|
||||
self.min_df = min_df
|
||||
self.n_jobs = n_jobs
|
||||
|
||||
def fit(self, X, y):
|
||||
self.count_vectorizer = CountVectorizer(min_df=self.min_df)
|
||||
X = self.count_vectorizer.fit_transform(X)
|
||||
|
||||
self.tf_vectorizer = TfidfTransformer(
|
||||
norm=None, use_idf=False, smooth_idf=False, sublinear_tf=self.sublinear_tf
|
||||
).fit(X)
|
||||
|
||||
if len(y.shape) == 1:
|
||||
y = np.expand_dims(y, axis=1)
|
||||
|
||||
nD, nC = y.shape
|
||||
nF = len(self.tf_vectorizer.get_feature_names_out())
|
||||
|
||||
if self.supervised_4cell_matrix is None:
|
||||
self.supervised_4cell_matrix = get_supervised_matrix(X, y, n_jobs=self.n_jobs)
|
||||
else:
|
||||
if self.supervised_4cell_matrix.shape != (nC, nF):
|
||||
raise ValueError("Shape of supervised information matrix is inconsistent with X and y")
|
||||
|
||||
tsr_matrix = get_tsr_matrix(self.supervised_4cell_matrix, self.tsr_function)
|
||||
|
||||
if self.global_policy == 'ave':
|
||||
self.global_tsr_vector = np.average(tsr_matrix, axis=0)
|
||||
elif self.global_policy == 'wave':
|
||||
category_prevalences = [sum(y[:,c])*1.0/nD for c in range(nC)]
|
||||
self.global_tsr_vector = np.average(tsr_matrix, axis=0, weights=category_prevalences)
|
||||
elif self.global_policy == 'sum':
|
||||
self.global_tsr_vector = np.sum(tsr_matrix, axis=0)
|
||||
elif self.global_policy == 'max':
|
||||
self.global_tsr_vector = np.amax(tsr_matrix, axis=0)
|
||||
return self
|
||||
|
||||
def fit_transform(self, X, y):
|
||||
return self.fit(X,y).transform(X)
|
||||
|
||||
def transform(self, X):
|
||||
if not hasattr(self, 'global_tsr_vector'): raise NameError('TSRweighting: transform method called before fit.')
|
||||
X = self.count_vectorizer.transform(X)
|
||||
tf_X = self.tf_vectorizer.transform(X).toarray()
|
||||
weighted_X = np.multiply(tf_X, self.global_tsr_vector)
|
||||
if self.norm is not None and self.norm!='none':
|
||||
weighted_X = sklearn.preprocessing.normalize(weighted_X, norm=self.norm, axis=1, copy=False)
|
||||
return csr_matrix(weighted_X)
|
||||
|
|
@ -0,0 +1,208 @@
|
|||
from scipy.sparse import issparse
|
||||
from sklearn.decomposition import TruncatedSVD
|
||||
from sklearn.feature_extraction.text import TfidfVectorizer, CountVectorizer
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
from sklearn.preprocessing import StandardScaler
|
||||
|
||||
import quapy as qp
|
||||
from data import LabelledCollection
|
||||
import numpy as np
|
||||
|
||||
from experimental_non_aggregative.custom_vectorizers import *
|
||||
from method._kdey import KDEBase
|
||||
from protocol import APP
|
||||
from quapy.method.aggregative import HDy, DistributionMatchingY
|
||||
from quapy.method.base import BaseQuantifier
|
||||
from scipy import optimize
|
||||
import pandas as pd
|
||||
import quapy.functional as F
|
||||
|
||||
|
||||
# TODO: explore the bernoulli (term presence/absence) variant
|
||||
# TODO: explore the multinomial (term frequency) variant
|
||||
# TODO: explore the multinomial + length normalization variant
|
||||
# TODO: consolidate the TSR-variant (e.g., using information gain) variant;
|
||||
# - works better with the idf?
|
||||
# - works better with length normalization?
|
||||
# - etc
|
||||
|
||||
class DxS(BaseQuantifier):
|
||||
def __init__(self, vectorizer=None, divergence='topsoe'):
|
||||
self.vectorizer = vectorizer
|
||||
self.divergence = divergence
|
||||
|
||||
# def __as_distribution(self, instances):
|
||||
# return np.asarray(instances.sum(axis=0) / instances.sum()).flatten()
|
||||
|
||||
def __as_distribution(self, instances):
|
||||
dist = instances.mean(axis=0)
|
||||
return np.asarray(dist).flatten()
|
||||
|
||||
def fit(self, text_instances, labels):
|
||||
|
||||
classes = np.unique(labels)
|
||||
|
||||
if self.vectorizer is not None:
|
||||
text_instances = self.vectorizer.fit_transform(text_instances, y=labels)
|
||||
|
||||
distributions = []
|
||||
for class_i in classes:
|
||||
distributions.append(self.__as_distribution(text_instances[labels == class_i]))
|
||||
|
||||
self.validation_distribution = np.asarray(distributions)
|
||||
|
||||
return self
|
||||
|
||||
def predict(self, text_instances):
|
||||
if self.vectorizer is not None:
|
||||
text_instances = self.vectorizer.transform(text_instances)
|
||||
|
||||
test_distribution = self.__as_distribution(text_instances)
|
||||
divergence = qp.functional.get_divergence(self.divergence)
|
||||
n_classes, n_feats = self.validation_distribution.shape
|
||||
|
||||
def match(prev):
|
||||
prev = np.expand_dims(prev, axis=0)
|
||||
mixture_distribution = (prev @ self.validation_distribution).flatten()
|
||||
return divergence(test_distribution, mixture_distribution)
|
||||
|
||||
# the initial point is set as the uniform distribution
|
||||
uniform_distribution = np.full(fill_value=1 / n_classes, shape=(n_classes,))
|
||||
|
||||
# solutions are bounded to those contained in the unit-simplex
|
||||
bounds = tuple((0, 1) for x in range(n_classes)) # values in [0,1]
|
||||
constraints = ({'type': 'eq', 'fun': lambda x: 1 - sum(x)}) # values summing up to 1
|
||||
r = optimize.minimize(match, x0=uniform_distribution, method='SLSQP', bounds=bounds, constraints=constraints)
|
||||
return r.x
|
||||
|
||||
|
||||
|
||||
class KDExML(BaseQuantifier, KDEBase):
|
||||
|
||||
def __init__(self, bandwidth=0.1, standardize=False):
|
||||
self._check_bandwidth(bandwidth)
|
||||
self.bandwidth = bandwidth
|
||||
self.standardize = standardize
|
||||
|
||||
def fit(self, X, y):
|
||||
classes = sorted(np.unique(y))
|
||||
|
||||
if self.standardize:
|
||||
self.scaler = StandardScaler()
|
||||
X = self.scaler.fit_transform(X)
|
||||
|
||||
if issparse(X):
|
||||
X = X.toarray()
|
||||
|
||||
self.mix_densities = self.get_mixture_components(X, y, classes, self.bandwidth)
|
||||
return self
|
||||
|
||||
def predict(self, X):
|
||||
"""
|
||||
Searches for the mixture model parameter (the sought prevalence values) that maximizes the likelihood
|
||||
of the data (i.e., that minimizes the negative log-likelihood)
|
||||
|
||||
:param X: instances in the sample
|
||||
:return: a vector of class prevalence estimates
|
||||
"""
|
||||
epsilon = 1e-10
|
||||
if issparse(X):
|
||||
X = X.toarray()
|
||||
n_classes = len(self.mix_densities)
|
||||
if self.standardize:
|
||||
X = self.scaler.transform(X)
|
||||
test_densities = [self.pdf(kde_i, X) for kde_i in self.mix_densities]
|
||||
|
||||
def neg_loglikelihood(prev):
|
||||
test_mixture_likelihood = sum(prev_i * dens_i for prev_i, dens_i in zip (prev, test_densities))
|
||||
test_loglikelihood = np.log(test_mixture_likelihood + epsilon)
|
||||
return -np.sum(test_loglikelihood)
|
||||
|
||||
return F.optim_minimize(neg_loglikelihood, n_classes)
|
||||
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
qp.environ['SAMPLE_SIZE'] = 250
|
||||
qp.environ['N_JOBS'] = -1
|
||||
min_df = 10
|
||||
# dataset = 'imdb'
|
||||
repeats = 10
|
||||
error = 'mae'
|
||||
|
||||
div = 'topsoe'
|
||||
|
||||
# generates tuples (dataset, method, method_name)
|
||||
# (the dataset is needed for methods that process the dataset differently)
|
||||
def gen_methods():
|
||||
|
||||
for dataset in qp.datasets.REVIEWS_SENTIMENT_DATASETS:
|
||||
|
||||
data = qp.datasets.fetch_reviews(dataset, tfidf=False)
|
||||
|
||||
# bernoulli_vectorizer = CountVectorizer(min_df=min_df, binary=True)
|
||||
# dxs = DxS(divergence=div, vectorizer=bernoulli_vectorizer)
|
||||
# yield data, dxs, 'DxS-Bernoulli'
|
||||
#
|
||||
# multinomial_vectorizer = CountVectorizer(min_df=min_df, binary=False)
|
||||
# dxs = DxS(divergence=div, vectorizer=multinomial_vectorizer)
|
||||
# yield data, dxs, 'DxS-multinomial'
|
||||
#
|
||||
# tf_vectorizer = TfidfVectorizer(sublinear_tf=False, use_idf=False, min_df=min_df, norm=None)
|
||||
# dxs = DxS(divergence=div, vectorizer=tf_vectorizer)
|
||||
# yield data, dxs, 'DxS-TF'
|
||||
#
|
||||
# logtf_vectorizer = TfidfVectorizer(sublinear_tf=True, use_idf=False, min_df=min_df, norm=None)
|
||||
# dxs = DxS(divergence=div, vectorizer=logtf_vectorizer)
|
||||
# yield data, dxs, 'DxS-logTF'
|
||||
#
|
||||
# tfidf_vectorizer = TfidfVectorizer(use_idf=True, min_df=min_df, norm=None)
|
||||
# dxs = DxS(divergence=div, vectorizer=tfidf_vectorizer)
|
||||
# yield data, dxs, 'DxS-TFIDF'
|
||||
#
|
||||
# tfidf_vectorizer = TfidfVectorizer(use_idf=True, min_df=min_df, norm='l2')
|
||||
# dxs = DxS(divergence=div, vectorizer=tfidf_vectorizer)
|
||||
# yield data, dxs, 'DxS-TFIDF-l2'
|
||||
|
||||
tsr_vectorizer = TSRweighting(tsr_function=information_gain, min_df=min_df, norm='l2')
|
||||
dxs = DxS(divergence=div, vectorizer=tsr_vectorizer)
|
||||
yield data, dxs, 'DxS-TFTSR-l2'
|
||||
|
||||
data = qp.datasets.fetch_reviews(dataset, tfidf=True, min_df=min_df)
|
||||
|
||||
kdex = KDExML()
|
||||
reduction = TruncatedSVD(n_components=100, random_state=0)
|
||||
red_data = qp.data.preprocessing.instance_transformation(data, transformer=reduction, inplace=False)
|
||||
yield red_data, kdex, 'KDEx'
|
||||
|
||||
hdy = HDy(LogisticRegression())
|
||||
yield data, hdy, 'HDy'
|
||||
|
||||
# dm = DistributionMatchingY(LogisticRegression(), divergence=div, nbins=5)
|
||||
# yield data, dm, 'DM-5b'
|
||||
#
|
||||
# dm = DistributionMatchingY(LogisticRegression(), divergence=div, nbins=10)
|
||||
# yield data, dm, 'DM-10b'
|
||||
|
||||
|
||||
|
||||
|
||||
result_path = 'results.csv'
|
||||
with open(result_path, 'wt') as csv:
|
||||
csv.write(f'Method\tDataset\tMAE\tMRAE\n')
|
||||
for data, quantifier, quant_name in gen_methods():
|
||||
quantifier.fit(*data.training.Xy)
|
||||
report = qp.evaluation.evaluation_report(quantifier, APP(data.test, repeats=repeats), error_metrics=['mae','mrae'], verbose=True)
|
||||
means = report.mean(numeric_only=True)
|
||||
csv.write(f'{quant_name}\t{data.name}\t{means["mae"]:.5f}\t{means["mrae"]:.5f}\n')
|
||||
|
||||
df = pd.read_csv(result_path, sep='\t')
|
||||
# print(df)
|
||||
|
||||
pv = df.pivot_table(index='Method', columns="Dataset", values=["MAE", "MRAE"])
|
||||
print(pv)
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -7,13 +7,17 @@ from . import functional
|
|||
from . import method
|
||||
from . import evaluation
|
||||
from . import protocol
|
||||
from . import plot
|
||||
from . import util
|
||||
from . import model_selection
|
||||
from . import classification
|
||||
import os
|
||||
|
||||
__version__ = '0.2.0'
|
||||
try:
|
||||
from . import plot
|
||||
except ImportError:
|
||||
plot = None
|
||||
|
||||
__version__ = '0.2.1'
|
||||
|
||||
|
||||
def _default_cls():
|
||||
|
|
@ -74,4 +78,3 @@ def _get_classifier(classifier):
|
|||
raise ValueError('neither classifier nor qp.environ["DEFAULT_CLS"] have been specified')
|
||||
return classifier
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1 +1,3 @@
|
|||
from . import svmperf
|
||||
from . import calibration
|
||||
from . import methods
|
||||
from . import svmperf
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
from copy import deepcopy
|
||||
|
||||
from abstention.calibration import NoBiasVectorScaling, TempScaling, VectorScaling
|
||||
from sklearn.base import BaseEstimator, clone
|
||||
from sklearn.model_selection import cross_val_predict, train_test_split
|
||||
from sklearn.preprocessing import LabelEncoder
|
||||
from sklearn.utils.validation import check_X_y
|
||||
import numpy as np
|
||||
|
||||
|
||||
|
|
@ -11,6 +12,17 @@ import numpy as np
|
|||
# see https://github.com/kundajelab/abstention
|
||||
|
||||
|
||||
def _require_abstention_calibration():
|
||||
try:
|
||||
from abstention.calibration import NoBiasVectorScaling, TempScaling, VectorScaling
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"Calibration methods in quapy.classification.calibration require the optional "
|
||||
"'abstention' package."
|
||||
) from exc
|
||||
return NoBiasVectorScaling, TempScaling, VectorScaling
|
||||
|
||||
|
||||
class RecalibratedProbabilisticClassifier:
|
||||
"""
|
||||
Abstract class for (re)calibration method from `abstention.calibration`, as defined in
|
||||
|
|
@ -142,6 +154,7 @@ class NBVSCalibration(RecalibratedProbabilisticClassifierBase):
|
|||
"""
|
||||
|
||||
def __init__(self, classifier, val_split=5, n_jobs=None, verbose=False):
|
||||
NoBiasVectorScaling, _, _ = _require_abstention_calibration()
|
||||
self.classifier = classifier
|
||||
self.calibrator = NoBiasVectorScaling(verbose=verbose)
|
||||
self.val_split = val_split
|
||||
|
|
@ -164,6 +177,7 @@ class BCTSCalibration(RecalibratedProbabilisticClassifierBase):
|
|||
"""
|
||||
|
||||
def __init__(self, classifier, val_split=5, n_jobs=None, verbose=False):
|
||||
_, TempScaling, _ = _require_abstention_calibration()
|
||||
self.classifier = classifier
|
||||
self.calibrator = TempScaling(verbose=verbose, bias_positions='all')
|
||||
self.val_split = val_split
|
||||
|
|
@ -186,6 +200,7 @@ class TSCalibration(RecalibratedProbabilisticClassifierBase):
|
|||
"""
|
||||
|
||||
def __init__(self, classifier, val_split=5, n_jobs=None, verbose=False):
|
||||
_, TempScaling, _ = _require_abstention_calibration()
|
||||
self.classifier = classifier
|
||||
self.calibrator = TempScaling(verbose=verbose)
|
||||
self.val_split = val_split
|
||||
|
|
@ -208,9 +223,84 @@ class VSCalibration(RecalibratedProbabilisticClassifierBase):
|
|||
"""
|
||||
|
||||
def __init__(self, classifier, val_split=5, n_jobs=None, verbose=False):
|
||||
_, _, VectorScaling = _require_abstention_calibration()
|
||||
self.classifier = classifier
|
||||
self.calibrator = VectorScaling(verbose=verbose)
|
||||
self.val_split = val_split
|
||||
self.n_jobs = n_jobs
|
||||
self.verbose = verbose
|
||||
|
||||
|
||||
class TemperatureScalingFromLogits(BaseEstimator):
|
||||
"""
|
||||
Calibrates a matrix of logits by learning a temperature-scaling mapping
|
||||
with the calibration methods from `abstention.calibration`.
|
||||
|
||||
This estimator is useful when the inputs are already logits produced by a
|
||||
pretrained classifier, and the goal is to transform them directly into
|
||||
calibrated posterior probabilities without retraining the underlying model.
|
||||
|
||||
:param bias_corrected: if True, uses Bias-Corrected Temperature Scaling
|
||||
(BCTS); otherwise, uses standard Temperature Scaling (TS)
|
||||
:param verbose: whether the underlying calibrator should display progress
|
||||
information
|
||||
"""
|
||||
|
||||
def __init__(self, bias_corrected=False, verbose=False):
|
||||
self.bias_corrected = bias_corrected
|
||||
self.verbose = verbose
|
||||
|
||||
def fit(self, X, y):
|
||||
"""
|
||||
Fits the logits calibrator.
|
||||
|
||||
:param X: array-like of shape `(n_samples, n_classes)` containing
|
||||
logits
|
||||
:param y: array-like of shape `(n_samples,)` containing class labels
|
||||
:return: self
|
||||
"""
|
||||
X, y = check_X_y(X, y)
|
||||
|
||||
self.label_encoder_ = LabelEncoder()
|
||||
y_enc = self.label_encoder_.fit_transform(y)
|
||||
self.classes_ = self.label_encoder_.classes_
|
||||
|
||||
n_classes = len(self.classes_)
|
||||
logits_dim = X.shape[1]
|
||||
if n_classes != logits_dim:
|
||||
raise ValueError(
|
||||
f'mismatch between the number of classes ({n_classes}) and the '
|
||||
f'dimensionality of the logits ({logits_dim})'
|
||||
)
|
||||
|
||||
_, TempScaling, _ = _require_abstention_calibration()
|
||||
calibrator = TempScaling(
|
||||
verbose=self.verbose,
|
||||
bias_positions='all' if self.bias_corrected else [],
|
||||
)
|
||||
self.calibrator_ = calibrator
|
||||
self.calibration_function_ = calibrator(X, np.eye(n_classes)[y_enc])
|
||||
return self
|
||||
|
||||
def predict_proba(self, X):
|
||||
"""
|
||||
Converts logits into calibrated posterior probabilities.
|
||||
|
||||
:param X: array-like of shape `(n_samples, n_classes)` containing
|
||||
logits
|
||||
:return: array-like of shape `(n_samples, n_classes)` with calibrated
|
||||
posterior probabilities
|
||||
"""
|
||||
return self.calibration_function_(X)
|
||||
|
||||
def predict(self, X):
|
||||
"""
|
||||
Predicts class labels after calibration.
|
||||
|
||||
:param X: array-like of shape `(n_samples, n_classes)` containing
|
||||
logits
|
||||
:return: array-like of shape `(n_samples,)` with class label
|
||||
predictions
|
||||
"""
|
||||
posteriors = self.predict_proba(X)
|
||||
return self.label_encoder_.inverse_transform(np.argmax(posteriors, axis=1))
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import numpy as np
|
||||
from sklearn.base import BaseEstimator
|
||||
from sklearn.decomposition import TruncatedSVD
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
|
@ -95,3 +96,23 @@ class LowRankLogisticRegression(BaseEstimator):
|
|||
if self.pca is None:
|
||||
return X
|
||||
return self.pca.transform(X)
|
||||
|
||||
|
||||
class MockClassifierFromPosteriors(BaseEstimator):
|
||||
"""
|
||||
Mock classifier that bypasses classifier training when the input instances
|
||||
are already posterior probabilities produced by a pretrained probabilistic
|
||||
classifier.
|
||||
|
||||
:param X: arrays of shape `(n_samples, n_classes)` are interpreted as posterior probabilities
|
||||
"""
|
||||
|
||||
def fit(self, X, y):
|
||||
self.classes_ = np.sort(np.unique(y))
|
||||
return self
|
||||
|
||||
def predict(self, X):
|
||||
return np.argmax(X, axis=1)
|
||||
|
||||
def predict_proba(self, X):
|
||||
return X
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import logging
|
||||
import os
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from pathlib import Path
|
||||
|
|
@ -42,7 +43,7 @@ class NeuralClassifierTrainer:
|
|||
batch_size=64,
|
||||
batch_size_test=512,
|
||||
padding_length=300,
|
||||
device='cuda',
|
||||
device='cpu',
|
||||
checkpointpath='../checkpoint/classifier_net.dat'):
|
||||
|
||||
super().__init__()
|
||||
|
|
@ -63,7 +64,7 @@ class NeuralClassifierTrainer:
|
|||
self.learner_hyperparams = self.net.get_params()
|
||||
self.checkpointpath = checkpointpath
|
||||
|
||||
print(f'[NeuralNetwork running on {device}]')
|
||||
logging.getLogger(__name__).info(f'NeuralNetwork running on {device}')
|
||||
os.makedirs(Path(checkpointpath).parent, exist_ok=True)
|
||||
|
||||
def reset_net_params(self, vocab_size, n_classes):
|
||||
|
|
@ -198,14 +199,15 @@ class NeuralClassifierTrainer:
|
|||
if self.early_stop.IMPROVED:
|
||||
torch.save(self.net.state_dict(), checkpoint)
|
||||
elif self.early_stop.STOP:
|
||||
print(f'training ended by patience exhasted; loading best model parameters in {checkpoint} '
|
||||
f'for epoch {self.early_stop.best_epoch}')
|
||||
logging.getLogger(__name__).info(
|
||||
f'training ended by patience exhausted; loading best model parameters in {checkpoint} '
|
||||
f'for epoch {self.early_stop.best_epoch}')
|
||||
self.net.load_state_dict(torch.load(checkpoint))
|
||||
break
|
||||
|
||||
print('performing one training pass over the validation set...')
|
||||
logging.getLogger(__name__).info('performing one training pass over the validation set...')
|
||||
self._train_epoch(valid_generator, self.status['tr'], pbar, epoch=0)
|
||||
print('[done]')
|
||||
logging.getLogger(__name__).info('done')
|
||||
|
||||
return self
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import logging
|
||||
import random
|
||||
import shutil
|
||||
import subprocess
|
||||
|
|
@ -66,8 +67,7 @@ class SVMperf(BaseEstimator, ClassifierMixin):
|
|||
# this would allow to run parallel instances of predict
|
||||
random_code = 'svmperfprocess'+'-'.join(str(local_random.randint(0, 1000000)) for _ in range(5))
|
||||
if self.host_folder is None:
|
||||
# tmp dir are removed after the fit terminates in multiprocessing...
|
||||
self.tmpdir = tempfile.TemporaryDirectory(suffix=random_code).name
|
||||
self.tmpdir = join(tempfile.gettempdir(), random_code)
|
||||
else:
|
||||
self.tmpdir = join(self.host_folder, '.' + random_code)
|
||||
makedirs(self.tmpdir, exist_ok=True)
|
||||
|
|
@ -79,14 +79,14 @@ class SVMperf(BaseEstimator, ClassifierMixin):
|
|||
|
||||
cmd = ' '.join([self.svmperf_learn, self.c_cmd, self.loss_cmd, traindat, self.model])
|
||||
if self.verbose:
|
||||
print('[Running]', cmd)
|
||||
p = subprocess.run(cmd.split(), stdout=PIPE, stderr=STDOUT)
|
||||
logging.getLogger(__name__).info(f'[Running] {cmd}')
|
||||
p = subprocess.run(cmd.split(), stdout=PIPE, stderr=PIPE)
|
||||
if not exists(self.model):
|
||||
print(p.stderr.decode('utf-8'))
|
||||
logging.getLogger(__name__).error(p.stderr.decode('utf-8'))
|
||||
remove(traindat)
|
||||
|
||||
if self.verbose:
|
||||
print(p.stdout.decode('utf-8'))
|
||||
logging.getLogger(__name__).info(p.stdout.decode('utf-8'))
|
||||
|
||||
return self
|
||||
|
||||
|
|
@ -125,11 +125,11 @@ class SVMperf(BaseEstimator, ClassifierMixin):
|
|||
|
||||
cmd = ' '.join([self.svmperf_classify, testdat, self.model, predictions_path])
|
||||
if self.verbose:
|
||||
print('[Running]', cmd)
|
||||
logging.getLogger(__name__).info(f'[Running] {cmd}')
|
||||
p = subprocess.run(cmd.split(), stdout=PIPE, stderr=STDOUT)
|
||||
|
||||
if self.verbose:
|
||||
print(p.stdout.decode('utf-8'))
|
||||
logging.getLogger(__name__).info(p.stdout.decode('utf-8'))
|
||||
|
||||
scores = np.loadtxt(predictions_path)
|
||||
remove(testdat)
|
||||
|
|
|
|||
|
|
@ -99,6 +99,9 @@ class SamplesFromDir(AbstractProtocol):
|
|||
sample, _ = self.load_fn(os.path.join(self.path_dir, f'{id}.txt'))
|
||||
yield sample, prevalence
|
||||
|
||||
def total(self):
|
||||
return len(self.true_prevs)
|
||||
|
||||
|
||||
class LabelledCollectionsFromDir(AbstractProtocol):
|
||||
|
||||
|
|
@ -113,6 +116,10 @@ class LabelledCollectionsFromDir(AbstractProtocol):
|
|||
lc = LabelledCollection.load(path=collection_path, loader_func=self.load_fn)
|
||||
yield lc
|
||||
|
||||
def total(self):
|
||||
return len(self.true_prevs)
|
||||
|
||||
|
||||
|
||||
class ResultSubmission:
|
||||
|
||||
|
|
@ -180,8 +187,7 @@ class ResultSubmission:
|
|||
try:
|
||||
df = pd.read_csv(path, index_col=0)
|
||||
except Exception as e:
|
||||
print(f'the file {path} does not seem to be a valid csv file. ')
|
||||
print(e)
|
||||
raise ValueError(f'the file {path} does not seem to be a valid csv file: {e}')
|
||||
return ResultSubmission.check_dataframe_format(df, path=path)
|
||||
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -33,7 +33,6 @@ class LabelledCollection:
|
|||
else:
|
||||
self.instances = np.asarray(instances)
|
||||
self.labels = np.asarray(labels)
|
||||
n_docs = len(self)
|
||||
if classes is None:
|
||||
self.classes_ = F.classes_from_labels(self.labels)
|
||||
else:
|
||||
|
|
@ -41,7 +40,13 @@ class LabelledCollection:
|
|||
self.classes_.sort()
|
||||
if len(set(self.labels).difference(set(classes))) > 0:
|
||||
raise ValueError(f'labels ({set(self.labels)}) contain values not included in classes_ ({set(classes)})')
|
||||
self.index = {class_: np.arange(n_docs)[self.labels == class_] for class_ in self.classes_}
|
||||
self._index = None
|
||||
|
||||
@property
|
||||
def index(self):
|
||||
if not hasattr(self, '_index') or self._index is None:
|
||||
self._index = {class_: np.arange(len(self))[self.labels == class_] for class_ in self.classes_}
|
||||
return self._index
|
||||
|
||||
@classmethod
|
||||
def load(cls, path: str, loader_func: callable, classes=None, **loader_kwargs):
|
||||
|
|
@ -324,7 +329,9 @@ class LabelledCollection:
|
|||
else:
|
||||
raise NotImplementedError('unsupported operation for collection types')
|
||||
labels = np.concatenate([lc.labels for lc in args])
|
||||
classes = np.unique(labels).sort()
|
||||
# union of each collection's own classes_, so a class declared but absent from
|
||||
# this particular join (e.g. an empty fold) is preserved at zero prevalence
|
||||
classes = np.unique(np.concatenate([lc.classes_ for lc in args]))
|
||||
return LabelledCollection(instances, labels, classes=classes)
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
import logging
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
import zipfile
|
||||
from os.path import join
|
||||
import pandas as pd
|
||||
from ucimlrepo import fetch_ucirepo
|
||||
from quapy.data.base import Dataset, LabelledCollection
|
||||
from quapy.data.preprocessing import text2tfidf, reduce_columns
|
||||
from quapy.data.preprocessing import standardize as standardizer
|
||||
|
|
@ -12,6 +12,17 @@ from quapy.util import download_file_if_not_exists, download_file, get_quapy_hom
|
|||
from sklearn.preprocessing import StandardScaler
|
||||
|
||||
|
||||
def _fetch_ucirepo(*args, **kwargs):
|
||||
try:
|
||||
from ucimlrepo import fetch_ucirepo
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"UCI dataset fetching requires the optional 'ucimlrepo' package. "
|
||||
"Install it to use fetch_UCIBinaryDataset or fetch_UCIMulticlassDataset."
|
||||
) from exc
|
||||
return fetch_ucirepo(*args, **kwargs)
|
||||
|
||||
|
||||
REVIEWS_SENTIMENT_DATASETS = ['hp', 'kindle', 'imdb']
|
||||
|
||||
TWITTER_SENTIMENT_DATASETS_TEST = [
|
||||
|
|
@ -109,6 +120,9 @@ LEQUA2024_SAMPLE_SIZE = {
|
|||
'T4': 250,
|
||||
}
|
||||
|
||||
IMAGE_DATASETS=['cifar10', 'cifar100', 'cifar100coarse', 'svhn', 'fashionmnist', 'mnist']
|
||||
IMAGE_EMBEDDINGS=['features', 'logits', 'predictions']
|
||||
|
||||
|
||||
def fetch_reviews(dataset_name, tfidf=False, min_df=None, data_home=None, pickle=False) -> Dataset:
|
||||
"""
|
||||
|
|
@ -199,8 +213,9 @@ def fetch_twitter(dataset_name, for_model_selection=False, min_df=None, data_hom
|
|||
if dataset_name in {'semeval13', 'semeval14', 'semeval15'}:
|
||||
trainset_name = 'semeval'
|
||||
testset_name = 'semeval' if for_model_selection else dataset_name
|
||||
print(f"the training and development sets for datasets 'semeval13', 'semeval14', 'semeval15' are common "
|
||||
f"(called 'semeval'); returning trainin-set='{trainset_name}' and test-set={testset_name}")
|
||||
logging.getLogger(__name__).info(
|
||||
f"the training and development sets for datasets 'semeval13', 'semeval14', 'semeval15' are common "
|
||||
f"(called 'semeval'); returning trainin-set='{trainset_name}' and test-set={testset_name}")
|
||||
else:
|
||||
if dataset_name == 'semeval' and for_model_selection==False:
|
||||
raise ValueError('dataset "semeval" can only be used for model selection. '
|
||||
|
|
@ -500,11 +515,11 @@ def fetch_UCIBinaryLabelledCollection(dataset_name, data_home=None, standardize=
|
|||
y = df["NSP"].astype(int).values
|
||||
elif group == "semeion":
|
||||
with download_tmp_file("semeion", "semeion.data") as tmp:
|
||||
df = pd.read_csv(tmp, header=None, sep='\s+')
|
||||
df = pd.read_csv(tmp, header=None, sep='\\s+')
|
||||
X = df.iloc[:, 0:256].astype(float).values
|
||||
y = df[263].values # 263 stands for digit 8 (labels are one-hot vectors from col 256-266)
|
||||
else:
|
||||
df = fetch_ucirepo(id=id)
|
||||
df = _fetch_ucirepo(id=id)
|
||||
X, y = df.data.features.to_numpy(), df.data.targets.to_numpy().squeeze()
|
||||
|
||||
# transform data when needed before returning (returned data will be pickled)
|
||||
|
|
@ -616,8 +631,8 @@ def fetch_UCIMulticlassDataset(
|
|||
are taken for training, and the rest (irrespective of `min_test_split`) is taken for test.
|
||||
:param max_train_instances: maximum number of instances to keep for training (defaults to 25000);
|
||||
set to -1 or None to avoid this check
|
||||
:param min_class_support: minimum number of istances per class. Classes with fewer instances
|
||||
are discarded (deafult is 100)
|
||||
:param min_class_support: integer or float, the minimum number or proportion of istances per class.
|
||||
Classes with fewer instances are discarded (deafult is 100).
|
||||
:param standardize: indicates whether the covariates should be standardized or not (default is True). If requested,
|
||||
standardization applies after the LabelledCollection is split, that is, the mean an std are computed only on the
|
||||
training portion of the data.
|
||||
|
|
@ -673,6 +688,11 @@ def fetch_UCIMulticlassLabelledCollection(dataset_name, data_home=None, min_clas
|
|||
f'Name {dataset_name} does not match any known dataset from the ' \
|
||||
f'UCI Machine Learning datasets repository (multiclass). ' \
|
||||
f'Valid ones are {UCI_MULTICLASS_DATASETS}'
|
||||
|
||||
assert (min_class_support is None or
|
||||
((isinstance(min_class_support, int) and min_class_support >= 0) or
|
||||
(isinstance(min_class_support, float) and 0. <= min_class_support < 1.))), \
|
||||
f'invalid value for {min_class_support=}; expected non negative integer or float in [0,1)'
|
||||
|
||||
if data_home is None:
|
||||
data_home = get_quapy_home()
|
||||
|
|
@ -739,26 +759,43 @@ def fetch_UCIMulticlassLabelledCollection(dataset_name, data_home=None, min_clas
|
|||
|
||||
file = join(data_home, 'uci_multiclass', dataset_name+'.pkl')
|
||||
|
||||
def dummify_categorical_features(df_features, dataset_id):
|
||||
categorical_features = {
|
||||
158: ["S1", "C1", "S2", "C2", "S3", "C3", "S4", "C4", "S5", "C5"], # poker_hand
|
||||
}
|
||||
|
||||
categorical = categorical_features.get(dataset_id, [])
|
||||
|
||||
X = df_features.copy()
|
||||
if categorical:
|
||||
X[categorical] = X[categorical].astype("category")
|
||||
X = pd.get_dummies(X, columns=categorical, drop_first=True)
|
||||
|
||||
return X
|
||||
|
||||
def download(id, name):
|
||||
df = fetch_ucirepo(id=id)
|
||||
df = _fetch_ucirepo(id=id)
|
||||
|
||||
df.data.features = pd.get_dummies(df.data.features, drop_first=True)
|
||||
X, y = df.data.features.to_numpy(dtype=np.float64), df.data.targets.to_numpy().squeeze()
|
||||
X_df = dummify_categorical_features(df.data.features, id)
|
||||
X = X_df.to_numpy(dtype=np.float64)
|
||||
y = df.data.targets.to_numpy().squeeze()
|
||||
|
||||
assert y.ndim == 1, 'more than one y'
|
||||
assert y.ndim == 1, f'error: the dataset {id=} {name=} has more than one target variable'
|
||||
|
||||
classes = np.sort(np.unique(y))
|
||||
y = np.searchsorted(classes, y)
|
||||
return LabelledCollection(X, y)
|
||||
|
||||
def filter_classes(data: LabelledCollection, min_ipc):
|
||||
if min_ipc is None:
|
||||
min_ipc = 0
|
||||
def filter_classes(data: LabelledCollection, min_class_support):
|
||||
if min_class_support is None or min_class_support == 0.:
|
||||
return data
|
||||
if isinstance(min_class_support, float):
|
||||
min_class_support = int(len(data) * min_class_support)
|
||||
classes = data.classes_
|
||||
# restrict classes to only those with at least min_ipc instances
|
||||
classes = classes[data.counts() >= min_ipc]
|
||||
# restrict classes to only those with at least min_class_support instances
|
||||
classes = classes[data.counts() >= min_class_support]
|
||||
# filter X and y keeping only datapoints belonging to valid classes
|
||||
filter_idx = np.in1d(data.y, classes)
|
||||
filter_idx = np.isin(data.y, classes)
|
||||
X, y = data.X[filter_idx], data.y[filter_idx]
|
||||
# map classes to range(len(classes))
|
||||
y = np.searchsorted(classes, y)
|
||||
|
|
@ -1030,3 +1067,90 @@ def fetch_IFCB(single_sample_train=True, for_model_selection=False, data_home=No
|
|||
return train, test_gen
|
||||
else:
|
||||
return train_gen, test_gen
|
||||
|
||||
|
||||
def _fetch_image_embedding_splits(dataset_name, embedding, data_home=None) -> tuple[LabelledCollection,LabelledCollection,LabelledCollection]:
|
||||
"""
|
||||
Loads a pre-generated embedding set (train, val, or test) of an image dataset from `Zenodo <https://zenodo.org/records/21131944>`_.
|
||||
|
||||
Embeddings were extracted using `this script <https://github.com/pglez82/visiondatasets_quapy>`_.
|
||||
|
||||
:param dataset_name: the name of the dataset: valid ones are 'cifar10', 'cifar100', 'cifar100coarse', 'svhn', 'fashionmnist', 'mnist'
|
||||
:param embedding: the type of embedding: valid ones are 'features' (next-to-last representations), 'logits' (pre-activation values), 'predictions' (posterior probabilities)
|
||||
:param data_home: specify the quapy home directory where collections will be dumped (leave empty to use the default
|
||||
~/quay_data/ directory)
|
||||
:return: a tuple (train, val, test) where each entry is an instance of :class:`quapy.data.base.LabelledCollection`
|
||||
"""
|
||||
assert dataset_name in IMAGE_DATASETS, \
|
||||
f'Name {dataset_name} does not match any known dataset. Valid ones are {IMAGE_DATASETS}'
|
||||
assert embedding in IMAGE_EMBEDDINGS, \
|
||||
f'Name {embedding} does not match any known type of embedding. Valid ones are {IMAGE_EMBEDDINGS}'
|
||||
if data_home is None:
|
||||
data_home = get_quapy_home()
|
||||
|
||||
dataset_network = {
|
||||
'cifar10': 'resnet18',
|
||||
'cifar100': 'resnet18',
|
||||
'cifar100coarse': 'resnet18',
|
||||
'svhn': 'resnet18',
|
||||
'fashionmnist': 'basiccnn',
|
||||
'mnist': 'basiccnn',
|
||||
}
|
||||
|
||||
trained_network = dataset_network[dataset_name]
|
||||
|
||||
def download_embedding_npz(dataset_name, trained_network, embedding):
|
||||
target_file = f'{dataset_name}_{trained_network}_{embedding}.npz'
|
||||
URL = f'https://zenodo.org/records/21131944/files/{target_file}'
|
||||
os.makedirs(join(data_home, 'image'), exist_ok=True)
|
||||
file_path = join(data_home, 'image', target_file)
|
||||
download_file_if_not_exists(URL, file_path)
|
||||
npz_file = np.load(file_path)
|
||||
return npz_file
|
||||
|
||||
embedding_dict = download_embedding_npz(dataset_name, trained_network, embedding=embedding)
|
||||
labels_dict = download_embedding_npz(dataset_name, trained_network, embedding='targets')
|
||||
|
||||
train = LabelledCollection(embedding_dict['train'], labels_dict['train'])
|
||||
val = LabelledCollection(embedding_dict['val'], labels_dict['val'], classes=train.classes)
|
||||
test = LabelledCollection(embedding_dict['test'], labels_dict['test'], classes=train.classes)
|
||||
|
||||
return train, val, test
|
||||
|
||||
|
||||
def fetch_image_embeddings(dataset_name, embedding, heldout_only=True, data_home=None) -> Dataset:
|
||||
"""
|
||||
Loads an image dataset with pre-generated embeddings. Available datasets include:
|
||||
|
||||
- 'cifar10', 'cifar100', 'cifar100coarse': see `Alex Krizhevsky and Geoffrey Hinton. Learning multiple layers of features from tiny images. Technical report, University of Toronto, Toronto, Ontario, 2009. <https://cave.cs.toronto.edu/kriz/learning-features-2009-TR.pdf>`_
|
||||
- 'mnist': `Yann LeCun, Corinna Cortes, and Christopher J. C. Burges. The MNIST database of handwritten digits. 1998. <http://yann.lecun.com/exdb/mnist/>`_
|
||||
- 'fashionmnist': `Han Xiao, Kashif Rasul, and Roland Vollgraf. Fashion-MNIST: a novel image dataset for benchmarking machine learning algorithms. arXiv preprint arXiv:1708.07747, 2017. <https://arxiv.org/abs/1708.07747>`_
|
||||
- 'svhn': `Yuval Netzer, Tao Wang, Adam Coates, Alessandro Bissacco, Baolin Wu, Andrew Y Ng, et al. Reading digits in natural images with unsupervised feature learning. In NIPS workshop on deep learning and unsupervised feature learning, volume 2011, page 4. Granada, 2011. <https://static.googleusercontent.com/media/research.google.com/es//pubs/archive/37648.pdf>`_
|
||||
|
||||
The image dataset are stored in `Zenodo <https://zenodo.org/records/21131944>`_ and were extracted using `this script <https://github.com/pglez82/visiondatasets_quapy>`_.
|
||||
|
||||
These embeddings were generated using a resnet18 or a simple cnn. In all cases, the network was trained using ~60% of the data, validated on ~25% of the data, and the remaining ~15% was used for test. Splits were created with stratification.
|
||||
Once the network is trained, it was used with frozen weights to generate embeddings for the training, validation, and test, in different formats (see below).
|
||||
It would therefore be convenient to use only heldout data (validation and test) for training and testing quantifiers (this is the default behavior), although the training+validation data can be accessed with `heldout_only=False`.
|
||||
|
||||
:param dataset_name: the name of the dataset: valid ones are 'cifar10', 'cifar100', 'cifar100coarse', 'svhn', 'fashionmnist', 'mnist'
|
||||
:param embedding: the type of embedding: valid ones are 'features' (next-to-last representations), 'logits' (pre-activation outputs), 'predictions' (post-softmax outputs, or predicted posterior probabilities)
|
||||
:param heldout_only: whether to discard the part of the training data used to train the neural model that generated the embeddings (default: True); set to False
|
||||
to obtain, as the training data, the original training+validation splits.
|
||||
:param data_home: specify the quapy home directory where collections will be dumped (leave empty to use the default
|
||||
~/quay_data/ directory)
|
||||
:return: an instance of :class:`quapy.data.base.Dataset`
|
||||
"""
|
||||
if data_home is None:
|
||||
data_home = get_quapy_home()
|
||||
|
||||
network_train, val, test = _fetch_image_embedding_splits(dataset_name, embedding, data_home)
|
||||
|
||||
if heldout_only:
|
||||
train = val
|
||||
else:
|
||||
train = network_train + val
|
||||
|
||||
return Dataset(train, test, name=dataset_name)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,37 @@ from quapy.util import map_parallel
|
|||
from .base import LabelledCollection
|
||||
|
||||
|
||||
def instance_transformation(dataset:Dataset, transformer, inplace=False):
|
||||
"""
|
||||
Transforms a :class:`quapy.data.base.Dataset` applying the `fit_transform` and `transform` functions
|
||||
of a (sklearn's) transformer.
|
||||
|
||||
:param dataset: a :class:`quapy.data.base.Dataset` where the instances of training and test collections are
|
||||
lists of str
|
||||
:param transformer: TransformerMixin implementing `fit_transform` and `transform` functions
|
||||
:param inplace: whether or not to apply the transformation inplace (True), or to a new copy (False, default)
|
||||
:return: a new :class:`quapy.data.base.Dataset` with transformed instances (if inplace=False) or a reference to the
|
||||
current Dataset (if inplace=True) where the instances have been transformed
|
||||
"""
|
||||
training_transformed = transformer.fit_transform(*dataset.training.Xy)
|
||||
test_transformed = transformer.transform(dataset.test.X)
|
||||
orig_name = dataset.name
|
||||
|
||||
if inplace:
|
||||
dataset.training = LabelledCollection(training_transformed, dataset.training.labels, dataset.classes_)
|
||||
dataset.test = LabelledCollection(test_transformed, dataset.test.labels, dataset.classes_)
|
||||
if hasattr(transformer, 'vocabulary_'):
|
||||
dataset.vocabulary = transformer.vocabulary_
|
||||
return dataset
|
||||
else:
|
||||
training = LabelledCollection(training_transformed, dataset.training.labels.copy(), dataset.classes_)
|
||||
test = LabelledCollection(test_transformed, dataset.test.labels.copy(), dataset.classes_)
|
||||
vocab = None
|
||||
if hasattr(transformer, 'vocabulary_'):
|
||||
vocab = transformer.vocabulary_
|
||||
return Dataset(training, test, vocabulary=vocab, name=orig_name)
|
||||
|
||||
|
||||
def text2tfidf(dataset:Dataset, min_df=3, sublinear_tf=True, inplace=False, **kwargs):
|
||||
"""
|
||||
Transforms a :class:`quapy.data.base.Dataset` of textual instances into a :class:`quapy.data.base.Dataset` of
|
||||
|
|
@ -29,18 +60,7 @@ def text2tfidf(dataset:Dataset, min_df=3, sublinear_tf=True, inplace=False, **kw
|
|||
__check_type(dataset.test.instances, np.ndarray, str)
|
||||
|
||||
vectorizer = TfidfVectorizer(min_df=min_df, sublinear_tf=sublinear_tf, **kwargs)
|
||||
training_documents = vectorizer.fit_transform(dataset.training.instances)
|
||||
test_documents = vectorizer.transform(dataset.test.instances)
|
||||
|
||||
if inplace:
|
||||
dataset.training = LabelledCollection(training_documents, dataset.training.labels, dataset.classes_)
|
||||
dataset.test = LabelledCollection(test_documents, dataset.test.labels, dataset.classes_)
|
||||
dataset.vocabulary = vectorizer.vocabulary_
|
||||
return dataset
|
||||
else:
|
||||
training = LabelledCollection(training_documents, dataset.training.labels.copy(), dataset.classes_)
|
||||
test = LabelledCollection(test_documents, dataset.test.labels.copy(), dataset.classes_)
|
||||
return Dataset(training, test, vectorizer.vocabulary_)
|
||||
return instance_transformation(dataset, vectorizer, inplace)
|
||||
|
||||
|
||||
def reduce_columns(dataset: Dataset, min_df=5, inplace=False):
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
import logging
|
||||
|
||||
import numpy as np
|
||||
from scipy.sparse import dok_matrix
|
||||
from tqdm import tqdm
|
||||
|
|
@ -30,7 +32,7 @@ def from_text(path, encoding='utf-8', verbose=1, class2int=True):
|
|||
all_sentences.append(sentence)
|
||||
all_labels.append(label)
|
||||
except ValueError:
|
||||
print(f'format error in {line}')
|
||||
logging.getLogger(__name__).warning(f'format error in {line}')
|
||||
return all_sentences, all_labels
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
import numpy as np
|
||||
from sklearn.metrics import f1_score
|
||||
import quapy as qp
|
||||
from quapy.functional import AitchisonDistance
|
||||
|
||||
|
||||
def from_name(err_name):
|
||||
|
|
@ -128,6 +129,82 @@ def se(prevs_true, prevs_hat):
|
|||
return ((prevs_hat - prevs_true) ** 2).mean(axis=-1)
|
||||
|
||||
|
||||
def sre(prevs_true, prevs_hat, prevs_train, eps=0.):
|
||||
"""
|
||||
Computes the squared ratio error between two prevalence vectors.
|
||||
The squared ratio error between prevalence vectors :math:`p` and
|
||||
:math:`\\hat{p}` with training prevalence :math:`p^{tr}` is:
|
||||
:math:`SRE(p,\\hat{p},p^{tr})=\\frac{1}{|\\mathcal{Y}|}\\sum_{i \\in \\mathcal{Y}}(w_i-\\hat{w}_i)^2`,
|
||||
where :math:`w_i=\\frac{p_i}{p^{tr}_i}`.
|
||||
|
||||
:param prevs_true: array-like with the true prevalence values
|
||||
:param prevs_hat: array-like with the predicted prevalence values
|
||||
:param prevs_train: array-like with the training prevalence values, or a single
|
||||
prevalence vector when all comparisons refer to the same training set
|
||||
:param eps: smoothing factor for the prevalence values (default 0, i.e., no smoothing)
|
||||
:return: squared ratio error
|
||||
"""
|
||||
prevs_true = np.asarray(prevs_true)
|
||||
prevs_hat = np.asarray(prevs_hat)
|
||||
prevs_train = np.asarray(prevs_train)
|
||||
assert prevs_true.shape == prevs_hat.shape, f'wrong shape {prevs_true.shape=} vs {prevs_hat.shape=}'
|
||||
assert prevs_true.shape[-1] == prevs_train.shape[-1], 'wrong shape for training prevalence'
|
||||
if prevs_true.ndim == 2 and prevs_train.ndim == 1:
|
||||
prevs_train = np.tile(prevs_train, reps=(prevs_true.shape[0], 1))
|
||||
if eps > 0:
|
||||
prevs_true = smooth(prevs_true, eps)
|
||||
prevs_hat = smooth(prevs_hat, eps)
|
||||
prevs_train = smooth(prevs_train, eps)
|
||||
|
||||
n_classes = prevs_true.shape[-1]
|
||||
w = prevs_true / prevs_train
|
||||
w_hat = prevs_hat / prevs_train
|
||||
return (1. / n_classes) * np.sum((w - w_hat) ** 2., axis=-1)
|
||||
|
||||
|
||||
def msre(prevs_true, prevs_hat, prevs_train, eps=0.):
|
||||
"""
|
||||
Computes the mean squared ratio error (see :meth:`quapy.error.sre`) across the sample pairs.
|
||||
|
||||
:param prevs_true: array-like of shape `(n_samples, n_classes,)` with the true prevalence values
|
||||
:param prevs_hat: array-like of shape equal to prevs_true with the predicted prevalence values
|
||||
:param prevs_train: array-like with the training prevalence values
|
||||
:param eps: smoothing factor (default 0, i.e., no smoothing)
|
||||
:return: mean squared ratio error
|
||||
"""
|
||||
return np.mean(sre(prevs_true, prevs_hat, prevs_train, eps))
|
||||
|
||||
|
||||
def aqe(prevs_true, prevs_hat):
|
||||
"""
|
||||
Computes the Aitchison distance between two prevalence vectors.
|
||||
The Aitchison distance between prevalence vectors :math:`p` and
|
||||
:math:`\\hat{p}` is computed as
|
||||
:math:`d_A(p,\\hat{p})=\\|\\mathrm{clr}(p)-\\mathrm{clr}(\\hat{p})\\|_2`,
|
||||
where :math:`\\mathrm{clr}(p)_i=\\log p_i-\\frac{1}{|\\mathcal{Y}|}
|
||||
\\sum_{j \\in \\mathcal{Y}} \\log p_j`.
|
||||
|
||||
:param prevs_true: array-like with the true prevalence values
|
||||
:param prevs_hat: array-like with the predicted prevalence values
|
||||
:return: Aitchison distance
|
||||
"""
|
||||
return AitchisonDistance(prevs_true, prevs_hat)
|
||||
|
||||
|
||||
def maqe(prevs_true, prevs_hat):
|
||||
"""
|
||||
Computes the mean Aitchison distance (see :meth:`quapy.error.aitchisondist`)
|
||||
across the sample pairs, i.e.,
|
||||
:math:`\\mathrm{mAitchisonDist}=\\frac{1}{n}\\sum_{i=1}^n
|
||||
d_A(p_i,\\hat{p}_i)`.
|
||||
|
||||
:param prevs_true: array-like with the true prevalence values
|
||||
:param prevs_hat: array-like with the predicted prevalence values
|
||||
:return: mean Aitchison distance
|
||||
"""
|
||||
return np.mean(aqe(prevs_true, prevs_hat))
|
||||
|
||||
|
||||
def mkld(prevs_true, prevs_hat, eps=None):
|
||||
"""Computes the mean Kullback-Leibler divergence (see :meth:`quapy.error.kld`) across the
|
||||
sample pairs. The distributions are smoothed using the `eps` factor
|
||||
|
|
@ -374,8 +451,8 @@ def __check_eps(eps=None):
|
|||
|
||||
|
||||
CLASSIFICATION_ERROR = {f1e, acce}
|
||||
QUANTIFICATION_ERROR = {mae, mnae, mrae, mnrae, mse, mkld, mnkld}
|
||||
QUANTIFICATION_ERROR_SINGLE = {ae, nae, rae, nrae, se, kld, nkld}
|
||||
QUANTIFICATION_ERROR = {mae, mnae, mrae, mnrae, mse, mkld, mnkld, msre, maqe}
|
||||
QUANTIFICATION_ERROR_SINGLE = {ae, nae, rae, nrae, se, kld, nkld, sre, aqe}
|
||||
QUANTIFICATION_ERROR_SMOOTH = {kld, nkld, rae, nrae, mkld, mnkld, mrae}
|
||||
CLASSIFICATION_ERROR_NAMES = {func.__name__ for func in CLASSIFICATION_ERROR}
|
||||
QUANTIFICATION_ERROR_NAMES = {func.__name__ for func in QUANTIFICATION_ERROR}
|
||||
|
|
@ -387,6 +464,9 @@ ERROR_NAMES = \
|
|||
f1_error = f1e
|
||||
acc_error = acce
|
||||
mean_absolute_error = mae
|
||||
squared_ratio_error = sre
|
||||
dist_aitchison = aqe
|
||||
mean_dist_aitchison = maqe
|
||||
absolute_error = ae
|
||||
mean_relative_absolute_error = mrae
|
||||
relative_absolute_error = rae
|
||||
|
|
|
|||
|
|
@ -1,11 +1,15 @@
|
|||
import warnings
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import defaultdict
|
||||
from functools import lru_cache
|
||||
from typing import Literal, Union, Callable
|
||||
from numpy.typing import ArrayLike
|
||||
|
||||
import scipy
|
||||
import numpy as np
|
||||
|
||||
import quapy as qp
|
||||
|
||||
|
||||
# ------------------------------------------------------------------------------------------
|
||||
# General utils
|
||||
|
|
@ -277,7 +281,7 @@ def l1_norm(prevalences: ArrayLike) -> np.ndarray:
|
|||
"""
|
||||
n_classes = prevalences.shape[-1]
|
||||
accum = prevalences.sum(axis=-1, keepdims=True)
|
||||
prevalences = np.true_divide(prevalences, accum, where=accum > 0)
|
||||
prevalences = np.true_divide(prevalences, accum, where=accum > 0, out=None)
|
||||
allzeros = accum.flatten() == 0
|
||||
if any(allzeros):
|
||||
if prevalences.ndim == 1:
|
||||
|
|
@ -389,6 +393,23 @@ def TopsoeDistance(P: np.ndarray, Q: np.ndarray, epsilon: float=1e-20):
|
|||
return np.sum(P*np.log((2*P+epsilon)/(P+Q+epsilon)) + Q*np.log((2*Q+epsilon)/(P+Q+epsilon)))
|
||||
|
||||
|
||||
def AitchisonDistance(prevs_true, prevs_hat):
|
||||
"""
|
||||
Computes the Aitchison distance between two prevalence vectors.
|
||||
The Aitchison distance between prevalence vectors :math:`p` and
|
||||
:math:`\\hat{p}` is computed as
|
||||
:math:`d_A(p,\\hat{p})=\\|\\mathrm{clr}(p)-\\mathrm{clr}(\\hat{p})\\|_2`,
|
||||
where :math:`\\mathrm{clr}(p)_i=\\log p_i-\\frac{1}{|\\mathcal{Y}|}
|
||||
\\sum_{j \\in \\mathcal{Y}} \\log p_j`.
|
||||
|
||||
:param prevs_true: array-like with the true prevalence values
|
||||
:param prevs_hat: array-like with the predicted prevalence values
|
||||
:return: Aitchison distance
|
||||
"""
|
||||
clr = CLRtransformation()
|
||||
return np.linalg.norm(clr(prevs_true) - clr(prevs_hat), axis=-1)
|
||||
|
||||
|
||||
def get_divergence(divergence: Union[str, Callable]):
|
||||
"""
|
||||
Guarantees that the divergence received as argument is a function. That is, if this argument is already
|
||||
|
|
@ -403,6 +424,8 @@ def get_divergence(divergence: Union[str, Callable]):
|
|||
return HellingerDistance
|
||||
elif divergence=='topsoe':
|
||||
return TopsoeDistance
|
||||
elif divergence=='aitchison':
|
||||
return AitchisonDistance
|
||||
else:
|
||||
raise ValueError(f'unknown divergence {divergence}')
|
||||
elif callable(divergence):
|
||||
|
|
@ -426,7 +449,7 @@ def argmin_prevalence(loss: Callable,
|
|||
:param method: string indicating the search strategy. Possible values are::
|
||||
'optim_minimize': uses scipy.optim
|
||||
'linear_search': carries out a linear search for binary problems in the space [0, 0.01, 0.02, ..., 1]
|
||||
'ternary_search': implements the ternary search (not yet implemented)
|
||||
'ternary_search': carries out a ternary search for binary problems in the interval [0,1]
|
||||
:return: np.ndarray, a prevalence vector
|
||||
"""
|
||||
if method == 'optim_minimize':
|
||||
|
|
@ -434,7 +457,7 @@ def argmin_prevalence(loss: Callable,
|
|||
elif method == 'linear_search':
|
||||
return linear_search(loss, n_classes)
|
||||
elif method == 'ternary_search':
|
||||
ternary_search(loss, n_classes)
|
||||
return ternary_search(loss, n_classes)
|
||||
else:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
|
@ -489,7 +512,32 @@ def linear_search(loss: Callable, n_classes: int):
|
|||
|
||||
|
||||
def ternary_search(loss: Callable, n_classes: int):
|
||||
raise NotImplementedError()
|
||||
"""
|
||||
Performs a ternary search for the best prevalence value in binary problems.
|
||||
This search assumes the loss is unimodal over the interval [0,1].
|
||||
|
||||
:param loss: (callable) the function to minimize
|
||||
:param n_classes: (int) the number of classes, i.e., the dimensionality of the prevalence vector
|
||||
:return: (ndarray) the best prevalence vector found
|
||||
"""
|
||||
assert n_classes == 2, 'ternary search is only available for binary problems'
|
||||
|
||||
left, right = 0., 1.
|
||||
tol = 1e-5
|
||||
while abs(right - left) >= tol:
|
||||
left_third = left + (right - left) / 3
|
||||
right_third = right - (right - left) / 3
|
||||
|
||||
left_loss = loss(np.asarray([1 - left_third, left_third]))
|
||||
right_loss = loss(np.asarray([1 - right_third, right_third]))
|
||||
|
||||
if left_loss < right_loss:
|
||||
right = right_third
|
||||
else:
|
||||
left = left_third
|
||||
|
||||
prev = (left + right) / 2
|
||||
return np.asarray([1 - prev, prev])
|
||||
|
||||
|
||||
# ------------------------------------------------------------------------------------------
|
||||
|
|
@ -621,7 +669,11 @@ def solve_adjustment(
|
|||
if method == "inversion":
|
||||
pass # We leave A and B unchanged
|
||||
elif method == "invariant-ratio":
|
||||
# Change the last equation to replace it with the normalization condition
|
||||
# Change the last equation to replace it with the normalization condition;
|
||||
# copy first so this does not mutate the caller's arrays (np.asarray above
|
||||
# returns the same object, not a copy, when the input is already float64)
|
||||
A = A.copy()
|
||||
B = B.copy()
|
||||
A[-1, :] = 1.0
|
||||
B[-1] = 1.0
|
||||
else:
|
||||
|
|
@ -649,3 +701,105 @@ def solve_adjustment(
|
|||
raise ValueError(f'unknown {solver=}')
|
||||
|
||||
|
||||
# ------------------------------------------------------------------------------------------
|
||||
# Transformations from Compositional analysis
|
||||
# ------------------------------------------------------------------------------------------
|
||||
|
||||
class CompositionalTransformation(ABC):
|
||||
"""
|
||||
Abstract class of transformations for compositional data.
|
||||
"""
|
||||
|
||||
EPSILON = 1e-12
|
||||
|
||||
@abstractmethod
|
||||
def __call__(self, X):
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def inverse(self, Z):
|
||||
...
|
||||
|
||||
|
||||
class CLRtransformation(CompositionalTransformation):
|
||||
"""
|
||||
Centered log-ratio (CLR) transformation.
|
||||
"""
|
||||
|
||||
def __call__(self, X):
|
||||
X = np.asarray(X)
|
||||
X = qp.error.smooth(X, self.EPSILON)
|
||||
geometric_mean = np.exp(np.mean(np.log(X), axis=-1, keepdims=True))
|
||||
return np.log(X / geometric_mean)
|
||||
|
||||
def inverse(self, Z):
|
||||
return scipy.special.softmax(Z, axis=-1)
|
||||
|
||||
|
||||
class ILRtransformation(CompositionalTransformation):
|
||||
"""
|
||||
Isometric log-ratio (ILR) transformation.
|
||||
"""
|
||||
|
||||
def __call__(self, X):
|
||||
X = np.asarray(X)
|
||||
X = qp.error.smooth(X, self.EPSILON)
|
||||
basis = self.get_V(X.shape[-1])
|
||||
return np.log(X) @ basis.T
|
||||
|
||||
def inverse(self, Z):
|
||||
Z = np.asarray(Z)
|
||||
basis = self.get_V(Z.shape[-1] + 1)
|
||||
logp = Z @ basis
|
||||
p = np.exp(logp)
|
||||
return p / np.sum(p, axis=-1, keepdims=True)
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_V(self, k):
|
||||
helmert = np.zeros((k, k))
|
||||
for i in range(1, k):
|
||||
helmert[i, :i] = 1
|
||||
helmert[i, i] = -i
|
||||
helmert[i] = helmert[i] / np.sqrt(i * (i + 1))
|
||||
return helmert[1:, :]
|
||||
|
||||
|
||||
def normalized_entropy(p):
|
||||
"""
|
||||
Computes the normalized Shannon entropy of a prevalence vector.
|
||||
|
||||
:param p: array-like prevalence vector summing to 1
|
||||
:return: float in [0,1]
|
||||
"""
|
||||
p = np.asarray(p)
|
||||
entropy = scipy.stats.entropy(p)
|
||||
max_entropy = np.log(len(p))
|
||||
return np.clip(entropy / max_entropy, 0, 1)
|
||||
|
||||
|
||||
def antagonistic_prevalence(p, strength=1):
|
||||
"""
|
||||
Reflects a prevalence vector in ILR space and maps it back to the simplex.
|
||||
|
||||
:param p: array-like prevalence vector
|
||||
:param strength: reflection strength in ILR space
|
||||
:return: prevalence vector in the simplex
|
||||
"""
|
||||
ilr = ILRtransformation()
|
||||
z = ilr(p)
|
||||
z_ant = -strength * z
|
||||
return ilr.inverse(z_ant)
|
||||
|
||||
|
||||
def in_simplex(x, atol=1e-8):
|
||||
"""
|
||||
Checks whether points lie in the probability simplex.
|
||||
|
||||
:param x: array-like of shape `(n_classes,)` or `(n_points, n_classes)`
|
||||
:param atol: numerical tolerance for the unit-sum check
|
||||
:return: boolean or boolean array
|
||||
"""
|
||||
x = np.asarray(x)
|
||||
non_negative = np.all(x >= 0, axis=-1)
|
||||
sum_to_one = np.isclose(x.sum(axis=-1), 1.0, atol=atol)
|
||||
return non_negative & sum_to_one
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from sklearn.exceptions import ConvergenceWarning
|
|||
warnings.simplefilter("ignore", ConvergenceWarning)
|
||||
|
||||
from . import confidence
|
||||
from . import _bayesian
|
||||
from . import base
|
||||
from . import aggregative
|
||||
from . import non_aggregative
|
||||
|
|
@ -14,6 +15,7 @@ AGGREGATIVE_METHODS = {
|
|||
aggregative.ACC,
|
||||
aggregative.PCC,
|
||||
aggregative.PACC,
|
||||
aggregative.RLLS,
|
||||
aggregative.EMQ,
|
||||
aggregative.HDy,
|
||||
aggregative.DyS,
|
||||
|
|
@ -24,11 +26,15 @@ AGGREGATIVE_METHODS = {
|
|||
aggregative.MS,
|
||||
aggregative.MS2,
|
||||
aggregative.DMy,
|
||||
aggregative.EDy,
|
||||
aggregative.KDEyML,
|
||||
aggregative.KDEyCS,
|
||||
aggregative.KDEyHD,
|
||||
# aggregative.OneVsAllAggregative,
|
||||
confidence.BayesianCC,
|
||||
_bayesian.BayesianKDEy,
|
||||
_bayesian.BayesianMAPLS,
|
||||
confidence.PQ,
|
||||
}
|
||||
|
||||
BINARY_METHODS = {
|
||||
|
|
@ -40,6 +46,7 @@ BINARY_METHODS = {
|
|||
aggregative.MAX,
|
||||
aggregative.MS,
|
||||
aggregative.MS2,
|
||||
confidence.PQ,
|
||||
}
|
||||
|
||||
MULTICLASS_METHODS = {
|
||||
|
|
@ -47,16 +54,22 @@ MULTICLASS_METHODS = {
|
|||
aggregative.ACC,
|
||||
aggregative.PCC,
|
||||
aggregative.PACC,
|
||||
aggregative.RLLS,
|
||||
aggregative.EMQ,
|
||||
aggregative.EDy,
|
||||
aggregative.KDEyML,
|
||||
aggregative.KDEyCS,
|
||||
aggregative.KDEyHD,
|
||||
confidence.BayesianCC
|
||||
confidence.BayesianCC,
|
||||
_bayesian.BayesianKDEy,
|
||||
_bayesian.BayesianMAPLS,
|
||||
non_aggregative.EDx,
|
||||
}
|
||||
|
||||
NON_AGGREGATIVE_METHODS = {
|
||||
non_aggregative.MaximumLikelihoodPrevalenceEstimation,
|
||||
non_aggregative.DMx
|
||||
non_aggregative.DMx,
|
||||
non_aggregative.EDx
|
||||
}
|
||||
|
||||
META_METHODS = {
|
||||
|
|
@ -68,5 +81,3 @@ QUANTIFICATION_METHODS = AGGREGATIVE_METHODS | NON_AGGREGATIVE_METHODS | META_ME
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,20 +1,61 @@
|
|||
"""
|
||||
Utility functions for `Bayesian quantification <https://arxiv.org/abs/2302.09159>`_ methods.
|
||||
Utilities and methods for Bayesian quantification.
|
||||
"""
|
||||
import contextlib
|
||||
import copy
|
||||
import importlib.resources
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import warnings
|
||||
from collections.abc import Iterable
|
||||
from numbers import Number, Real
|
||||
|
||||
import numpy as np
|
||||
from joblib import Parallel, delayed
|
||||
from sklearn.base import BaseEstimator
|
||||
from tqdm import tqdm
|
||||
|
||||
import quapy as qp
|
||||
import quapy.functional as F
|
||||
from quapy.data import LabelledCollection
|
||||
from quapy.method._kdey import KDEBase
|
||||
from quapy.method.aggregative import AggregativeSoftQuantifier
|
||||
from quapy.method.confidence import ConfidenceRegionABC, WithConfidenceABC
|
||||
from quapy.protocol import AbstractProtocol
|
||||
|
||||
# stan's plugin discovery (stan.plugins.get_plugins) calls pkg_resources.iter_entry_points()
|
||||
# on every model build, each of which re-emits setuptools' pkg_resources deprecation notice;
|
||||
# this is upstream pystan noise, not actionable in quapy, so it is silenced here.
|
||||
warnings.filterwarnings(
|
||||
"ignore",
|
||||
message=r".*pkg_resources is deprecated.*",
|
||||
category=UserWarning,
|
||||
)
|
||||
|
||||
try:
|
||||
import jax
|
||||
import jax.numpy as jnp
|
||||
import jax.random as jrandom
|
||||
from jax.scipy.special import logsumexp as jax_logsumexp
|
||||
import numpyro
|
||||
import numpyro.distributions as dist
|
||||
from numpyro.infer import MCMC, NUTS
|
||||
import stan
|
||||
import stan.common
|
||||
|
||||
DEPENDENCIES_INSTALLED = True
|
||||
except ImportError:
|
||||
except ImportError as e:
|
||||
logging.getLogger(__name__).warning(f'Bayesian dependencies failed to import: {e!r}')
|
||||
jax = None
|
||||
jnp = None
|
||||
jrandom = None
|
||||
jax_logsumexp = None
|
||||
numpyro = None
|
||||
dist = None
|
||||
MCMC = None
|
||||
NUTS = None
|
||||
stan = None
|
||||
|
||||
DEPENDENCIES_INSTALLED = False
|
||||
|
||||
|
|
@ -24,28 +65,124 @@ P_TEST_C: str = "P_test(C)"
|
|||
P_C_COND_Y: str = "P(C|Y)"
|
||||
|
||||
|
||||
def model(n_c_unlabeled: np.ndarray, n_y_and_c_labeled: np.ndarray) -> None:
|
||||
"""
|
||||
Defines a probabilistic model in `NumPyro <https://num.pyro.ai/>`_.
|
||||
def _require_bayesian_dependencies():
|
||||
if not DEPENDENCIES_INSTALLED:
|
||||
raise ImportError(
|
||||
"Auxiliary dependencies are required. "
|
||||
"Run `$ pip install quapy[bayes]` to install them."
|
||||
)
|
||||
|
||||
:param n_c_unlabeled: a `np.ndarray` of shape `(n_predicted_classes,)`
|
||||
with entry `c` being the number of instances predicted as class `c`.
|
||||
:param n_y_and_c_labeled: a `np.ndarray` of shape `(n_classes, n_predicted_classes)`
|
||||
with entry `(y, c)` being the number of instances labeled as class `y` and predicted as class `c`.
|
||||
|
||||
def _resolve_dirichlet_prior(prior, n_classes, *, allow_mapls_priors=False, n_test=None, map_prev=None, map_lambda=None):
|
||||
if isinstance(prior, str):
|
||||
if prior == 'uniform':
|
||||
return np.ones(n_classes, dtype=float)
|
||||
if allow_mapls_priors and prior in {'map', 'map2'}:
|
||||
if n_test is None or map_prev is None:
|
||||
raise ValueError('MAPLS priors require n_test and map_prev')
|
||||
if prior == 'map':
|
||||
lam = map_lambda
|
||||
else:
|
||||
lam = get_lambda(
|
||||
test_probs=map_prev["test_probs"],
|
||||
pz=map_prev["train_prev"],
|
||||
q_prior=map_prev["map_estimate"],
|
||||
dvg=kl_div,
|
||||
)
|
||||
alpha_0 = alpha0_from_lambda(lam, n_test=n_test, n_classes=n_classes)
|
||||
return np.full(n_classes, alpha_0, dtype=float)
|
||||
raise ValueError(f"unknown prior specification {prior!r}")
|
||||
if isinstance(prior, Number):
|
||||
return np.full(n_classes, float(prior), dtype=float)
|
||||
|
||||
alpha = np.asarray(prior, dtype=float)
|
||||
if alpha.ndim != 1 or len(alpha) != n_classes:
|
||||
raise ValueError(f'wrong shape for prior; expected {n_classes} values, found shape {alpha.shape}')
|
||||
return alpha
|
||||
|
||||
|
||||
def _validate_temperature(temperature):
|
||||
if not isinstance(temperature, Real) or temperature <= 0:
|
||||
raise ValueError(f'expected a positive real value for temperature; found {temperature!r}')
|
||||
return float(temperature)
|
||||
|
||||
|
||||
def model_bayesianCC(
|
||||
n_c_unlabeled: np.ndarray,
|
||||
n_y_and_c_labeled: np.ndarray,
|
||||
temperature: float,
|
||||
alpha: np.ndarray,
|
||||
) -> None:
|
||||
"""
|
||||
NumPyro model for BayesianCC.
|
||||
"""
|
||||
n_y_labeled = n_y_and_c_labeled.sum(axis=1)
|
||||
|
||||
K = len(n_c_unlabeled)
|
||||
L = len(n_y_labeled)
|
||||
n_pred_classes = len(n_c_unlabeled)
|
||||
n_classes = len(n_y_labeled)
|
||||
|
||||
pi_ = numpyro.sample(P_TEST_Y, dist.Dirichlet(jnp.ones(L)))
|
||||
p_c_cond_y = numpyro.sample(P_C_COND_Y, dist.Dirichlet(jnp.ones(K).repeat(L).reshape(L, K)))
|
||||
pi_ = numpyro.sample(P_TEST_Y, dist.Dirichlet(jnp.asarray(alpha, dtype=jnp.float32)))
|
||||
p_c_cond_y = numpyro.sample(
|
||||
P_C_COND_Y,
|
||||
dist.Dirichlet(jnp.ones(n_pred_classes).repeat(n_classes).reshape(n_classes, n_pred_classes)),
|
||||
)
|
||||
|
||||
with numpyro.plate('plate', L):
|
||||
numpyro.sample('F_yc', dist.Multinomial(n_y_labeled, p_c_cond_y), obs=n_y_and_c_labeled)
|
||||
if temperature == 1.0:
|
||||
with numpyro.plate('plate', n_classes):
|
||||
numpyro.sample('F_yc', dist.Multinomial(n_y_labeled, p_c_cond_y), obs=n_y_and_c_labeled)
|
||||
|
||||
p_c = numpyro.deterministic(P_TEST_C, jnp.einsum("yc,y->c", p_c_cond_y, pi_))
|
||||
numpyro.sample('N_c', dist.Multinomial(jnp.sum(n_c_unlabeled), p_c), obs=n_c_unlabeled)
|
||||
return
|
||||
|
||||
with numpyro.plate('plate_y', n_classes):
|
||||
logp_f = dist.Multinomial(n_y_labeled, p_c_cond_y).log_prob(n_y_and_c_labeled)
|
||||
|
||||
numpyro.factor('F_yc_loglik', jnp.sum(logp_f) / temperature)
|
||||
|
||||
p_c = numpyro.deterministic(P_TEST_C, jnp.einsum("yc,y->c", p_c_cond_y, pi_))
|
||||
numpyro.sample('N_c', dist.Multinomial(jnp.sum(n_c_unlabeled), p_c), obs=n_c_unlabeled)
|
||||
logp_n = dist.Multinomial(jnp.sum(n_c_unlabeled), p_c).log_prob(n_c_unlabeled)
|
||||
numpyro.factor('N_c_loglik', logp_n / temperature)
|
||||
|
||||
|
||||
def model(n_c_unlabeled: np.ndarray, n_y_and_c_labeled: np.ndarray) -> None:
|
||||
"""
|
||||
Backward-compatible BayesianCC model with a uniform prior.
|
||||
"""
|
||||
alpha = np.ones(n_y_and_c_labeled.shape[0], dtype=float)
|
||||
return model_bayesianCC(n_c_unlabeled, n_y_and_c_labeled, temperature=1.0, alpha=alpha)
|
||||
|
||||
|
||||
def sample_posterior_bayesianCC(
|
||||
n_c_unlabeled: np.ndarray,
|
||||
n_y_and_c_labeled: np.ndarray,
|
||||
num_warmup: int,
|
||||
num_samples: int,
|
||||
alpha: np.ndarray,
|
||||
temperature: float = 1.0,
|
||||
seed: int = 0,
|
||||
) -> dict:
|
||||
"""
|
||||
Samples from the BayesianCC posterior using NumPyro.
|
||||
"""
|
||||
_require_bayesian_dependencies()
|
||||
temperature = _validate_temperature(temperature)
|
||||
|
||||
mcmc = numpyro.infer.MCMC(
|
||||
numpyro.infer.NUTS(model_bayesianCC),
|
||||
num_warmup=num_warmup,
|
||||
num_samples=num_samples,
|
||||
progress_bar=False,
|
||||
)
|
||||
rng_key = jax.random.PRNGKey(seed)
|
||||
mcmc.run(
|
||||
rng_key,
|
||||
n_c_unlabeled=n_c_unlabeled,
|
||||
n_y_and_c_labeled=n_y_and_c_labeled,
|
||||
temperature=temperature,
|
||||
alpha=alpha,
|
||||
)
|
||||
return mcmc.get_samples()
|
||||
|
||||
|
||||
def sample_posterior(
|
||||
|
|
@ -56,24 +193,691 @@ def sample_posterior(
|
|||
seed: int = 0,
|
||||
) -> dict:
|
||||
"""
|
||||
Samples from the Bayesian quantification model in NumPyro using the
|
||||
`NUTS <https://arxiv.org/abs/1111.4246>`_ sampler.
|
||||
|
||||
:param n_c_unlabeled: a `np.ndarray` of shape `(n_predicted_classes,)`
|
||||
with entry `c` being the number of instances predicted as class `c`.
|
||||
:param n_y_and_c_labeled: a `np.ndarray` of shape `(n_classes, n_predicted_classes)`
|
||||
with entry `(y, c)` being the number of instances labeled as class `y` and predicted as class `c`.
|
||||
:param num_warmup: the number of warmup steps.
|
||||
:param num_samples: the number of samples to draw.
|
||||
:seed: the random seed.
|
||||
:return: a `dict` with the samples. The keys are the names of the latent variables.
|
||||
Backward-compatible wrapper around BayesianCC sampling.
|
||||
"""
|
||||
mcmc = numpyro.infer.MCMC(
|
||||
numpyro.infer.NUTS(model),
|
||||
alpha = np.ones(n_y_and_c_labeled.shape[0], dtype=float)
|
||||
return sample_posterior_bayesianCC(
|
||||
n_c_unlabeled=n_c_unlabeled,
|
||||
n_y_and_c_labeled=n_y_and_c_labeled,
|
||||
num_warmup=num_warmup,
|
||||
num_samples=num_samples,
|
||||
progress_bar=False
|
||||
alpha=alpha,
|
||||
temperature=1.0,
|
||||
seed=seed,
|
||||
)
|
||||
rng_key = jax.random.PRNGKey(seed)
|
||||
mcmc.run(rng_key, n_c_unlabeled=n_c_unlabeled, n_y_and_c_labeled=n_y_and_c_labeled)
|
||||
return mcmc.get_samples()
|
||||
|
||||
|
||||
def load_stan_file():
|
||||
return importlib.resources.files('quapy.method').joinpath('stan/pq.stan').read_text(encoding='utf-8')
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _suppress_stan_logging():
|
||||
with open(os.devnull, "w") as devnull:
|
||||
old_stderr = sys.stderr
|
||||
sys.stderr = devnull
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
sys.stderr = old_stderr
|
||||
|
||||
|
||||
def pq_stan(stan_code, n_bins, pos_hist, neg_hist, test_hist, number_of_samples, num_warmup, stan_seed):
|
||||
"""
|
||||
Samples posterior prevalences for PQ from a Stan model.
|
||||
"""
|
||||
_require_bayesian_dependencies()
|
||||
logging.getLogger("stan.common").setLevel(logging.ERROR)
|
||||
|
||||
stan_data = {
|
||||
'n_bucket': n_bins,
|
||||
'train_neg': neg_hist.tolist(),
|
||||
'train_pos': pos_hist.tolist(),
|
||||
'test': test_hist.tolist(),
|
||||
'posterior': 1,
|
||||
}
|
||||
|
||||
with _suppress_stan_logging():
|
||||
stan_model = stan.build(stan_code, data=stan_data, random_seed=stan_seed)
|
||||
fit = stan_model.sample(num_chains=1, num_samples=number_of_samples, num_warmup=num_warmup)
|
||||
|
||||
return fit['prev']
|
||||
|
||||
|
||||
class BayesianKDEy(AggregativeSoftQuantifier, KDEBase, WithConfidenceABC):
|
||||
"""
|
||||
Bayesian version of KDEy.
|
||||
|
||||
This method relies on extra dependencies, which have to be installed via:
|
||||
`$ pip install quapy[bayes]`
|
||||
|
||||
:param classifier: a scikit-learn's BaseEstimator, or None, in which case
|
||||
the classifier is taken to be the one indicated in
|
||||
`qp.environ['DEFAULT_CLS']`
|
||||
:param fit_classifier: whether to train the classifier, or consider it
|
||||
already fit
|
||||
:param val_split: specifies the data used for generating classifier
|
||||
predictions. This specification can be made as float in (0, 1)
|
||||
indicating the proportion of stratified held-out validation set to be
|
||||
extracted from the training set; or as an integer (default 5),
|
||||
indicating that the predictions are to be generated in a `k`-fold
|
||||
cross-validation manner (with this integer indicating the value for
|
||||
`k`); or as a tuple `(X,y)` defining the specific set of data to use
|
||||
for validation. Set to None when the method does not require any
|
||||
validation data, in order to avoid that some portion of the training
|
||||
data be wasted.
|
||||
:param kernel: kernel function for KDE. Available kernels include
|
||||
{'gaussian', 'aitchison', 'ilr'} (default 'gaussian')
|
||||
:param bandwidth: bandwidth for the kernel (default 0.1)
|
||||
:param shrinkage: regularization strength for Aitchison/ILR kernels
|
||||
(default 0.0)
|
||||
:param num_warmup: number of warmup iterations for the MCMC sampler
|
||||
(default 500)
|
||||
:param num_samples: number of samples to draw from the posterior
|
||||
(default 1000)
|
||||
:param mcmc_seed: random seed for the MCMC sampler (default 0)
|
||||
:param confidence_level: float in [0,1] to construct a confidence region
|
||||
around the point estimate (default 0.95)
|
||||
:param region: string, set to `intervals` for constructing confidence
|
||||
intervals (default), or to `ellipse` for constructing an ellipse in
|
||||
the probability simplex, or to `ellipse-clr` for constructing an
|
||||
ellipse in the Centered-Log Ratio (CLR) unconstrained space.
|
||||
:param bonferroni: bool (default False), whether to apply Bonferroni
|
||||
correction when `region='intervals'`. This parameter has no effect
|
||||
for ellipse-based regions.
|
||||
:param temperature: temperature (>0) for posterior calibration
|
||||
(default 1.)
|
||||
:param prior: an array-like with the alpha parameters of a Dirichlet
|
||||
prior, a scalar real value to be broadcast to all classes, or the
|
||||
string 'uniform' for a uniform, uninformative prior (default)
|
||||
:param verbose: bool, whether to display the progress bar
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
classifier: BaseEstimator = None,
|
||||
fit_classifier=True,
|
||||
val_split: int = 5,
|
||||
kernel='gaussian',
|
||||
bandwidth=0.1,
|
||||
shrinkage=0.0,
|
||||
num_warmup: int = 500,
|
||||
num_samples: int = 1_000,
|
||||
mcmc_seed: int = 0,
|
||||
confidence_level: float = 0.95,
|
||||
region: str = 'intervals',
|
||||
bonferroni: bool = False,
|
||||
temperature: float = 1.0,
|
||||
prior='uniform',
|
||||
verbose: bool = False,
|
||||
):
|
||||
_require_bayesian_dependencies()
|
||||
if num_warmup <= 0:
|
||||
raise ValueError(f'parameter {num_warmup=} must be a positive integer')
|
||||
if num_samples <= 0:
|
||||
raise ValueError(f'parameter {num_samples=} must be a positive integer')
|
||||
|
||||
self.kernel = KDEBase._check_kernel(kernel)
|
||||
self.bandwidth = KDEBase._check_bandwidth(bandwidth, self.kernel)
|
||||
assert 0 <= shrinkage < 1, 'shrinkage must be in [0,1)'
|
||||
assert self.kernel != 'gaussian' or shrinkage == 0, \
|
||||
'shrinkage is only supported for Aitchison/ILR kernels'
|
||||
|
||||
super().__init__(classifier, fit_classifier, val_split)
|
||||
self.shrinkage = float(shrinkage)
|
||||
self.num_warmup = num_warmup
|
||||
self.num_samples = num_samples
|
||||
self.mcmc_seed = mcmc_seed
|
||||
self.confidence_level = confidence_level
|
||||
self.region = region
|
||||
self.bonferroni = bonferroni
|
||||
self.temperature = _validate_temperature(temperature)
|
||||
self.prior = prior
|
||||
self.verbose = verbose
|
||||
self.prevalence_samples = None
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
self.mix_densities = self.get_mixture_components(
|
||||
classif_predictions,
|
||||
labels,
|
||||
self.classes_,
|
||||
self.bandwidth,
|
||||
self.kernel,
|
||||
)
|
||||
return self
|
||||
|
||||
def sample_from_posterior(self, classif_predictions):
|
||||
test_log_densities = np.asarray(
|
||||
[self.pdf(kde_i, classif_predictions, self.kernel, log_densities=True) for kde_i in self.mix_densities]
|
||||
)
|
||||
alpha = _resolve_dirichlet_prior(self.prior, len(self.mix_densities))
|
||||
|
||||
mcmc = MCMC(
|
||||
NUTS(self._numpyro_model),
|
||||
num_warmup=self.num_warmup,
|
||||
num_samples=self.num_samples,
|
||||
num_chains=1,
|
||||
progress_bar=self.verbose,
|
||||
)
|
||||
mcmc.run(jrandom.PRNGKey(self.mcmc_seed), test_log_densities=test_log_densities, alpha=alpha)
|
||||
self.prevalence_samples = np.asarray(mcmc.get_samples()["prev"])
|
||||
return self.prevalence_samples
|
||||
|
||||
def aggregate(self, classif_predictions: np.ndarray):
|
||||
return self.sample_from_posterior(classif_predictions).mean(axis=0)
|
||||
|
||||
def predict_conf(self, instances, confidence_level=None) -> (np.ndarray, ConfidenceRegionABC):
|
||||
confidence_level = self.confidence_level if confidence_level is None else confidence_level
|
||||
classif_predictions = self.classify(instances)
|
||||
point_estimate = self.aggregate(classif_predictions)
|
||||
region = WithConfidenceABC.construct_region(
|
||||
self.prevalence_samples,
|
||||
confidence_level=confidence_level,
|
||||
method=self.region,
|
||||
bonferroni=self.bonferroni,
|
||||
)
|
||||
return point_estimate, region
|
||||
|
||||
def _numpyro_model(self, test_log_densities, alpha):
|
||||
prev = numpyro.sample("prev", dist.Dirichlet(jnp.asarray(alpha)))
|
||||
log_likelihood = jnp.sum(jax_logsumexp(jnp.log(prev)[:, None] + test_log_densities, axis=0))
|
||||
numpyro.factor("loglik", (1.0 / self.temperature) * log_likelihood)
|
||||
|
||||
|
||||
class _JaxILRTransformation(F.CompositionalTransformation):
|
||||
"""
|
||||
JAX-backed ILR transform used inside Bayesian MAPLS.
|
||||
"""
|
||||
|
||||
def __call__(self, X):
|
||||
X = jnp.asarray(X)
|
||||
X = qp.error.smooth(np.asarray(X), self.EPSILON)
|
||||
X = jnp.asarray(X)
|
||||
basis = jnp.asarray(self.get_V(X.shape[-1]))
|
||||
return jnp.log(X) @ basis.T
|
||||
|
||||
def inverse(self, Z):
|
||||
Z = jnp.asarray(Z)
|
||||
basis = jnp.asarray(self.get_V(Z.shape[-1] + 1))
|
||||
logp = Z @ basis
|
||||
p = jnp.exp(logp)
|
||||
return p / jnp.sum(p, axis=-1, keepdims=True)
|
||||
|
||||
def get_V(self, k):
|
||||
return F.ILRtransformation().get_V(k)
|
||||
|
||||
|
||||
class BayesianMAPLS(AggregativeSoftQuantifier, WithConfidenceABC):
|
||||
"""
|
||||
Bayesian variant of the MLLS/EMQ method proposed by
|
||||
Ye, Changkun, et al. "Label shift estimation for class-imbalance problem:
|
||||
A bayesian approach." Proceedings of the IEEE/CVF Winter Conference on
|
||||
Applications of Computer Vision. 2024.
|
||||
|
||||
Code adapted from:
|
||||
https://github.com/ChangkunYe/MAPLS/blob/main/label_shift/mapls.py
|
||||
|
||||
This method relies on extra dependencies, which have to be installed via:
|
||||
`$ pip install quapy[bayes]`
|
||||
|
||||
:param classifier: a scikit-learn's BaseEstimator, or None, in which case
|
||||
the classifier is taken to be the one indicated in
|
||||
`qp.environ['DEFAULT_CLS']`
|
||||
:param fit_classifier: whether to train the classifier, or consider it
|
||||
already fit
|
||||
:param val_split: specifies the data used for generating classifier
|
||||
predictions. This specification can be made as float in (0, 1)
|
||||
indicating the proportion of stratified held-out validation set to be
|
||||
extracted from the training set; or as an integer (default 5),
|
||||
indicating that the predictions are to be generated in a `k`-fold
|
||||
cross-validation manner (with this integer indicating the value for
|
||||
`k`); or as a tuple `(X,y)` defining the specific set of data to use
|
||||
for validation. Set to None when the method does not require any
|
||||
validation data, in order to avoid that some portion of the training
|
||||
data be wasted.
|
||||
:param exact_train_prev: set to True (default) for using the true training
|
||||
prevalence as the initial observation; set to False for computing the
|
||||
training prevalence as an estimate of it, i.e., as the expected value
|
||||
of the posterior probabilities of the training instances.
|
||||
:param num_warmup: number of warmup iterations for the MCMC sampler
|
||||
(default 500)
|
||||
:param num_samples: number of samples to draw from the posterior
|
||||
(default 1000)
|
||||
:param mcmc_seed: random seed for the MCMC sampler (default 0)
|
||||
:param confidence_level: float in [0,1] to construct a confidence region
|
||||
around the point estimate (default 0.95)
|
||||
:param region: string, set to `intervals` for constructing confidence
|
||||
intervals (default), or to `ellipse` for constructing an ellipse in
|
||||
the probability simplex, or to `ellipse-clr` for constructing an
|
||||
ellipse in the Centered-Log Ratio (CLR) unconstrained space.
|
||||
:param bonferroni: bool (default False), whether to apply Bonferroni
|
||||
correction when `region='intervals'`. This parameter has no effect
|
||||
for ellipse-based regions.
|
||||
:param temperature: temperature (>0) for posterior calibration
|
||||
(default 1.)
|
||||
:param prior: an array-like with the alpha parameters of a Dirichlet
|
||||
prior, a scalar real value to be broadcast to all classes, or one of
|
||||
{'uniform', 'map', 'map2'} (default 'uniform')
|
||||
:param mapls_chain_init: whether to initialize the Markov chain with a
|
||||
preliminary EM point estimate (default True)
|
||||
:param verbose: bool, whether to display the progress bar
|
||||
(default False)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
classifier: BaseEstimator = None,
|
||||
fit_classifier=True,
|
||||
val_split: int = 5,
|
||||
exact_train_prev=True,
|
||||
num_warmup: int = 500,
|
||||
num_samples: int = 1_000,
|
||||
mcmc_seed: int = 0,
|
||||
confidence_level: float = 0.95,
|
||||
region: str = 'intervals',
|
||||
bonferroni: bool = False,
|
||||
temperature: float = 1.0,
|
||||
prior='uniform',
|
||||
mapls_chain_init=True,
|
||||
verbose=False,
|
||||
):
|
||||
_require_bayesian_dependencies()
|
||||
if num_warmup <= 0:
|
||||
raise ValueError(f'parameter {num_warmup=} must be a positive integer')
|
||||
if num_samples <= 0:
|
||||
raise ValueError(f'parameter {num_samples=} must be a positive integer')
|
||||
if not (
|
||||
(isinstance(prior, str) and prior in {'uniform', 'map', 'map2'})
|
||||
or isinstance(prior, Number)
|
||||
or (isinstance(prior, Iterable) and all(isinstance(v, Number) for v in prior))
|
||||
):
|
||||
raise ValueError(
|
||||
f'wrong type for {prior=}; expected one of {{"uniform", "map", "map2"}}, '
|
||||
'a real scalar, or an array-like of real values'
|
||||
)
|
||||
|
||||
super().__init__(classifier, fit_classifier, val_split)
|
||||
self.exact_train_prev = exact_train_prev
|
||||
self.num_warmup = num_warmup
|
||||
self.num_samples = num_samples
|
||||
self.mcmc_seed = mcmc_seed
|
||||
self.confidence_level = confidence_level
|
||||
self.region = region
|
||||
self.bonferroni = bonferroni
|
||||
self.temperature = _validate_temperature(temperature)
|
||||
self.prior = prior
|
||||
self.mapls_chain_init = mapls_chain_init
|
||||
self.verbose = verbose
|
||||
self.prevalence_samples = None
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
self.train_post = classif_predictions
|
||||
if self.exact_train_prev:
|
||||
self.train_prevalence = F.prevalence_from_labels(labels, classes=self.classes_)
|
||||
else:
|
||||
self.train_prevalence = F.prevalence_from_probabilities(classif_predictions)
|
||||
self.ilr = _JaxILRTransformation()
|
||||
return self
|
||||
|
||||
def sample_from_posterior(self, classif_predictions):
|
||||
n_test, n_classes = classif_predictions.shape
|
||||
map_estimate, lam = mapls(
|
||||
self.train_post,
|
||||
test_probs=classif_predictions,
|
||||
pz=self.train_prevalence,
|
||||
return_lambda=True,
|
||||
)
|
||||
|
||||
z0 = self.ilr(map_estimate)
|
||||
if isinstance(self.prior, str) and self.prior in {'map', 'map2'}:
|
||||
prior_context = {
|
||||
"test_probs": classif_predictions,
|
||||
"train_prev": self.train_prevalence,
|
||||
"map_estimate": map_estimate,
|
||||
}
|
||||
alpha = _resolve_dirichlet_prior(
|
||||
self.prior,
|
||||
n_classes,
|
||||
allow_mapls_priors=True,
|
||||
n_test=n_test,
|
||||
map_prev=prior_context,
|
||||
map_lambda=lam,
|
||||
)
|
||||
else:
|
||||
alpha = _resolve_dirichlet_prior(self.prior, n_classes)
|
||||
|
||||
mcmc = MCMC(
|
||||
NUTS(self._numpyro_model),
|
||||
num_warmup=self.num_warmup,
|
||||
num_samples=self.num_samples,
|
||||
num_chains=1,
|
||||
progress_bar=self.verbose,
|
||||
)
|
||||
mcmc.run(
|
||||
jrandom.PRNGKey(self.mcmc_seed),
|
||||
test_posteriors=classif_predictions,
|
||||
alpha=alpha,
|
||||
init_params={"z": z0} if self.mapls_chain_init else None,
|
||||
)
|
||||
|
||||
samples = mcmc.get_samples()["z"]
|
||||
self.prevalence_samples = np.asarray(self.ilr.inverse(samples))
|
||||
return self.prevalence_samples
|
||||
|
||||
def aggregate(self, classif_predictions: np.ndarray):
|
||||
return self.sample_from_posterior(classif_predictions).mean(axis=0)
|
||||
|
||||
def predict_conf(self, instances, confidence_level=None) -> (np.ndarray, ConfidenceRegionABC):
|
||||
confidence_level = self.confidence_level if confidence_level is None else confidence_level
|
||||
classif_predictions = self.classify(instances)
|
||||
point_estimate = self.aggregate(classif_predictions)
|
||||
region = WithConfidenceABC.construct_region(
|
||||
self.prevalence_samples,
|
||||
confidence_level=confidence_level,
|
||||
method=self.region,
|
||||
bonferroni=self.bonferroni,
|
||||
)
|
||||
return point_estimate, region
|
||||
|
||||
def _log_likelihood(self, test_classif, test_prev, train_prev):
|
||||
log_w = jnp.log(test_prev) - jnp.log(train_prev)
|
||||
return jnp.sum(jax_logsumexp(jnp.log(test_classif) + log_w, axis=-1))
|
||||
|
||||
def _numpyro_model(self, test_posteriors, alpha):
|
||||
test_posteriors = jnp.asarray(test_posteriors)
|
||||
n_classes = test_posteriors.shape[1]
|
||||
|
||||
z = numpyro.sample("z", dist.Normal(jnp.zeros(n_classes - 1), 1.0))
|
||||
prev = self.ilr.inverse(z)
|
||||
train_prev = jnp.asarray(self.train_prevalence)
|
||||
alpha = jnp.asarray(alpha)
|
||||
|
||||
numpyro.factor("dirichlet_prior", dist.Dirichlet(alpha).log_prob(prev))
|
||||
numpyro.factor(
|
||||
"likelihood",
|
||||
(1.0 / self.temperature) * self._log_likelihood(test_posteriors, test_prev=prev, train_prev=train_prev),
|
||||
)
|
||||
|
||||
|
||||
def mapls(
|
||||
train_probs: np.ndarray,
|
||||
test_probs: np.ndarray,
|
||||
pz: np.ndarray,
|
||||
qy_mode: str = 'soft',
|
||||
max_iter: int = 100,
|
||||
init_mode: str = 'identical',
|
||||
lam: float = None,
|
||||
dvg_name='kl',
|
||||
return_lambda=False,
|
||||
):
|
||||
cls_num = len(pz)
|
||||
assert test_probs.shape[-1] == cls_num
|
||||
if not isinstance(max_iter, int) or max_iter < 0:
|
||||
raise ValueError(f'expected a non-negative integer for max_iter; found {max_iter!r}')
|
||||
|
||||
if dvg_name == 'kl':
|
||||
dvg = kl_div
|
||||
elif dvg_name == 'js':
|
||||
dvg = js_div
|
||||
else:
|
||||
raise ValueError(f'Unsupported distribution distance measure {dvg_name!r}; expected "kl" or "js"')
|
||||
|
||||
q_prior = np.ones(cls_num) / cls_num
|
||||
if lam is None:
|
||||
lam = get_lambda(test_probs, pz, q_prior, dvg=dvg, max_iter=max_iter)
|
||||
|
||||
qz = mapls_em(
|
||||
test_probs,
|
||||
pz,
|
||||
lam,
|
||||
q_prior,
|
||||
cls_num,
|
||||
init_mode=init_mode,
|
||||
max_iter=max_iter,
|
||||
qy_mode=qy_mode,
|
||||
)
|
||||
return (qz, lam) if return_lambda else qz
|
||||
|
||||
|
||||
def mapls_em(probs, pz, lam, q_prior, cls_num, init_mode='identical', max_iter=100, qy_mode='soft'):
|
||||
pz = np.asarray(pz, dtype=float)
|
||||
pz = pz / np.sum(pz)
|
||||
if init_mode == 'uniform':
|
||||
qz = np.ones(cls_num) / cls_num
|
||||
elif init_mode == 'identical':
|
||||
qz = pz.copy()
|
||||
else:
|
||||
raise ValueError('init_mode should be either "uniform" or "identical"')
|
||||
|
||||
w = qz / pz
|
||||
for _ in range(max_iter):
|
||||
mapls_probs = normalized(probs * w, axis=-1, order=1)
|
||||
if qy_mode == 'hard':
|
||||
pred = np.argmax(mapls_probs, axis=-1)
|
||||
qz_new = np.bincount(pred.reshape(-1), minlength=cls_num)
|
||||
elif qy_mode == 'soft':
|
||||
qz_new = np.mean(mapls_probs, axis=0)
|
||||
else:
|
||||
raise ValueError('qy_mode should be either "soft" or "hard"')
|
||||
|
||||
qz = lam * qz_new + (1 - lam) * q_prior
|
||||
qz /= qz.sum()
|
||||
w = qz / pz
|
||||
|
||||
return qz
|
||||
|
||||
|
||||
def get_lambda(test_probs, pz, q_prior, dvg, max_iter=50):
|
||||
n_classes = len(pz)
|
||||
qz_pred = mapls_em(test_probs, pz, 1, 0, n_classes, max_iter=max_iter)
|
||||
|
||||
tu_div = dvg(qz_pred, q_prior)
|
||||
ts_div = dvg(qz_pred, pz)
|
||||
su_div = dvg(pz, q_prior)
|
||||
|
||||
su_conf = 1 - lambda_forward(su_div, lambda_inverse(dpq=0.5, lam=0.2))
|
||||
tu_conf = lambda_forward(tu_div, lambda_inverse(dpq=0.5, lam=su_conf))
|
||||
ts_conf = lambda_forward(ts_div, lambda_inverse(dpq=0.5, lam=su_conf))
|
||||
|
||||
confs = np.array([tu_conf, 1 - ts_conf])
|
||||
weights = np.array([0.9, 0.1])
|
||||
return np.sum(weights * confs)
|
||||
|
||||
|
||||
def lambda_inverse(dpq, lam):
|
||||
return (1 / (1 - lam) - 1) / dpq
|
||||
|
||||
|
||||
def lambda_forward(dpq, gamma):
|
||||
return gamma * dpq / (1 + gamma * dpq)
|
||||
|
||||
|
||||
def get_lamda(test_probs, pz, q_prior, dvg, max_iter=50):
|
||||
return get_lambda(test_probs, pz, q_prior, dvg, max_iter=max_iter)
|
||||
|
||||
|
||||
def lam_inv(dpq, lam):
|
||||
return lambda_inverse(dpq, lam)
|
||||
|
||||
|
||||
def lam_forward(dpq, gamma):
|
||||
return lambda_forward(dpq, gamma)
|
||||
|
||||
|
||||
def kl_div(p, q, eps=1e-12):
|
||||
p = np.asarray(p, dtype=float)
|
||||
q = np.asarray(q, dtype=float)
|
||||
|
||||
mask = p > 0
|
||||
return np.sum(p[mask] * np.log(p[mask] / (q[mask] + eps)))
|
||||
|
||||
|
||||
def js_div(p, q):
|
||||
assert (np.abs(np.sum(p) - 1) < 1e-6) and (np.abs(np.sum(q) - 1) < 1e-6)
|
||||
m = (p + q) / 2
|
||||
return kl_div(p, m) / 2 + kl_div(q, m) / 2
|
||||
|
||||
|
||||
def normalized(a, axis=-1, order=2):
|
||||
l2 = np.atleast_1d(np.linalg.norm(a, order, axis))
|
||||
l2[l2 == 0] = 1
|
||||
return a / np.expand_dims(l2, axis)
|
||||
|
||||
|
||||
def alpha0_from_lambda(lam, n_test, n_classes):
|
||||
return 1 + n_test * (1 - lam) / (lam * n_classes)
|
||||
|
||||
|
||||
def alpha0_from_lamda(lam, n_test, n_classes):
|
||||
return alpha0_from_lambda(lam, n_test, n_classes)
|
||||
|
||||
|
||||
def calibrate_temperature(
|
||||
method: WithConfidenceABC,
|
||||
train: LabelledCollection,
|
||||
val_prot: AbstractProtocol,
|
||||
temp_grid=(0.5, 1.0, 1.5, 2.0, 5.0, 10.0, 100.0),
|
||||
nominal_coverage: float = 0.95,
|
||||
amplitude_threshold=1.0,
|
||||
criterion: str = 'winkler',
|
||||
n_jobs: int = 1,
|
||||
verbose: bool = True,
|
||||
):
|
||||
"""
|
||||
Calibrates the temperature parameter of a Bayesian quantifier with
|
||||
confidence regions by selecting the value that yields the best validation
|
||||
trade-off between nominal coverage and region sharpness.
|
||||
|
||||
The method is first fitted on ``train``. For each candidate temperature,
|
||||
the fitted quantifier is deep-copied, its ``temperature`` attribute is
|
||||
replaced, and it is evaluated on the samples generated by ``val_prot``.
|
||||
Candidate temperatures whose average region amplitude exceeds
|
||||
``amplitude_threshold`` are discarded.
|
||||
|
||||
When ``criterion='winkler'``, the surviving candidate with minimum mean
|
||||
Winkler score is selected. When ``criterion='auto'``, the selected
|
||||
temperature is the one whose empirical coverage is closest to
|
||||
``nominal_coverage``.
|
||||
|
||||
:param method: a quantifier implementing :class:`WithConfidenceABC` and
|
||||
exposing a writable ``temperature`` attribute
|
||||
:param train: training set used to fit the quantifier
|
||||
:param val_prot: validation protocol yielding pairs ``(sample, true_prev)``
|
||||
:param temp_grid: candidate temperatures to evaluate
|
||||
:param nominal_coverage: target confidence level used by the Winkler score
|
||||
and coverage selection
|
||||
:param amplitude_threshold: maximum allowed average simplex proportion of
|
||||
the region. It can also be set to ``'auto'`` to use a heuristic based
|
||||
on the number of classes
|
||||
:param criterion: either ``'winkler'`` (default) or ``'auto'``
|
||||
:param n_jobs: number of parallel jobs across candidate temperatures
|
||||
:param verbose: whether to display progress information
|
||||
:return: the selected temperature value
|
||||
"""
|
||||
if not hasattr(method, 'temperature'):
|
||||
raise ValueError(f'{method.__class__.__name__} does not expose a temperature attribute')
|
||||
if not isinstance(method, WithConfidenceABC):
|
||||
raise TypeError(f'{method.__class__.__name__} is not an instance of WithConfidenceABC')
|
||||
if not 0 < nominal_coverage < 1:
|
||||
raise ValueError(f'{nominal_coverage=} must be in the interval (0,1)')
|
||||
if criterion not in {'auto', 'winkler'}:
|
||||
raise ValueError(f'unknown {criterion=}; valid ones are "auto" or "winkler"')
|
||||
if amplitude_threshold != 'auto':
|
||||
if not isinstance(amplitude_threshold, Real) or amplitude_threshold > 1.0:
|
||||
raise ValueError(
|
||||
f'wrong value for {amplitude_threshold=}; it must either be "auto" or a real value <= 1.0'
|
||||
)
|
||||
temperatures = sorted(_validate_temperature(temp) for temp in temp_grid)
|
||||
|
||||
if amplitude_threshold == 'auto':
|
||||
amplitude_threshold = 0.1 / np.log(train.n_classes + 1)
|
||||
|
||||
if amplitude_threshold > 0.1:
|
||||
print(f'warning: the {amplitude_threshold=} is too large; this may lead to uninformative regions')
|
||||
|
||||
def _evaluate_temperature_job(job_id, temperature):
|
||||
local_method = copy.deepcopy(method)
|
||||
local_method.temperature = temperature
|
||||
|
||||
coverage = 0
|
||||
amplitudes = []
|
||||
winklers = []
|
||||
errors = []
|
||||
|
||||
pbar = tqdm(
|
||||
enumerate(val_prot()),
|
||||
position=job_id,
|
||||
total=val_prot.total(),
|
||||
disable=not verbose,
|
||||
)
|
||||
|
||||
for i, (sample, prev) in pbar:
|
||||
point_estim, conf_region = local_method.predict_conf(sample)
|
||||
|
||||
if prev in conf_region:
|
||||
coverage += 1
|
||||
|
||||
amplitudes.append(conf_region.montecarlo_proportion(n_trials=50_000))
|
||||
if criterion == 'winkler':
|
||||
winklers.append(conf_region.mean_winkler_score(true_prev=prev, alpha=1 - nominal_coverage))
|
||||
errors.append(qp.error.mae(prev, point_estim))
|
||||
|
||||
description = (
|
||||
f'job={job_id} T={temperature}: '
|
||||
f'MAE={np.mean(errors):.6f} '
|
||||
f'coverage={coverage / (i + 1) * 100:.2f}% '
|
||||
f'amplitude={np.mean(amplitudes) * 100:.4f}% '
|
||||
)
|
||||
if criterion == 'winkler':
|
||||
description += f'winkler={np.mean(winklers):.4f}'
|
||||
pbar.set_description(description)
|
||||
|
||||
mean_coverage = coverage / val_prot.total()
|
||||
mean_amplitude = float(np.mean(amplitudes))
|
||||
mean_winkler = float(np.mean(winklers)) if criterion == 'winkler' else None
|
||||
mean_error = float(np.mean(errors))
|
||||
return temperature, mean_coverage, mean_amplitude, mean_winkler, mean_error
|
||||
|
||||
method.fit(*train.Xy)
|
||||
raw_results = Parallel(n_jobs=n_jobs, backend="loky")(
|
||||
delayed(_evaluate_temperature_job)(job_id, temperature)
|
||||
for job_id, temperature in tqdm(enumerate(temperatures), disable=not verbose)
|
||||
)
|
||||
filtered_results = [
|
||||
(temperature, coverage, amplitude, winkler, error)
|
||||
for temperature, coverage, amplitude, winkler, error in raw_results
|
||||
if amplitude < amplitude_threshold
|
||||
]
|
||||
|
||||
chosen_temperature = 1.0
|
||||
chosen_coverage = chosen_amplitude = chosen_winkler = chosen_error = None
|
||||
|
||||
if filtered_results:
|
||||
if criterion == 'winkler':
|
||||
chosen_temperature, chosen_coverage, chosen_amplitude, chosen_winkler, chosen_error = min(
|
||||
filtered_results, key=lambda item: item[3]
|
||||
)
|
||||
else:
|
||||
chosen_temperature, chosen_coverage, chosen_amplitude, chosen_winkler, chosen_error = min(
|
||||
filtered_results, key=lambda item: abs(item[1] - nominal_coverage)
|
||||
)
|
||||
|
||||
if verbose and chosen_coverage is not None:
|
||||
message = (
|
||||
f'\nChosen_temperature={chosen_temperature:.2f} got '
|
||||
f'MAE={chosen_error:.6f} '
|
||||
f'coverage={chosen_coverage * 100:.2f}% '
|
||||
f'amplitude={chosen_amplitude * 100:.4f}% '
|
||||
)
|
||||
if criterion == 'winkler':
|
||||
message += f'winkler={chosen_winkler:.4f}'
|
||||
print(message)
|
||||
|
||||
return chosen_temperature
|
||||
|
||||
|
||||
def temp_calibration(*args, **kwargs):
|
||||
"""
|
||||
Backward-compatible alias for :func:`calibrate_temperature`.
|
||||
"""
|
||||
return calibrate_temperature(*args, **kwargs)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,144 @@
|
|||
from typing import Callable, Union
|
||||
|
||||
import numpy as np
|
||||
from sklearn.metrics.pairwise import euclidean_distances, manhattan_distances
|
||||
|
||||
import quapy as qp
|
||||
import quapy.functional as F
|
||||
from quapy.method._helper import _get_quadprog
|
||||
|
||||
|
||||
class _EnergyDistanceCore:
|
||||
"""Shared numerical core for energy-distance quantifiers."""
|
||||
|
||||
def _check_ed_init_parameters(self):
|
||||
self.distance = self._resolve_distance_function(self.distance)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_distance_function(distance):
|
||||
if isinstance(distance, str):
|
||||
if distance == 'manhattan':
|
||||
return manhattan_distances
|
||||
if distance == 'euclidean':
|
||||
return euclidean_distances
|
||||
raise ValueError(
|
||||
f"unknown distance {distance!r}; valid aliases are 'manhattan' and 'euclidean'"
|
||||
)
|
||||
if not hasattr(distance, '__call__'):
|
||||
raise ValueError('distance must be a valid string alias or a callable function')
|
||||
return distance
|
||||
|
||||
def _is_pd(self, m):
|
||||
"""Check whether a symmetric matrix is positive definite."""
|
||||
return self._dpofa(m)[0] == 0
|
||||
|
||||
def _dpofa(self, m):
|
||||
"""Factor a symmetric positive definite matrix."""
|
||||
r = np.array(m, copy=True)
|
||||
n = len(r)
|
||||
for k in range(n):
|
||||
s = 0.0
|
||||
if k >= 1:
|
||||
for i in range(k):
|
||||
t = r[i, k]
|
||||
if i > 0:
|
||||
t = t - np.sum(r[0:i, i] * r[0:i, k])
|
||||
t = t / r[i, i]
|
||||
r[i, k] = t
|
||||
s = s + t * t
|
||||
s = r[k, k] - s
|
||||
if s <= 0.0:
|
||||
return k + 1, r
|
||||
r[k, k] = np.sqrt(s)
|
||||
return 0, r
|
||||
|
||||
def _nearest_pd(self, A):
|
||||
"""Project a matrix onto the cone of positive-definite matrices."""
|
||||
B = (A + A.T) / 2
|
||||
_, s, V = np.linalg.svd(B)
|
||||
H = V.T @ np.diag(s) @ V
|
||||
A2 = (B + H) / 2
|
||||
A3 = (A2 + A2.T) / 2
|
||||
|
||||
if self._is_pd(A3):
|
||||
return A3
|
||||
|
||||
spacing = np.spacing(np.linalg.norm(A))
|
||||
identity_matrix = np.eye(A.shape[0])
|
||||
k = 1
|
||||
while not self._is_pd(A3):
|
||||
mineig = np.min(np.real(np.linalg.eigvals(A3)))
|
||||
A3 += identity_matrix * (-mineig * k ** 2 + spacing)
|
||||
k += 1
|
||||
|
||||
return A3
|
||||
|
||||
def _compute_ed_param_train(self, distance_func, train_distrib, n_cls_i):
|
||||
"""Pre-compute the training-side terms of the ED optimization problem."""
|
||||
n_classes = len(train_distrib)
|
||||
K = np.zeros((n_classes, n_classes), dtype=float)
|
||||
for i in range(n_classes):
|
||||
K[i, i] = distance_func(train_distrib[i], train_distrib[i]).sum()
|
||||
for j in range(i + 1, n_classes):
|
||||
K[i, j] = distance_func(train_distrib[i], train_distrib[j]).sum()
|
||||
K[j, i] = K[i, j]
|
||||
|
||||
K = K / np.dot(n_cls_i, n_cls_i.T)
|
||||
|
||||
B = np.zeros((n_classes - 1, n_classes - 1), dtype=float)
|
||||
for i in range(n_classes - 1):
|
||||
B[i, i] = -K[i, i] - K[-1, -1] + 2 * K[i, -1]
|
||||
for j in range(n_classes - 1):
|
||||
if j == i:
|
||||
continue
|
||||
B[i, j] = -K[i, j] - K[-1, -1] + K[i, -1] + K[j, -1]
|
||||
|
||||
G = 2 * B
|
||||
if not self._is_pd(G):
|
||||
G = self._nearest_pd(G)
|
||||
|
||||
C = -np.vstack([np.ones((1, n_classes - 1)), -np.eye(n_classes - 1)]).T
|
||||
b = -np.array([1] + [0] * (n_classes - 1), dtype=float)
|
||||
|
||||
return K, G, C, b
|
||||
|
||||
def _compute_ed_param_test(self, distance_func, train_distrib, test_distrib, K, n_cls_i):
|
||||
"""Compute the test-dependent linear term of the ED objective."""
|
||||
n_classes = len(train_distrib)
|
||||
Kt = np.zeros(n_classes, dtype=float)
|
||||
for i in range(n_classes):
|
||||
Kt[i] = distance_func(train_distrib[i], test_distrib).sum()
|
||||
|
||||
Kt = Kt / (n_cls_i.squeeze() * float(test_distrib.shape[0]))
|
||||
return 2 * (-Kt[:-1] + K[:-1, -1] + Kt[-1] - K[-1, -1])
|
||||
|
||||
def _solve_ed(self, G, a, C, b):
|
||||
"""Solve the energy-distance quadratic program."""
|
||||
quadprog = _get_quadprog()
|
||||
sol = quadprog.solve_qp(G=G, a=a, C=C, b=b)
|
||||
prevalences = sol[0]
|
||||
prevalences = np.append(prevalences, 1 - prevalences.sum())
|
||||
return F.normalize_prevalence(prevalences, method='clip')
|
||||
|
||||
def _fit_energy_model(self, train_distrib):
|
||||
self.train_distrib_ = tuple(train_distrib)
|
||||
self.train_n_cls_i_ = np.asarray(
|
||||
[[distrib.shape[0]] for distrib in self.train_distrib_],
|
||||
dtype=float,
|
||||
)
|
||||
self.K_, self.G_, self.C_, self.b_ = self._compute_ed_param_train(
|
||||
self.distance,
|
||||
self.train_distrib_,
|
||||
self.train_n_cls_i_,
|
||||
)
|
||||
return self
|
||||
|
||||
def _predict_energy(self, test_distrib):
|
||||
self.a_ = self._compute_ed_param_test(
|
||||
self.distance,
|
||||
self.train_distrib_,
|
||||
test_distrib,
|
||||
self.K_,
|
||||
self.train_n_cls_i_,
|
||||
)
|
||||
return self._solve_ed(G=self.G_, a=self.a_, C=self.C_, b=self.b_)
|
||||
|
|
@ -0,0 +1,117 @@
|
|||
"""
|
||||
Internal helper utilities shared by quantification methods.
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
from sklearn.metrics import confusion_matrix
|
||||
from sklearn.preprocessing import LabelEncoder
|
||||
|
||||
|
||||
def _get_abstention_calibrators():
|
||||
try:
|
||||
from abstention.calibration import NoBiasVectorScaling, TempScaling, VectorScaling
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"Posterior calibration for EMQ requires the optional 'abstention' package."
|
||||
) from exc
|
||||
return {
|
||||
'nbvs': NoBiasVectorScaling(),
|
||||
'bcts': TempScaling(bias_positions='all'),
|
||||
'ts': TempScaling(),
|
||||
'vs': VectorScaling(),
|
||||
}
|
||||
|
||||
|
||||
def _get_cvxpy():
|
||||
try:
|
||||
import cvxpy as cp
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"RLLS requires the optional 'cvxpy' package."
|
||||
) from exc
|
||||
return cp
|
||||
|
||||
|
||||
def _get_quadprog():
|
||||
try:
|
||||
import quadprog
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"EDy requires the optional 'quadprog' package."
|
||||
) from exc
|
||||
return quadprog
|
||||
|
||||
|
||||
def _labels_to_indices(labels, classes):
|
||||
encoder = LabelEncoder().fit(classes)
|
||||
return encoder.transform(labels)
|
||||
|
||||
|
||||
def _rlls_check_mode(mode):
|
||||
valid = {'soft', 'hard'}
|
||||
if mode not in valid:
|
||||
raise ValueError(f'unknown mode {mode!r}; valid ones are {valid}')
|
||||
return mode
|
||||
|
||||
|
||||
def _rlls_joint_distribution(posteriors, labels, classes, mode='soft'):
|
||||
mode = _rlls_check_mode(mode)
|
||||
posteriors = np.asarray(posteriors, dtype=float)
|
||||
labels = np.asarray(labels)
|
||||
n_samples, n_classes = posteriors.shape
|
||||
assert n_classes == len(classes), 'wrong number of posterior columns'
|
||||
|
||||
if mode == 'hard':
|
||||
pred = np.argmax(posteriors, axis=1)
|
||||
encoded_labels = _labels_to_indices(labels, classes)
|
||||
joint = confusion_matrix(encoded_labels, pred, labels=np.arange(n_classes)).T.astype(float)
|
||||
return joint / n_samples
|
||||
|
||||
joint = np.zeros((n_classes, n_classes), dtype=float)
|
||||
for class_index, class_ in enumerate(classes):
|
||||
idx = labels == class_
|
||||
if idx.any():
|
||||
joint[:, class_index] = posteriors[idx].sum(axis=0)
|
||||
return joint / n_samples
|
||||
|
||||
|
||||
def _rlls_predicted_marginal(posteriors, mode='soft'):
|
||||
mode = _rlls_check_mode(mode)
|
||||
posteriors = np.asarray(posteriors, dtype=float)
|
||||
if mode == 'soft':
|
||||
return posteriors.mean(axis=0)
|
||||
|
||||
pred = np.argmax(posteriors, axis=1)
|
||||
counts = np.bincount(pred, minlength=posteriors.shape[1]).astype(float)
|
||||
return counts / counts.sum()
|
||||
|
||||
|
||||
def _rlls_compute_3deltaC(n_classes, n_train, delta):
|
||||
return 3 * (
|
||||
2 * np.log(2 * n_classes / delta) / (3 * n_train)
|
||||
+ np.sqrt(2 * np.log(2 * n_classes / delta) / n_train)
|
||||
)
|
||||
|
||||
|
||||
def _rlls_compute_weights(C_zy, qz, pz, rho, clip=False):
|
||||
cp = _get_cvxpy()
|
||||
|
||||
n_classes = C_zy.shape[0]
|
||||
theta = cp.Variable(n_classes)
|
||||
b = qz - pz
|
||||
objective = cp.Minimize(cp.pnorm(C_zy @ theta - b) + rho * cp.pnorm(theta))
|
||||
constraints = [-1 <= theta]
|
||||
problem = cp.Problem(objective, constraints)
|
||||
|
||||
try:
|
||||
problem.solve(verbose=False, solver=cp.SCS)
|
||||
except cp.error.SolverError:
|
||||
problem.solve(verbose=False, solver=cp.SCS, use_indirect=True)
|
||||
|
||||
if theta.value is None:
|
||||
raise RuntimeError('RLLS optimization failed to produce a solution')
|
||||
|
||||
w = 1 + np.asarray(theta.value, dtype=float)
|
||||
if clip and np.any(w < 0):
|
||||
w = np.clip(w, 0, None)
|
||||
return w
|
||||
|
|
@ -1,11 +1,13 @@
|
|||
import numpy as np
|
||||
from numbers import Real
|
||||
from sklearn.base import BaseEstimator
|
||||
from sklearn.neighbors import KernelDensity
|
||||
|
||||
import quapy as qp
|
||||
from quapy.method._helper import _labels_to_indices
|
||||
from quapy.method.aggregative import AggregativeSoftQuantifier
|
||||
import quapy.functional as F
|
||||
|
||||
from scipy.special import logsumexp
|
||||
from sklearn.metrics.pairwise import rbf_kernel
|
||||
|
||||
|
||||
|
|
@ -15,44 +17,57 @@ class KDEBase:
|
|||
"""
|
||||
|
||||
BANDWIDTH_METHOD = ['scott', 'silverman']
|
||||
KERNELS = ['gaussian', 'aitchison', 'ilr']
|
||||
|
||||
@classmethod
|
||||
def _check_bandwidth(cls, bandwidth):
|
||||
def _check_bandwidth(cls, bandwidth, kernel):
|
||||
"""
|
||||
Checks that the bandwidth parameter is correct
|
||||
|
||||
:param bandwidth: either a string (see BANDWIDTH_METHOD) or a float
|
||||
:return: the bandwidth if the check is passed, or raises an exception for invalid values
|
||||
"""
|
||||
assert bandwidth in KDEBase.BANDWIDTH_METHOD or isinstance(bandwidth, float), \
|
||||
assert bandwidth in KDEBase.BANDWIDTH_METHOD or isinstance(bandwidth, Real), \
|
||||
f'invalid bandwidth, valid ones are {KDEBase.BANDWIDTH_METHOD} or float values'
|
||||
if isinstance(bandwidth, float):
|
||||
assert 0 < bandwidth < 1, \
|
||||
"the bandwith for KDEy should be in (0,1), since this method models the unit simplex"
|
||||
if isinstance(bandwidth, Real):
|
||||
bandwidth = float(bandwidth)
|
||||
return bandwidth
|
||||
|
||||
def get_kde_function(self, X, bandwidth):
|
||||
@classmethod
|
||||
def _check_kernel(cls, kernel):
|
||||
assert kernel in KDEBase.KERNELS, f'unknown {kernel=}'
|
||||
return kernel
|
||||
|
||||
def get_kde_function(self, X, bandwidth, kernel):
|
||||
"""
|
||||
Wraps the KDE function from scikit-learn.
|
||||
|
||||
:param X: data for which the density function is to be estimated
|
||||
:param bandwidth: the bandwidth of the kernel
|
||||
:param kernel: the kernel family
|
||||
:return: a scikit-learn's KernelDensity object
|
||||
"""
|
||||
X = self.transform_posteriors(X, kernel)
|
||||
bandwidth = self.effective_bandwidth(bandwidth, kernel)
|
||||
return KernelDensity(bandwidth=bandwidth).fit(X)
|
||||
|
||||
def pdf(self, kde, X):
|
||||
def pdf(self, kde, X, kernel, log_densities=False):
|
||||
"""
|
||||
Wraps the density evalution of scikit-learn's KDE. Scikit-learn returns log-scores (s), so this
|
||||
function returns :math:`e^{s}`
|
||||
|
||||
:param kde: a previously fit KDE function
|
||||
:param X: the data for which the density is to be estimated
|
||||
:param kernel: the kernel family
|
||||
:return: np.ndarray with the densities
|
||||
"""
|
||||
return np.exp(kde.score_samples(X))
|
||||
X = self.transform_posteriors(X, kernel)
|
||||
log_density = kde.score_samples(X)
|
||||
if log_densities:
|
||||
return log_density
|
||||
return np.exp(log_density)
|
||||
|
||||
def get_mixture_components(self, X, y, classes, bandwidth):
|
||||
def get_mixture_components(self, X, y, classes, bandwidth, kernel):
|
||||
"""
|
||||
Returns an array containing the mixture components, i.e., the KDE functions for each class.
|
||||
|
||||
|
|
@ -60,22 +75,57 @@ class KDEBase:
|
|||
:param y: the class labels
|
||||
:param n_classes: integer, the number of classes
|
||||
:param bandwidth: float, the bandwidth of the kernel
|
||||
:param kernel: the kernel family
|
||||
:return: a list of KernelDensity objects, each fitted with the corresponding class-specific covariates
|
||||
"""
|
||||
class_cond_X = []
|
||||
for cat in classes:
|
||||
selX = X[y==cat]
|
||||
if selX.size==0:
|
||||
selX = [F.uniform_prevalence(len(classes))]
|
||||
raise ValueError(f'empty class {cat}')
|
||||
class_cond_X.append(selX)
|
||||
return [self.get_kde_function(X_cond_yi, bandwidth) for X_cond_yi in class_cond_X]
|
||||
return [self.get_kde_function(X_cond_yi, bandwidth, kernel) for X_cond_yi in class_cond_X]
|
||||
|
||||
def transform_posteriors(self, X, kernel):
|
||||
if kernel in {'aitchison', 'ilr'}:
|
||||
X = self.shrink_posteriors(X)
|
||||
if kernel == 'aitchison':
|
||||
return self.clr_transform(X)
|
||||
if kernel == 'ilr':
|
||||
return self.ilr_transform(X)
|
||||
return X
|
||||
|
||||
def shrink_posteriors(self, X):
|
||||
shrinkage = getattr(self, 'shrinkage', 0.0)
|
||||
if shrinkage <= 0:
|
||||
return X
|
||||
X = np.asarray(X)
|
||||
n_classes = X.shape[-1]
|
||||
uniform = np.full(n_classes, 1.0 / n_classes, dtype=X.dtype)
|
||||
return (1.0 - shrinkage) * X + shrinkage * uniform
|
||||
|
||||
def effective_bandwidth(self, bandwidth, kernel):
|
||||
shrinkage = getattr(self, 'shrinkage', 0.0)
|
||||
if shrinkage > 0 and kernel in {'aitchison', 'ilr'} and isinstance(bandwidth, Real):
|
||||
return (1.0 - shrinkage) * float(bandwidth)
|
||||
return bandwidth
|
||||
|
||||
def clr_transform(self, X):
|
||||
if not hasattr(self, 'clr'):
|
||||
self.clr = F.CLRtransformation()
|
||||
return self.clr(X)
|
||||
|
||||
def ilr_transform(self, X):
|
||||
if not hasattr(self, 'ilr'):
|
||||
self.ilr = F.ILRtransformation()
|
||||
return self.ilr(X)
|
||||
|
||||
|
||||
class KDEyML(AggregativeSoftQuantifier, KDEBase):
|
||||
"""
|
||||
Kernel Density Estimation model for quantification (KDEy) relying on the Kullback-Leibler divergence (KLD) as
|
||||
the divergence measure to be minimized. This method was first proposed in the paper
|
||||
`Kernel Density Estimation for Multiclass Quantification <https://arxiv.org/abs/2401.00490>`_, in which
|
||||
`Kernel Density Estimation for Multiclass Quantification <https://link.springer.com/article/10.1007/s10994-024-06726-5>`_ (`arXiv <https://arxiv.org/abs/2401.00490>`_), in which
|
||||
the authors show that minimizing the distribution mathing criterion for KLD is akin to performing
|
||||
maximum likelihood (ML).
|
||||
|
||||
|
|
@ -107,17 +157,31 @@ class KDEyML(AggregativeSoftQuantifier, KDEBase):
|
|||
are to be generated in a `k`-fold cross-validation manner (with this integer indicating the value
|
||||
for `k`); or as a tuple (X,y) defining the specific set of data to use for validation.
|
||||
:param bandwidth: float, the bandwidth of the Kernel
|
||||
:param kernel: kernel of KDE, valid ones are in KDEBase.KERNELS
|
||||
:param shrinkage: amount of shrinkage towards the uniform distribution to apply before
|
||||
Aitchison/ILR transformations. Must be in ``[0,1)``.
|
||||
:param random_state: a seed to be set before fitting any base quantifier (default None)
|
||||
"""
|
||||
|
||||
def __init__(self, classifier: BaseEstimator=None, fit_classifier=True, val_split=5, bandwidth=0.1,
|
||||
random_state=None):
|
||||
kernel='gaussian', shrinkage=0.0, random_state=None):
|
||||
super().__init__(classifier, fit_classifier, val_split)
|
||||
self.bandwidth = KDEBase._check_bandwidth(bandwidth)
|
||||
self.bandwidth = KDEBase._check_bandwidth(bandwidth, kernel)
|
||||
self.kernel = self._check_kernel(kernel)
|
||||
assert 0 <= shrinkage < 1, 'shrinkage must be in [0,1)'
|
||||
assert self.kernel != 'gaussian' or shrinkage == 0, \
|
||||
'shrinkage is only supported for Aitchison/ILR kernels'
|
||||
self.shrinkage = float(shrinkage)
|
||||
self.random_state=random_state
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
self.mix_densities = self.get_mixture_components(classif_predictions, labels, self.classes_, self.bandwidth)
|
||||
self.mix_densities = self.get_mixture_components(
|
||||
classif_predictions,
|
||||
labels,
|
||||
self.classes_,
|
||||
self.bandwidth,
|
||||
self.kernel,
|
||||
)
|
||||
return self
|
||||
|
||||
def aggregate(self, posteriors: np.ndarray):
|
||||
|
|
@ -129,14 +193,25 @@ class KDEyML(AggregativeSoftQuantifier, KDEBase):
|
|||
:return: a vector of class prevalence estimates
|
||||
"""
|
||||
with qp.util.temp_seed(self.random_state):
|
||||
epsilon = 1e-10
|
||||
epsilon = 1e-12
|
||||
n_classes = len(self.mix_densities)
|
||||
test_densities = [self.pdf(kde_i, posteriors) for kde_i in self.mix_densities]
|
||||
if (self.kernel != 'gaussian' and n_classes >= 20) or n_classes >= 30:
|
||||
test_log_densities = [
|
||||
self.pdf(kde_i, posteriors, self.kernel, log_densities=True)
|
||||
for kde_i in self.mix_densities
|
||||
]
|
||||
|
||||
def neg_loglikelihood(prev):
|
||||
test_mixture_likelihood = sum(prev_i * dens_i for prev_i, dens_i in zip (prev, test_densities))
|
||||
test_loglikelihood = np.log(test_mixture_likelihood + epsilon)
|
||||
return -np.sum(test_loglikelihood)
|
||||
def neg_loglikelihood(prev):
|
||||
prev = qp.error.smooth(prev, eps=epsilon)
|
||||
test_loglikelihood = logsumexp(np.log(prev)[:, None] + test_log_densities, axis=0)
|
||||
return -np.sum(test_loglikelihood)
|
||||
else:
|
||||
test_densities = [self.pdf(kde_i, posteriors, self.kernel) for kde_i in self.mix_densities]
|
||||
|
||||
def neg_loglikelihood(prev):
|
||||
test_mixture_likelihood = prev @ test_densities
|
||||
test_loglikelihood = np.log(test_mixture_likelihood + epsilon)
|
||||
return -np.sum(test_loglikelihood)
|
||||
|
||||
return F.optim_minimize(neg_loglikelihood, n_classes)
|
||||
|
||||
|
|
@ -145,7 +220,7 @@ class KDEyHD(AggregativeSoftQuantifier, KDEBase):
|
|||
"""
|
||||
Kernel Density Estimation model for quantification (KDEy) relying on the squared Hellinger Disntace (HD) as
|
||||
the divergence measure to be minimized. This method was first proposed in the paper
|
||||
`Kernel Density Estimation for Multiclass Quantification <https://arxiv.org/abs/2401.00490>`_, in which
|
||||
`Kernel Density Estimation for Multiclass Quantification <https://link.springer.com/article/10.1007/s10994-024-06726-5>`_ (`arXiv <https://arxiv.org/abs/2401.00490>`_), in which
|
||||
the authors proposed a Monte Carlo approach for minimizing the divergence.
|
||||
|
||||
The distribution matching optimization problem comes down to solving:
|
||||
|
|
@ -191,18 +266,22 @@ class KDEyHD(AggregativeSoftQuantifier, KDEBase):
|
|||
|
||||
super().__init__(classifier, fit_classifier, val_split)
|
||||
self.divergence = divergence
|
||||
self.bandwidth = KDEBase._check_bandwidth(bandwidth)
|
||||
self.bandwidth = KDEBase._check_bandwidth(bandwidth, kernel='gaussian')
|
||||
self.random_state=random_state
|
||||
self.montecarlo_trials = montecarlo_trials
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
self.mix_densities = self.get_mixture_components(classif_predictions, labels, self.classes_, self.bandwidth)
|
||||
self.mix_densities = self.get_mixture_components(
|
||||
classif_predictions, labels, self.classes_, self.bandwidth, 'gaussian'
|
||||
)
|
||||
|
||||
N = self.montecarlo_trials
|
||||
rs = self.random_state
|
||||
n = len(self.classes_)
|
||||
self.reference_samples = np.vstack([kde_i.sample(N//n, random_state=rs) for kde_i in self.mix_densities])
|
||||
self.reference_classwise_densities = np.asarray([self.pdf(kde_j, self.reference_samples) for kde_j in self.mix_densities])
|
||||
self.reference_classwise_densities = np.asarray(
|
||||
[self.pdf(kde_j, self.reference_samples, 'gaussian') for kde_j in self.mix_densities]
|
||||
)
|
||||
self.reference_density = np.mean(self.reference_classwise_densities, axis=0) # equiv. to (uniform @ self.reference_classwise_densities)
|
||||
|
||||
return self
|
||||
|
|
@ -212,8 +291,8 @@ class KDEyHD(AggregativeSoftQuantifier, KDEBase):
|
|||
# apply importance sampling (IS). In this version we compute D(p_alpha||q) with IS
|
||||
n_classes = len(self.mix_densities)
|
||||
|
||||
test_kde = self.get_kde_function(posteriors, self.bandwidth)
|
||||
test_densities = self.pdf(test_kde, self.reference_samples)
|
||||
test_kde = self.get_kde_function(posteriors, self.bandwidth, 'gaussian')
|
||||
test_densities = self.pdf(test_kde, self.reference_samples, 'gaussian')
|
||||
|
||||
def f_squared_hellinger(u):
|
||||
return (np.sqrt(u)-1)**2
|
||||
|
|
@ -243,7 +322,7 @@ class KDEyCS(AggregativeSoftQuantifier):
|
|||
"""
|
||||
Kernel Density Estimation model for quantification (KDEy) relying on the Cauchy-Schwarz divergence (CS) as
|
||||
the divergence measure to be minimized. This method was first proposed in the paper
|
||||
`Kernel Density Estimation for Multiclass Quantification <https://arxiv.org/abs/2401.00490>`_, in which
|
||||
`Kernel Density Estimation for Multiclass Quantification <https://link.springer.com/article/10.1007/s10994-024-06726-5>`_ (`arXiv <https://arxiv.org/abs/2401.00490>`_), in which
|
||||
the authors proposed a Monte Carlo approach for minimizing the divergence.
|
||||
|
||||
The distribution matching optimization problem comes down to solving:
|
||||
|
|
@ -278,7 +357,7 @@ class KDEyCS(AggregativeSoftQuantifier):
|
|||
|
||||
def __init__(self, classifier: BaseEstimator=None, fit_classifier=True, val_split=5, bandwidth=0.1):
|
||||
super().__init__(classifier, fit_classifier, val_split)
|
||||
self.bandwidth = KDEBase._check_bandwidth(bandwidth)
|
||||
self.bandwidth = KDEBase._check_bandwidth(bandwidth, kernel='gaussian')
|
||||
|
||||
def gram_matrix_mix_sum(self, X, Y=None):
|
||||
# this adapts the output of the rbf_kernel function (pairwise evaluations of Gaussian kernels k(x,y))
|
||||
|
|
@ -296,13 +375,11 @@ class KDEyCS(AggregativeSoftQuantifier):
|
|||
|
||||
P, y = classif_predictions, labels
|
||||
n = len(self.classes_)
|
||||
|
||||
assert all(sorted(np.unique(y)) == np.arange(n)), \
|
||||
'label name gaps not allowed in current implementation'
|
||||
y = _labels_to_indices(y, self.classes_)
|
||||
|
||||
# counts_inv keeps track of the relative weight of each datapoint within its class
|
||||
# (i.e., the weight in its KDE model)
|
||||
counts_inv = 1 / (F.counts_from_labels(y, classes=self.classes_))
|
||||
counts_inv = 1 / (F.counts_from_labels(y, classes=np.arange(n)))
|
||||
|
||||
# tr_tr_sums corresponds to symbol \overline{B} in the paper
|
||||
tr_tr_sums = np.zeros(shape=(n,n), dtype=float)
|
||||
|
|
@ -353,4 +430,3 @@ class KDEyCS(AggregativeSoftQuantifier):
|
|||
return partA + partB #+ partC
|
||||
|
||||
return F.optim_minimize(divergence, n)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
import random
|
||||
|
|
@ -173,7 +174,7 @@ class QuaNetTrainer(BaseQuantifier):
|
|||
order_by=0 if data.binary else None,
|
||||
**self.quanet_params
|
||||
).to(self.device)
|
||||
print(self.quanet)
|
||||
logging.getLogger(__name__).debug(self.quanet)
|
||||
|
||||
self.optim = torch.optim.Adam(self.quanet.parameters(), lr=self.lr)
|
||||
early_stop = EarlyStop(self.patience, lower_is_better=True)
|
||||
|
|
@ -188,8 +189,9 @@ class QuaNetTrainer(BaseQuantifier):
|
|||
if early_stop.IMPROVED:
|
||||
torch.save(self.quanet.state_dict(), checkpoint)
|
||||
elif early_stop.STOP:
|
||||
print(f'training ended by patience exhausted; loading best model parameters in {checkpoint} '
|
||||
f'for epoch {early_stop.best_epoch}')
|
||||
logging.getLogger(__name__).info(
|
||||
f'training ended by patience exhausted; loading best model parameters in {checkpoint} '
|
||||
f'for epoch {early_stop.best_epoch}')
|
||||
self.quanet.load_state_dict(torch.load(checkpoint))
|
||||
break
|
||||
|
||||
|
|
|
|||
|
|
@ -110,10 +110,10 @@ class ThresholdOptimization(BinaryAggregativeQuantifier):
|
|||
TN = np.logical_and(y == y_, y == self.neg_label).sum()
|
||||
return TP, FP, FN, TN
|
||||
|
||||
def _compute_tpr(self, TP, FP):
|
||||
if TP + FP == 0:
|
||||
def _compute_tpr(self, TP, FN):
|
||||
if TP + FN == 0:
|
||||
return 1
|
||||
return TP / (TP + FP)
|
||||
return TP / (TP + FN)
|
||||
|
||||
def _compute_fpr(self, FP, TN):
|
||||
if FP + TN == 0:
|
||||
|
|
|
|||
|
|
@ -1,10 +1,8 @@
|
|||
import warnings
|
||||
from abc import ABC, abstractmethod
|
||||
from argparse import ArgumentError
|
||||
from copy import deepcopy
|
||||
from typing import Callable, Literal, Union
|
||||
import numpy as np
|
||||
from abstention.calibration import NoBiasVectorScaling, TempScaling, VectorScaling
|
||||
from numpy.f2py.crackfortran import true_intent_list
|
||||
from sklearn.base import BaseEstimator
|
||||
from sklearn.calibration import CalibratedClassifierCV
|
||||
from sklearn.exceptions import NotFittedError
|
||||
|
|
@ -18,11 +16,18 @@ from quapy.functional import get_divergence
|
|||
from quapy.classification.svmperf import SVMperf
|
||||
from quapy.data import LabelledCollection
|
||||
from quapy.method.base import BaseQuantifier, BinaryQuantifier, OneVsAllGeneric
|
||||
from quapy.method import _bayesian
|
||||
from quapy.method._energy import _EnergyDistanceCore
|
||||
from quapy.method._helper import (
|
||||
_get_abstention_calibrators,
|
||||
_get_cvxpy,
|
||||
_rlls_check_mode,
|
||||
_rlls_joint_distribution,
|
||||
_rlls_predicted_marginal,
|
||||
_rlls_compute_3deltaC,
|
||||
_rlls_compute_weights,
|
||||
_labels_to_indices,
|
||||
)
|
||||
|
||||
# import warnings
|
||||
# from sklearn.exceptions import ConvergenceWarning
|
||||
# warnings.filterwarnings("ignore", category=ConvergenceWarning)
|
||||
|
||||
|
||||
# Abstract classes
|
||||
|
|
@ -78,9 +83,9 @@ class AggregativeQuantifier(BaseQuantifier, ABC):
|
|||
(f'when {val_split=} is indicated as an integer, it represents the number of folds in a kFCV '
|
||||
f'and must thus be >1')
|
||||
if val_split==5 and not fit_classifier:
|
||||
print(f'Warning: {val_split=} will be ignored when the classifier is already trained '
|
||||
f'({fit_classifier=}). Parameter {self.val_split=} will be set to None. Set {val_split=} '
|
||||
f'to None to avoid this warning.')
|
||||
warnings.warn(f'{val_split=} will be ignored when the classifier is already trained '
|
||||
f'({fit_classifier=}). Parameter {self.val_split=} will be set to None. Set {val_split=} '
|
||||
f'to None to avoid this warning.')
|
||||
self.val_split=None
|
||||
if val_split!=5:
|
||||
assert fit_classifier, (f'Parameter {val_split=} has been modified, but {fit_classifier=} '
|
||||
|
|
@ -339,8 +344,8 @@ class AggregativeSoftQuantifier(AggregativeQuantifier, ABC):
|
|||
"""
|
||||
if not hasattr(self.classifier, self._classifier_method()):
|
||||
if adapt_if_necessary:
|
||||
print(f'warning: The learner {self.classifier.__class__.__name__} does not seem to be '
|
||||
f'probabilistic. The learner will be calibrated (using CalibratedClassifierCV).')
|
||||
warnings.warn(f'The learner {self.classifier.__class__.__name__} does not seem to be '
|
||||
f'probabilistic. The learner will be calibrated (using CalibratedClassifierCV).')
|
||||
self.classifier = CalibratedClassifierCV(self.classifier, cv=5)
|
||||
else:
|
||||
raise AssertionError(f'error: The learner {self.classifier.__class__.__name__} does not '
|
||||
|
|
@ -367,8 +372,13 @@ class BinaryAggregativeQuantifier(AggregativeQuantifier, BinaryQuantifier):
|
|||
# ------------------------------------
|
||||
class CC(AggregativeCrispQuantifier):
|
||||
"""
|
||||
The most basic Quantification method. One that simply classifies all instances and counts how many have been
|
||||
attributed to each of the classes in order to compute class prevalence estimates.
|
||||
`Classify & Count` (CC), the most basic quantification method, one that
|
||||
simply classifies all instances and counts how many have been attributed to
|
||||
each class in order to compute class prevalence estimates. This baseline is
|
||||
the unadjusted estimator discussed, among others, in
|
||||
`Forman, G. (2008). Quantifying counts and costs via classification.
|
||||
Data Mining and Knowledge Discovery, 17, 164-206
|
||||
<https://link.springer.com/article/10.1007/s10618-008-0097-y>`_.
|
||||
|
||||
:param classifier: a sklearn's Estimator that generates a classifier
|
||||
"""
|
||||
|
|
@ -396,14 +406,19 @@ class CC(AggregativeCrispQuantifier):
|
|||
|
||||
class PCC(AggregativeSoftQuantifier):
|
||||
"""
|
||||
`Probabilistic Classify & Count <https://ieeexplore.ieee.org/abstract/document/5694031>`_,
|
||||
the probabilistic variant of CC that relies on the posterior probabilities returned by a probabilistic classifier.
|
||||
`Probabilistic Classify & Count` (PCC), the probabilistic variant of CC
|
||||
that relies on the posterior probabilities returned by a probabilistic
|
||||
classifier, introduced in
|
||||
`Bella, A., Ferri, C., Hernández-Orallo, J., and Ramírez-Quintana, M.J.
|
||||
(2010). Quantification via probability estimators. In Proceedings of the
|
||||
2010 IEEE International Conference on Data Mining (ICDM 2010)
|
||||
<https://ieeexplore.ieee.org/abstract/document/5694031>`_.
|
||||
|
||||
:param classifier: a sklearn's Estimator that generates a classifier
|
||||
"""
|
||||
|
||||
def __init__(self, classifier: BaseEstimator = None, fit_classifier: bool = True):
|
||||
super().__init__(classifier, fit_classifier, val_split=None)
|
||||
def __init__(self, classifier: BaseEstimator = None, fit_classifier: bool = True, val_split=None):
|
||||
super().__init__(classifier, fit_classifier, val_split=val_split)
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
"""
|
||||
|
|
@ -420,9 +435,12 @@ class PCC(AggregativeSoftQuantifier):
|
|||
|
||||
class ACC(AggregativeCrispQuantifier):
|
||||
"""
|
||||
`Adjusted Classify & Count <https://link.springer.com/article/10.1007/s10618-008-0097-y>`_,
|
||||
the "adjusted" variant of :class:`CC`, that corrects the predictions of CC
|
||||
according to the `misclassification rates`.
|
||||
`Adjusted Classify & Count` (ACC), the "adjusted" variant of :class:`CC`
|
||||
that corrects the predictions of CC according to the
|
||||
misclassification rates, originally proposed in
|
||||
`Forman, G. (2008). Quantifying counts and costs via classification.
|
||||
Data Mining and Knowledge Discovery, 17, 164-206
|
||||
<https://link.springer.com/article/10.1007/s10618-008-0097-y>`_.
|
||||
|
||||
:param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be
|
||||
the one indicated in `qp.environ['DEFAULT_CLS']`
|
||||
|
|
@ -564,8 +582,13 @@ class ACC(AggregativeCrispQuantifier):
|
|||
|
||||
class PACC(AggregativeSoftQuantifier):
|
||||
"""
|
||||
`Probabilistic Adjusted Classify & Count <https://ieeexplore.ieee.org/abstract/document/5694031>`_,
|
||||
the probabilistic variant of ACC that relies on the posterior probabilities returned by a probabilistic classifier.
|
||||
`Probabilistic Adjusted Classify & Count` (PACC), the probabilistic
|
||||
variant of ACC that relies on the posterior probabilities returned by a
|
||||
probabilistic classifier, introduced in
|
||||
`Bella, A., Ferri, C., Hernández-Orallo, J., and Ramírez-Quintana, M.J.
|
||||
(2010). Quantification via probability estimators. In Proceedings of the
|
||||
2010 IEEE International Conference on Data Mining (ICDM 2010)
|
||||
<https://ieeexplore.ieee.org/abstract/document/5694031>`_.
|
||||
|
||||
:param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be
|
||||
the one indicated in `qp.environ['DEFAULT_CLS']`
|
||||
|
|
@ -670,6 +693,108 @@ class PACC(AggregativeSoftQuantifier):
|
|||
return confusion.T
|
||||
|
||||
|
||||
class RLLS(AggregativeSoftQuantifier):
|
||||
"""
|
||||
`Regularized Learning for Domain Adaptation under Label Shifts
|
||||
<https://arxiv.org/abs/1903.09734>`_, used here as an aggregative
|
||||
quantifier.
|
||||
|
||||
This implementation ports the regularized weight-estimation component of
|
||||
RLLS to QuaPy's aggregative interface. It estimates label-shift ratios from
|
||||
validation posteriors and source labels, then rescales the source
|
||||
prevalence to obtain target prevalence estimates.
|
||||
|
||||
This method relies on the optional `cvxpy` dependency.
|
||||
|
||||
:param classifier: a scikit-learn's BaseEstimator, or None, in which case
|
||||
the classifier is taken to be the one indicated in
|
||||
`qp.environ['DEFAULT_CLS']`
|
||||
:param fit_classifier: whether to train the learner (default is True). Set
|
||||
to False if the learner has been trained outside the quantifier.
|
||||
:param val_split: specifies the data used for generating classifier
|
||||
predictions. This specification can be made as float in (0, 1)
|
||||
indicating the proportion of stratified held-out validation set to be
|
||||
extracted from the training set; or as an integer (default 5),
|
||||
indicating that the predictions are to be generated in a `k`-fold
|
||||
cross-validation manner; or as a tuple `(X, y)` defining the specific
|
||||
set of data to use for validation. This method requires source
|
||||
predictions and therefore needs `val_split` whenever
|
||||
`fit_classifier=True`.
|
||||
:param mode: whether source- and target-domain quantities are estimated
|
||||
from posterior probabilities (`soft`, default) or from argmax
|
||||
predictions (`hard`)
|
||||
:param alpha: multiplicative factor for the regularization level (default
|
||||
0.01)
|
||||
:param delta: confidence parameter used in the finite-sample regularizer
|
||||
(default 0.05)
|
||||
:param clip_weights: if True, clips negative importance weights to zero
|
||||
before converting them into prevalence estimates
|
||||
:param norm: the normalization method passed to
|
||||
:func:`quapy.functional.normalize_prevalence`
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
classifier: BaseEstimator = None,
|
||||
fit_classifier=True,
|
||||
val_split=5,
|
||||
mode: Literal['soft', 'hard'] = 'soft',
|
||||
alpha: float = 0.01,
|
||||
delta: float = 0.05,
|
||||
clip_weights: bool = True,
|
||||
norm: Literal['clip', 'mapsimplex', 'condsoftmax'] = 'clip',
|
||||
):
|
||||
super().__init__(classifier, fit_classifier, val_split)
|
||||
self.mode = mode
|
||||
self.alpha = alpha
|
||||
self.delta = delta
|
||||
self.clip_weights = clip_weights
|
||||
self.norm = norm
|
||||
self.last_w_ = None
|
||||
|
||||
def _check_init_parameters(self):
|
||||
_get_cvxpy()
|
||||
_rlls_check_mode(self.mode)
|
||||
if not isinstance(self.alpha, (int, float)) or self.alpha < 0:
|
||||
raise ValueError(f'expected a non-negative real value for alpha; found {self.alpha!r}')
|
||||
if not isinstance(self.delta, (int, float)) or not (0 < self.delta < 1):
|
||||
raise ValueError(f'expected delta to be in (0,1); found {self.delta!r}')
|
||||
if self.norm not in ACC.NORMALIZATIONS:
|
||||
raise ValueError(f"unknown normalization; valid ones are {ACC.NORMALIZATIONS}")
|
||||
if self.fit_classifier and self.val_split is None:
|
||||
raise ValueError(
|
||||
'RLLS requires validation predictions for aggregation_fit; '
|
||||
'please set val_split to an integer, float, or validation tuple'
|
||||
)
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
if classif_predictions is None or labels is None:
|
||||
raise ValueError('RLLS requires source posterior predictions and source labels')
|
||||
|
||||
self.train_prevalence_ = F.prevalence_from_labels(labels, classes=self.classes_)
|
||||
self.C_zy_ = _rlls_joint_distribution(
|
||||
classif_predictions,
|
||||
labels,
|
||||
self.classes_,
|
||||
mode=self.mode,
|
||||
)
|
||||
self.pz_ = _rlls_predicted_marginal(classif_predictions, mode=self.mode)
|
||||
self.rho_ = _rlls_compute_3deltaC(len(self.classes_), len(labels), self.delta)
|
||||
|
||||
def aggregate(self, classif_posteriors):
|
||||
qz = _rlls_predicted_marginal(classif_posteriors, mode=self.mode)
|
||||
w = _rlls_compute_weights(
|
||||
self.C_zy_,
|
||||
qz,
|
||||
self.pz_,
|
||||
rho=self.alpha * self.rho_,
|
||||
clip=self.clip_weights,
|
||||
)
|
||||
self.last_w_ = w
|
||||
estimate = self.train_prevalence_ * w
|
||||
return F.normalize_prevalence(estimate, method=self.norm)
|
||||
|
||||
|
||||
class EMQ(AggregativeSoftQuantifier):
|
||||
"""
|
||||
`Expectation Maximization for Quantification <https://ieeexplore.ieee.org/abstract/document/6789744>`_ (EMQ),
|
||||
|
|
@ -732,7 +857,7 @@ class EMQ(AggregativeSoftQuantifier):
|
|||
self.exact_train_prev = exact_train_prev
|
||||
self.calib = calib
|
||||
self.on_calib_error = on_calib_error
|
||||
self.n_jobs = n_jobs
|
||||
self.n_jobs = qp._get_njobs(n_jobs)
|
||||
|
||||
@classmethod
|
||||
def EMQ_BCTS(cls, classifier: BaseEstimator, fit_classifier=True, val_split=5, on_calib_error="raise", n_jobs=None):
|
||||
|
|
@ -769,15 +894,15 @@ class EMQ(AggregativeSoftQuantifier):
|
|||
def _check_init_parameters(self):
|
||||
if self.val_split is not None:
|
||||
if self.exact_train_prev and self.calib is None:
|
||||
raise RuntimeWarning(f'The parameter {self.val_split=} was specified for EMQ, while the parameters '
|
||||
f'{self.exact_train_prev=} and {self.calib=}. This has no effect and causes an '
|
||||
f'unnecessary overload.')
|
||||
warnings.warn(f'The parameter {self.val_split=} was specified for EMQ, while the parameters '
|
||||
f'{self.exact_train_prev=} and {self.calib=}. This has no effect and causes an '
|
||||
f'unnecessary overload.', RuntimeWarning)
|
||||
else:
|
||||
if self.calib is not None:
|
||||
print(f'[warning] The parameter {self.calib=} requires the val_split be different from None. '
|
||||
f'This parameter will be set to 5. To avoid this warning, set this value to a float value '
|
||||
f'indicating the proportion of training data to be used as validation, or to an integer '
|
||||
f'indicating the number of folds for kFCV.')
|
||||
warnings.warn(f'The parameter {self.calib=} requires the val_split be different from None. '
|
||||
f'This parameter will be set to 5. To avoid this warning, set this value to a float value '
|
||||
f'indicating the proportion of training data to be used as validation, or to an integer '
|
||||
f'indicating the number of folds for kFCV.')
|
||||
self.val_split = 5
|
||||
|
||||
def classify(self, X):
|
||||
|
|
@ -839,22 +964,17 @@ class EMQ(AggregativeSoftQuantifier):
|
|||
requires_predictions = (self.calib is not None) or (not self.exact_train_prev)
|
||||
if P is None and requires_predictions:
|
||||
# classifier predictions were not generated because val_split=None
|
||||
raise ArgumentError(self.val_split, self.__class__.__name__ +
|
||||
": Classifier predictions for the aggregative fit were not generated because "
|
||||
"val_split=None. This usually happens when you enable calibrations or heuristics "
|
||||
"during model selection but left val_split set to its default value (None). "
|
||||
"Please provide one of the following values for val_split: (i) an integer >1 "
|
||||
"(e.g. val_split=5) for k-fold cross-validation; (ii) a float in (0,1) (e.g. "
|
||||
"val_split=0.3) for a proportion split; or (iii) a tuple (X, y) with explicit "
|
||||
"validation data")
|
||||
raise ValueError(self.__class__.__name__ +
|
||||
": Classifier predictions for the aggregative fit were not generated because "
|
||||
"val_split=None. This usually happens when you enable calibrations or heuristics "
|
||||
"during model selection but left val_split set to its default value (None). "
|
||||
"Please provide one of the following values for val_split: (i) an integer >1 "
|
||||
"(e.g. val_split=5) for k-fold cross-validation; (ii) a float in (0,1) (e.g. "
|
||||
"val_split=0.3) for a proportion split; or (iii) a tuple (X, y) with explicit "
|
||||
"validation data")
|
||||
|
||||
if self.calib is not None:
|
||||
calibrator = {
|
||||
'nbvs': NoBiasVectorScaling(),
|
||||
'bcts': TempScaling(bias_positions='all'),
|
||||
'ts': TempScaling(),
|
||||
'vs': VectorScaling()
|
||||
}.get(self.calib, None)
|
||||
calibrator = _get_abstention_calibrators().get(self.calib, None)
|
||||
|
||||
if calibrator is None:
|
||||
raise ValueError(f'invalid value for {self.calib=}; valid ones are {EMQ.CALIB_OPTIONS}')
|
||||
|
|
@ -922,7 +1042,7 @@ class EMQ(AggregativeSoftQuantifier):
|
|||
s += 1
|
||||
|
||||
if not converged:
|
||||
print('[warning] the method has reached the maximum number of iterations; it might have not converged')
|
||||
warnings.warn('the method has reached the maximum number of iterations; it might have not converged')
|
||||
|
||||
return qs, ps
|
||||
|
||||
|
|
@ -937,6 +1057,10 @@ class HDy(AggregativeSoftQuantifier, BinaryAggregativeQuantifier):
|
|||
class-conditional distributions of the posterior probabilities returned for the positive and negative validation
|
||||
examples, respectively. The parameters of the mixture thus represent the estimates of the class prevalence values.
|
||||
|
||||
This dedicated class is kept for backward compatibility as the historical
|
||||
HDy implementation. The same historical preset is also available as
|
||||
:meth:`DMy.HDy`.
|
||||
|
||||
:param classifier: a scikit-learn's BaseEstimator, or None, in which case the classifier is taken to be
|
||||
the one indicated in `qp.environ['DEFAULT_CLS']`
|
||||
|
||||
|
|
@ -983,9 +1107,6 @@ class HDy(AggregativeSoftQuantifier, BinaryAggregativeQuantifier):
|
|||
Px = classif_posteriors[:, self.pos_label] # takes only the P(y=+1|x)
|
||||
|
||||
prev_estimations = []
|
||||
# for bins in np.linspace(10, 110, 11, dtype=int): #[10, 20, 30, ..., 100, 110]
|
||||
# Pxy0_density, _ = np.histogram(self.Pxy0, bins=bins, range=(0, 1), density=True)
|
||||
# Pxy1_density, _ = np.histogram(self.Pxy1, bins=bins, range=(0, 1), density=True)
|
||||
for bins in self.bins:
|
||||
Pxy0_density = self.Pxy0_density[bins]
|
||||
Pxy1_density = self.Pxy1_density[bins]
|
||||
|
|
@ -995,13 +1116,12 @@ class HDy(AggregativeSoftQuantifier, BinaryAggregativeQuantifier):
|
|||
# the authors proposed to search for the prevalence yielding the best matching as a linear search
|
||||
# at small steps (modern implementations resort to an optimization procedure,
|
||||
# see class DistributionMatching)
|
||||
prev_selected, min_dist = None, None
|
||||
for prev in F.prevalence_linspace(grid_points=101, repeats=1, smooth_limits_epsilon=0.0):
|
||||
Px_train = prev * Pxy1_density + (1 - prev) * Pxy0_density
|
||||
hdy = F.HellingerDistance(Px_train, Px_test)
|
||||
if prev_selected is None or hdy < min_dist:
|
||||
prev_selected, min_dist = prev, hdy
|
||||
prev_estimations.append(prev_selected)
|
||||
def loss(prev):
|
||||
class1_prev = prev[1]
|
||||
Px_train = class1_prev * Pxy1_density + (1 - class1_prev) * Pxy0_density
|
||||
return F.HellingerDistance(Px_train, Px_test)
|
||||
|
||||
prev_estimations.append(F.linear_search(loss, n_classes=2)[1])
|
||||
|
||||
class1_prev = np.median(prev_estimations)
|
||||
return F.as_binary_prevalence(class1_prev)
|
||||
|
|
@ -1042,7 +1162,7 @@ class DyS(AggregativeSoftQuantifier, BinaryAggregativeQuantifier):
|
|||
self.tol = tol
|
||||
self.divergence = divergence
|
||||
self.n_bins = n_bins
|
||||
self.n_jobs = n_jobs
|
||||
self.n_jobs = qp._get_njobs(n_jobs)
|
||||
|
||||
def _ternary_search(self, f, left, right, tol):
|
||||
"""
|
||||
|
|
@ -1160,6 +1280,10 @@ class DMy(AggregativeSoftQuantifier):
|
|||
|
||||
:param cdf: whether to use CDF instead of PDF (default False)
|
||||
|
||||
:param search: string indicating the search strategy used to estimate the prevalence values.
|
||||
Valid options are `optim_minimize` (default, works for binary and multiclass problems),
|
||||
`linear_search` (binary only), and `ternary_search` (binary only)
|
||||
|
||||
:param n_jobs: number of parallel workers (default None)
|
||||
"""
|
||||
|
||||
|
|
@ -1170,15 +1294,36 @@ class DMy(AggregativeSoftQuantifier):
|
|||
self.divergence = divergence
|
||||
self.cdf = cdf
|
||||
self.search = search
|
||||
self.n_jobs = n_jobs
|
||||
self.n_jobs = qp._get_njobs(n_jobs)
|
||||
|
||||
# @classmethod
|
||||
# def HDy(cls, classifier, val_split=5, n_jobs=None):
|
||||
# from quapy.method.meta import MedianEstimator
|
||||
#
|
||||
# hdy = DMy(classifier=classifier, val_split=val_split, search='linear_search', divergence='HD')
|
||||
# hdy = AggregativeMedianEstimator(hdy, param_grid={'nbins': np.linspace(10, 110, 11).astype(int)}, n_jobs=n_jobs)
|
||||
# return hdy
|
||||
@classmethod
|
||||
def HDy(cls, classifier: BaseEstimator = None, fit_classifier=True, val_split=5, n_jobs=None):
|
||||
"""
|
||||
Historical HDy preset expressed as a configuration of :class:`DMy`.
|
||||
|
||||
This preset reproduces the original HDy setup by using Hellinger
|
||||
distance, PDF matching, linear search, and a median sweep over
|
||||
`nbins` in `[10, 20, ..., 110]`.
|
||||
|
||||
:param classifier: a scikit-learn's BaseEstimator, or None
|
||||
:param fit_classifier: whether to train the learner
|
||||
:param val_split: validation specification for generating posteriors
|
||||
:param n_jobs: number of parallel workers
|
||||
:return: an instance of :class:`AggregativeMedianEstimator` configured
|
||||
to reproduce the historical HDy preset
|
||||
"""
|
||||
base = cls(
|
||||
classifier=classifier,
|
||||
fit_classifier=fit_classifier,
|
||||
val_split=val_split,
|
||||
nbins=10,
|
||||
divergence='HD',
|
||||
cdf=False,
|
||||
search='linear_search',
|
||||
n_jobs=n_jobs,
|
||||
)
|
||||
param_grid = {'nbins': np.linspace(10, 110, 11, dtype=int)}
|
||||
return AggregativeMedianEstimator(base_quantifier=base, param_grid=param_grid, n_jobs=n_jobs)
|
||||
|
||||
def _get_distributions(self, posteriors):
|
||||
histograms = []
|
||||
|
|
@ -1210,7 +1355,9 @@ class DMy(AggregativeSoftQuantifier):
|
|||
:param labels: array-like with the true labels associated to each posterior
|
||||
"""
|
||||
posteriors, true_labels = classif_predictions, labels
|
||||
n_classes = len(self.classifier.classes_)
|
||||
classes = self.classifier.classes_
|
||||
n_classes = len(classes)
|
||||
true_labels = _labels_to_indices(true_labels, classes)
|
||||
|
||||
self.validation_distribution = qp.util.parallel(
|
||||
func=self._get_distributions,
|
||||
|
|
@ -1321,9 +1468,9 @@ def newSVMKLD(svmperf_base=None, C=1):
|
|||
return newELM(svmperf_base, loss='kld', C=C)
|
||||
|
||||
|
||||
def newSVMKLD(svmperf_base=None, C=1):
|
||||
def newSVMNKLD(svmperf_base=None, C=1):
|
||||
"""
|
||||
SVM(KLD) is an Explicit Loss Minimization (ELM) quantifier set to optimize for the Kullback-Leibler Divergence
|
||||
SVM(NKLD) is an Explicit Loss Minimization (ELM) quantifier set to optimize for the Kullback-Leibler Divergence
|
||||
normalized via the logistic function, as proposed by
|
||||
`Esuli et al. 2015 <https://dl.acm.org/doi/abs/10.1145/2700406>`_.
|
||||
Equivalent to:
|
||||
|
|
@ -1450,7 +1597,7 @@ class OneVsAllAggregative(OneVsAllGeneric, AggregativeQuantifier):
|
|||
return F.normalize_prevalence(prevalences)
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
self._parallel(self._delayed_binary_aggregate_fit(c, classif_predictions, labels))
|
||||
self._parallel(self._delayed_binary_aggregate_fit, classif_predictions, labels)
|
||||
return self
|
||||
|
||||
def _delayed_binary_classification(self, c, X):
|
||||
|
|
@ -1462,7 +1609,7 @@ class OneVsAllAggregative(OneVsAllGeneric, AggregativeQuantifier):
|
|||
|
||||
def _delayed_binary_aggregate_fit(self, c, classif_predictions, labels):
|
||||
# trains the aggregation function of the cth quantifier
|
||||
return self.dict_binary_quantifiers[c].aggregate_fit(classif_predictions[:, c], labels)
|
||||
return self.dict_binary_quantifiers[c].aggregation_fit(classif_predictions[:, c], labels == c)
|
||||
|
||||
|
||||
class AggregativeMedianEstimator(BinaryQuantifier):
|
||||
|
|
@ -1528,8 +1675,7 @@ class AggregativeMedianEstimator(BinaryQuantifier):
|
|||
((params, X, y) for params in cls_configs),
|
||||
seed=qp.environ.get('_R_SEED', None),
|
||||
n_jobs=self.n_jobs,
|
||||
asarray=False,
|
||||
backend='threading'
|
||||
asarray=False
|
||||
)
|
||||
else:
|
||||
model = self.base_quantifier
|
||||
|
|
@ -1541,8 +1687,7 @@ class AggregativeMedianEstimator(BinaryQuantifier):
|
|||
self._delayed_fit_aggregation,
|
||||
itertools.product(models_preds, q_configs),
|
||||
seed=qp.environ.get('_R_SEED', None),
|
||||
n_jobs=self.n_jobs,
|
||||
backend='threading'
|
||||
n_jobs=self.n_jobs
|
||||
)
|
||||
else:
|
||||
configs = qp.model_selection.expand_grid(self.param_grid)
|
||||
|
|
@ -1550,8 +1695,7 @@ class AggregativeMedianEstimator(BinaryQuantifier):
|
|||
self._delayed_fit,
|
||||
((params, X, y) for params in configs),
|
||||
seed=qp.environ.get('_R_SEED', None),
|
||||
n_jobs=self.n_jobs,
|
||||
backend='threading'
|
||||
n_jobs=self.n_jobs
|
||||
)
|
||||
return self
|
||||
|
||||
|
|
@ -1564,12 +1708,111 @@ class AggregativeMedianEstimator(BinaryQuantifier):
|
|||
self._delayed_predict,
|
||||
((model, instances) for model in self.models),
|
||||
seed=qp.environ.get('_R_SEED', None),
|
||||
n_jobs=self.n_jobs,
|
||||
backend='threading'
|
||||
n_jobs=self.n_jobs
|
||||
)
|
||||
return np.median(prev_preds, axis=0)
|
||||
|
||||
|
||||
class EDy(_EnergyDistanceCore, AggregativeSoftQuantifier):
|
||||
"""
|
||||
Energy Distance y (EDy), a posterior-space distribution-matching quantifier
|
||||
based on energy distance.
|
||||
|
||||
The method represents each class by the posterior-probability vectors
|
||||
produced by a probabilistic classifier on validation data, and estimates the
|
||||
test prevalence vector by matching the test posterior distribution against
|
||||
the class-conditional validation distributions through an energy-distance
|
||||
objective solved as a quadratic program. The method is therefore another
|
||||
instance of the general mixture-matching view of quantification, but it
|
||||
operates directly on posterior vectors rather than on histogram summaries.
|
||||
|
||||
This implementation works for binary and multiclass single-label
|
||||
quantification and relies on the optional ``quadprog`` dependency. It was
|
||||
adapted to QuaPy's current aggregative API from the original implementation
|
||||
available in `quantificationlib <https://github.com/AICGijon/quantificationlib>`_,
|
||||
and now shares its numerical core with the classifier-free
|
||||
:class:`quapy.method.non_aggregative.EDx` variant.
|
||||
|
||||
The current implementation follows the energy-distance formulation discussed
|
||||
in:
|
||||
|
||||
* Alberto Castaño, Laura Morán-Fernández, Jaime Alonso,
|
||||
Verónica Bolón-Canedo, Amparo Alonso-Betanzos, and Juan José del Coz.
|
||||
*An analysis of quantification methods based on matching distributions*.
|
||||
* Hideko Kawakubo, Marthinus Christoffel du Plessis, and Masashi Sugiyama
|
||||
(2016). *Computationally efficient class-prior estimation under class
|
||||
balance change using energy distance*. IEICE Transactions on Information
|
||||
and Systems, 99(1):176-186.
|
||||
|
||||
:param classifier: a scikit-learn ``BaseEstimator``, or ``None`` to use
|
||||
``qp.environ['DEFAULT_CLS']``
|
||||
:param fit_classifier: whether to train the learner (default ``True``).
|
||||
Set to ``False`` if the learner has already been trained outside the
|
||||
quantifier
|
||||
:param val_split: specification of the data used for generating validation
|
||||
posterior probabilities. This can be an integer (default ``5``) for
|
||||
k-fold cross-validation, a float in ``(0, 1)`` for a held-out split,
|
||||
or a tuple ``(X, y)`` with explicit validation data
|
||||
:param distance: distance used to compare posterior vectors. Valid string
|
||||
aliases are ``'manhattan'`` (default) and ``'euclidean'``; a custom
|
||||
callable compatible with pairwise-distance signatures can also be used
|
||||
:param n_jobs: number of parallel workers (default ``None``, meaning the
|
||||
value is taken from the environment)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
classifier: BaseEstimator = None,
|
||||
fit_classifier: bool = True,
|
||||
val_split=5,
|
||||
distance: Union[str, Callable] = 'manhattan',
|
||||
n_jobs=None,
|
||||
):
|
||||
super().__init__(classifier, fit_classifier, val_split)
|
||||
self.distance = distance
|
||||
self.n_jobs = qp._get_njobs(n_jobs)
|
||||
self.train_n_cls_i_ = None
|
||||
self.train_distrib_ = None
|
||||
self.K_ = None
|
||||
self.G_ = None
|
||||
self.C_ = None
|
||||
self.b_ = None
|
||||
self.a_ = None
|
||||
|
||||
def _check_init_parameters(self):
|
||||
self._check_ed_init_parameters()
|
||||
|
||||
def aggregation_fit(self, classif_predictions, labels):
|
||||
"""
|
||||
Estimate the class-conditional posterior distributions on validation
|
||||
data and pre-compute the quadratic-program parameters that depend only
|
||||
on the training side.
|
||||
|
||||
In EDy, the validation posteriors are not discretized into histograms.
|
||||
Instead, each class is represented by the cloud of posterior vectors
|
||||
observed for that class, and these clouds are then compared through the
|
||||
selected pairwise distance.
|
||||
|
||||
:param classif_predictions: posterior probabilities returned by the
|
||||
classifier on validation data
|
||||
:param labels: true labels associated to each posterior vector
|
||||
"""
|
||||
posteriors = np.asarray(classif_predictions, dtype=float)
|
||||
labels = np.asarray(labels)
|
||||
train_distrib = [posteriors[labels == class_] for class_ in self.classes_]
|
||||
return self._fit_energy_model(train_distrib)
|
||||
|
||||
def aggregate(self, posteriors: np.ndarray):
|
||||
"""Estimate the prevalence vector for a test sample.
|
||||
|
||||
:param posteriors: posterior probabilities returned by the classifier
|
||||
for the instances in the test sample
|
||||
:return: a prevalence vector of shape ``(n_classes,)``
|
||||
"""
|
||||
posteriors = np.asarray(posteriors, dtype=float)
|
||||
return self._predict_energy(posteriors)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# imports
|
||||
# ---------------------------------------------------------------
|
||||
|
|
@ -1588,6 +1831,7 @@ KDEyML = _kdey.KDEyML
|
|||
KDEyHD = _kdey.KDEyHD
|
||||
KDEyCS = _kdey.KDEyCS
|
||||
|
||||
|
||||
# ---------------------------------------------------------------
|
||||
# aliases
|
||||
# ---------------------------------------------------------------
|
||||
|
|
@ -1599,6 +1843,8 @@ ProbabilisticAdjustedClassifyAndCount = PACC
|
|||
ExpectationMaximizationQuantifier = EMQ
|
||||
SLD = EMQ
|
||||
DistributionMatchingY = DMy
|
||||
EnergyDistanceY = EDy
|
||||
HellingerDistanceY = HDy
|
||||
HistoricalHDy = DMy.HDy
|
||||
MedianSweep = MS
|
||||
MedianSweep2 = MS2
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import warnings
|
||||
from abc import ABCMeta, abstractmethod
|
||||
from copy import deepcopy
|
||||
|
||||
|
|
@ -84,8 +85,8 @@ class OneVsAllGeneric(OneVsAll, BaseQuantifier):
|
|||
assert isinstance(binary_quantifier, BaseQuantifier), \
|
||||
f'{binary_quantifier} does not seem to be a Quantifier'
|
||||
if isinstance(binary_quantifier, qp.method.aggregative.AggregativeQuantifier):
|
||||
print('[warning] the quantifier seems to be an instance of qp.method.aggregative.AggregativeQuantifier; '
|
||||
f'you might prefer instantiating {qp.method.aggregative.OneVsAllAggregative.__name__}')
|
||||
warnings.warn('the quantifier seems to be an instance of qp.method.aggregative.AggregativeQuantifier; '
|
||||
f'you might prefer instantiating {qp.method.aggregative.OneVsAllAggregative.__name__}')
|
||||
self.binary_quantifier = binary_quantifier
|
||||
self.n_jobs = qp._get_njobs(n_jobs)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import itertools
|
||||
import logging
|
||||
from copy import deepcopy
|
||||
from typing import Union, List
|
||||
import numpy as np
|
||||
|
|
@ -26,65 +27,6 @@ else:
|
|||
QuaNet = "QuaNet is not available due to missing torch package"
|
||||
|
||||
|
||||
class MedianEstimator2(BinaryQuantifier):
|
||||
"""
|
||||
This method is a meta-quantifier that returns, as the estimated class prevalence values, the median of the
|
||||
estimation returned by differently (hyper)parameterized base quantifiers.
|
||||
The median of unit-vectors is only guaranteed to be a unit-vector for n=2 dimensions,
|
||||
i.e., in cases of binary quantification.
|
||||
|
||||
:param base_quantifier: the base, binary quantifier
|
||||
:param random_state: a seed to be set before fitting any base quantifier (default None)
|
||||
:param param_grid: the grid or parameters towards which the median will be computed
|
||||
:param n_jobs: number of parllel workes
|
||||
"""
|
||||
def __init__(self, base_quantifier: BinaryQuantifier, param_grid: dict, random_state=None, n_jobs=None):
|
||||
self.base_quantifier = base_quantifier
|
||||
self.param_grid = param_grid
|
||||
self.random_state = random_state
|
||||
self.n_jobs = qp._get_njobs(n_jobs)
|
||||
|
||||
def get_params(self, deep=True):
|
||||
return self.base_quantifier.get_params(deep)
|
||||
|
||||
def set_params(self, **params):
|
||||
self.base_quantifier.set_params(**params)
|
||||
|
||||
def _delayed_fit(self, args):
|
||||
with qp.util.temp_seed(self.random_state):
|
||||
params, X, y = args
|
||||
model = deepcopy(self.base_quantifier)
|
||||
model.set_params(**params)
|
||||
model.fit(X, y)
|
||||
return model
|
||||
|
||||
def fit(self, X, y):
|
||||
self._check_binary(y, self.__class__.__name__)
|
||||
|
||||
configs = qp.model_selection.expand_grid(self.param_grid)
|
||||
self.models = qp.util.parallel(
|
||||
self._delayed_fit,
|
||||
((params, X, y) for params in configs),
|
||||
seed=qp.environ.get('_R_SEED', None),
|
||||
n_jobs=self.n_jobs
|
||||
)
|
||||
return self
|
||||
|
||||
def _delayed_predict(self, args):
|
||||
model, instances = args
|
||||
return model.predict(instances)
|
||||
|
||||
def predict(self, X):
|
||||
prev_preds = qp.util.parallel(
|
||||
self._delayed_predict,
|
||||
((model, X) for model in self.models),
|
||||
seed=qp.environ.get('_R_SEED', None),
|
||||
n_jobs=self.n_jobs
|
||||
)
|
||||
prev_preds = np.asarray(prev_preds)
|
||||
return np.median(prev_preds, axis=0)
|
||||
|
||||
|
||||
class MedianEstimator(BinaryQuantifier):
|
||||
"""
|
||||
This method is a meta-quantifier that returns, as the estimated class prevalence values, the median of the
|
||||
|
|
@ -213,7 +155,7 @@ class Ensemble(BaseQuantifier):
|
|||
|
||||
def _sout(self, msg):
|
||||
if self.verbose:
|
||||
print('[Ensemble]' + msg)
|
||||
logging.getLogger(__name__).info('[Ensemble] ' + msg)
|
||||
|
||||
def fit(self, X, y):
|
||||
|
||||
|
|
@ -402,7 +344,7 @@ def _select_k(elements, order, k):
|
|||
def _delayed_new_instance(args):
|
||||
base_quantifier, data, val_split, prev, posteriors, keep_samples, verbose, sample_size = args
|
||||
if verbose:
|
||||
print(f'\tfit-start for prev {F.strprev(prev)}, sample_size={sample_size}')
|
||||
logging.getLogger(__name__).info(f'fit-start for prev {F.strprev(prev)}, sample_size={sample_size}')
|
||||
model = deepcopy(base_quantifier)
|
||||
|
||||
if val_split is not None:
|
||||
|
|
@ -422,7 +364,7 @@ def _delayed_new_instance(args):
|
|||
tr_distribution = get_probability_distribution(posteriors[sample_index]) if (posteriors is not None) else None
|
||||
|
||||
if verbose:
|
||||
print(f'\t--fit-ended for prev {F.strprev(prev)}')
|
||||
logging.getLogger(__name__).info(f'fit-ended for prev {F.strprev(prev)}')
|
||||
|
||||
return (model, tr_prevalence, tr_distribution, sample if keep_samples else None)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,11 +1,20 @@
|
|||
from typing import Union, Callable
|
||||
from itertools import product
|
||||
from tqdm import tqdm
|
||||
from typing import Union, Callable, Counter
|
||||
import numpy as np
|
||||
from sklearn.feature_extraction.text import CountVectorizer
|
||||
from sklearn.utils import resample
|
||||
from sklearn.preprocessing import normalize
|
||||
|
||||
from quapy.method.confidence import WithConfidenceABC, ConfidenceRegionABC
|
||||
from quapy.functional import get_divergence
|
||||
from quapy.data import LabelledCollection
|
||||
from quapy.method.base import BaseQuantifier, BinaryQuantifier
|
||||
from quapy.method._helper import _labels_to_indices
|
||||
from quapy.method._energy import _EnergyDistanceCore
|
||||
import quapy.functional as F
|
||||
from scipy.optimize import lsq_linear
|
||||
from scipy import sparse
|
||||
import quapy as qp
|
||||
|
||||
|
||||
class MaximumLikelihoodPrevalenceEstimation(BaseQuantifier):
|
||||
|
|
@ -52,6 +61,9 @@ class DMx(BaseQuantifier):
|
|||
or a callable function taking two ndarrays of the same dimension as input (default "HD", meaning Hellinger
|
||||
Distance)
|
||||
:param cdf: whether to use CDF instead of PDF (default False)
|
||||
:param search: string indicating the search strategy used to estimate the prevalence values.
|
||||
Valid options are `optim_minimize` (default, works for binary and multiclass problems),
|
||||
`linear_search` (binary only), and `ternary_search` (binary only)
|
||||
:param n_jobs: number of parallel workers (default None)
|
||||
"""
|
||||
|
||||
|
|
@ -87,7 +99,6 @@ class DMx(BaseQuantifier):
|
|||
return hdx
|
||||
|
||||
def __get_distributions(self, X):
|
||||
|
||||
histograms = []
|
||||
for feat_idx in range(self.nfeats):
|
||||
feature = X[:, feat_idx]
|
||||
|
|
@ -116,7 +127,9 @@ class DMx(BaseQuantifier):
|
|||
"""
|
||||
self.nfeats = X.shape[1]
|
||||
self.feat_ranges = _get_features_range(X)
|
||||
n_classes = len(np.unique(y))
|
||||
classes = np.unique(y)
|
||||
y = _labels_to_indices(y, classes)
|
||||
n_classes = len(classes)
|
||||
|
||||
self.validation_distribution = np.asarray(
|
||||
[self.__get_distributions(X[y==cat]) for cat in range(n_classes)]
|
||||
|
|
@ -149,53 +162,231 @@ class DMx(BaseQuantifier):
|
|||
return F.argmin_prevalence(loss, n_classes, method=self.search)
|
||||
|
||||
|
||||
# class ReadMe(BaseQuantifier):
|
||||
#
|
||||
# def __init__(self, bootstrap_trials=100, bootstrap_range=100, bagging_trials=100, bagging_range=25, **vectorizer_kwargs):
|
||||
# raise NotImplementedError('under development ...')
|
||||
# self.bootstrap_trials = bootstrap_trials
|
||||
# self.bootstrap_range = bootstrap_range
|
||||
# self.bagging_trials = bagging_trials
|
||||
# self.bagging_range = bagging_range
|
||||
# self.vectorizer_kwargs = vectorizer_kwargs
|
||||
#
|
||||
# def fit(self, data: LabelledCollection):
|
||||
# X, y = data.Xy
|
||||
# self.vectorizer = CountVectorizer(binary=True, **self.vectorizer_kwargs)
|
||||
# X = self.vectorizer.fit_transform(X)
|
||||
# self.class_conditional_X = {i: X[y==i] for i in range(data.classes_)}
|
||||
#
|
||||
# def predict(self, X):
|
||||
# X = self.vectorizer.transform(X)
|
||||
#
|
||||
# # number of features
|
||||
# num_docs, num_feats = X.shape
|
||||
#
|
||||
# # bootstrap
|
||||
# p_boots = []
|
||||
# for _ in range(self.bootstrap_trials):
|
||||
# docs_idx = np.random.choice(num_docs, size=self.bootstra_range, replace=False)
|
||||
# class_conditional_X = {i: X[docs_idx] for i, X in self.class_conditional_X.items()}
|
||||
# Xboot = X[docs_idx]
|
||||
#
|
||||
# # bagging
|
||||
# p_bags = []
|
||||
# for _ in range(self.bagging_trials):
|
||||
# feat_idx = np.random.choice(num_feats, size=self.bagging_range, replace=False)
|
||||
# class_conditional_Xbag = {i: X[:, feat_idx] for i, X in class_conditional_X.items()}
|
||||
# Xbag = Xboot[:,feat_idx]
|
||||
# p = self.std_constrained_linear_ls(Xbag, class_conditional_Xbag)
|
||||
# p_bags.append(p)
|
||||
# p_boots.append(np.mean(p_bags, axis=0))
|
||||
#
|
||||
# p_mean = np.mean(p_boots, axis=0)
|
||||
# p_std = np.std(p_bags, axis=0)
|
||||
#
|
||||
# return p_mean
|
||||
#
|
||||
#
|
||||
# def std_constrained_linear_ls(self, X, class_cond_X: dict):
|
||||
# pass
|
||||
class EDx(_EnergyDistanceCore, BaseQuantifier):
|
||||
"""
|
||||
Energy Distance x (EDx), a covariate-space distribution-matching
|
||||
quantifier based on energy distance.
|
||||
|
||||
EDx is the classifier-free counterpart of :class:`quapy.method.aggregative.EDy`.
|
||||
Instead of representing each class through posterior-probability vectors, it
|
||||
represents each class by the cloud of raw feature vectors observed in the
|
||||
training set and estimates the test prevalence vector by solving the same
|
||||
energy-distance quadratic program directly in feature space.
|
||||
|
||||
This implementation works for binary and multiclass single-label
|
||||
quantification and relies on the optional ``quadprog`` dependency. The
|
||||
current QuaPy adaptation shares its numerical core with EDy and keeps
|
||||
credit to the original implementation available in
|
||||
`quantificationlib <https://github.com/AICGijon/quantificationlib>`_.
|
||||
|
||||
The formulation follows the same references as EDy, namely:
|
||||
|
||||
* Alberto Castaño, Laura Morán-Fernández, Jaime Alonso,
|
||||
Verónica Bolón-Canedo, Amparo Alonso-Betanzos, and Juan José del Coz.
|
||||
*An analysis of quantification methods based on matching distributions*.
|
||||
* Hideko Kawakubo, Marthinus Christoffel du Plessis, and Masashi Sugiyama
|
||||
(2016). *Computationally efficient class-prior estimation under class
|
||||
balance change using energy distance*. IEICE Transactions on Information
|
||||
and Systems, 99(1):176-186.
|
||||
|
||||
:param distance: distance used to compare feature vectors. Valid string
|
||||
aliases are ``'manhattan'`` (default) and ``'euclidean'``; a custom
|
||||
callable compatible with pairwise-distance signatures can also be used
|
||||
:param n_jobs: number of parallel workers (default ``None``, meaning the
|
||||
value is taken from the environment)
|
||||
"""
|
||||
|
||||
def __init__(self, distance: Union[str, Callable] = 'manhattan', n_jobs=None):
|
||||
self.distance = distance
|
||||
self.n_jobs = qp._get_njobs(n_jobs)
|
||||
self.classes_ = None
|
||||
self.n_features_in_ = None
|
||||
self.train_distrib_ = None
|
||||
self.train_n_cls_i_ = None
|
||||
self.K_ = None
|
||||
self.G_ = None
|
||||
self.C_ = None
|
||||
self.b_ = None
|
||||
self.a_ = None
|
||||
|
||||
def fit(self, X, y):
|
||||
"""Fit class-conditional feature-space distributions from training data."""
|
||||
self._check_ed_init_parameters()
|
||||
labels = np.asarray(y)
|
||||
self.classes_ = np.unique(labels)
|
||||
self.n_features_in_ = X.shape[1]
|
||||
train_distrib = [X[labels == class_] for class_ in self.classes_]
|
||||
return self._fit_energy_model(train_distrib)
|
||||
|
||||
def predict(self, X):
|
||||
"""Estimate class prevalences for a test sample of raw instances."""
|
||||
assert X.shape[1] == self.n_features_in_, (
|
||||
f'wrong shape; expected {self.n_features_in_}, found {X.shape[1]}'
|
||||
)
|
||||
return self._predict_energy(X)
|
||||
|
||||
|
||||
class ReadMe(BaseQuantifier, WithConfidenceABC):
|
||||
"""
|
||||
ReadMe is a non-aggregative quantification system proposed by
|
||||
`Daniel Hopkins and Gary King, 2007. A method of automated nonparametric content analysis for
|
||||
social science. American Journal of Political Science, 54(1):229–247.
|
||||
<https://onlinelibrary.wiley.com/doi/abs/10.1111/j.1540-5907.2009.00428.x>`_.
|
||||
The idea is to estimate `Q(Y=i)` directly from:
|
||||
|
||||
:math:`Q(X)=\\sum_{i=1} Q(X|Y=i) Q(Y=i)`
|
||||
|
||||
via least-squares regression, i.e., without incurring the cost of computing posterior probabilities.
|
||||
However, this poses a very difficult representation in which the vector `Q(X)` and the matrix `Q(X|Y=i)`
|
||||
can be of very high dimensions. In order to render the problem tracktable, ReadMe performs bagging in
|
||||
the feature space. ReadMe also combines bagging with bootstrap in order to derive confidence intervals
|
||||
around point estimations.
|
||||
|
||||
We use the same default parameters as in the official
|
||||
`R implementation <https://github.com/iqss-research/ReadMeV1/blob/master/R/prototype.R>`_.
|
||||
|
||||
:param prob_model: str ('naive', or 'full'), selects the modality in which the probabilities `Q(X)` and
|
||||
`Q(X|Y)` are to be modelled. Options include "full", which corresponds to the original formulation of
|
||||
ReadMe, in which X is constrained to be a binary matrix (e.g., of term presence/absence) and in which
|
||||
`Q(X)` and `Q(X|Y)` are modelled, respectively, as matrices of `(2^K, 1)` and `(2^K, n)` values, where
|
||||
`K` is the number of columns in the data matrix (i.e., `bagging_range`), and `n` is the number of classes.
|
||||
Of course, this approach is computationally prohibited for large `K`, so the authors advised against computing it
|
||||
for matrices with `K>25` (although we recommend even smaller values of `K`). A much faster model is "naive", which
|
||||
considers the `Q(X)` and `Q(X|Y)` be multinomial distributions under the `bag-of-words` perspective. In this
|
||||
case, `bagging_range` can be set to much larger values. Default is "full" (i.e., original ReadMe behavior).
|
||||
:param bootstrap_trials: int, number of bootstrap trials (default 300)
|
||||
:param bagging_trials: int, number of bagging trials (default 300)
|
||||
:param bagging_range: int, number of features to keep for each bagging trial (default 15)
|
||||
:param confidence_level: float, a value in (0,1) reflecting the desired confidence level (default 0.95)
|
||||
:param region: str in 'intervals', 'ellipse', 'ellipse-clr'; indicates the preferred method for
|
||||
defining the confidence region (see :class:`WithConfidenceABC`)
|
||||
:param bonferroni: bool (default False), whether to apply Bonferroni correction when
|
||||
`region='intervals'`. This parameter has no effect for ellipse-based regions.
|
||||
:param random_state: int or None, allows replicability (default None)
|
||||
:param verbose: bool, whether to display information during the process (default False)
|
||||
"""
|
||||
|
||||
MAX_FEATURES_FOR_EMPIRICAL_ESTIMATION = 25
|
||||
PROBABILISTIC_MODELS = ["naive", "full"]
|
||||
|
||||
def __init__(self,
|
||||
prob_model="full",
|
||||
bootstrap_trials=300,
|
||||
bagging_trials=300,
|
||||
bagging_range=15,
|
||||
confidence_level=0.95,
|
||||
region='intervals',
|
||||
bonferroni=False,
|
||||
random_state=None,
|
||||
verbose=False):
|
||||
assert prob_model in ReadMe.PROBABILISTIC_MODELS, \
|
||||
f'unknown {prob_model=}, valid ones are {ReadMe.PROBABILISTIC_MODELS=}'
|
||||
self.prob_model = prob_model
|
||||
self.bootstrap_trials = bootstrap_trials
|
||||
self.bagging_trials = bagging_trials
|
||||
self.bagging_range = bagging_range
|
||||
self.confidence_level = confidence_level
|
||||
self.region = region
|
||||
self.bonferroni = bonferroni
|
||||
self.random_state = random_state
|
||||
self.verbose = verbose
|
||||
|
||||
def fit(self, X, y):
|
||||
self._check_matrix(X)
|
||||
|
||||
self.rng = np.random.default_rng(self.random_state)
|
||||
self.classes_ = np.unique(y)
|
||||
|
||||
Xsize = X.shape[0]
|
||||
|
||||
# Bootstrap loop
|
||||
self.Xboots, self.yboots = [], []
|
||||
for _ in range(self.bootstrap_trials):
|
||||
idx = self.rng.choice(Xsize, size=Xsize, replace=True)
|
||||
self.Xboots.append(X[idx])
|
||||
self.yboots.append(y[idx])
|
||||
|
||||
return self
|
||||
|
||||
def predict_conf(self, X, confidence_level=None) -> (np.ndarray, ConfidenceRegionABC):
|
||||
self._check_matrix(X)
|
||||
if confidence_level is None:
|
||||
confidence_level = self.confidence_level
|
||||
|
||||
n_features = X.shape[1]
|
||||
boots_prevalences = []
|
||||
for Xboots, yboots in tqdm(
|
||||
zip(self.Xboots, self.yboots),
|
||||
desc='bootstrap predictions', total=self.bootstrap_trials, disable=not self.verbose
|
||||
):
|
||||
bagging_estimates = []
|
||||
for _ in range(self.bagging_trials):
|
||||
feat_idx = self.rng.choice(n_features, size=self.bagging_range, replace=False)
|
||||
Xboots_bagging = Xboots[:, feat_idx]
|
||||
X_boots_bagging = X[:, feat_idx]
|
||||
bagging_prev = self._quantify_iteration(Xboots_bagging, yboots, X_boots_bagging)
|
||||
bagging_estimates.append(bagging_prev)
|
||||
|
||||
boots_prevalences.append(np.mean(bagging_estimates, axis=0))
|
||||
|
||||
conf = WithConfidenceABC.construct_region(boots_prevalences, confidence_level, method=self.region, bonferroni=self.bonferroni)
|
||||
prev_estim = conf.point_estimate()
|
||||
|
||||
return prev_estim, conf
|
||||
|
||||
def predict(self, X):
|
||||
prev_estim, _ = self.predict_conf(X)
|
||||
return prev_estim
|
||||
|
||||
def _quantify_iteration(self, Xtr, ytr, Xte):
|
||||
"""Single ReadMe estimate."""
|
||||
PX_given_Y = np.asarray([self._compute_P(Xtr[ytr == c]) for i,c in enumerate(self.classes_)])
|
||||
PX = self._compute_P(Xte)
|
||||
|
||||
res = lsq_linear(A=PX_given_Y.T, b=PX, bounds=(0, 1))
|
||||
pY = np.maximum(res.x, 0)
|
||||
return pY / pY.sum()
|
||||
|
||||
def _check_matrix(self, X):
|
||||
"""the "full" model requires estimating empirical distributions; due to the high computational cost,
|
||||
this function is only made available for binary matrices"""
|
||||
if self.prob_model == 'full' and not self._is_binary_matrix(X):
|
||||
raise ValueError('the empirical distribution can only be computed efficiently on binary matrices')
|
||||
|
||||
def _is_binary_matrix(self, X):
|
||||
data = X.data if sparse.issparse(X) else X
|
||||
return np.all((data == 0) | (data == 1))
|
||||
|
||||
def _compute_P(self, X):
|
||||
if self.prob_model == 'naive':
|
||||
return self._multinomial_distribution(X)
|
||||
elif self.prob_model == 'full':
|
||||
return self._empirical_distribution(X)
|
||||
else:
|
||||
raise ValueError(f'unknown {self.prob_model}; valid ones are {ReadMe.PROBABILISTIC_MODELS=}')
|
||||
|
||||
def _empirical_distribution(self, X):
|
||||
|
||||
if X.shape[1] > self.MAX_FEATURES_FOR_EMPIRICAL_ESTIMATION:
|
||||
raise ValueError(f'the empirical distribution can only be computed efficiently for dimensions '
|
||||
f'less or equal than {self.MAX_FEATURES_FOR_EMPIRICAL_ESTIMATION}')
|
||||
|
||||
# we first convert every binary row (e.g., 0 0 1 0 1) into the equivalent number (e.g., 5);
|
||||
# this will speed up subsequent comparisons a lot
|
||||
K = X.shape[1]
|
||||
binary_powers = 1 << np.arange(K-1, -1, -1) # (2^K, ..., 32, 16, 8, 4, 2, 1)
|
||||
X_as_binary_numbers = X @ binary_powers # e.g., [0 0 1 0 1] @ [16, 8, 4, 2, 1] = 5
|
||||
|
||||
# count occurrences and compute probs
|
||||
counts = np.bincount(X_as_binary_numbers, minlength=2 ** K).astype(float)
|
||||
probs = counts / counts.sum()
|
||||
|
||||
return probs
|
||||
|
||||
def _multinomial_distribution(self, X):
|
||||
PX = np.asarray(X.sum(axis=0))
|
||||
PX = normalize(PX, norm='l1', axis=1)
|
||||
return PX.ravel()
|
||||
|
||||
|
||||
def _get_features_range(X):
|
||||
|
|
@ -211,4 +402,8 @@ def _get_features_range(X):
|
|||
# aliases
|
||||
#---------------------------------------------------------------
|
||||
|
||||
DistributionMatchingX = DMx
|
||||
|
||||
HDx = DMx.HDx
|
||||
DistributionMatchingX = DMx
|
||||
EnergyDistanceX = EDx
|
||||
HellingerDistanceX = HDx
|
||||
|
|
@ -0,0 +1,39 @@
|
|||
data {
|
||||
int<lower=0> n_bucket;
|
||||
array[n_bucket] int<lower=0> train_pos;
|
||||
array[n_bucket] int<lower=0> train_neg;
|
||||
array[n_bucket] int<lower=0> test;
|
||||
int<lower=0,upper=1> posterior;
|
||||
}
|
||||
|
||||
transformed data{
|
||||
row_vector<lower=0>[n_bucket] train_pos_rv;
|
||||
row_vector<lower=0>[n_bucket] train_neg_rv;
|
||||
row_vector<lower=0>[n_bucket] test_rv;
|
||||
real n_test;
|
||||
|
||||
train_pos_rv = to_row_vector( train_pos );
|
||||
train_neg_rv = to_row_vector( train_neg );
|
||||
test_rv = to_row_vector( test );
|
||||
n_test = sum( test );
|
||||
}
|
||||
|
||||
parameters {
|
||||
simplex[n_bucket] p_neg;
|
||||
simplex[n_bucket] p_pos;
|
||||
real<lower=0,upper=1> prev_prior;
|
||||
}
|
||||
|
||||
model {
|
||||
if( posterior ) {
|
||||
target += train_neg_rv * log( p_neg );
|
||||
target += train_pos_rv * log( p_pos );
|
||||
target += test_rv * log( p_neg * ( 1 - prev_prior) + p_pos * prev_prior );
|
||||
}
|
||||
}
|
||||
|
||||
generated quantities {
|
||||
real<lower=0,upper=1> prev;
|
||||
prev = sum( binomial_rng(test, 1 / ( 1 + (p_neg./p_pos) *(1-prev_prior)/prev_prior ) ) ) / n_test;
|
||||
}
|
||||
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
import itertools
|
||||
import logging
|
||||
import signal
|
||||
from copy import deepcopy
|
||||
from enum import Enum
|
||||
|
|
@ -91,7 +92,7 @@ class GridSearchQ(BaseQuantifier):
|
|||
|
||||
def _sout(self, msg):
|
||||
if self.verbose:
|
||||
print(f'[{self.__class__.__name__}:{self.model.__class__.__name__}]: {msg}')
|
||||
logging.getLogger(__name__).info(f'[{self.__class__.__name__}:{self.model.__class__.__name__}]: {msg}')
|
||||
|
||||
def __check_error_measure(self, error):
|
||||
if error in qp.error.QUANTIFICATION_ERROR:
|
||||
|
|
|
|||
222
quapy/plot.py
|
|
@ -1,12 +1,15 @@
|
|||
from collections import defaultdict
|
||||
import matplotlib.pyplot as plt
|
||||
from matplotlib.pyplot import get_cmap
|
||||
import numpy as np
|
||||
from matplotlib import cm
|
||||
from scipy.stats import ttest_ind_from_stats
|
||||
from matplotlib.ticker import ScalarFormatter
|
||||
import math
|
||||
|
||||
from matplotlib import cm
|
||||
import matplotlib.colors as mcolors
|
||||
import matplotlib.pyplot as plt
|
||||
from matplotlib.colors import LinearSegmentedColormap, ListedColormap
|
||||
from matplotlib.pyplot import get_cmap
|
||||
from matplotlib.ticker import ScalarFormatter
|
||||
import numpy as np
|
||||
from scipy.stats import ttest_ind_from_stats
|
||||
|
||||
import quapy as qp
|
||||
|
||||
plt.rcParams['figure.figsize'] = [10, 6]
|
||||
|
|
@ -480,7 +483,6 @@ def brokenbar_supremacy_by_drift(method_names, true_prevs, estim_prevs, tr_prevs
|
|||
best_bucket_methods.append(method_order[method_index])
|
||||
best_methods.append(best_bucket_methods)
|
||||
salient_methods.update(best_bucket_methods)
|
||||
print(best_bucket_methods)
|
||||
|
||||
if binning=='isomerous':
|
||||
fig, axes = plt.subplots(2, 1, gridspec_kw={'height_ratios': [0.2, 1]}, figsize=(20, len(salient_methods)))
|
||||
|
|
@ -560,6 +562,123 @@ def brokenbar_supremacy_by_drift(method_names, true_prevs, estim_prevs, tr_prevs
|
|||
return fig, ax
|
||||
|
||||
|
||||
def plot_simplex(
|
||||
point_layers=None,
|
||||
region_layers=None,
|
||||
density_function=None,
|
||||
density_color='#1f77b4',
|
||||
density_alpha=1.0,
|
||||
resolution=400,
|
||||
class_names=None,
|
||||
title=None,
|
||||
show_legend=True,
|
||||
legend_loc='lower center',
|
||||
legend_bbox_to_anchor=(0.5, -0.08),
|
||||
legend_ncol=2,
|
||||
figsize=(6.8, 6.2),
|
||||
class_name_fontsize=10,
|
||||
title_fontsize=11,
|
||||
legend_fontsize=9,
|
||||
ax=None,
|
||||
savepath=None):
|
||||
"""
|
||||
Plots data on the ternary simplex for three-class quantification problems.
|
||||
|
||||
This utility is convenient for visualising prevalence vectors, posterior triplets,
|
||||
confidence regions, or any other points that lie on the 2-dimensional probability
|
||||
simplex. The plot can combine three optional layer types:
|
||||
|
||||
* `point_layers`: scatter layers for one or more prevalence clouds or reference points
|
||||
* `region_layers`: shaded regions defined by callables on prevalence vectors
|
||||
* `density_function`: a scalar function evaluated on the simplex and rendered as a heatmap
|
||||
|
||||
Each entry in `point_layers` is a dictionary with a mandatory `points` field
|
||||
containing an array-like of shape `(n_points, 3)` or `(3,)`. Optional fields are
|
||||
`label` for the legend and `style` for matplotlib scatter keyword arguments.
|
||||
|
||||
Each entry in `region_layers` is a dictionary with a mandatory `fn` field containing
|
||||
a callable that receives prevalence vectors and returns region-membership scores.
|
||||
Optional fields are `label`, `color`, and `alpha`.
|
||||
|
||||
:param point_layers: optional list of point-layer dictionaries
|
||||
:param region_layers: optional list of region-layer dictionaries
|
||||
:param density_function: optional callable receiving prevalence vectors and returning
|
||||
scalar values
|
||||
:param density_color: color used for the density heatmap
|
||||
:param density_alpha: opacity for the density heatmap
|
||||
:param resolution: number of grid steps per axis used for rendering regions and densities
|
||||
:param class_names: optional list or tuple with the three class names
|
||||
:param title: optional plot title
|
||||
:param show_legend: whether to display the legend
|
||||
:param legend_loc: location string passed to matplotlib for the legend
|
||||
:param legend_bbox_to_anchor: optional legend anchor box
|
||||
:param legend_ncol: number of legend columns
|
||||
:param figsize: figure size used when `ax` is not provided
|
||||
:param class_name_fontsize: fontsize used for simplex vertex labels
|
||||
:param title_fontsize: fontsize used for the optional title
|
||||
:param legend_fontsize: fontsize used for the legend
|
||||
:param ax: optional matplotlib axes object; if not provided, a new figure is created
|
||||
:param savepath: path where to save the plot; if not indicated, the plot is shown when
|
||||
`ax` is not provided
|
||||
:return: returns `(fig, ax)` matplotlib objects for eventual customisation
|
||||
"""
|
||||
if class_names is None:
|
||||
class_names = ('Y=1', 'Y=2', 'Y=3')
|
||||
if len(class_names) != 3:
|
||||
raise ValueError(f'expected exactly 3 class names, got {len(class_names)}')
|
||||
|
||||
own_figure = ax is None
|
||||
if own_figure:
|
||||
fig, ax = plt.subplots(figsize=figsize)
|
||||
else:
|
||||
fig = ax.figure
|
||||
|
||||
if density_function is not None:
|
||||
_plot_simplex_density(ax, density_function, resolution, density_color, density_alpha)
|
||||
|
||||
if region_layers:
|
||||
_plot_simplex_regions(ax, region_layers, resolution)
|
||||
|
||||
if point_layers:
|
||||
_plot_simplex_points(ax, point_layers)
|
||||
|
||||
simplex_ymax = np.sqrt(3) / 2
|
||||
triangle = np.array([
|
||||
[0.0, 0.0],
|
||||
[1.0, 0.0],
|
||||
[0.5, simplex_ymax],
|
||||
[0.0, 0.0],
|
||||
])
|
||||
ax.plot(triangle[:, 0], triangle[:, 1], color='black')
|
||||
|
||||
ax.text(-0.05, -0.05, class_names[0], ha='right', va='top', fontsize=class_name_fontsize)
|
||||
ax.text(1.05, -0.05, class_names[1], ha='left', va='top', fontsize=class_name_fontsize)
|
||||
ax.text(0.5, simplex_ymax + 0.05, class_names[2], ha='center', va='bottom', fontsize=class_name_fontsize)
|
||||
|
||||
if title is not None:
|
||||
ax.set_title(title, fontsize=title_fontsize)
|
||||
|
||||
ax.set_aspect('equal')
|
||||
ax.set_xlim(-0.1, 1.1)
|
||||
ax.set_ylim(-0.1, simplex_ymax + 0.1)
|
||||
ax.axis('off')
|
||||
|
||||
if show_legend:
|
||||
_, labels = ax.get_legend_handles_labels()
|
||||
if labels:
|
||||
ax.legend(loc=legend_loc, bbox_to_anchor=legend_bbox_to_anchor, ncol=legend_ncol, fontsize=legend_fontsize, frameon=False)
|
||||
|
||||
fig.tight_layout(pad=0.6)
|
||||
|
||||
if savepath is not None:
|
||||
qp.util.create_parent_dir(savepath)
|
||||
fig.savefig(savepath, bbox_inches='tight')
|
||||
elif own_figure:
|
||||
plt.show()
|
||||
|
||||
return fig, ax
|
||||
|
||||
|
||||
def _merge(method_names, true_prevs, estim_prevs):
|
||||
ndims = true_prevs[0].shape[1]
|
||||
data = defaultdict(lambda: {'true': np.empty(shape=(0, ndims)), 'estim': np.empty(shape=(0, ndims))})
|
||||
|
|
@ -612,6 +731,94 @@ def _join_data_by_drift(method_names, true_prevs, estim_prevs, tr_prevs, x_error
|
|||
return data
|
||||
|
||||
|
||||
def _simplex_to_cartesian(prevalences):
|
||||
prevalences = np.asarray(prevalences, dtype=float)
|
||||
prevalences = np.atleast_2d(prevalences)
|
||||
if prevalences.shape[1] != 3:
|
||||
raise ValueError(f'plot_simplex expects prevalence vectors of shape (_, 3); found {prevalences.shape}')
|
||||
x = prevalences[:, 1] + 0.5 * prevalences[:, 2]
|
||||
y = prevalences[:, 2] * (np.sqrt(3) / 2)
|
||||
return x, y
|
||||
|
||||
|
||||
def _barycentric_from_xy(x, y):
|
||||
p3 = 2 * y / np.sqrt(3)
|
||||
p2 = x - 0.5 * p3
|
||||
p1 = 1 - p2 - p3
|
||||
return np.stack([p1, p2, p3], axis=-1)
|
||||
|
||||
|
||||
def _simplex_mesh(resolution):
|
||||
simplex_ymax = np.sqrt(3) / 2
|
||||
xs = np.linspace(0, 1, resolution)
|
||||
ys = np.linspace(0, simplex_ymax, resolution)
|
||||
grid_x, grid_y = np.meshgrid(xs, ys)
|
||||
pts_bary = _barycentric_from_xy(grid_x, grid_y)
|
||||
mask = np.all(pts_bary >= 0, axis=-1)
|
||||
return xs, ys, pts_bary, mask
|
||||
|
||||
|
||||
def _evaluate_simplex_function(function, points):
|
||||
points = np.asarray(points, dtype=float)
|
||||
try:
|
||||
values = np.asarray(function(points), dtype=float)
|
||||
if values.shape == (points.shape[0],):
|
||||
return values
|
||||
if values.shape == points.shape[:-1]:
|
||||
return values.reshape(-1)
|
||||
except Exception:
|
||||
pass
|
||||
return np.asarray([function(point) for point in points], dtype=float)
|
||||
|
||||
|
||||
def _region_colormap(color='blue', alpha=0.35):
|
||||
return ListedColormap([
|
||||
(1.0, 1.0, 1.0, 0.0),
|
||||
(*mcolors.to_rgb(color), alpha),
|
||||
])
|
||||
|
||||
|
||||
def _plot_simplex_points(ax, point_layers):
|
||||
for layer in point_layers:
|
||||
points = np.asarray(layer['points'], dtype=float)
|
||||
style = {'s': 25, 'alpha': 0.8}
|
||||
style.update(layer.get('style', {}))
|
||||
ax.scatter(*_simplex_to_cartesian(points), label=layer.get('label'), **style)
|
||||
|
||||
|
||||
def _plot_simplex_regions(ax, region_layers, resolution):
|
||||
xs, ys, pts_bary, simplex_mask = _simplex_mesh(resolution)
|
||||
valid_points = pts_bary[simplex_mask]
|
||||
|
||||
for layer in region_layers:
|
||||
mask = np.zeros(simplex_mask.shape, dtype=float)
|
||||
values = _evaluate_simplex_function(layer['fn'], valid_points)
|
||||
mask[simplex_mask] = values
|
||||
ax.pcolormesh(
|
||||
xs,
|
||||
ys,
|
||||
mask,
|
||||
shading='auto',
|
||||
cmap=_region_colormap(layer.get('color', 'blue'), layer.get('alpha', 0.35)),
|
||||
)
|
||||
if layer.get('label') is not None:
|
||||
ax.scatter([], [], color=layer.get('color', 'blue'), alpha=layer.get('alpha', 0.35), label=layer['label'])
|
||||
|
||||
|
||||
def _plot_simplex_density(ax, density_function, resolution, color, alpha):
|
||||
xs, ys, pts_bary, simplex_mask = _simplex_mesh(resolution)
|
||||
valid_points = pts_bary[simplex_mask]
|
||||
density = np.full(simplex_mask.shape, np.nan, dtype=float)
|
||||
values = _evaluate_simplex_function(density_function, valid_points)
|
||||
min_v, max_v = np.min(values), np.max(values)
|
||||
if max_v > min_v:
|
||||
values = (values - min_v) / (max_v - min_v)
|
||||
density[simplex_mask] = values
|
||||
|
||||
cmap = LinearSegmentedColormap.from_list('simplex_density', ['white', color])
|
||||
ax.pcolormesh(xs, ys, density, shading='auto', cmap=cmap, alpha=alpha)
|
||||
|
||||
|
||||
def calibration_plot(prob_classifier, X, y, nbins=10, savepath=None):
|
||||
posteriors = prob_classifier.predict_proba(X)
|
||||
assert posteriors.ndim==2, 'calibration plot only works for binary problems'
|
||||
|
|
@ -619,7 +826,6 @@ def calibration_plot(prob_classifier, X, y, nbins=10, savepath=None):
|
|||
pred_y = posteriors>=0.5
|
||||
bins = np.linspace(0, 1, nbins + 1)
|
||||
binned_values = np.digitize(posteriors, bins, right=False)
|
||||
print(np.unique(binned_values))
|
||||
correct = pred_y == y
|
||||
bin_centers = (bins[:-1] + bins[1:]) / 2
|
||||
bins_names = np.arange(nbins)
|
||||
|
|
|
|||
|
|
@ -8,8 +8,8 @@ from contextlib import ExitStack
|
|||
from abc import ABCMeta, abstractmethod
|
||||
from quapy.data import LabelledCollection
|
||||
import quapy.functional as F
|
||||
from os.path import exists
|
||||
from glob import glob
|
||||
from collections.abc import Iterable
|
||||
from numbers import Number
|
||||
|
||||
|
||||
class AbstractProtocol(metaclass=ABCMeta):
|
||||
|
|
@ -171,7 +171,7 @@ class AbstractStochasticSeededProtocol(AbstractProtocol):
|
|||
return sample
|
||||
|
||||
|
||||
class OnLabelledCollectionProtocol:
|
||||
class OnLabelledCollectionProtocol(AbstractStochasticSeededProtocol):
|
||||
"""
|
||||
Protocols that generate samples from a :class:`qp.data.LabelledCollection` object.
|
||||
"""
|
||||
|
|
@ -229,8 +229,17 @@ class OnLabelledCollectionProtocol:
|
|||
elif return_type=='index':
|
||||
return lambda lc,params:params
|
||||
|
||||
def sample(self, index):
|
||||
"""
|
||||
Realizes the sample given the index of the instances.
|
||||
|
||||
class APP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
||||
:param index: indexes of the instances to select
|
||||
:return: an instance of :class:`qp.data.LabelledCollection`
|
||||
"""
|
||||
return self.data.sampling_from_index(index)
|
||||
|
||||
|
||||
class APP(OnLabelledCollectionProtocol):
|
||||
"""
|
||||
Implementation of the artificial prevalence protocol (APP).
|
||||
The APP consists of exploring a grid of prevalence values containing `n_prevalences` points (e.g.,
|
||||
|
|
@ -311,15 +320,6 @@ class APP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
|||
indexes.append(index)
|
||||
return indexes
|
||||
|
||||
def sample(self, index):
|
||||
"""
|
||||
Realizes the sample given the index of the instances.
|
||||
|
||||
:param index: indexes of the instances to select
|
||||
:return: an instance of :class:`qp.data.LabelledCollection`
|
||||
"""
|
||||
return self.data.sampling_from_index(index)
|
||||
|
||||
def total(self):
|
||||
"""
|
||||
Returns the number of samples that will be generated
|
||||
|
|
@ -329,7 +329,7 @@ class APP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
|||
return F.num_prevalence_combinations(self.n_prevalences, self.data.n_classes, self.repeats)
|
||||
|
||||
|
||||
class NPP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
||||
class NPP(OnLabelledCollectionProtocol):
|
||||
"""
|
||||
A generator of samples that implements the natural prevalence protocol (NPP). The NPP consists of drawing
|
||||
samples uniformly at random, therefore approximately preserving the natural prevalence of the collection.
|
||||
|
|
@ -365,15 +365,6 @@ class NPP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
|||
indexes.append(index)
|
||||
return indexes
|
||||
|
||||
def sample(self, index):
|
||||
"""
|
||||
Realizes the sample given the index of the instances.
|
||||
|
||||
:param index: indexes of the instances to select
|
||||
:return: an instance of :class:`qp.data.LabelledCollection`
|
||||
"""
|
||||
return self.data.sampling_from_index(index)
|
||||
|
||||
def total(self):
|
||||
"""
|
||||
Returns the number of samples that will be generated (equals to "repeats")
|
||||
|
|
@ -383,7 +374,7 @@ class NPP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
|||
return self.repeats
|
||||
|
||||
|
||||
class UPP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
||||
class UPP(OnLabelledCollectionProtocol):
|
||||
"""
|
||||
A variant of :class:`APP` that, instead of using a grid of equidistant prevalence values,
|
||||
relies on the Kraemer algorithm for sampling unit (k-1)-simplex uniformly at random, with
|
||||
|
|
@ -423,14 +414,63 @@ class UPP(AbstractStochasticSeededProtocol, OnLabelledCollectionProtocol):
|
|||
indexes.append(index)
|
||||
return indexes
|
||||
|
||||
def sample(self, index):
|
||||
def total(self):
|
||||
"""
|
||||
Realizes the sample given the index of the instances.
|
||||
Returns the number of samples that will be generated (equals to "repeats")
|
||||
|
||||
:param index: indexes of the instances to select
|
||||
:return: an instance of :class:`qp.data.LabelledCollection`
|
||||
:return: int
|
||||
"""
|
||||
return self.data.sampling_from_index(index)
|
||||
return self.repeats
|
||||
|
||||
|
||||
class DirichletProtocol(OnLabelledCollectionProtocol):
|
||||
"""
|
||||
A protocol that establishes a prior Dirichlet distribution for the prevalence of the samples.
|
||||
Note that providing an all-ones vector of Dirichlet parameters is equivalent to invoking the
|
||||
APP protocol (although each protocol will generate a different series of samples given a
|
||||
fixed seed, since the implementation is different).
|
||||
|
||||
:param data: a `LabelledCollection` from which the samples will be drawn
|
||||
:param alpha: an array-like of shape (n_classes,) with the parameters of the Dirichlet distribution
|
||||
:param sample_size: integer, the number of instances in each sample; if None (default) then it is taken from
|
||||
qp.environ["SAMPLE_SIZE"]. If this is not set, a ValueError exception is raised.
|
||||
:param repeats: the number of samples to generate. Default is 100.
|
||||
:param random_state: allows replicating samples across runs (default 0, meaning that the sequence of samples
|
||||
will be the same every time the protocol is called)
|
||||
:param return_type: set to "sample_prev" (default) to get the pairs of (sample, prevalence) at each iteration, or
|
||||
to "labelled_collection" to get instead instances of LabelledCollection
|
||||
"""
|
||||
|
||||
def __init__(self, data: LabelledCollection, alpha, sample_size=None, repeats=100, random_state=0,
|
||||
return_type='sample_prev'):
|
||||
n_classes = data.n_classes
|
||||
if isinstance(alpha, str) and alpha == 'uniform':
|
||||
self.alpha = np.ones(n_classes, dtype=float)
|
||||
elif isinstance(alpha, Number):
|
||||
self.alpha = np.full(n_classes, float(alpha), dtype=float)
|
||||
else:
|
||||
self.alpha = np.asarray(alpha, dtype=float)
|
||||
if self.alpha.ndim != 1 or len(self.alpha) != n_classes:
|
||||
raise ValueError(
|
||||
f'wrong shape for alpha; expected {n_classes} values, found shape {self.alpha.shape}'
|
||||
)
|
||||
|
||||
super(DirichletProtocol, self).__init__(random_state)
|
||||
self.data = data
|
||||
self.sample_size = qp._get_sample_size(sample_size)
|
||||
self.repeats = repeats
|
||||
self.random_state = random_state
|
||||
self.collator = OnLabelledCollectionProtocol.get_collator(return_type)
|
||||
|
||||
def samples_parameters(self):
|
||||
"""
|
||||
Return all the necessary parameters to replicate the samples.
|
||||
|
||||
:return: a list of indexes that realize the sampling
|
||||
"""
|
||||
prevs = np.random.dirichlet(self.alpha, size=self.repeats)
|
||||
indexes = [self.data.sampling_index(self.sample_size, *prevs_i) for prevs_i in prevs]
|
||||
return indexes
|
||||
|
||||
def total(self):
|
||||
"""
|
||||
|
|
@ -450,7 +490,7 @@ class DomainMixer(AbstractStochasticSeededProtocol):
|
|||
:param sample_size: integer, the number of instances in each sample; if None (default) then it is taken from
|
||||
qp.environ["SAMPLE_SIZE"]. If this is not set, a ValueError exception is raised.
|
||||
:param repeats: int, number of samples to draw for every mixture rate
|
||||
:param prevalence: the prevalence to preserv along the mixtures. If specified, should be an array containing
|
||||
:param prevalence: the prevalence to preserve along the mixtures. If specified, should be an array containing
|
||||
one prevalence value (positive float) for each class and summing up to one. If not specified, the prevalence
|
||||
will be taken from the domain A (default).
|
||||
:param mixture_points: an integer indicating the number of points to take from a linear scale (e.g., 21 will
|
||||
|
|
|
|||
|
|
@ -0,0 +1,10 @@
|
|||
"""
|
||||
Fast unit tests live in ``test_*.py`` and are intended to run without network
|
||||
access or large external resources.
|
||||
|
||||
Slow integration tests that download datasets or depend on optional stacks live
|
||||
in ``integration_*.py`` and are meant to be run explicitly, e.g.:
|
||||
|
||||
python -m unittest quapy.tests.integration_datasets
|
||||
python -m unittest quapy.tests.integration_methods
|
||||
"""
|
||||
|
|
@ -0,0 +1,48 @@
|
|||
import numpy as np
|
||||
from sklearn.datasets import make_classification
|
||||
|
||||
from quapy.data import LabelledCollection
|
||||
from quapy.data.base import Dataset
|
||||
|
||||
|
||||
def make_labelled_collection(
|
||||
n_samples=200,
|
||||
n_features=12,
|
||||
n_classes=2,
|
||||
class_sep=1.5,
|
||||
random_state=0,
|
||||
):
|
||||
n_informative = min(n_features, max(4, n_classes * 2))
|
||||
X, y = make_classification(
|
||||
n_samples=n_samples,
|
||||
n_features=n_features,
|
||||
n_informative=n_informative,
|
||||
n_redundant=0,
|
||||
n_repeated=0,
|
||||
n_classes=n_classes,
|
||||
n_clusters_per_class=1,
|
||||
class_sep=class_sep,
|
||||
random_state=random_state,
|
||||
)
|
||||
classes = np.arange(n_classes)
|
||||
return LabelledCollection(X, y, classes=classes)
|
||||
|
||||
|
||||
def make_dataset(
|
||||
n_train=150,
|
||||
n_test=80,
|
||||
n_features=12,
|
||||
n_classes=2,
|
||||
class_sep=1.5,
|
||||
random_state=0,
|
||||
name='synthetic',
|
||||
):
|
||||
data = make_labelled_collection(
|
||||
n_samples=n_train + n_test,
|
||||
n_features=n_features,
|
||||
n_classes=n_classes,
|
||||
class_sep=class_sep,
|
||||
random_state=random_state,
|
||||
)
|
||||
training, test = data.split_stratified(train_prop=n_train / (n_train + n_test), random_state=random_state)
|
||||
return Dataset(training, test, name=name)
|
||||
|
|
@ -1,3 +1,10 @@
|
|||
"""
|
||||
Integration tests for dataset fetchers and large external resources.
|
||||
|
||||
This module is intentionally excluded from default ``unittest`` discovery by
|
||||
using an ``integration_*.py`` filename.
|
||||
"""
|
||||
|
||||
import os
|
||||
import unittest
|
||||
|
||||
|
|
@ -5,11 +12,11 @@ from sklearn.feature_extraction.text import TfidfVectorizer
|
|||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
import quapy.functional as F
|
||||
from quapy.method.aggregative import PCC
|
||||
from quapy.data.datasets import *
|
||||
from quapy.method.aggregative import PCC
|
||||
|
||||
|
||||
class TestDatasets(unittest.TestCase):
|
||||
class IntegrationDatasetsTest(unittest.TestCase):
|
||||
|
||||
def new_quantifier(self):
|
||||
return PCC(LogisticRegression(C=0.001, max_iter=100))
|
||||
|
|
@ -17,13 +24,11 @@ class TestDatasets(unittest.TestCase):
|
|||
def _check_dataset(self, dataset):
|
||||
train, test = dataset.reduce().train_test
|
||||
q = self.new_quantifier()
|
||||
print(f'testing method {q} in {dataset.name}...', end='')
|
||||
if len(train)>500:
|
||||
if len(train) > 500:
|
||||
train = train.sampling(500)
|
||||
q.fit(*dataset.training.Xy)
|
||||
estim_prevalences = q.predict(dataset.test.instances)
|
||||
self.assertTrue(F.check_prevalence_vector(estim_prevalences))
|
||||
print(f'[done]')
|
||||
|
||||
def _check_samples(self, gen, q, max_samples_test=5, vectorizer=None):
|
||||
for X, p in gen():
|
||||
|
|
@ -37,54 +42,37 @@ class TestDatasets(unittest.TestCase):
|
|||
|
||||
def test_reviews(self):
|
||||
for dataset_name in REVIEWS_SENTIMENT_DATASETS:
|
||||
print(f'loading dataset {dataset_name}...', end='')
|
||||
dataset = fetch_reviews(dataset_name, tfidf=True, min_df=10)
|
||||
dataset.stats()
|
||||
dataset.reduce()
|
||||
print(f'[done]')
|
||||
self._check_dataset(dataset)
|
||||
|
||||
def test_twitter(self):
|
||||
# all the datasets are contained in the same resource; if the first one
|
||||
# works, there is no need to test for the rest
|
||||
for dataset_name in TWITTER_SENTIMENT_DATASETS_TEST[:1]:
|
||||
print(f'loading dataset {dataset_name}...', end='')
|
||||
dataset = fetch_twitter(dataset_name, min_df=10)
|
||||
dataset.stats()
|
||||
dataset.reduce()
|
||||
print(f'[done]')
|
||||
self._check_dataset(dataset)
|
||||
|
||||
def test_UCIBinaryDataset(self):
|
||||
for dataset_name in UCI_BINARY_DATASETS:
|
||||
print(f'loading dataset {dataset_name}...', end='')
|
||||
dataset = fetch_UCIBinaryDataset(dataset_name)
|
||||
dataset.stats()
|
||||
dataset.reduce()
|
||||
print(f'[done]')
|
||||
self._check_dataset(dataset)
|
||||
|
||||
def test_UCIMultiDataset(self):
|
||||
for dataset_name in UCI_MULTICLASS_DATASETS:
|
||||
print(f'loading dataset {dataset_name}...', end='')
|
||||
dataset = fetch_UCIMulticlassDataset(dataset_name)
|
||||
dataset.stats()
|
||||
n_classes = dataset.n_classes
|
||||
uniform_prev = F.uniform_prevalence(n_classes)
|
||||
dataset.training = dataset.training.sampling(100, *uniform_prev)
|
||||
dataset.test = dataset.test.sampling(100, *uniform_prev)
|
||||
print(f'[done]')
|
||||
self._check_dataset(dataset)
|
||||
|
||||
def test_lequa2022(self):
|
||||
if os.environ.get('QUAPY_TESTS_OMIT_LARGE_DATASETS'):
|
||||
print("omitting test_lequa2022 because QUAPY_TESTS_OMIT_LARGE_DATASETS is set")
|
||||
return
|
||||
|
||||
for dataset_name in LEQUA2022_VECTOR_TASKS:
|
||||
print(f'LeQu2022: loading dataset {dataset_name}...', end='')
|
||||
train, gen_val, gen_test = fetch_lequa2022(dataset_name)
|
||||
train.stats()
|
||||
n_classes = train.n_classes
|
||||
train = train.sampling(100, *F.uniform_prevalence(n_classes))
|
||||
q = self.new_quantifier()
|
||||
|
|
@ -93,9 +81,7 @@ class TestDatasets(unittest.TestCase):
|
|||
self._check_samples(gen_test, q, max_samples_test=5)
|
||||
|
||||
for dataset_name in LEQUA2022_TEXT_TASKS:
|
||||
print(f'LeQu2022: loading dataset {dataset_name}...', end='')
|
||||
train, gen_val, gen_test = fetch_lequa2022(dataset_name)
|
||||
train.stats()
|
||||
n_classes = train.n_classes
|
||||
train = train.sampling(100, *F.uniform_prevalence(n_classes))
|
||||
tfidf = TfidfVectorizer()
|
||||
|
|
@ -107,13 +93,10 @@ class TestDatasets(unittest.TestCase):
|
|||
|
||||
def test_lequa2024(self):
|
||||
if os.environ.get('QUAPY_TESTS_OMIT_LARGE_DATASETS'):
|
||||
print("omitting test_lequa2024 because QUAPY_TESTS_OMIT_LARGE_DATASETS is set")
|
||||
return
|
||||
|
||||
for task in LEQUA2024_TASKS:
|
||||
print(f'LeQu2024: loading task {task}...', end='')
|
||||
train, gen_val, gen_test = fetch_lequa2024(task, merge_T3=True)
|
||||
train.stats()
|
||||
n_classes = train.n_classes
|
||||
train = train.sampling(100, *F.uniform_prevalence(n_classes))
|
||||
q = self.new_quantifier()
|
||||
|
|
@ -121,16 +104,12 @@ class TestDatasets(unittest.TestCase):
|
|||
self._check_samples(gen_val, q, max_samples_test=5)
|
||||
self._check_samples(gen_test, q, max_samples_test=5)
|
||||
|
||||
|
||||
def test_IFCB(self):
|
||||
if os.environ.get('QUAPY_TESTS_OMIT_LARGE_DATASETS'):
|
||||
print("omitting test_IFCB because QUAPY_TESTS_OMIT_LARGE_DATASETS is set")
|
||||
return
|
||||
|
||||
print(f'loading dataset IFCB.')
|
||||
for mod_sel in [False, True]:
|
||||
train, gen = fetch_IFCB(single_sample_train=True, for_model_selection=mod_sel)
|
||||
train.stats()
|
||||
n_classes = train.n_classes
|
||||
train = train.sampling(100, *F.uniform_prevalence(n_classes))
|
||||
q = self.new_quantifier()
|
||||
|
|
@ -0,0 +1,42 @@
|
|||
"""
|
||||
Integration tests for optional or resource-heavy method end-to-end checks.
|
||||
|
||||
This module is intentionally excluded from default ``unittest`` discovery by
|
||||
using an ``integration_*.py`` filename.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
|
||||
import quapy as qp
|
||||
from quapy.functional import check_prevalence_vector
|
||||
|
||||
|
||||
class IntegrationMethodsTest(unittest.TestCase):
|
||||
|
||||
def test_quanet(self):
|
||||
try:
|
||||
import quapy.classification.neural
|
||||
except ModuleNotFoundError:
|
||||
print('the torch package is not installed; skipping integration test for QuaNet')
|
||||
return
|
||||
|
||||
qp.environ['SAMPLE_SIZE'] = 10
|
||||
|
||||
dataset = qp.datasets.fetch_reviews('kindle', pickle=True).reduce()
|
||||
qp.data.preprocessing.index(dataset, min_df=5, inplace=True)
|
||||
|
||||
from quapy.classification.neural import CNNnet
|
||||
from quapy.classification.neural import NeuralClassifierTrainer
|
||||
from quapy.method.meta import QuaNet
|
||||
|
||||
cnn = CNNnet(dataset.vocabulary_size, dataset.n_classes)
|
||||
learner = NeuralClassifierTrainer(cnn, device='cpu')
|
||||
model = QuaNet(learner, device='cpu', n_epochs=2, tr_iter_per_poch=10, va_iter_per_poch=10, patience=2)
|
||||
|
||||
model.fit(*dataset.training.Xy)
|
||||
estim_prevalences = model.predict(dataset.test.instances)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
@ -0,0 +1,69 @@
|
|||
import unittest
|
||||
|
||||
import numpy as np
|
||||
|
||||
from quapy.method._bayesian import (
|
||||
_validate_temperature,
|
||||
_resolve_dirichlet_prior,
|
||||
kl_div,
|
||||
js_div,
|
||||
normalized,
|
||||
lambda_inverse,
|
||||
lambda_forward,
|
||||
)
|
||||
|
||||
|
||||
class TestBayesianUtils(unittest.TestCase):
|
||||
"""
|
||||
Smoke tests for the pure-numpy helper functions in method/_bayesian.py. These do not require the
|
||||
optional jax/stan/numpyro dependencies (unlike BayesianKDEy/BayesianMAPLS themselves), so they can
|
||||
run regardless of whether `quapy[bayes]` is installed.
|
||||
"""
|
||||
|
||||
def test_validate_temperature(self):
|
||||
self.assertEqual(_validate_temperature(2), 2.0)
|
||||
self.assertEqual(_validate_temperature(0.5), 0.5)
|
||||
with self.assertRaises(ValueError):
|
||||
_validate_temperature(0)
|
||||
with self.assertRaises(ValueError):
|
||||
_validate_temperature(-1)
|
||||
with self.assertRaises(ValueError):
|
||||
_validate_temperature('not-a-number')
|
||||
|
||||
def test_resolve_dirichlet_prior(self):
|
||||
np.testing.assert_array_equal(_resolve_dirichlet_prior('uniform', n_classes=3), np.ones(3))
|
||||
np.testing.assert_array_equal(_resolve_dirichlet_prior(2.5, n_classes=3), np.full(3, 2.5))
|
||||
np.testing.assert_array_equal(_resolve_dirichlet_prior([1, 2, 3], n_classes=3), np.array([1., 2., 3.]))
|
||||
with self.assertRaises(ValueError):
|
||||
_resolve_dirichlet_prior([1, 2], n_classes=3) # wrong shape
|
||||
with self.assertRaises(ValueError):
|
||||
_resolve_dirichlet_prior('unknown-prior', n_classes=3)
|
||||
|
||||
def test_kl_div_and_js_div(self):
|
||||
p = np.array([0.5, 0.5])
|
||||
self.assertAlmostEqual(kl_div(p, p), 0.0, places=6)
|
||||
self.assertAlmostEqual(js_div(p, p), 0.0, places=6)
|
||||
|
||||
q = np.array([0.9, 0.1])
|
||||
self.assertGreater(kl_div(p, q), 0.0)
|
||||
self.assertGreater(js_div(p, q), 0.0)
|
||||
# JS divergence is symmetric, unlike KL
|
||||
self.assertAlmostEqual(js_div(p, q), js_div(q, p), places=6)
|
||||
|
||||
def test_normalized(self):
|
||||
# normalized() operates on a batch of row vectors, shape (n_samples, n_features)
|
||||
a = np.array([[3.0, 4.0], [0.0, 0.0]])
|
||||
normed = normalized(a)
|
||||
self.assertAlmostEqual(np.linalg.norm(normed[0]), 1.0, places=6)
|
||||
# an all-zero row should not raise (guarded division), and stays all-zero
|
||||
np.testing.assert_array_equal(normed[1], np.array([0.0, 0.0]))
|
||||
|
||||
def test_lambda_inverse_forward_roundtrip(self):
|
||||
dpq, lam = 0.3, 0.4
|
||||
gamma = lambda_inverse(dpq, lam)
|
||||
recovered_lam = lambda_forward(dpq, gamma)
|
||||
self.assertAlmostEqual(recovered_lam, lam, places=6)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
@ -0,0 +1,35 @@
|
|||
import unittest
|
||||
|
||||
import numpy as np
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
from quapy.classification.calibration import NBVSCalibration, BCTSCalibration, TSCalibration, VSCalibration
|
||||
from quapy.tests._synthetic import make_labelled_collection
|
||||
|
||||
|
||||
class TestCalibration(unittest.TestCase):
|
||||
|
||||
data = make_labelled_collection(n_samples=200, n_features=10, n_classes=3, random_state=23)
|
||||
|
||||
def test_calibration_methods_fit_predict_proba(self):
|
||||
X, y = self.data.Xy
|
||||
for calib_cls in [NBVSCalibration, BCTSCalibration, TSCalibration, VSCalibration]:
|
||||
model = calib_cls(LogisticRegression(max_iter=2000), val_split=5)
|
||||
model.fit(X, y)
|
||||
posteriors = model.predict_proba(X)
|
||||
self.assertEqual(posteriors.shape, (len(y), self.data.n_classes))
|
||||
np.testing.assert_allclose(posteriors.sum(axis=1), 1.0, rtol=1e-5,
|
||||
err_msg=f'{calib_cls.__name__} posteriors do not sum to 1')
|
||||
predictions = model.predict(X)
|
||||
self.assertEqual(len(predictions), len(y))
|
||||
|
||||
def test_calibration_with_float_val_split(self):
|
||||
X, y = self.data.Xy
|
||||
model = BCTSCalibration(LogisticRegression(max_iter=2000), val_split=0.3)
|
||||
model.fit(X, y)
|
||||
posteriors = model.predict_proba(X)
|
||||
self.assertEqual(posteriors.shape, (len(y), self.data.n_classes))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
@ -0,0 +1,89 @@
|
|||
import unittest
|
||||
|
||||
import numpy as np
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
from quapy.method.aggregative import PACC
|
||||
from quapy.method.confidence import ConfidenceIntervals, ConfidenceEllipseSimplex, AggregativeBootstrap, WithConfidenceABC
|
||||
from quapy.tests._synthetic import make_dataset
|
||||
|
||||
|
||||
def _dirichlet_samples(n_classes=3, n_samples=300, random_state=0):
|
||||
rng = np.random.RandomState(random_state)
|
||||
return rng.dirichlet(np.ones(n_classes) * 5, size=n_samples)
|
||||
|
||||
|
||||
class TestConfidenceRegions(unittest.TestCase):
|
||||
|
||||
def test_confidence_intervals_contain_own_mean(self):
|
||||
samples = _dirichlet_samples()
|
||||
region = ConfidenceIntervals(samples)
|
||||
point_estimate = region.point_estimate()
|
||||
self.assertEqual(region.coverage(point_estimate), 1.)
|
||||
self.assertEqual(region.n_dim, 3)
|
||||
|
||||
def test_confidence_ellipse_simplex_contains_own_mean(self):
|
||||
samples = _dirichlet_samples()
|
||||
region = ConfidenceEllipseSimplex(samples)
|
||||
point_estimate = region.point_estimate()
|
||||
self.assertIn(point_estimate, region)
|
||||
|
||||
def test_construct_region_applies_bonferroni_only_to_intervals(self):
|
||||
samples = _dirichlet_samples()
|
||||
region_plain = WithConfidenceABC.construct_region(samples, confidence_level=0.9, method='intervals')
|
||||
region_bonf = WithConfidenceABC.construct_region(samples, confidence_level=0.9, method='intervals', bonferroni=True)
|
||||
ellipse = WithConfidenceABC.construct_region(samples, confidence_level=0.9, method='ellipse', bonferroni=True)
|
||||
|
||||
self.assertAlmostEqual(region_plain.alpha, 0.1)
|
||||
self.assertAlmostEqual(region_bonf.alpha, 0.1 / samples.shape[1])
|
||||
self.assertIsInstance(ellipse, ConfidenceEllipseSimplex)
|
||||
|
||||
def test_simplex_portion_is_cached_and_consistent(self):
|
||||
# regression test for the @lru_cache-on-bound-method memory leak fix:
|
||||
# results must still be memoized per instance, and two instances must not share state
|
||||
region1 = ConfidenceEllipseSimplex(_dirichlet_samples(random_state=1))
|
||||
region2 = ConfidenceEllipseSimplex(_dirichlet_samples(random_state=2))
|
||||
|
||||
p1_first = region1.simplex_portion()
|
||||
p1_second = region1.simplex_portion()
|
||||
self.assertEqual(p1_first, p1_second)
|
||||
|
||||
p2 = region2.simplex_portion()
|
||||
self.assertTrue(hasattr(region1, '_simplex_portion_cache'))
|
||||
self.assertTrue(hasattr(region2, '_simplex_portion_cache'))
|
||||
# each instance keeps its own cached value
|
||||
self.assertEqual(region1._simplex_portion_cache, p1_first)
|
||||
self.assertEqual(region2._simplex_portion_cache, p2)
|
||||
|
||||
def test_aggregative_bootstrap_end_to_end(self):
|
||||
dataset = make_dataset(n_train=150, n_test=50, n_classes=3, n_features=12, random_state=5)
|
||||
learner = LogisticRegression(max_iter=2000)
|
||||
learner.fit(*dataset.training.Xy)
|
||||
quantifier = AggregativeBootstrap(
|
||||
PACC(learner, fit_classifier=False), n_test_samples=50, confidence_level=0.9
|
||||
)
|
||||
quantifier.fit(*dataset.training.Xy)
|
||||
point_estimate, region = quantifier.predict_conf(dataset.test.X)
|
||||
self.assertEqual(len(point_estimate), 3)
|
||||
self.assertEqual(region.coverage(point_estimate), 1.)
|
||||
|
||||
def test_aggregative_bootstrap_exposes_bonferroni(self):
|
||||
dataset = make_dataset(n_train=150, n_test=50, n_classes=3, n_features=12, random_state=7)
|
||||
learner = LogisticRegression(max_iter=2000)
|
||||
learner.fit(*dataset.training.Xy)
|
||||
quantifier = AggregativeBootstrap(
|
||||
PACC(learner, fit_classifier=False),
|
||||
n_test_samples=20,
|
||||
confidence_level=0.9,
|
||||
region='intervals',
|
||||
bonferroni=True,
|
||||
random_state=0,
|
||||
)
|
||||
quantifier.fit(*dataset.training.Xy)
|
||||
_, region = quantifier.predict_conf(dataset.test.X)
|
||||
|
||||
self.assertAlmostEqual(region.alpha, 0.1 / 3)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
@ -1,43 +1,41 @@
|
|||
import inspect
|
||||
import unittest
|
||||
|
||||
import numpy as np
|
||||
|
||||
import quapy as qp
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
from time import time
|
||||
|
||||
import numpy as np
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
import quapy as qp
|
||||
from quapy.error import QUANTIFICATION_ERROR_SINGLE_NAMES
|
||||
from quapy.method.aggregative import EMQ, PCC
|
||||
from quapy.method.base import BaseQuantifier
|
||||
from quapy.tests._synthetic import make_dataset
|
||||
|
||||
|
||||
class EvalTestCase(unittest.TestCase):
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.data = make_dataset(n_train=140, n_test=90, n_classes=2, random_state=7, name='eval')
|
||||
|
||||
def test_eval_speedup(self):
|
||||
"""
|
||||
Checks whether the speed-up heuristics used by qp.evaluation work, i.e., actually save time
|
||||
"""
|
||||
|
||||
data = qp.datasets.fetch_reviews('hp', tfidf=True, min_df=10, pickle=True)
|
||||
train, test = data.training, data.test
|
||||
|
||||
protocol = qp.protocol.APP(test, sample_size=1000, n_prevalences=11, repeats=1, random_state=1)
|
||||
train, test = self.data.training, self.data.test
|
||||
protocol = qp.protocol.APP(test, sample_size=30, n_prevalences=5, repeats=1, random_state=1)
|
||||
|
||||
class SlowLR(LogisticRegression):
|
||||
def predict_proba(self, X):
|
||||
import time
|
||||
time.sleep(1)
|
||||
import time as _time
|
||||
_time.sleep(0.05)
|
||||
return super().predict_proba(X)
|
||||
|
||||
emq = EMQ(SlowLR()).fit(*train.Xy)
|
||||
emq = EMQ(SlowLR(max_iter=1000)).fit(*train.Xy)
|
||||
|
||||
tinit = time()
|
||||
score = qp.evaluation.evaluate(emq, protocol, error_metric='mae', verbose=True, aggr_speedup='force')
|
||||
tend_optim = time()-tinit
|
||||
print(f'evaluation (with optimization) took {tend_optim}s [MAE={score:.4f}]')
|
||||
score = qp.evaluation.evaluate(emq, protocol, error_metric='mae', aggr_speedup='force')
|
||||
tend_optim = time() - tinit
|
||||
self.assertTrue(isinstance(score, float))
|
||||
|
||||
class NonAggregativeEMQ(BaseQuantifier):
|
||||
|
||||
def __init__(self, cls):
|
||||
self.emq = EMQ(cls)
|
||||
|
||||
|
|
@ -48,31 +46,32 @@ class EvalTestCase(unittest.TestCase):
|
|||
self.emq.fit(X, y)
|
||||
return self
|
||||
|
||||
emq = NonAggregativeEMQ(SlowLR()).fit(*train.Xy)
|
||||
emq = NonAggregativeEMQ(SlowLR(max_iter=1000)).fit(*train.Xy)
|
||||
|
||||
tinit = time()
|
||||
score = qp.evaluation.evaluate(emq, protocol, error_metric='mae', verbose=True)
|
||||
score = qp.evaluation.evaluate(emq, protocol, error_metric='mae')
|
||||
tend_no_optim = time() - tinit
|
||||
print(f'evaluation (w/o optimization) took {tend_no_optim}s [MAE={score:.4f}]')
|
||||
|
||||
self.assertEqual(tend_no_optim>(tend_optim/2), True)
|
||||
self.assertTrue(isinstance(score, float))
|
||||
self.assertGreater(tend_no_optim, tend_optim)
|
||||
|
||||
def test_evaluation_output(self):
|
||||
"""
|
||||
Checks the evaluation functions return correct types for different error_metrics
|
||||
"""
|
||||
train, test = self.data.training, self.data.test
|
||||
qp.environ['SAMPLE_SIZE'] = 30
|
||||
protocol = qp.protocol.APP(test, sample_size=30, n_prevalences=5, repeats=1, random_state=0)
|
||||
q = PCC(LogisticRegression(max_iter=1000)).fit(*train.Xy)
|
||||
|
||||
data = qp.datasets.fetch_reviews('hp', tfidf=True, min_df=10, pickle=True).reduce(n_train=100, n_test=100)
|
||||
train, test = data.training, data.test
|
||||
def supports_evaluation(err):
|
||||
required = [
|
||||
p for p in inspect.signature(err).parameters.values()
|
||||
if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD) and p.default is inspect._empty
|
||||
]
|
||||
return len(required) <= 2
|
||||
|
||||
qp.environ['SAMPLE_SIZE']=100
|
||||
|
||||
protocol = qp.protocol.APP(test, random_state=0)
|
||||
|
||||
q = PCC(LogisticRegression()).fit(*train.Xy)
|
||||
|
||||
single_errors = list(QUANTIFICATION_ERROR_SINGLE_NAMES)
|
||||
averaged_errors = ['m'+e for e in single_errors]
|
||||
single_errors = [
|
||||
e for e in QUANTIFICATION_ERROR_SINGLE_NAMES
|
||||
if supports_evaluation(qp.error.from_name(e))
|
||||
]
|
||||
averaged_errors = ['m' + e for e in single_errors]
|
||||
single_errors = single_errors + [qp.error.from_name(e) for e in single_errors]
|
||||
averaged_errors = averaged_errors + [qp.error.from_name(e) for e in averaged_errors]
|
||||
for error_metric, averaged_error_metric in zip(single_errors, averaged_errors):
|
||||
|
|
@ -81,7 +80,6 @@ class EvalTestCase(unittest.TestCase):
|
|||
|
||||
scores = qp.evaluation.evaluate(q, protocol, error_metric=error_metric)
|
||||
self.assertTrue(isinstance(scores, np.ndarray))
|
||||
|
||||
self.assertEqual(scores.mean(), score)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,27 @@
|
|||
import unittest
|
||||
|
||||
import numpy as np
|
||||
|
||||
import quapy.functional as F
|
||||
|
||||
|
||||
class TestFunctional(unittest.TestCase):
|
||||
|
||||
def test_ternary_search_binary(self):
|
||||
def loss(prev):
|
||||
return (prev[1] - 0.37) ** 2
|
||||
|
||||
result = F.argmin_prevalence(loss, n_classes=2, method='ternary_search')
|
||||
self.assertTrue(np.allclose(result.sum(), 1.0))
|
||||
self.assertAlmostEqual(result[1], 0.37, places=3)
|
||||
|
||||
def test_ternary_search_multiclass_not_supported(self):
|
||||
def loss(prev):
|
||||
return np.sum((prev - np.array([0.2, 0.3, 0.5])) ** 2)
|
||||
|
||||
with self.assertRaises(AssertionError):
|
||||
F.argmin_prevalence(loss, n_classes=3, method='ternary_search')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
@ -1,130 +1,226 @@
|
|||
import itertools
|
||||
import inspect
|
||||
import unittest
|
||||
|
||||
import numpy as np
|
||||
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
import quapy as qp
|
||||
from quapy.method.aggregative import ACC
|
||||
from quapy.method.meta import Ensemble
|
||||
from quapy.method import AGGREGATIVE_METHODS, BINARY_METHODS, NON_AGGREGATIVE_METHODS
|
||||
from quapy.method.non_aggregative import DMx, EDx, HDx
|
||||
from quapy.method.aggregative import ACC, DMy, EDy, KDEyCS, RLLS
|
||||
from quapy.method.meta import Ensemble
|
||||
from quapy.functional import check_prevalence_vector
|
||||
from quapy.tests._synthetic import make_dataset
|
||||
import quapy as qp
|
||||
|
||||
# a random selection of composed methods to test the qunfold integration
|
||||
from quapy.method.composable import check_compatible_qunfold_version
|
||||
OPTIONAL_AGGREGATIVE_METHODS = {
|
||||
'BayesianCC',
|
||||
'BayesianKDEy',
|
||||
'BayesianMAPLS',
|
||||
'PQ',
|
||||
'RLLS',
|
||||
'EDy',
|
||||
}
|
||||
|
||||
from quapy.method.composable import (
|
||||
ComposableQuantifier,
|
||||
LeastSquaresLoss,
|
||||
HellingerSurrogateLoss,
|
||||
ClassRepresentation,
|
||||
HistogramRepresentation,
|
||||
CVClassifier
|
||||
)
|
||||
OPTIONAL_NON_AGGREGATIVE_METHODS = {
|
||||
'EDx',
|
||||
}
|
||||
|
||||
COMPOSABLE_METHODS = [
|
||||
ComposableQuantifier( # ACC
|
||||
LeastSquaresLoss(),
|
||||
ClassRepresentation(CVClassifier(LogisticRegression()))
|
||||
),
|
||||
ComposableQuantifier( # HDy
|
||||
HellingerSurrogateLoss(),
|
||||
HistogramRepresentation(
|
||||
3, # 3 bins per class
|
||||
preprocessor = ClassRepresentation(CVClassifier(LogisticRegression()))
|
||||
)
|
||||
),
|
||||
]
|
||||
|
||||
class TestMethods(unittest.TestCase):
|
||||
|
||||
tiny_dataset_multiclass = qp.datasets.fetch_UCIMulticlassDataset('academic-success').reduce(n_test=10)
|
||||
tiny_dataset_binary = qp.datasets.fetch_UCIBinaryDataset('ionosphere').reduce(n_test=10)
|
||||
tiny_dataset_multiclass = make_dataset(
|
||||
n_train=140, n_test=40, n_classes=3, n_features=12, random_state=11, name='synthetic-multiclass'
|
||||
)
|
||||
tiny_dataset_binary = make_dataset(
|
||||
n_train=140, n_test=40, n_classes=2, n_features=12, random_state=13, name='synthetic-binary'
|
||||
)
|
||||
datasets = [tiny_dataset_binary, tiny_dataset_multiclass]
|
||||
|
||||
def test_aggregative(self):
|
||||
for dataset in TestMethods.datasets:
|
||||
learner = LogisticRegression()
|
||||
learner = LogisticRegression(max_iter=2000)
|
||||
learner.fit(*dataset.training.Xy)
|
||||
|
||||
for model in AGGREGATIVE_METHODS:
|
||||
if model.__name__ in OPTIONAL_AGGREGATIVE_METHODS:
|
||||
continue
|
||||
if not dataset.binary and model in BINARY_METHODS:
|
||||
print(f'skipping the test of binary model {model.__name__} on multiclass dataset {dataset.name}')
|
||||
continue
|
||||
|
||||
q = model(learner, fit_classifier=False)
|
||||
print('testing', q)
|
||||
kwargs = {'fit_classifier': False}
|
||||
if 'val_split' in inspect.signature(model.__init__).parameters:
|
||||
kwargs['val_split'] = None
|
||||
q = model(learner, **kwargs)
|
||||
q.fit(*dataset.training.Xy)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
def test_non_aggregative(self):
|
||||
for dataset in TestMethods.datasets:
|
||||
|
||||
for model in NON_AGGREGATIVE_METHODS:
|
||||
if model.__name__ in OPTIONAL_NON_AGGREGATIVE_METHODS:
|
||||
continue
|
||||
if not dataset.binary and model in BINARY_METHODS:
|
||||
print(f'skipping the test of binary model {model.__name__} on multiclass dataset {dataset.name}')
|
||||
continue
|
||||
|
||||
q = model()
|
||||
print(f'testing {q} on dataset {dataset.name}')
|
||||
q.fit(*dataset.training.Xy)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
def test_ensembles(self):
|
||||
qp.environ['SAMPLE_SIZE'] = 10
|
||||
qp.environ['SAMPLE_SIZE'] = 20
|
||||
|
||||
base_quantifier = ACC(LogisticRegression())
|
||||
def policy_supported(policy):
|
||||
if policy in {'ave', 'ptr', 'ds'}:
|
||||
return True
|
||||
err = qp.error.from_name(policy)
|
||||
required = [
|
||||
p for p in inspect.signature(err).parameters.values()
|
||||
if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD) and p.default is inspect._empty
|
||||
]
|
||||
return len(required) <= 2
|
||||
|
||||
base_quantifier = ACC(LogisticRegression(max_iter=2000))
|
||||
for dataset, policy in itertools.product(TestMethods.datasets, Ensemble.VALID_POLICIES):
|
||||
if not policy_supported(policy):
|
||||
continue
|
||||
if not dataset.binary and policy == 'ds':
|
||||
print(f'skipping the test of binary policy ds on non-binary dataset {dataset}')
|
||||
continue
|
||||
|
||||
print(f'testing {base_quantifier} on dataset {dataset.name} with {policy=}')
|
||||
ensemble = Ensemble(quantifier=base_quantifier, size=3, policy=policy, n_jobs=-1)
|
||||
ensemble = Ensemble(quantifier=base_quantifier, size=3, policy=policy, n_jobs=1)
|
||||
ensemble.fit(*dataset.training.Xy)
|
||||
estim_prevalences = ensemble.predict(dataset.test.instances)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
def test_quanet(self):
|
||||
def test_composable(self):
|
||||
try:
|
||||
import quapy.classification.neural
|
||||
except ModuleNotFoundError:
|
||||
print('the torch package is not installed; skipping unit test for QuaNet')
|
||||
from quapy.method.composable import check_compatible_qunfold_version
|
||||
from quapy.method.composable import (
|
||||
ComposableQuantifier,
|
||||
LeastSquaresLoss,
|
||||
HellingerSurrogateLoss,
|
||||
ClassRepresentation,
|
||||
HistogramRepresentation,
|
||||
CVClassifier,
|
||||
)
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
qp.environ['SAMPLE_SIZE'] = 10
|
||||
composable_methods = [
|
||||
ComposableQuantifier(
|
||||
LeastSquaresLoss(),
|
||||
ClassRepresentation(CVClassifier(LogisticRegression()))
|
||||
),
|
||||
ComposableQuantifier(
|
||||
HellingerSurrogateLoss(),
|
||||
HistogramRepresentation(
|
||||
3,
|
||||
preprocessor=ClassRepresentation(CVClassifier(LogisticRegression()))
|
||||
)
|
||||
),
|
||||
]
|
||||
|
||||
# load the kindle dataset as text, and convert words to numerical indexes
|
||||
dataset = qp.datasets.fetch_reviews('kindle', pickle=True).reduce()
|
||||
qp.data.preprocessing.index(dataset, min_df=5, inplace=True)
|
||||
|
||||
from quapy.classification.neural import CNNnet
|
||||
cnn = CNNnet(dataset.vocabulary_size, dataset.n_classes)
|
||||
|
||||
from quapy.classification.neural import NeuralClassifierTrainer
|
||||
learner = NeuralClassifierTrainer(cnn, device='cpu')
|
||||
|
||||
from quapy.method.meta import QuaNet
|
||||
model = QuaNet(learner, device='cpu', n_epochs=2, tr_iter_per_poch=10, va_iter_per_poch=10, patience=2)
|
||||
|
||||
model.fit(*dataset.training.Xy)
|
||||
estim_prevalences = model.predict(dataset.test.instances)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
def test_composable(self):
|
||||
if check_compatible_qunfold_version():
|
||||
for dataset in TestMethods.datasets:
|
||||
for q in COMPOSABLE_METHODS:
|
||||
print('testing', q)
|
||||
for q in composable_methods:
|
||||
q.fit(*dataset.training.Xy)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
print(estim_prevalences)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
else:
|
||||
from quapy.method.composable import __old_version_message
|
||||
print(__old_version_message)
|
||||
|
||||
def test_rlls(self):
|
||||
try:
|
||||
import cvxpy # noqa: F401
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
dataset = TestMethods.tiny_dataset_multiclass
|
||||
q = RLLS(LogisticRegression(max_iter=2000), val_split=3)
|
||||
q.fit(*dataset.training.Xy)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
|
||||
def test_edy(self):
|
||||
try:
|
||||
import quadprog # noqa: F401
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
dataset = TestMethods.tiny_dataset_multiclass
|
||||
q = EDy(LogisticRegression(max_iter=2000), val_split=3)
|
||||
q.fit(*dataset.training.Xy)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
|
||||
def test_edx(self):
|
||||
try:
|
||||
import quadprog # noqa: F401
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
dataset = TestMethods.tiny_dataset_multiclass
|
||||
q = EDx()
|
||||
q.fit(*dataset.training.Xy)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
|
||||
|
||||
def test_dmy_noncanonical_labels(self):
|
||||
dataset = TestMethods.tiny_dataset_multiclass
|
||||
label_names = np.asarray(['class-a', 'class-c', 'class-z'])
|
||||
y_train = label_names[dataset.training.y]
|
||||
y_test = label_names[dataset.test.y]
|
||||
|
||||
q = DMy(LogisticRegression(max_iter=2000), val_split=3)
|
||||
q.fit(dataset.training.X, y_train)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
self.assertEqual(len(estim_prevalences), len(np.unique(y_test)))
|
||||
|
||||
|
||||
def test_dmx_noncanonical_labels(self):
|
||||
dataset = TestMethods.tiny_dataset_multiclass
|
||||
label_names = np.asarray(['class-a', 'class-c', 'class-z'])
|
||||
y_train = label_names[dataset.training.y]
|
||||
|
||||
q = DMx()
|
||||
q.fit(dataset.training.X, y_train)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
self.assertEqual(len(estim_prevalences), len(np.unique(y_train)))
|
||||
|
||||
def test_kdeycs_noncanonical_labels(self):
|
||||
dataset = TestMethods.tiny_dataset_multiclass
|
||||
label_names = np.asarray(['class-a', 'class-c', 'class-z'])
|
||||
y_train = label_names[dataset.training.y]
|
||||
|
||||
q = KDEyCS(LogisticRegression(max_iter=2000), val_split=3)
|
||||
q.fit(dataset.training.X, y_train)
|
||||
estim_prevalences = q.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences))
|
||||
self.assertEqual(len(estim_prevalences), len(np.unique(y_train)))
|
||||
|
||||
|
||||
def test_historical_distribution_matching_presets(self):
|
||||
dataset = TestMethods.tiny_dataset_binary
|
||||
|
||||
hdy = DMy.HDy(LogisticRegression(max_iter=2000), val_split=3)
|
||||
hdy.fit(*dataset.training.Xy)
|
||||
prev_hdy = hdy.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(prev_hdy))
|
||||
|
||||
hdx = HDx()
|
||||
hdx.fit(*dataset.training.Xy)
|
||||
prev_hdx = hdx.predict(dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(prev_hdx))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
|
|||
|
|
@ -1,104 +1,87 @@
|
|||
import time
|
||||
import unittest
|
||||
|
||||
import numpy as np
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
import quapy as qp
|
||||
from quapy.method.aggregative import PACC
|
||||
from quapy.model_selection import GridSearchQ
|
||||
from quapy.protocol import APP
|
||||
import time
|
||||
from quapy.tests._synthetic import make_dataset
|
||||
|
||||
|
||||
class ModselTestCase(unittest.TestCase):
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
data = make_dataset(
|
||||
n_train=220,
|
||||
n_test=120,
|
||||
n_classes=2,
|
||||
n_features=16,
|
||||
class_sep=1.8,
|
||||
random_state=1,
|
||||
name='modsel',
|
||||
)
|
||||
cls.training, cls.validation = data.training.split_stratified(0.7, random_state=1)
|
||||
|
||||
def test_modsel(self):
|
||||
"""
|
||||
Checks whether a model selection exploration takes a good hyperparameter
|
||||
Checks whether a model selection exploration picks the better hyperparameter.
|
||||
"""
|
||||
|
||||
q = PACC(LogisticRegression(random_state=1, max_iter=5000))
|
||||
|
||||
data = qp.datasets.fetch_reviews('imdb', tfidf=True, min_df=10).reduce(random_state=1)
|
||||
training, validation = data.training.split_stratified(0.7, random_state=1)
|
||||
|
||||
param_grid = {'classifier__C': [0.000001, 10.]}
|
||||
app = APP(validation, sample_size=100, random_state=1)
|
||||
param_grid = {'classifier__C': [0.000001, 10.0]}
|
||||
app = APP(self.validation, sample_size=30, n_prevalences=5, repeats=1, random_state=1)
|
||||
q = GridSearchQ(
|
||||
q, param_grid, protocol=app, error='mae', refit=False, timeout=-1, verbose=True, n_jobs=-1
|
||||
).fit(*training.Xy)
|
||||
print('best params', q.best_params_)
|
||||
print('best score', q.best_score_)
|
||||
q, param_grid, protocol=app, error='mae', refit=False, timeout=-1, verbose=False, n_jobs=-1
|
||||
).fit(*self.training.Xy)
|
||||
|
||||
self.assertEqual(q.best_params_['classifier__C'], 10.0)
|
||||
self.assertEqual(q.best_model().get_params()['classifier__C'], 10.0)
|
||||
|
||||
def test_modsel_parallel(self):
|
||||
"""
|
||||
Checks whether a parallelized model selection actually is faster than a sequential exploration but
|
||||
obtains the same optimal parameters
|
||||
Checks whether sequential and parallel model selection agree on the best parameters.
|
||||
"""
|
||||
|
||||
q = PACC(LogisticRegression(random_state=1, max_iter=3000))
|
||||
|
||||
data = qp.datasets.fetch_reviews('imdb', tfidf=True, min_df=50)
|
||||
training, validation = data.training.split_stratified(0.7, random_state=1)
|
||||
|
||||
param_grid = {'classifier__C': np.logspace(-3,3,7), 'classifier__class_weight': ['balanced', None]}
|
||||
app = APP(validation, sample_size=100, random_state=1)
|
||||
param_grid = {'classifier__C': np.logspace(-3, 3, 7), 'classifier__class_weight': ['balanced', None]}
|
||||
app = APP(self.validation, sample_size=30, n_prevalences=5, repeats=1, random_state=1)
|
||||
|
||||
def do_gridsearch(n_jobs):
|
||||
print('starting model selection in sequential exploration')
|
||||
t_init = time.time()
|
||||
modsel = GridSearchQ(
|
||||
q, param_grid, protocol=app, error='mae', refit=False, timeout=-1, n_jobs=n_jobs, verbose=True
|
||||
).fit(*training.Xy)
|
||||
t_end = time.time()-t_init
|
||||
best_c = modsel.best_params_['classifier__C']
|
||||
print(f'[done] took {t_end:.2f}s best C = {best_c}')
|
||||
return t_end, best_c
|
||||
q, param_grid, protocol=app, error='mae', refit=False, timeout=-1, n_jobs=n_jobs, verbose=False
|
||||
).fit(*self.training.Xy)
|
||||
t_end = time.time() - t_init
|
||||
return t_end, modsel.best_params_
|
||||
|
||||
tend_seq, best_c_seq = do_gridsearch(n_jobs=1)
|
||||
tend_par, best_c_par = do_gridsearch(n_jobs=-1)
|
||||
|
||||
print(tend_seq, best_c_seq)
|
||||
print(tend_par, best_c_par)
|
||||
|
||||
self.assertEqual(best_c_seq, best_c_par)
|
||||
self.assertLess(tend_par, tend_seq)
|
||||
_, best_seq = do_gridsearch(n_jobs=1)
|
||||
_, best_par = do_gridsearch(n_jobs=-1)
|
||||
|
||||
self.assertEqual(best_seq, best_par)
|
||||
|
||||
def test_modsel_timeout(self):
|
||||
|
||||
class SlowLR(LogisticRegression):
|
||||
def fit(self, X, y, sample_weight=None):
|
||||
import time
|
||||
time.sleep(10)
|
||||
super(SlowLR, self).fit(X, y, sample_weight)
|
||||
time.sleep(2)
|
||||
return super().fit(X, y, sample_weight)
|
||||
|
||||
q = PACC(SlowLR())
|
||||
q = PACC(SlowLR(max_iter=1000))
|
||||
param_grid = {'classifier__C': np.logspace(-1, 1, 3)}
|
||||
app = APP(self.validation, sample_size=30, n_prevalences=5, repeats=1, random_state=1)
|
||||
|
||||
data = qp.datasets.fetch_reviews('imdb', tfidf=True, min_df=10).reduce(random_state=1)
|
||||
training, validation = data.training.split_stratified(0.7, random_state=1)
|
||||
|
||||
param_grid = {'classifier__C': np.logspace(-1,1,3)}
|
||||
app = APP(validation, sample_size=100, random_state=1)
|
||||
|
||||
print('Expecting TimeoutError to be raised')
|
||||
modsel = GridSearchQ(
|
||||
q, param_grid, protocol=app, timeout=3, n_jobs=-1, verbose=True, raise_errors=True
|
||||
q, param_grid, protocol=app, timeout=1, n_jobs=-1, verbose=False, raise_errors=True
|
||||
)
|
||||
with self.assertRaises(TimeoutError):
|
||||
modsel.fit(*training.Xy)
|
||||
modsel.fit(*self.training.Xy)
|
||||
|
||||
print('Expecting ValueError to be raised')
|
||||
modsel = GridSearchQ(
|
||||
q, param_grid, protocol=app, timeout=3, n_jobs=-1, verbose=True, raise_errors=False
|
||||
q, param_grid, protocol=app, timeout=1, n_jobs=-1, verbose=False, raise_errors=False
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
# this exception is not raised because of the timeout, but because no combination of hyperparams
|
||||
# succedded (in this case, a ValueError is raised, regardless of "raise_errors"
|
||||
modsel.fit(*training.Xy)
|
||||
modsel.fit(*self.training.Xy)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
|
|
|||
|
|
@ -0,0 +1,41 @@
|
|||
import unittest
|
||||
|
||||
try:
|
||||
import matplotlib
|
||||
matplotlib.use('Agg')
|
||||
import matplotlib.pyplot as plt
|
||||
HAS_MATPLOTLIB = True
|
||||
except ImportError:
|
||||
plt = None
|
||||
HAS_MATPLOTLIB = False
|
||||
|
||||
import numpy as np
|
||||
|
||||
import quapy as qp
|
||||
|
||||
|
||||
@unittest.skipUnless(HAS_MATPLOTLIB and qp.plot is not None, 'matplotlib is not available')
|
||||
class TestPlot(unittest.TestCase):
|
||||
|
||||
def test_plot_simplex_smoke(self):
|
||||
rng = np.random.default_rng(0)
|
||||
true_prev = np.array([0.2, 0.3, 0.5])
|
||||
cloud = rng.dirichlet(alpha=30 * true_prev, size=50)
|
||||
|
||||
fig, ax = plt.subplots(figsize=(5, 5))
|
||||
fig, ax = qp.plot.plot_simplex(
|
||||
point_layers=[
|
||||
{'points': cloud, 'label': 'cloud', 'style': {'s': 8, 'alpha': 0.2}},
|
||||
{'points': true_prev, 'label': 'target', 'style': {'s': 50, 'color': 'black'}},
|
||||
],
|
||||
region_layers=[
|
||||
{'fn': lambda p: p[:, 2] >= 0.4, 'label': 'high class-3', 'color': 'green', 'alpha': 0.2},
|
||||
],
|
||||
density_function=lambda p: np.exp(-25 * np.sum((p - true_prev) ** 2, axis=1)),
|
||||
class_names=['A', 'B', 'C'],
|
||||
ax=ax,
|
||||
)
|
||||
|
||||
self.assertIs(fig, ax.figure)
|
||||
self.assertGreaterEqual(len(ax.collections), 2)
|
||||
plt.close(fig)
|
||||
|
|
@ -3,7 +3,7 @@ import numpy as np
|
|||
|
||||
import quapy.functional
|
||||
from quapy.data import LabelledCollection
|
||||
from quapy.protocol import APP, NPP, UPP, DomainMixer, AbstractStochasticSeededProtocol
|
||||
from quapy.protocol import APP, NPP, UPP, DomainMixer, AbstractStochasticSeededProtocol, DirichletProtocol
|
||||
|
||||
|
||||
def mock_labelled_collection(prefix=''):
|
||||
|
|
@ -138,6 +138,31 @@ class TestProtocols(unittest.TestCase):
|
|||
|
||||
self.assertNotEqual(samples1, samples2)
|
||||
|
||||
def test_dirichlet_replicate(self):
|
||||
data = mock_labelled_collection()
|
||||
p = DirichletProtocol(data, alpha=[1, 2, 3, 4], sample_size=5, repeats=10, random_state=42)
|
||||
|
||||
samples1 = samples_to_str(p)
|
||||
samples2 = samples_to_str(p)
|
||||
|
||||
self.assertEqual(samples1, samples2)
|
||||
|
||||
p = DirichletProtocol(data, alpha=[1, 2, 3, 4], sample_size=5, repeats=10, random_state=0)
|
||||
|
||||
samples1 = samples_to_str(p)
|
||||
samples2 = samples_to_str(p)
|
||||
|
||||
self.assertEqual(samples1, samples2)
|
||||
|
||||
def test_dirichlet_not_replicate(self):
|
||||
data = mock_labelled_collection()
|
||||
p = DirichletProtocol(data, alpha=[1, 2, 3, 4], sample_size=5, repeats=10, random_state=None)
|
||||
|
||||
samples1 = samples_to_str(p)
|
||||
samples2 = samples_to_str(p)
|
||||
|
||||
self.assertNotEqual(samples1, samples2)
|
||||
|
||||
def test_covariate_shift_replicate(self):
|
||||
dataA = mock_labelled_collection('domA')
|
||||
dataB = mock_labelled_collection('domB')
|
||||
|
|
|
|||
|
|
@ -0,0 +1,56 @@
|
|||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
import numpy as np
|
||||
|
||||
from quapy.data.reader import from_text, from_sparse, from_csv, reindex_labels, binarize
|
||||
|
||||
|
||||
class TestReader(unittest.TestCase):
|
||||
|
||||
def _write_tmp(self, content, suffix='.txt'):
|
||||
fd, path = tempfile.mkstemp(suffix=suffix)
|
||||
with os.fdopen(fd, 'w') as f:
|
||||
f.write(content)
|
||||
self.addCleanup(os.remove, path)
|
||||
return path
|
||||
|
||||
def test_from_text(self):
|
||||
path = self._write_tmp('1\tthis is positive\n0\tthis is negative\n')
|
||||
sentences, labels = from_text(path, verbose=0)
|
||||
self.assertEqual(sentences, ['this is positive', 'this is negative'])
|
||||
self.assertEqual(labels, [1, 0])
|
||||
|
||||
def test_from_text_skips_malformed_lines(self):
|
||||
# a line without a tab separator should be skipped (and warned about), not raise
|
||||
path = self._write_tmp('1\tgood line\nthis line has no label\n0\tanother good line\n')
|
||||
sentences, labels = from_text(path, verbose=0)
|
||||
self.assertEqual(sentences, ['good line', 'another good line'])
|
||||
self.assertEqual(labels, [1, 0])
|
||||
|
||||
def test_from_sparse(self):
|
||||
# format: <label> <col:val> <col:val> ... (1-indexed columns)
|
||||
path = self._write_tmp('1 1:0.5 2:1.0\n-1 2:2.0\n', suffix='.dat')
|
||||
X, y = from_sparse(path)
|
||||
self.assertEqual(X.shape[0], 2)
|
||||
np.testing.assert_array_equal(y, np.array([2, 0])) # labels shifted by +1
|
||||
|
||||
def test_from_csv(self):
|
||||
path = self._write_tmp('a,1.0,2.0\nb,3.0,4.0\n', suffix='.csv')
|
||||
X, y = from_csv(path)
|
||||
np.testing.assert_array_equal(X, np.array([[1.0, 2.0], [3.0, 4.0]]))
|
||||
np.testing.assert_array_equal(y, np.array(['a', 'b']))
|
||||
|
||||
def test_reindex_labels(self):
|
||||
indexed, classnames = reindex_labels(['B', 'B', 'A', 'C'])
|
||||
np.testing.assert_array_equal(indexed, np.array([1, 1, 0, 2]))
|
||||
np.testing.assert_array_equal(classnames, np.array(['A', 'B', 'C']))
|
||||
|
||||
def test_binarize(self):
|
||||
binarized = binarize([1, 2, 3, 1, 1, 0], pos_class=2)
|
||||
np.testing.assert_array_equal(binarized, np.array([0, 1, 0, 0, 0, 0]))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
@ -1,19 +1,29 @@
|
|||
import unittest
|
||||
|
||||
import numpy as np
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
import quapy as qp
|
||||
import quapy.functional as F
|
||||
from quapy.data import LabelledCollection
|
||||
from quapy.functional import strprev
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
import numpy as np
|
||||
from quapy.method.aggregative import PACC
|
||||
import quapy.functional as F
|
||||
from quapy.tests._synthetic import make_dataset
|
||||
|
||||
|
||||
class TestReplicability(unittest.TestCase):
|
||||
|
||||
def test_prediction_replicability(self):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.binary_dataset = make_dataset(
|
||||
n_train=180, n_test=80, n_classes=2, n_features=10, random_state=21, name='rep-binary'
|
||||
)
|
||||
cls.multiclass_dataset = make_dataset(
|
||||
n_train=180, n_test=80, n_classes=3, n_features=12, random_state=23, name='rep-multiclass'
|
||||
)
|
||||
|
||||
dataset = qp.datasets.fetch_UCIBinaryDataset('yeast')
|
||||
train, test = dataset.train_test
|
||||
def test_prediction_replicability(self):
|
||||
train, test = self.binary_dataset.train_test
|
||||
|
||||
with qp.util.temp_seed(0):
|
||||
lr = LogisticRegression(random_state=0, max_iter=10000)
|
||||
|
|
@ -29,7 +39,6 @@ class TestReplicability(unittest.TestCase):
|
|||
|
||||
self.assertEqual(str_prev1, str_prev2)
|
||||
|
||||
|
||||
def test_samping_replicability(self):
|
||||
|
||||
def equal_collections(c1, c2, value=True):
|
||||
|
|
@ -60,53 +69,33 @@ class TestReplicability(unittest.TestCase):
|
|||
sample2 = data.sampling(50, *[0.7, 0.3])
|
||||
equal_collections(sample1, sample2, True)
|
||||
|
||||
sample1 = data.sampling(50, *[0.7, 0.3], random_state=0)
|
||||
sample2 = data.sampling(50, *[0.7, 0.3], random_state=0)
|
||||
equal_collections(sample1, sample2, True)
|
||||
|
||||
sample1_tr, sample1_te = data.split_stratified(train_prop=0.7, random_state=0)
|
||||
sample2_tr, sample2_te = data.split_stratified(train_prop=0.7, random_state=0)
|
||||
equal_collections(sample1_tr, sample2_tr, True)
|
||||
equal_collections(sample1_te, sample2_te, True)
|
||||
|
||||
with qp.util.temp_seed(0):
|
||||
sample1_tr, sample1_te = data.split_stratified(train_prop=0.7)
|
||||
with qp.util.temp_seed(0):
|
||||
sample2_tr, sample2_te = data.split_stratified(train_prop=0.7)
|
||||
equal_collections(sample1_tr, sample2_tr, True)
|
||||
equal_collections(sample1_te, sample2_te, True)
|
||||
|
||||
|
||||
def test_parallel_replicability(self):
|
||||
|
||||
train, test = qp.datasets.fetch_UCIMulticlassDataset('dry-bean').reduce().train_test
|
||||
|
||||
test = test.sampling(500, *[0.1, 0.0, 0.1, 0.1, 0.2, 0.5, 0.0])
|
||||
train, test = self.multiclass_dataset.train_test
|
||||
test = test.sampling(60, *[0.2, 0.3, 0.5], random_state=4)
|
||||
|
||||
with qp.util.temp_seed(10):
|
||||
pacc = PACC(LogisticRegression(), val_split=.5, n_jobs=2)
|
||||
pacc = PACC(LogisticRegression(max_iter=5000), val_split=.5, n_jobs=2)
|
||||
pacc.fit(*train.Xy)
|
||||
prev1 = F.strprev(pacc.predict(test.instances))
|
||||
|
||||
with qp.util.temp_seed(0):
|
||||
pacc = PACC(LogisticRegression(), val_split=.5, n_jobs=2)
|
||||
pacc = PACC(LogisticRegression(max_iter=5000), val_split=.5, n_jobs=2)
|
||||
pacc.fit(*train.Xy)
|
||||
prev2 = F.strprev(pacc.predict(test.instances))
|
||||
|
||||
with qp.util.temp_seed(0):
|
||||
pacc = PACC(LogisticRegression(), val_split=.5, n_jobs=2)
|
||||
pacc = PACC(LogisticRegression(max_iter=5000), val_split=.5, n_jobs=2)
|
||||
pacc.fit(*train.Xy)
|
||||
prev3 = F.strprev(pacc.predict(test.instances))
|
||||
|
||||
print(prev1)
|
||||
print(prev2)
|
||||
print(prev3)
|
||||
|
||||
self.assertNotEqual(prev1, prev2)
|
||||
self.assertEqual(prev2, prev3)
|
||||
|
||||
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,33 @@
|
|||
import unittest
|
||||
|
||||
from sklearn.linear_model import LogisticRegression
|
||||
|
||||
from quapy.functional import check_prevalence_vector
|
||||
from quapy.method.aggregative import T50, MAX, X, MS, MS2
|
||||
from quapy.tests._synthetic import make_dataset
|
||||
|
||||
|
||||
class TestThresholdOptim(unittest.TestCase):
|
||||
|
||||
dataset = make_dataset(n_train=140, n_test=40, n_classes=2, n_features=12, random_state=17, name='synthetic-binary')
|
||||
|
||||
def test_compute_tpr_fpr_edge_cases(self):
|
||||
# regression test for the TP/FN vs TP/FP parameter-naming mix-up in _compute_tpr
|
||||
model = T50()
|
||||
self.assertEqual(model._compute_tpr(TP=5, FN=5), 0.5)
|
||||
self.assertEqual(model._compute_tpr(TP=0, FN=0), 1) # guarded division by zero
|
||||
self.assertEqual(model._compute_fpr(FP=3, TN=7), 0.3)
|
||||
self.assertEqual(model._compute_fpr(FP=0, TN=0), 0) # guarded division by zero
|
||||
|
||||
def test_threshold_methods_fit_predict(self):
|
||||
learner = LogisticRegression(max_iter=2000)
|
||||
learner.fit(*self.dataset.training.Xy)
|
||||
for model_cls in [T50, MAX, X, MS, MS2]:
|
||||
model = model_cls(learner, fit_classifier=False, val_split=None)
|
||||
model.fit(*self.dataset.training.Xy)
|
||||
estim_prevalences = model.predict(self.dataset.test.X)
|
||||
self.assertTrue(check_prevalence_vector(estim_prevalences), f'{model_cls.__name__} failed')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||