Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 40 additions & 6 deletions src/nessai/proposal/flowproposal/flowproposal.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import datetime
import logging
import math
import warnings

import numpy as np
Expand All @@ -26,6 +27,21 @@
logger = logging.getLogger(__name__)


def _clip_weights(weights: np.ndarray) -> np.ndarray:
"""Clip the largest ``ceil(sqrt(N))`` weights to their mean."""
weights = weights.copy()
num_clip = math.ceil(math.sqrt(len(weights)))
if not num_clip:
return weights
if num_clip >= len(weights):
clip_idx = np.arange(len(weights))
else:
clip_idx = np.argpartition(-weights, num_clip - 1)[:num_clip]
weights[clip_idx] = weights[clip_idx].mean()
weights /= weights.mean()
return weights


class FlowProposal(BaseFlowProposal):
"""Proposal that samples in latent space using the trained flow.

Expand Down Expand Up @@ -120,6 +136,7 @@ def __init__(
truncation_method=None,
truncation_methods=None,
truncation_kwargs=None,
clip_population_weights=False,
**kwargs,
):
super().__init__(model, poolsize=poolsize, **kwargs)
Expand All @@ -130,6 +147,7 @@ def __init__(
drawsize,
latent_prior=latent_prior,
latent_temperature=latent_temperature,
clip_population_weights=clip_population_weights,
)

self._truncation_scheme = TruncationScheme()
Expand Down Expand Up @@ -250,6 +268,7 @@ def configure_population(
drawsize,
latent_prior=None,
latent_temperature=None,
clip_population_weights=False,
) -> None:
"""Configure settings related to population."""
if drawsize is None:
Expand All @@ -272,6 +291,17 @@ def configure_population(
self.drawsize = drawsize
self.latent_prior = "flow"
self.latent_temperature = latent_temperature
self.clip_population_weights = clip_population_weights

def _get_population_log_weights(self, log_weights) -> np.ndarray:
"""Return log-weights used in the rejection step during population."""
log_weights = np.asarray(log_weights, dtype=float)
log_weights = log_weights - np.nanmax(log_weights)
if not self.clip_population_weights:
return log_weights

weights = _clip_weights(np.exp(log_weights))
return np.log(weights) - np.log(weights.max())

def configure_truncation(
self,
Expand Down Expand Up @@ -424,7 +454,6 @@ def populate(
log_n_expected = -np.inf
n_proposed = 0
log_weights = np.empty(0)
log_constant = -np.inf
n_accepted = 0
accept = None

Expand Down Expand Up @@ -471,8 +500,10 @@ def populate(
if self.accumulate_weights:
samples = np.concatenate([samples, x])
log_weights = np.concatenate([log_weights, log_w])
log_constant = max(np.nanmax(log_w), log_constant)
log_n_expected = logsumexp(log_weights - log_constant)
log_weights_rejection = self._get_population_log_weights(
log_weights
)
log_n_expected = logsumexp(log_weights_rejection)

logger.debug(
"Drawn %s - n expected: %s / %s",
Expand All @@ -483,13 +514,13 @@ def populate(

if log_n_expected >= log_n:
log_u = np.log(self.rng.random(len(log_weights)))
accept = (log_weights - log_constant) > log_u
accept = log_weights_rejection > log_u
n_accepted = np.sum(accept)
if n_proposed > max_samples:
logger.warning("Reached max samples (%s)", max_samples)
break
else:
log_w -= log_w.max()
log_w = self._get_population_log_weights(log_w)
log_u = np.log(self.rng.random(len(log_w)))
accept = log_w > log_u
n_accept_batch = accept.sum()
Expand All @@ -503,8 +534,11 @@ def populate(

if self.accumulate_weights:
if accept is None or len(accept) != len(samples):
log_weights_rejection = self._get_population_log_weights(
log_weights
)
log_u = np.log(self.rng.random(len(log_weights)))
accept = (log_weights - log_constant) > log_u
accept = log_weights_rejection > log_u
logger.debug("Total number of samples: %s", samples.size)
n_accepted = np.sum(accept)
self.x = samples[accept][:n_samples]
Expand Down
74 changes: 74 additions & 0 deletions src/nessai/proposal/flowproposal/truncation.py
Original file line number Diff line number Diff line change
Expand Up @@ -395,6 +395,34 @@ def apply_after_backward(self, proposal, x, log_q, z):
return get_subset_arrays(keep, x, log_q, z)


class LogProposalThresholdTruncation(BaseTruncationRule):
"""Truncate samples using the current log-proposal threshold."""

name = "log_proposal_threshold"
_transient_defaults = {"_threshold": np.nan}

def __init__(self, quantile: float = 0.05) -> None:
self.quantile = float(quantile)

@property
def threshold(self) -> float:
return self._threshold

def prepare(self, proposal, worst_point, radius=None):
live_points = proposal.training_data.copy()
log_q = proposal.forward_pass(live_points)[1]
self._threshold = float(np.quantile(log_q, self.quantile))

def apply_after_backward(self, proposal, x, log_q, z):
keep = log_q > self.threshold
logger.debug(
"Accepting %s / %s samples above log-proposal threshold",
int(keep.sum()),
len(x),
)
return get_subset_arrays(keep, x, log_q, z)


class LikelihoodThresholdTruncation(BaseTruncationRule):
"""Truncate samples using the current likelihood threshold."""

Expand Down Expand Up @@ -429,10 +457,56 @@ def apply_after_likelihood(self, proposal, x, log_q, z):
return get_subset_arrays(keep, x, log_q, z)


class LogWeightThresholdTruncation(BaseTruncationRule):
"""Truncate samples using the current log-weight threshold.

The threshold is set based on the current live points.
"""

name = "log_weight_threshold"
_transient_defaults = {"_threshold": np.nan}

def __init__(
self, quantile: float = 0.05, enlargement: float = 0.0
) -> None:
super().__init__()
self.quantile = quantile
self.enlargement = enlargement

@property
def threshold(self) -> float:
return self._threshold

def prepare(self, proposal, worst_point, radius=None):
live_points = proposal.training_data.copy()
log_q = proposal.forward_pass(live_points)[1]
log_w = proposal.compute_weights(
live_points,
log_q=log_q,
)

self._threshold = np.quantile(log_w, self.quantile) - self.enlargement

def apply_after_backward(self, proposal, x, log_q, z):
log_w = proposal.compute_weights(
x,
log_q=log_q,
)
keep = log_w > self.threshold
logger.debug(
"Accepting %s / %s samples above log-weight threshold",
int(keep.sum()),
len(x),
)
return get_subset_arrays(keep, x, log_q, z)


TRUNCATION_REGISTRY = {
"latent_radius": LatentRadiusTruncation,
"min_log_q": MinLogQTruncation,
"likelihood_threshold": LikelihoodThresholdTruncation,
"log_proposal_threshold": LogProposalThresholdTruncation,
"weights_threshold": LogWeightThresholdTruncation,
}


Expand Down
Loading