Source code for fastbnns.simulation.generators
"""Collections of data generators to, e.g., aid dataset creation."""
from collections.abc import Callable
from typing import Any
import torch
[docs]
def sample_simulator(
simulator: Callable, simulator_kwargs_generator: dict, simulator_kwargs: dict
) -> dict:
"""Call simulator with randomly sampled inputs.
Args:
simulator: Callable that accepts keyword arguments and returns simulated
data.
simulator_kwargs_generator: Dictionary whose keys define keyword
arguments of `simulator` and whose values are Callable and return
valid values of associated keyword arguments.
simulator_kwargs: Fixed keyword arguments to be merged with arguments
generated by `simulator_kwargs_generator` before passing to `simulator`.
"""
generated_kwargs = {
key: value_gen() for key, value_gen in simulator_kwargs_generator.items()
}
simulator_kwargs = simulator_kwargs | generated_kwargs
return {"input": simulator_kwargs, "output": simulator(**simulator_kwargs)}
[docs]
class Generator(torch.nn.Module):
"""Generator to allow calling simulator with stochastic inputs."""
def __init__(
self,
simulator: Callable,
simulator_kwargs: dict,
simulator_kwargs_generator: dict,
) -> None:
"""Initializer for stochastic simulator dataset.
Args:
simulator: Callable that accepts keyword arguments and returns simulated
data.
simulator_kwargs: Fixed keyword arguments to be merged with arguments
generated by `simulator_kwargs_generator` before passing to `simulator`.
simulator_kwargs_generator: Dictionary whose keys define keyword
arguments of `simulator` and whose values are Callable and return
valid values of associated keyword arguments.
"""
super().__init__()
self.simulator = simulator
self.simulator_kwargs = simulator_kwargs
self.simulator_kwargs_generator = simulator_kwargs_generator
[docs]
def forward(self) -> Any:
return sample_simulator(
simulator=self.simulator,
simulator_kwargs=self.simulator_kwargs,
simulator_kwargs_generator=self.simulator_kwargs_generator,
)
if __name__ == "__main__":
import matplotlib.pyplot as plt
import numpy as np
import polynomials
# Example usage of sample_simulator(): simulate random samples from a polynomial.
simulator = polynomials.polynomial
simulator_kwargs_generator = {"x": lambda: np.random.random() - 0.5}
simulator_kwargs = {"order": 2, "coefficients": (0.0, 0.0, 1.0)}
sim_gen = Generator(
simulator=simulator,
simulator_kwargs=simulator_kwargs,
simulator_kwargs_generator=simulator_kwargs_generator,
)
n_samples = 100
x = np.empty(n_samples)
y = np.empty(n_samples)
for n in range(n_samples):
data_dict = sim_gen()
x[n] = data_dict["input"]["x"]
y[n] = data_dict["output"]
fig, ax = plt.subplots()
ax.plot(x, y, ".")
plt.show()