Implement histogram based kl divergence computation
This commit is contained in:
@@ -87,7 +87,24 @@ def divergence(p, q, type='kl', n_components=0):
|
|||||||
d (float): approximate divergence
|
d (float): approximate divergence
|
||||||
"""
|
"""
|
||||||
if type == 'kl':
|
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)
|
pm = torch.mean(p)
|
||||||
qm = torch.mean(q)
|
qm = torch.mean(q)
|
||||||
pv = torch.var(p)
|
pv = torch.var(p)
|
||||||
@@ -162,3 +179,22 @@ def nanmean(v, *args, inplace=False, **kwargs):
|
|||||||
v[is_nan] = 0
|
v[is_nan] = 0
|
||||||
result = v.sum(*args, **kwargs) / (~is_nan).float().sum(*args, **kwargs)
|
result = v.sum(*args, **kwargs) / (~is_nan).float().sum(*args, **kwargs)
|
||||||
return result
|
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
|
||||||
|
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import torch
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from src import metrics
|
from src import metrics
|
||||||
from src.metrics import divergence
|
from src.metrics import divergence, evaluate_histogram
|
||||||
from sklearn.model_selection import GridSearchCV
|
from sklearn.model_selection import GridSearchCV
|
||||||
from sklearn.neighbors import KernelDensity
|
from sklearn.neighbors import KernelDensity
|
||||||
|
|
||||||
@@ -20,39 +20,43 @@ def test_kl_divergence():
|
|||||||
d3 = divergence(p, q, type='kl', n_components=3)
|
d3 = divergence(p, q, type='kl', n_components=3)
|
||||||
print(d3)
|
print(d3)
|
||||||
assert isinstance(d3, float)
|
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):
|
def test_evaluate_histogram():
|
||||||
idx = np.digitize(x, bin_edges)
|
N = 10000
|
||||||
mask = np.logical_and(np.less(0, idx), np.less(idx, len(bin_edges)))
|
p = torch.randn(N)
|
||||||
r = np.zeros_like(x)
|
q = 1.0 + 2.0 * torch.randn(2*N)
|
||||||
r[mask] = hist[idx[mask] - 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 = 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__':
|
if __name__ == '__main__':
|
||||||
# p = torch.randn(10)
|
p = torch.randn(10)
|
||||||
# q = 1.0 + 2.0 * torch.randn(10)
|
q = 1.0 + 2.0 * torch.randn(10)
|
||||||
|
|
||||||
# p = p.unsqueeze(-1)
|
p = p.unsqueeze(-1)
|
||||||
# q = q.unsqueeze(-1)
|
q = q.unsqueeze(-1)
|
||||||
# p_hist, p_edges = np.histogram(p.unsqueeze(-1), bins='auto', density=True)
|
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)
|
q_hist, q_edges = np.histogram(q.unsqueeze(-1), bins='auto', density=True)
|
||||||
# # px = p_hist[np.digitize(p, p_edges) - 1]
|
# px = p_hist[np.digitize(p, p_edges) - 1]
|
||||||
|
|
||||||
# # qx = q_hist[np.digitize(p, q_edges) - 1]
|
# qx = q_hist[np.digitize(p, q_edges) - 1]
|
||||||
# px = evaluate_histogram(p, p_hist, p_edges)
|
px = evaluate_histogram(p, p_hist, p_edges)
|
||||||
# qx = evaluate_histogram(q, 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()
|
|
||||||
|
|||||||
Reference in New Issue
Block a user