| """Generate cumulative regret plot for the Volterra memory example.""" |
|
|
| import functools |
| import numpy as np |
|
|
| import flax.linen as nn |
| import jax.numpy as jnp |
| import jax.random as jr |
| from jax import config, jit, value_and_grad |
| from jax.lax import scan |
| from jax.scipy.linalg import expm |
| import matplotlib.pyplot as plt |
| import optax |
| from tqdm import tqdm |
|
|
| from mpm import get_system_params, legsval, whitesignal, wray_and_green_output |
|
|
|
|
|
|
| def log_downsample(x, y, num=2000): |
| idx = jnp.unique(jnp.logspace(0, jnp.log10(len(x)-1), num).astype(int)) |
| return x[idx], y[idx] |
|
|
|
|
| def initialize_predictor(n: int, *, seed: int = 0): |
| """Initialize parameters and optimizer for the linear/quadratic predictors.""" |
|
|
| key = jr.PRNGKey(seed) |
| key1, key2 = jr.split(key) |
| params = { |
| "w1": 1e-2 * jr.normal(key1, (n,)), |
| "w2": 0.0 * jr.normal(key2, (n, n)), |
| "b": jnp.zeros(1), |
| } |
| optimizer = optax.sgd(learning_rate=3e-2) |
| opt_state = optimizer.init(params) |
| return params, optimizer, opt_state |
|
|
|
|
| def initialize_mlp_predictor(n: int, *, hidden_dim: int = 128, seed: int = 0): |
| """Initialize parameters and optimizer for the MLP predictor.""" |
| key = jr.PRNGKey(seed) |
| model = MlpReadout(hidden_dim=hidden_dim) |
| params = model.init(key, jnp.zeros((n,))) |
| optimizer = optax.sgd(learning_rate=1e-2) |
| opt_state = optimizer.init(params) |
| return model, params, optimizer, opt_state |
|
|
|
|
| class MlpReadout(nn.Module): |
| hidden_dim: int = 128 |
|
|
| @nn.compact |
| def __call__(self, state): |
| hidden = nn.Dense(self.hidden_dim)(state) |
| hidden = jnp.tanh(hidden) |
| pred = nn.Dense(1)(hidden) |
| return jnp.squeeze(pred, axis=-1) |
|
|
|
|
| def loss_func_linear(params, state, out_signal): |
| w1, b = params["w1"], params["b"] |
| pred = b[0] + jnp.inner(state, w1) |
| return jnp.mean((pred - out_signal) ** 2), pred |
|
|
|
|
| def loss_func_quadratic(params, state, out_signal): |
| w1, w2, b = params["w1"], params["w2"], params["b"] |
| pred = b[0] + jnp.inner(state, w1) |
| pred = pred + jnp.sum(w2 * jnp.outer(state, state)) |
| return jnp.mean((pred - out_signal) ** 2), pred |
|
|
|
|
| def loss_func_mlp(params, state, out_signal, *, model): |
| pred = model.apply(params, state) |
| return jnp.mean((pred - out_signal) ** 2), pred |
|
|
|
|
| def run_regret_experiment(): |
| dt = 1e-2 |
|
|
| |
| n = 64 |
| measure = "legs" |
| params, _, _ = get_system_params(measure, n) |
| A, b = params |
| timescale = 3 * dt / 0.08 |
| A, b = A / timescale, b / timescale |
| A_d = expm(dt * A) |
| b_d = jnp.linalg.solve(A, A_d @ b - b) |
| A, b = jnp.asarray(A_d), jnp.asarray(b_d) |
|
|
| loss_funcs = [loss_func_linear, loss_func_quadratic, None] |
| labels = ["Linear", "Quadratic", "MLP"] |
| colors = ["peru", "mediumseagreen", "steelblue"] |
|
|
| fig = plt.figure(figsize=(7, 2.5)) |
| ax_regret = fig.add_subplot(1, 3, 1) |
| plt.sca(ax_regret) |
|
|
| quadratic_params = None |
| num_trials = 5 |
| all_cumulative = {label: [] for label in labels} |
|
|
| pbar = tqdm(range(num_trials)) |
|
|
| for trial_idx in pbar: |
| np.random.seed(trial_idx) |
| input_data = jnp.asarray(whitesignal(1e7 * dt, dt, freq=10)) |
| output_data = jnp.asarray(wray_and_green_output(np.asarray(input_data))) |
| signals = jnp.stack([input_data, output_data], axis=1) |
|
|
| for iter_idx, (loss_func, label, color) in enumerate(zip(loss_funcs, labels, colors)): |
| seed = trial_idx * len(labels) + iter_idx |
| if label == "MLP": |
| model, params, optimizer, opt_state = initialize_mlp_predictor(n, seed=seed) |
| loss_func_i = functools.partial(loss_func_mlp, model=model) |
| else: |
| params, optimizer, opt_state = initialize_predictor(n, seed=seed) |
| loss_func_i = loss_func |
| val_grad_loss = jit(value_and_grad(loss_func_i, has_aux=True)) |
|
|
| def step(carry, signals): |
| in_x, out_x = signals |
| state, params, opt_state = carry |
|
|
| |
| state = A @ state + in_x * b |
|
|
| |
| (loss_val, pred), grads = val_grad_loss(params, state, out_x) |
| updates, opt_state = optimizer.update(grads, opt_state) |
| params = optax.apply_updates(params, updates) |
|
|
| return (state, params, opt_state), (loss_val, pred) |
|
|
| initial = (jnp.zeros(n), params, opt_state) |
| (_, params, _), (loss_vals, _) = scan(step, initial, signals) |
| all_cumulative[label].append(jnp.cumsum(loss_vals)) |
|
|
| if label == "Quadratic" and trial_idx == 0: |
| quadratic_params = params |
|
|
| steps = jnp.arange(len(all_cumulative[labels[0]][0])) |
| for label, color in zip(labels, colors): |
| curves = jnp.stack(all_cumulative[label]) |
| mean = jnp.mean(curves, axis=0) |
| sem = jnp.std(curves, axis=0) / jnp.sqrt(num_trials) |
|
|
| x_ds, mean_ds = log_downsample(steps, mean) |
| _, sem_ds = log_downsample(steps, sem) |
|
|
| ax_regret.loglog(x_ds, mean_ds, label=label, c=color) |
| ax_regret.fill_between(x_ds, mean_ds - sem_ds, mean_ds + sem_ds, |
| color=color, alpha=0.3) |
|
|
| plt.xlabel("Step") |
| plt.xlim(10, None) |
| plt.ylim(1e-1, None) |
| plt.ylabel("Cumulative Error") |
| plt.legend() |
|
|
| |
| ax = fig.add_subplot(1, 3, 2) |
| plt.sca(ax) |
| tau_vals = dt * jnp.arange(50) |
| a, m, k = 2.0, 0.3, 0.08 |
| mu = lambda t: a / m * jnp.exp(-k * t) * jnp.sin(m * t) |
| true_filter = mu(jnp.arange(50)) |
| true_kernel = 4e-3 * jnp.outer(true_filter, true_filter) |
| vmax = 0.06 |
| plt.imshow( |
| true_kernel, |
| cmap="coolwarm", |
| extent=[0, 49, 0, 49], |
| origin="lower", |
| vmax=vmax, |
| vmin=-vmax, |
| ) |
| ax.set_title("True Kernel") |
| ax.set_xlabel(r"$\tau_1$") |
| ax.set_ylabel(r"$\tau_2$") |
|
|
| |
| ax = fig.add_subplot(1, 3, 3) |
| plt.sca(ax) |
| if quadratic_params is not None: |
| w2 = quadratic_params["w2"] |
| p_tau = dt * 1.0 / timescale * jnp.exp(-1.0 / timescale * tau_vals) |
|
|
| legsvals = jnp.zeros((len(tau_vals), n)) |
| for i in range(n): |
| c = np.zeros(n) |
| c[i] = 1.0 |
| legsvals = legsvals.at[:, i].set(legsval(np.asarray(tau_vals / timescale), c)) |
|
|
| legsvals = legsvals[None, :, None] * legsvals[:, None, :, None] |
| legsvals = ( |
| legsvals |
| * w2[None, None] |
| * p_tau.reshape(1, -1, 1, 1) |
| * p_tau.reshape(-1, 1, 1, 1) |
| ) |
| kernel = jnp.sum(legsvals, axis=(-1, -2)) |
| else: |
| kernel = jnp.zeros((50, 50)) |
|
|
| plt.imshow( |
| kernel, |
| cmap="coolwarm", |
| extent=[0, 49, 0, 49], |
| origin="lower", |
| vmax=vmax, |
| vmin=-vmax, |
| ) |
| ax.set_title("Inferred Kernel") |
| ax.set_xlabel(r"$\tau_1$") |
| ax.set_ylabel(r"$\tau_2$") |
|
|
| plt.tight_layout() |
| plt.savefig("volterra_regret.png") |
| plt.savefig("volterra_regret.pdf") |
| plt.close("all") |
|
|
|
|
| if __name__ == "__main__": |
| run_regret_experiment() |
|
|