calpit 0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
calpit/__init__.py ADDED
@@ -0,0 +1,9 @@
1
+ <<<<<<< before updating
2
+ from .diagnostics_and_calibration import CalPit
3
+
4
+ __all__ = ["CalPit"]
5
+ =======
6
+ from .example_module import greetings, meaning
7
+
8
+ __all__ = ["greetings", "meaning"]
9
+ >>>>>>> after updating
calpit/_version.py ADDED
@@ -0,0 +1,16 @@
1
+ # file generated by setuptools_scm
2
+ # don't change, don't track in version control
3
+ TYPE_CHECKING = False
4
+ if TYPE_CHECKING:
5
+ from typing import Tuple, Union
6
+ VERSION_TUPLE = Tuple[Union[int, str], ...]
7
+ else:
8
+ VERSION_TUPLE = object
9
+
10
+ version: str
11
+ __version__: str
12
+ __version_tuple__: VERSION_TUPLE
13
+ version_tuple: VERSION_TUPLE
14
+
15
+ __version__ = version = '0.1'
16
+ __version_tuple__ = version_tuple = (0, 1)
calpit/datasets.py ADDED
File without changes
@@ -0,0 +1,305 @@
1
+ from pathlib import Path
2
+ import numpy as np
3
+ import torch
4
+ from torch.utils.data import TensorDataset, DataLoader
5
+ from scipy.interpolate import PchipInterpolator
6
+ from tqdm import trange
7
+
8
+
9
+ from calpit.nn.models import MLP
10
+ from calpit.nn.utils import count_parameters, RandomDataset, EarlyStopping
11
+ from calpit.metrics import probability_integral_transform
12
+ from calpit.utils import trapz_grid
13
+
14
+
15
+ class CalPit:
16
+ def __init__(self, model, input_dim=None, hidden_layers=None, **args):
17
+ """
18
+ Initializes an instance of the CalPit Class.
19
+
20
+ Args:
21
+ model (str or torch.nn.Module): The model to be used to learn the conditional PIT.
22
+ Can be any pytorch model that outputs a value between 0 and 1.
23
+ A string with the name of an inbuilt can also be provided. Currently supports: `mlp`
24
+
25
+ input_dim (int, optional): The input dimension of the model. Defaults to None.
26
+ hidden_layers (list, optional): A list of hidden layer sizes for the MLP models. Defaults to None.
27
+ """
28
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
29
+ if model == "mlp":
30
+ self.model = MLP(input_dim + 1, hidden_layers, 1).to(self.device)
31
+ else:
32
+ self.model = model.to(self.device)
33
+
34
+ count_parameters(self.model)
35
+
36
+ self.training_loss = None
37
+ self.validation_bce = None
38
+ self.val_loss_min = None
39
+
40
+ def fit(
41
+ self,
42
+ x_calib,
43
+ y_calib=None,
44
+ cde_calib=None,
45
+ y_grid=None,
46
+ pit_calib=None,
47
+ oversample=1,
48
+ n_cov_val=201,
49
+ patience=20,
50
+ n_epochs=1000,
51
+ lr=0.001,
52
+ weight_decay=1e-5,
53
+ batch_size=2048,
54
+ frac_train=0.9,
55
+ lr_decay=0.99,
56
+ trace_func=print,
57
+ seed=299792458,
58
+ num_workers=1,
59
+ checkpt_path="_results/checkpoint_.pt",
60
+ ):
61
+ """
62
+ Train the model using the calibration data.
63
+
64
+ Args:
65
+ x_calib (numpy.ndarray): The input features for calibration data.
66
+ y_calib (numpy.ndarray, optional): The target values for calibration data.
67
+ cde_calib (numpy.ndarray, optional): The conditional density estimates for calibration.
68
+ y_grid (numpy.ndarray, optional): The grid of target values for calibration.
69
+ pit_calib (numpy.ndarray, optional): The probability integral transforms for the given CDEs evaluated at y_calib.
70
+ Either pit_calib or y_calib, cde_calib and y_grid must be provided.
71
+ oversample (int, optional): The oversampling factor for the training data. Default is 1.
72
+ This is used to upsample the number of coverage values used for training.
73
+ n_cov_val (int, optional): The number of coverage values to use for validation. Default is 201.
74
+ patience (int, optional): The number of epochs to wait for improvement in validation loss before early stopping. Default is 20.
75
+ n_epochs (int, optional): The maximum number of epochs for training. Default is 1000.
76
+ lr (float, optional): The initial learning rate for the optimizer (AdamW). Default is 0.001.
77
+ weight_decay (float, optional): The weight decay for the optimizer. Default is 1e-5.
78
+ batch_size (int, optional): The batch size for training and validation. Default is 2048.
79
+ frac_train (float, optional): The fraction of data to use for training.
80
+ The rest is used for the validation set used to determine when to stop training. Default is 0.9.
81
+ lr_decay (float, optional): The learning rate decay factor for the rule,
82
+ learning_rate(epoch) = lr*lr_decay ** epoch. Default is 0.99.
83
+ trace_func (function, optional): The function used for printing training progress. Default is print.
84
+ seed (int, optional): The random seed for reproducibility. Default is 299792458.
85
+ num_workers (int, optional): The number of CPU worker threads for data loading. Default is 1.
86
+ checkpt_path (str, optional): The path to save the checkpoint of the best model. Default is "_results/checkpoint_.pt".
87
+
88
+ Returns:
89
+ torch.nn.Module: The trained model.
90
+ """
91
+ # method implementation
92
+ if pit_calib is None:
93
+ if y_calib is None or cde_calib is None or y_grid is None:
94
+ raise ValueError("Either pit_calib or, y_calib, cde_calib and y_grid must be provided")
95
+ pit_calib = probability_integral_transform(cde_calib, y_grid, y_calib)
96
+
97
+ cov_grid = np.linspace(0.001, 0.999, n_cov_val)
98
+ # Split into train and valid sets
99
+ train_size = int(frac_train * len(x_calib))
100
+ valid_size = len(x_calib) - train_size
101
+
102
+ rnd_idx = np.random.default_rng(seed=seed).permutation(len(x_calib))
103
+ x_train_rnd = x_calib[rnd_idx[:train_size]]
104
+ x_val_rnd = x_calib[rnd_idx[train_size:]]
105
+ pit_train_rand = pit_calib[rnd_idx[:train_size]]
106
+ pit_val_rand = pit_calib[rnd_idx[train_size:]]
107
+
108
+ # Creat randomized Data set for training
109
+ trainset = RandomDataset(x_train_rnd, pit_train_rand, oversample=oversample)
110
+
111
+ # Create static dataset for validation
112
+ feature_val = torch.cat(
113
+ [
114
+ torch.Tensor(np.repeat(cov_grid, len(x_val_rnd)))[:, None],
115
+ torch.Tensor(np.tile(x_val_rnd, (n_cov_val, 1))),
116
+ ],
117
+ dim=-1,
118
+ )
119
+ target_val = torch.Tensor(
120
+ np.tile(pit_val_rand, n_cov_val) <= np.repeat(cov_grid, len(x_val_rnd))
121
+ ).float()[:, None]
122
+
123
+ validset = TensorDataset(feature_val, target_val)
124
+
125
+ # Create Data loader
126
+ train_dataloader = DataLoader(trainset, batch_size=batch_size, shuffle=True, num_workers=num_workers)
127
+ valid_dataloader = DataLoader(validset, batch_size=batch_size, shuffle=False, num_workers=num_workers)
128
+
129
+ # Initialize the Model and optimizer, etc.
130
+ training_loss = []
131
+ validation_bce = []
132
+
133
+ # Optimizer
134
+ optimizer = torch.optim.AdamW(self.model.parameters(), lr=lr, weight_decay=weight_decay)
135
+ # Use lr decay
136
+ schedule_rule = lambda epoch: lr_decay**epoch
137
+ scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=schedule_rule)
138
+ checkpt_path = Path(checkpt_path)
139
+ checkpt_path.parent.mkdir(parents=True, exist_ok=True)
140
+ early_stopping = EarlyStopping(
141
+ patience=patience, verbose=True, path=checkpt_path, trace_func=trace_func
142
+ )
143
+
144
+ # Training loop for all epochs
145
+ for epoch in range(1, n_epochs + 1):
146
+ training_loss_batch = []
147
+ validation_mse_batch = []
148
+ validation_bce_batch = []
149
+
150
+ # Training loop per epoch
151
+ self.model.train() # prep model for training
152
+ for batch, (feature, target) in enumerate(train_dataloader, start=1):
153
+ feature = feature.to(self.device)
154
+ target = target.to(self.device)
155
+
156
+ # Zero your gradients for every batch!
157
+ optimizer.zero_grad()
158
+ # Make predictions for this batch
159
+ output = self.model(feature.float())
160
+
161
+ # Compute the loss and its gradients
162
+
163
+ loss_fn = torch.nn.BCELoss(reduction="sum")
164
+ loss = loss_fn(torch.clamp(torch.squeeze(output), min=0.0, max=1.0), torch.squeeze(target))
165
+
166
+ loss.backward()
167
+ # Adjust learning weights
168
+ optimizer.step()
169
+
170
+ # record training loss
171
+ training_loss_batch.append(loss.item())
172
+
173
+ # Validation
174
+ self.model.eval() # prep model for evaluation
175
+
176
+ for feature, target in valid_dataloader:
177
+
178
+ feature = feature.to(self.device)
179
+ target = target.to(self.device)
180
+
181
+ # forward pass: compute predicted outputs by passing inputs to the model
182
+ output = self.model(feature.float())
183
+
184
+ # calculate the loss
185
+ mse = ((output - target.float()) ** 2).sum()
186
+ # record validation loss
187
+ validation_mse_batch.append(mse.item())
188
+
189
+ criterion = torch.nn.BCELoss(reduction="sum")
190
+ bce = criterion(torch.clamp(torch.squeeze(output), min=0, max=1), torch.squeeze(target))
191
+ validation_bce_batch.append(bce.item())
192
+
193
+ # calculate average loss over an epoch
194
+ train_loss_epoch = np.sum(training_loss_batch) / (train_size * oversample)
195
+ valid_bce_epoch = np.sum(validation_bce_batch) / (valid_size * n_cov_val)
196
+ training_loss.append(train_loss_epoch)
197
+ validation_bce.append(valid_bce_epoch)
198
+
199
+ epoch_len = len(str(n_epochs))
200
+ # print training/validation statistics
201
+ msg = (
202
+ f"[{epoch:>{epoch_len}}/{n_epochs:>{epoch_len}}] | "
203
+ + f"train_loss: {train_loss_epoch:.5f} |"
204
+ + f"valid_bce: {valid_bce_epoch:.5f} | "
205
+ )
206
+
207
+ trace_func(msg)
208
+
209
+ # change the lr
210
+ scheduler.step()
211
+
212
+ # early_stopping needs the validation loss to check if it has decresed,
213
+ # and if it has, it will make a checkpoint of the current model
214
+ early_stopping(valid_bce_epoch, self.model)
215
+
216
+ if early_stopping.early_stop:
217
+ print("Early stopping")
218
+ break
219
+
220
+ # # load the last checkpoint with the best model
221
+ self.model.load_state_dict(torch.load(checkpt_path))
222
+ self.training_loss = np.array(training_loss)
223
+ self.validation_bce = np.array(validation_bce)
224
+ self.val_loss_min = early_stopping.val_loss_min
225
+ return self.model
226
+
227
+ def predict(self, x_test, cov_grid, batch_size=2048):
228
+ """
229
+ Predicts the conditional PIT values for the given test data and coverage grid.
230
+
231
+ Args:
232
+ x_test (numpy.ndarray): The input features of the test data.
233
+ cov_grid (numpy.ndarray): The coverage grid at which the PIT values are to be evaluated.
234
+ batch_size (int, optional): The batch size for prediction. Defaults to 2048.
235
+
236
+ Returns:
237
+ numpy.ndarray: The predicted conditional PIT values.
238
+ """
239
+ self.model.eval()
240
+ self.model.to(self.device)
241
+
242
+ pred_pit = []
243
+ n_test = len(x_test)
244
+ n_cov = len(cov_grid)
245
+ n_batches = (n_test - 1) // batch_size + 1
246
+
247
+ for i in trange(n_batches):
248
+ x = x_test[i * batch_size : (i + 1) * batch_size]
249
+ if cov_grid.ndim == 1:
250
+ with torch.no_grad():
251
+ pred_pit_batch = (
252
+ self.model(
253
+ torch.Tensor(
254
+ np.hstack([np.repeat(cov_grid, len(x))[:, None], np.tile(x, (n_cov, 1))])
255
+ ).to(self.device)
256
+ )
257
+ .detach()
258
+ .cpu()
259
+ .numpy()
260
+ .reshape(n_cov, -1)
261
+ .T
262
+ )
263
+ elif cov_grid.ndim == 2:
264
+ c = cov_grid[i * batch_size : (i + 1) * batch_size]
265
+ with torch.no_grad():
266
+ pred_pit_batch = (
267
+ self.model(
268
+ torch.Tensor(
269
+ np.hstack([np.ravel(c)[:, None], np.repeat(x, c.shape[1], axis=0)])
270
+ ).to(self.device)
271
+ )
272
+ .detach()
273
+ .cpu()
274
+ .numpy()
275
+ .reshape(len(x), -1)
276
+ )
277
+
278
+ pred_pit_batch[pred_pit_batch < 0] = 0
279
+ pred_pit_batch[pred_pit_batch > 1] = 1
280
+ pred_pit.extend(pred_pit_batch)
281
+ return np.array(pred_pit)
282
+
283
+ def transform(self, x_test, cde_test, y_grid, batch_size=2048):
284
+ """
285
+ Transforms the input CDEs for a test data set to calibrated CDEs.
286
+
287
+ Args:
288
+ x_test (array-like): The input features of the test data.
289
+ cde_test (array-like): The initial CDEs for the test data that is to be transformed.
290
+ y_grid (array-like): The grid of values for the CDEs.
291
+ batch_size (int, optional): The batch size for prediction. Defaults to 2048.
292
+
293
+ Returns:
294
+ numpy.ndarray: The transformed CDEs for the given.
295
+ """
296
+ cdf_test = trapz_grid(cde_test, y_grid)
297
+ cdf_test_new = self.predict(x_test, cov_grid=cdf_test, batch_size=batch_size)
298
+ cdf_funct = PchipInterpolator(y_grid, cdf_test_new, extrapolate=True, axis=1)
299
+ pdf_func = cdf_funct.derivative(1)
300
+ cde_test_new = pdf_func(y_grid)
301
+ return cde_test_new
302
+
303
+ def fit_transform(self, **args):
304
+ """Fit the model and transform the data in one go"""
305
+ raise NotImplementedError
@@ -0,0 +1,14 @@
1
+ """An example module containing simplistic methods under benchmarking."""
2
+
3
+ import random
4
+ import time
5
+
6
+
7
+ def runtime_computation():
8
+ """Runtime computation consuming between 0 and 5 seconds."""
9
+ time.sleep(random.uniform(0, 5))
10
+
11
+
12
+ def memory_computation():
13
+ """Memory computation for a random list up to 512 samples."""
14
+ return [0] * random.randint(0, 512)
calpit/metrics.py ADDED
@@ -0,0 +1,150 @@
1
+ import numpy as np
2
+
3
+
4
+ def cde_loss(cde_estimates: np.ndarray, y_grid: np.ndarray, y_test: np.ndarray) -> tuple:
5
+ """
6
+ Calculates conditional density estimation loss on holdout data.
7
+
8
+ Args:
9
+ cde_estimates (numpy.array): An array where each row is a density estimate on y_grid.
10
+ z_grid (numpy.array): An array of the grid points at which cde_estimates is evaluated.
11
+ z_test (numpy.array): An array of the true y values corresponding to the rows of cde_estimates.
12
+
13
+ Returns:
14
+ tuple: A tuple containing the loss and the standard error of the loss.
15
+
16
+ Raises:
17
+ ValueError: If the dimensions of the input tensors are not compatible.
18
+ """
19
+
20
+ if len(y_test.shape) == 1:
21
+ y_test = y_test.reshape(-1, 1)
22
+ if len(y_grid.shape) == 1:
23
+ y_grid = y_grid.reshape(-1, 1)
24
+
25
+ n_obs, n_grid = cde_estimates.shape
26
+ n_samples, feats_samples = y_test.shape
27
+ n_grid_points, feats_grid = y_grid.shape
28
+
29
+ if n_obs != n_samples:
30
+ raise ValueError(
31
+ f"Number of samples in CDEs should be the same as in z_test.Currently {n_obs} and {n_samples}."
32
+ )
33
+ if n_grid != n_grid_points:
34
+ raise ValueError(
35
+ f"Number of grid points in CDEs should be the same as in z_grid. Currently {n_grid} and {n_grid_points}."
36
+ )
37
+
38
+ if feats_samples != feats_grid:
39
+ raise ValueError(
40
+ f"Dimensionality of test points and grid points need to coincise. Currently {feats_samples} and {feats_grid}."
41
+ )
42
+
43
+ integrals = np.trapz(cde_estimates**2, np.squeeze(y_grid), axis=1)
44
+
45
+ nn_ids = np.argmin(np.abs(y_grid - y_test.T), axis=0)
46
+ likeli = cde_estimates[(tuple(np.arange(n_samples)), tuple(nn_ids))]
47
+
48
+ losses = integrals - 2 * likeli
49
+ loss = np.mean(losses)
50
+ se_error = np.std(losses, axis=0) / (n_obs**0.5)
51
+
52
+ return loss, se_error
53
+
54
+
55
+ def kolmogorov_smirnov_statistic(cdf_test: np.ndarray, cdf_ref: np.ndarray) -> np.ndarray:
56
+ """
57
+ Calculate the Kolmogorov-Smirnov statistic between two cumulative distribution functions (CDFs).
58
+
59
+ Parameters:
60
+ cdf_test (np.ndarray): CDF of the test distribution.
61
+ cdf_ref (np.ndarray): CDF of the reference distribution on the same grid.
62
+
63
+ Returns:
64
+ np.ndarray: The Kolmogorov-Smirnov statistic.
65
+
66
+ """
67
+ ks = np.max(np.abs(cdf_test - cdf_ref), axis=-1)
68
+
69
+ return ks
70
+
71
+
72
+ def cramer_von_mises_statistic(cdf_test: np.ndarray, cdf_ref: np.ndarray) -> np.ndarray:
73
+ """
74
+ Calculates the Cramer-von Mises statistic between two cumulative distribution functions (CDFs).
75
+
76
+ Args:
77
+ cdf_test (np.ndarray): CDF of the test distribution.
78
+ cdf_ref (np.ndarray): CDF of the reference distribution on the same grid.
79
+
80
+ Returns:
81
+ np.ndarray: The Cramer-von Mises statistic.
82
+
83
+ """
84
+ diff = (cdf_test - cdf_ref) ** 2
85
+
86
+ cvm2 = np.trapz(diff, cdf_ref, axis=-1)
87
+ return np.sqrt(cvm2)
88
+
89
+
90
+ def anderson_darling_statistic(cdf_test: np.ndarray, cdf_ref: np.ndarray, n_tot: int = 1) -> np.ndarray:
91
+ """
92
+ Calculates the Anderson-Darling statistic between two cumulative distribution functions (CDFs).
93
+
94
+ Args:
95
+ cdf_test (np.ndarray): CDF of the test distribution (1D array).
96
+ cdf_ref (np.ndarray): CDF of the reference distribution on the same grid (1D array).
97
+ n_tot (int): Scaling factor equal to the number of PDFs used to construct ECDF.
98
+
99
+ Returns:
100
+ np.ndarray: The Anderson-Darling statistic.
101
+
102
+ """
103
+ num = (cdf_test - cdf_ref) ** 2
104
+ den = cdf_ref * (1 - cdf_ref)
105
+
106
+ ad2 = n_tot * np.trapz((num / den), cdf_ref, axis=-1)
107
+ return np.sqrt(ad2)
108
+
109
+
110
+ def probability_integral_transform(cde: np.ndarray, y_grid: np.ndarray, y_test: np.ndarray) -> np.ndarray:
111
+ """
112
+ Calculates the Probability Integral Transform (PIT) based on Conditional Density Estimates (CDE).
113
+
114
+ Args:
115
+ cde (np.ndarray): A numpy array of conditional density estimates.
116
+ Each row corresponds to an observation, each column corresponds to a grid point.
117
+ y_grid (np.ndarray): A numpy array of the grid points at which cde is evaluated.
118
+ y_test (np.ndarray): A numpy array of the true y values corresponding to the rows of cde.
119
+
120
+ Returns:
121
+ np.ndarray: A numpy array of PIT values.
122
+
123
+ Raises:
124
+ ValueError: If the number of samples in cde is not the same as in y_test,
125
+ or if the number of grid points in cde is not the same as in y_grid.
126
+
127
+ """
128
+ # flatten the input arrays to 1D
129
+ y_grid = np.ravel(y_grid)
130
+ y_test = np.ravel(y_test)
131
+
132
+ # Sanity checks
133
+ nrow_cde, ncol_cde = cde.shape
134
+ n_samples = y_test.shape[0]
135
+ n_grid_points = y_grid.shape[0]
136
+
137
+ if nrow_cde != n_samples:
138
+ raise ValueError(
139
+ f"Number of samples in CDEs should be the same as in z_test. Currently {nrow_cde} and {n_samples}."
140
+ )
141
+ if ncol_cde != n_grid_points:
142
+ raise ValueError(
143
+ f"Number of grid points in CDEs should be the same as in z_grid. Currently {nrow_cde} and {n_grid_points}."
144
+ )
145
+
146
+ # Vectorized implementation using masked arrays
147
+ pit = np.ma.masked_array(cde, (y_grid > y_test[:, np.newaxis]))
148
+ pit = np.trapz(pit, y_grid)
149
+
150
+ return np.array(pit)
calpit/nn/__init__.py ADDED
@@ -0,0 +1,2 @@
1
+ from .umnn import MonotonicNN # noqa
2
+ from .ispline_nn import IsplineNN # noqa
@@ -0,0 +1,102 @@
1
+ import torch
2
+ import torch.nn as nn
3
+ from splinebasis import ISplineBasis
4
+
5
+
6
+ class ISplineLayer(nn.Module):
7
+ def __init__(self, in_features, num_basis,dropout_p=0):
8
+ super().__init__()
9
+ self.in_features = in_features
10
+ self.num_basis = num_basis
11
+ self.coefs = nn.Sequential(nn.Linear(in_features, num_basis), nn.Softmax(dim=-1),nn.Dropout(p=dropout_p))
12
+ self.grid = torch.linspace(0, 1, 1000)
13
+ self.basis_vectors = ISplineBasis(
14
+ order=3, num_basis=num_basis, lower=0, upper=1, grid=self.grid
15
+ ).basis_vectors
16
+ self.basis_vectors = torch.from_numpy(self.basis_vectors)
17
+
18
+ # def init_weights(m):
19
+ # if isinstance(m, nn.Linear):
20
+ # torch.nn.init.kaiming_normal_(m.weight)
21
+ # m.bias.data.fill_(0.01)
22
+
23
+ # self.coefs.apply(init_weights)
24
+
25
+ def interp1d(self, x, y, x_new):
26
+ # 2. Find where in the original data, the values to interpolate
27
+ # would be inserted.
28
+ # Note: If x_new[n] == x[m], then m is returned by searchsorted.
29
+ # y = torch.moveaxis(y,axis,0)
30
+ # y = y.reshape((y.shape[0],-1))
31
+
32
+ x_new_indices = torch.searchsorted(x, x_new)
33
+
34
+ # 3. Clip x_new_indices so that they are within the range of
35
+ # self.x indices and at least 1. Removes mis-interpolation
36
+ # of x_new[n] = x[0]
37
+ x_new_indices = x_new_indices.clip(1, len(x) - 1)
38
+
39
+ # 4. Calculate the slope of regions that each x_new value falls in.
40
+ lo = x_new_indices - 1
41
+ hi = x_new_indices
42
+
43
+ x_lo = x[lo]
44
+ x_hi = x[hi]
45
+ y_lo = y[lo]
46
+ y_hi = y[hi]
47
+
48
+ # Note that the following two expressions rely on the specifics of the
49
+ # broadcasting semantics.
50
+ slope = (y_hi - y_lo) / (x_hi - x_lo)[:, None]
51
+
52
+ # 5. Calculate the actual value for each entry in x_new.
53
+ y_new = slope * (x_new - x_lo)[:, None] + y_lo
54
+
55
+ return y_new
56
+
57
+ def forward(self, x, alpha):
58
+ grid = self.grid.to(alpha)
59
+ basis_vectors = self.basis_vectors.to(alpha)
60
+ basis = self.interp1d(grid, basis_vectors, alpha)
61
+
62
+ # print(basis.shape)
63
+ # print(self.coefs(x).shape)
64
+ # print(self.coefs(x))
65
+ weighted_basis = self.coefs(x) * basis
66
+ # print(weighted_basis.shape)
67
+ return weighted_basis.sum(axis=-1)
68
+
69
+
70
+ class IsplineNN(nn.Module):
71
+ def __init__(self, input_dim, hidden_layers=[512, 512, 512],dropout_p=0.5, num_basis=10):
72
+ super().__init__()
73
+ self.all_layers = [input_dim + 1]
74
+ self.hidden_layers = hidden_layers
75
+ self.all_layers.extend(hidden_layers)
76
+ self.num_basis = num_basis
77
+ self.dropout_p = dropout_p
78
+ self.spline_layer = ISplineLayer(in_features=self.hidden_layers[-1], num_basis=self.num_basis,dropout_p=self.dropout_p)
79
+
80
+ self.mlp_layer_list = []
81
+ for i in range(len(self.all_layers) - 1):
82
+ self.mlp_layer_list.append(nn.Linear(self.all_layers[i], self.all_layers[i + 1]))
83
+ self.mlp_layer_list.append(nn.PReLU())
84
+
85
+ # self.mlp_layer_list.append(nn.Dropout(p=dropout_p))
86
+ self.mlp_layers = nn.Sequential(*self.mlp_layer_list)
87
+
88
+ def init_weights(m):
89
+ if isinstance(m, nn.Linear):
90
+ torch.nn.init.kaiming_normal_(m.weight)
91
+ m.bias.data.fill_(0.01)
92
+
93
+ self.mlp_layers.apply(init_weights)
94
+
95
+ def forward(self, x):
96
+ alpha = x[:, 0]
97
+
98
+ res = self.mlp_layers(x)
99
+
100
+ res = self.spline_layer(res, alpha)
101
+
102
+ return res
calpit/nn/models.py ADDED
@@ -0,0 +1,52 @@
1
+ import torch.nn as nn
2
+
3
+
4
+ class MLP(nn.Module):
5
+ """
6
+ Multi-Layer Perceptron (MLP) neural network model.
7
+
8
+ Args:
9
+ input_dim (int): The number of input features.
10
+ hidden_layers (list): A list of integers representing the number of units in each hidden layer.
11
+ output_dim (int): The number of output units. Defaults to 1.
12
+ sigmoid (bool): Whether to apply a sigmoid activation function to the output layer. Defaults to True.
13
+
14
+ Methods:
15
+ forward(x): Performs a forward pass through the MLP.
16
+
17
+ """
18
+
19
+ def __init__(self, input_dim, hidden_layers, output_dim=1, sigmoid=True):
20
+ super().__init__()
21
+ self.all_layers = [input_dim]
22
+ self.all_layers.extend(hidden_layers)
23
+ self.all_layers.append(output_dim)
24
+ self.layer_list = []
25
+
26
+ for i in range(len(self.all_layers) - 1):
27
+ self.layer_list.append(nn.Linear(self.all_layers[i], self.all_layers[i + 1]))
28
+ self.layer_list.append(nn.PReLU())
29
+ self.layer_list.pop()
30
+ if sigmoid:
31
+ self.layer_list.append(nn.Sigmoid())
32
+ self.layers = nn.Sequential(*self.layer_list)
33
+
34
+ def init_weights(m):
35
+ if isinstance(m, nn.Linear):
36
+ nn.init.kaiming_normal_(m.weight)
37
+ m.bias.data.fill_(0.01)
38
+
39
+ self.layers.apply(init_weights)
40
+
41
+ def forward(self, x):
42
+ """
43
+ Performs a forward pass through the MLP.
44
+
45
+ Args:
46
+ x (torch.Tensor): The input tensor.
47
+
48
+ Returns:
49
+ torch.Tensor: The output tensor.
50
+
51
+ """
52
+ return self.layers(x)