Example: Neural Network Surrogate#
Model for subcutaneous autoinjectors#
We reproduce the surrogate-modeling example of Sree et al. (2023) using an NN surrogate of an expensive biomechanical model.
Background#
Autoinjectors are drug delivery devices that work kind of like an automated syringe. Basically, you put the device against the patient’s skin, press a button, and a spring within the device pushes the needle and drug is released into the patient.
Fig. 3 Generic autoinjector components. This original conceptual schematic is not to scale.#
To efficiently design an optimal autoinjector, you need a computational model of the injection process. To this end, Sree et al. (2023) present a biomechanical model for injections into subcutaneous tissue (i.e., the tissue right below your skin) via an autoinjector.
Fig. 4 Coupled mechanical and finite-element model of the injection process. Adapted from Sree et al. (2023).#
Fig. 5 Conceptual injection sequence. This original schematic shows no simulation results and is not to scale.#
The model, however, is expensive to evaluate—one evaluation takes multiple hours on several parallel processors. This is too slow for things like design optimization or uncertainty quantification. So we will speed it up by training a surrogate. Let’s get started.
Dataset#
To train the surrogate, we’ll use a dataset with thousands of evaluations of the expensive model (data provided by the authors of Sree et al. (2023)). Let’s import the data:
The model has 10 inputs and 4 outputs. The inputs are things like drug viscosity, needle size, tissue biomechanical properties, etc. The outputs are the maximum acceleration/deceleration of fluid in syringe, the total time of injection, and the depth of needle insertion at the onset of drug delivery.
Preprocessing#
Ideally, we want the input and output data to be standardized and more-or-less evenly distributed. This helps avoid numerical difficulties in training.
Let’s visualize the distribution of each input:
array([<Axes: title={'center': 'Viscosity'}>,
<Axes: title={'center': 'Fill volume'}>,
<Axes: title={'center': 'Air gap'}>,
<Axes: title={'center': 'Needle length'}>,
<Axes: title={'center': 'Needle diameter'}>,
<Axes: title={'center': 'Spring force'}>,
<Axes: title={'center': 'Spring constant'}>,
<Axes: title={'center': '$\\kappa_5$'}>,
<Axes: title={'center': '$\\kappa_6$'}>,
<Axes: title={'center': '$\\kappa_7$'}>], dtype=object)
They are evenly distributed, but not standardized (i.e., the means and standard deviations are not 0 and 1, respectively). Let’s standardize the input data:
input_transform = StandardScaler(X_train.values)
X_train_scaled = vmap(input_transform.forward)(X_train.values)
X_test_scaled = vmap(input_transform.forward)(X_test.values)
Now let’s visualize the distribution of each output:
Some of the outputs are skewed. Let’s apply a log transformation, then standardize the output data:
def partial_log_transform(x):
return jnp.hstack([x[0], jnp.log(x[1]), jnp.log(x[2]), x[3]])
def partial_exp_transform(y):
return jnp.array([y[0], jnp.exp(y[1]), jnp.exp(y[2]), y[3]])
output_transform = StandardScaler(y_train.values, pretransform_forward=partial_log_transform, pretransform_inverse=partial_exp_transform)
y_train_scaled = vmap(output_transform.forward)(y_train.values)
y_test_scaled = vmap(output_transform.forward)(y_test.values)
Okay, now let’s visualize the transformed inputs and outputs:
array([<Axes: title={'center': 'Needle displacement'}>,
<Axes: title={'center': 'Injection time'}>,
<Axes: title={'center': 'Maximum acceleration'}>,
<Axes: title={'center': 'Maximum deceleration'}>], dtype=object)
Nice! The input/output data are all more-or-less evenly distributed and standardized. Now on to building the neural network surrogate.
Creating the NN surrogate#
Let’s set aside some of the training data for on-the-fly validation during training:
X_train_scaled_subset, X_val_scaled = jnp.split(X_train_scaled, [int(0.9 * X_train_scaled.shape[0])])
y_train_scaled_subset, y_val_scaled = jnp.split(y_train_scaled, [int(0.9 * y_train_scaled.shape[0])])
The best choice of neural network architecture is problem dependent. We’ll just use a simple multilayer perceptron (MLP) with 5 hidden layers and 100 neurons per layer.
Define the training loop:
Here is how to train the surrogate model using 500 points from the training set:
LOSS_FREQ, SAVE_FREQ, PRINT_FREQ = 20, 100, 1000
MAX_ITERS = 5000
key, subkey = jrandom.split(key)
trained_mlp, train_losses, val_losses = train(
batch_size=128,
max_epochs=None,
max_iters=MAX_ITERS,
learning_rate=1e-3,
X_train=X_train_scaled_subset[:500],
y_train=y_train_scaled_subset[:500],
X_val=X_val_scaled,
y_val=y_val_scaled,
loss_freq=LOSS_FREQ,
save_freq=SAVE_FREQ,
print_freq=PRINT_FREQ,
key=subkey
)
Epoch 0 , iter 0 , train loss 0.9905, val loss: 1.0175
Epoch 250 , iter 1000 , train loss 0.0009, val loss: 0.0238
Epoch 500 , iter 2000 , train loss 0.0003, val loss: 0.0234
Epoch 750 , iter 3000 , train loss 0.0003, val loss: 0.0233
Epoch 1000, iter 4000 , train loss 0.0001, val loss: 0.0231
Epoch 1250, iter 5000 , train loss 0.0002, val loss: 0.0230
array([<Axes: xlabel='Iteration', ylabel='Loss'>], dtype=object)
We now have a trained surrogate model.
Surrogate diagnostics#
We use the test data to assess the surrogate’s accuracy.
Parity plot#
Let’s plot the predicted output against the true output (for the test dataset), along with the RMSE:
# Parity plot
y_train_pred = vmap(trained_mlp)(X_train_scaled)
y_test_pred = vmap(trained_mlp)(X_test_scaled)
# Root mean square error
rmse = lambda y, y_hat: jnp.sqrt(jnp.mean((y - y_hat)**2))
The closer the points are to the identity line, the more accurate the surrogate is. This looks pretty good for only training on 500 points. Let’s try to improve it with more training data.
Convergence with training-set size#
The more data you have available to train the surrogate, the better. In fact, you should see the prediction error go to zero as you increase the amount of training data (so long as the neural network is expressive enough and you train for long enough).
Let’s observe this convergence behavior by training the same neural network with different dataset sizes:
dataset_sizes = [100, 1000, 2500, 5000, 9000]
trained_mlps_N = {}
train_losses_N = {}
val_losses_N = {}
for N in dataset_sizes:
print(f"Training NN with dataset size {N}.\n")
key, subkey = jrandom.split(key)
trained_mlps_N[N], train_losses_N[N], val_losses_N[N] = train(
batch_size=128,
max_epochs=None,
max_iters=MAX_ITERS,
learning_rate=1e-3,
X_train=X_train_scaled_subset[:N],
y_train=y_train_scaled_subset[:N],
X_val=X_val_scaled,
y_val=y_val_scaled,
loss_freq=LOSS_FREQ,
save_freq=SAVE_FREQ,
print_freq=PRINT_FREQ,
key=subkey
)
print("\nTraining complete.\n\n")
Training NN with dataset size 100.
Epoch 0 , iter 0 , train loss 0.9435, val loss: 1.0165
Epoch 1000, iter 1000 , train loss 0.0000, val loss: 0.0940
Epoch 2000, iter 2000 , train loss 0.0000, val loss: 0.0934
Epoch 3000, iter 3000 , train loss 0.0000, val loss: 0.0925
Epoch 4000, iter 4000 , train loss 0.0000, val loss: 0.0917
Epoch 5000, iter 5000 , train loss 0.0001, val loss: 0.0917
Training complete.
Training NN with dataset size 1000.
Epoch 0 , iter 0 , train loss 0.9553, val loss: 1.0205
Epoch 125 , iter 1000 , train loss 0.0025, val loss: 0.0139
Epoch 250 , iter 2000 , train loss 0.0013, val loss: 0.0125
Epoch 375 , iter 3000 , train loss 0.0009, val loss: 0.0124
Epoch 500 , iter 4000 , train loss 0.0005, val loss: 0.0123
Epoch 625 , iter 5000 , train loss 0.0005, val loss: 0.0128
Training complete.
Training NN with dataset size 2500.
Epoch 0 , iter 0 , train loss 0.9168, val loss: 1.0283
Epoch 50 , iter 1000 , train loss 0.0052, val loss: 0.0088
Epoch 100 , iter 2000 , train loss 0.0025, val loss: 0.0076
Epoch 150 , iter 3000 , train loss 0.0019, val loss: 0.0067
Epoch 200 , iter 4000 , train loss 0.0018, val loss: 0.0063
Epoch 250 , iter 5000 , train loss 0.0009, val loss: 0.0060
Training complete.
Training NN with dataset size 5000.
Epoch 0 , iter 0 , train loss 1.0080, val loss: 1.0217
Epoch 25 , iter 1000 , train loss 0.0063, val loss: 0.0083
Epoch 50 , iter 2000 , train loss 0.0034, val loss: 0.0051
Epoch 75 , iter 3000 , train loss 0.0028, val loss: 0.0046
Epoch 100 , iter 4000 , train loss 0.0022, val loss: 0.0040
Epoch 125 , iter 5000 , train loss 0.0020, val loss: 0.0038
Training complete.
Training NN with dataset size 9000.
Epoch 0 , iter 0 , train loss 1.0497, val loss: 1.0190
Epoch 14 , iter 1000 , train loss 0.0053, val loss: 0.0072
Epoch 28 , iter 2000 , train loss 0.0033, val loss: 0.0050
Epoch 42 , iter 3000 , train loss 0.0025, val loss: 0.0039
Epoch 56 , iter 4000 , train loss 0.0018, val loss: 0.0034
Epoch 70 , iter 5000 , train loss 0.0019, val loss: 0.0032
Training complete.
Let’s repeat the parity plots for different amounts of training data:
Again, the more data, the better. The fit for \(N=9,000\) is quite good.
Let’s plot the RMSE against the dataset size:
rmse_i = []
for N in dataset_sizes:
y_test_pred = vmap(trained_mlps_N[N])(X_test_scaled)
rmse_i.append(vmap(rmse, 1)(y_test_scaled, y_test_pred))
rmse_i = jnp.array(rmse_i)
array([<Axes: xlabel='Dataset size', ylabel='RMSE'>], dtype=object)
For this problem, the surrogate is pretty good when trained on 2,000 (or more) points. We could probably further improve accuracy by training for longer or using a different NN architecture.
Sensitivity analysis with surrogate#
Now that we have a trained and tested surrogate, there are many useful things we can do with it. As an example, we’ll demonstrate using the surrogate to do Sobol sensitivity analysis.
First, let’s create a helper function for evaluating the unscaled surrogate model.
def _surrogate(x, nn_model):
x_scaled = input_transform.forward(x)
y_scaled = nn_model(x_scaled)
y = output_transform.inverse(y_scaled)
return y
# This is the unscaled surrogate.
surrogate = partial(_surrogate, nn_model=trained_mlps_N[9000])
Now, let’s create bounds for the inputs (based on physical intuition and/or literature values):
import SALib.sample.sobol as sobol
import SALib.analyze.sobol as analyze_sobol
input_bounds_dict = {
'viscosity_cP': [1.0, 20.0],
'fill_volume_mL': [1.0, 1.05],
'air_gap_height_mm': [4.0, 5.0],
'needle_length_mm': [8.0, 15.9],
'needle_diameter_mm': [0.133, 0.21],
'spring_force_N': [18.0, 36.0],
'spring_constant_N_per_mm': [150.0, 250.0],
'kappa5': [0.0, 3.0],
'kappa6': [0.0, 4.0],
'kappa7': [0.0, 0.1]
}
input_bounds = np.array([[input_bounds_dict[i][0], input_bounds_dict[i][1]] for i in INPUT_NAMES])
print('The input bounds are:\n', input_bounds)
The input bounds are:
[[1.00e+00 2.00e+01]
[1.00e+00 1.05e+00]
[4.00e+00 5.00e+00]
[8.00e+00 1.59e+01]
[1.33e-01 2.10e-01]
[1.80e+01 3.60e+01]
[1.50e+02 2.50e+02]
[0.00e+00 3.00e+00]
[0.00e+00 4.00e+00]
[0.00e+00 1.00e-01]]
Next, we create Sobol samples of the inputs and pass them through the surrogate:
problem = {
'num_vars': len(INPUT_NAMES),
'names': np.array(INPUT_NAMES),
'bounds': input_bounds
}
# The number of samples to generate (should be a power of 2).
N = 512
# Generate the samples.
sobol_samples = sobol.sample(problem, N, calc_second_order=False)
# Evaluate the surrogate model at the Sobol samples.
sobol_outputs = vmap(surrogate)(sobol_samples)
Finally, we calculate and plot the Sobol indices:
sobol_indices = {}
for i, o_name in enumerate(OUTPUT_NAMES):
sobol_indices[o_name] = analyze_sobol.analyze(problem, sobol_outputs[:, i], calc_second_order=False, print_to_console=False)
Excellent! We now know to which inputs the outputs are most sensitive. For example, we can see that the injection time is very sensitive to the drug viscosity and needle diameter, somewhat sensitive to needle length and injector spring force, and not sensitive to any other inputs. This information can be used, for example, to further study and understand the physics behind the model, or to investigate how identifiable each parameter is given an experimental dataset.
Remember that sensitivity analysis would not have been feasible without a surrogate model—the true physical model was just too expensive. With a surrogate however, we can do sensitivity analysis, design optimization, uncertainty quantification, etc. all at a reasonable computational cost.