From f468b3b7a4463b025bb8028da0114e6e03a53c3a Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Mon, 2 Aug 2021 19:14:40 +0200 Subject: [PATCH] Implement histogram based kl divergence computation --- src/metrics.py | 40 +++++++++++++++++++++++++-- tests/test_metrics.py | 64 +++++++++++++++++++++++-------------------- 2 files changed, 72 insertions(+), 32 deletions(-) diff --git a/src/metrics.py b/src/metrics.py index 175d574..fc28f7f 100644 --- a/src/metrics.py +++ b/src/metrics.py @@ -87,7 +87,24 @@ def divergence(p, q, type='kl', n_components=0): d (float): approximate divergence """ if type == 'kl': - if n_components == 0: + if n_components < 0: + # Use histogram binning to discretize sampled distributions + p_hist, p_edges = np.histogram(p.unsqueeze(-1), bins='auto', density=True) + q_hist, q_edges = np.histogram(q.unsqueeze(-1), bins='auto', density=True) + px = evaluate_histogram(p, p_hist, p_edges) + qx = evaluate_histogram(p, q_hist, q_edges) + p_supp = ~np.isclose(px, 0.0) + q_supp = ~np.isclose(qx, 0.0) + if np.any(np.logical_and(p_supp, ~q_supp)): + # if not support(p) subset support(q) + return np.nan + elif ~np.any(p_supp): + # if p is zero everywhere + return 0. + d = np.mean(np.log(px[p_supp] / qx[p_supp])) + return d + elif n_components == 0: + # Assume p and q to be Gaussian pm = torch.mean(p) qm = torch.mean(q) pv = torch.var(p) @@ -161,4 +178,23 @@ def nanmean(v, *args, inplace=False, **kwargs): is_nan = torch.isnan(v) v[is_nan] = 0 result = v.sum(*args, **kwargs) / (~is_nan).float().sum(*args, **kwargs) - return result \ No newline at end of file + return result + + +def evaluate_histogram(x, hist, bin_edges): + """ + Evaluate a histogram + Args: + x (array) : points at which to evaluate the histogram + hist (array): histogram values in terms of number of occurrences or probability + bin_edges (array): edges of histogram bins + e.g. from hist, bin_edges = np.histogram(p, bins='auto', density=True) + Return: + r: tensor: (batch,) kl between each sample + """ + idx = np.digitize(x, bin_edges) + mask = np.logical_and(np.less(0, idx), np.less(idx, len(bin_edges))) + r = np.zeros_like(x) + r[mask] = hist[idx[mask] - 1] + return r + diff --git a/tests/test_metrics.py b/tests/test_metrics.py index b219bf2..bef990d 100644 --- a/tests/test_metrics.py +++ b/tests/test_metrics.py @@ -2,7 +2,7 @@ import torch import numpy as np from src import metrics -from src.metrics import divergence +from src.metrics import divergence, evaluate_histogram from sklearn.model_selection import GridSearchCV from sklearn.neighbors import KernelDensity @@ -20,39 +20,43 @@ def test_kl_divergence(): d3 = divergence(p, q, type='kl', n_components=3) print(d3) assert isinstance(d3, float) + d4 = divergence(p, q, type='kl', n_components=-1) + assert isinstance(d4, float) + d5 = divergence(q, p, type='kl', n_components=-1) + assert isinstance(d5, float) + p = 1000000 * torch.randn(1000) + q = torch.randn(1000) + d6 = divergence(p, q, type='kl', n_components=-1) + assert np.isnan(d6) -def evaluate_histogram(x, hist, bin_edges): - idx = np.digitize(x, bin_edges) - mask = np.logical_and(np.less(0, idx), np.less(idx, len(bin_edges))) - r = np.zeros_like(x) - r[mask] = hist[idx[mask] - 1] +def test_evaluate_histogram(): + N = 10000 + p = torch.randn(N) + q = 1.0 + 2.0 * torch.randn(2*N) + + p_hist, p_edges = np.histogram(p.unsqueeze(-1), bins='auto', density=True) + q_hist, q_edges = np.histogram(q.unsqueeze(-1), bins='auto', density=True) + px = evaluate_histogram(p, p_hist, p_edges) + assert px.shape == p.shape + qx = evaluate_histogram(q, p_hist, p_edges) + assert qx.shape == q.shape + px = evaluate_histogram(p, q_hist, q_edges) + assert px.shape == p.shape + qx = evaluate_histogram(q, q_hist, q_edges) + assert qx.shape == q.shape if __name__ == '__main__': - # p = torch.randn(10) - # q = 1.0 + 2.0 * torch.randn(10) + p = torch.randn(10) + q = 1.0 + 2.0 * torch.randn(10) - # p = p.unsqueeze(-1) - # q = q.unsqueeze(-1) - # p_hist, p_edges = np.histogram(p.unsqueeze(-1), bins='auto', density=True) - # q_hist, q_edges = np.histogram(q.unsqueeze(-1), bins='auto', density=True) - # # px = p_hist[np.digitize(p, p_edges) - 1] + p = p.unsqueeze(-1) + q = q.unsqueeze(-1) + p_hist, p_edges = np.histogram(p.unsqueeze(-1), bins='auto', density=True) + q_hist, q_edges = np.histogram(q.unsqueeze(-1), bins='auto', density=True) + # px = p_hist[np.digitize(p, p_edges) - 1] - # # qx = q_hist[np.digitize(p, q_edges) - 1] - # px = evaluate_histogram(p, p_hist, p_edges) - # qx = evaluate_histogram(q, p_hist, p_edges) - - # use grid search cross-validation to optimize the bandwidth - params = {'bandwidth': np.logspace(-1, 1, 3)} - grid = GridSearchCV(KernelDensity(), params) - grid.fit(p) - p_kde = grid.best_estimator_ - grid = GridSearchCV(KernelDensity(), params) - grid.fit(q) - q_kde = grid.best_estimator_ - px = p_kde.score_samples(p) - qx = q_kde.score_samples(p) - d = np.mean(px - qx).item() - print(d) - test_kl_divergence() \ No newline at end of file + # qx = q_hist[np.digitize(p, q_edges) - 1] + px = evaluate_histogram(p, p_hist, p_edges) + qx = evaluate_histogram(q, p_hist, p_edges)