Using PredictionsBasedSampler for Dataset Sampling in VERONA
The PredictionsBasedSampler filters VERONA datasets by model predictions, selecting either correctly or incorrectly classified samples based on the sample_correct_predictions parameter.
The VERONA framework (ada-research/verona) provides a modular architecture for neural network robustness analysis. Its sampling subsystem allows you to create targeted subsets of data using the PredictionsBasedSampler class, which implements the abstract DatasetSampler interface to select data points according to a model's inference results.
What is PredictionsBasedSampler?
PredictionsBasedSampler is the default concrete implementation of the DatasetSampler abstract base class defined in ada_verona/dataset_sampler/dataset_sampler.py#L22. It inherits from DatasetSampler (predictions_based_sampler.py#L23) and implements the required sample(network, dataset) method to return an ExperimentDataset containing only the selected data points.
The sampler is designed to work with any class implementing the Network interface (such as ONNXNetwork or PyTorch wrappers) and any dataset adhering to the ExperimentDataset contract.
Core Sampling Logic
The sample method (predictions_based_sampler.py#L40‑L75) executes an eight-step algorithm to filter your dataset:
- Device selection – Automatically selects CUDA if available via
torch.device. - Model initialization – Calls
network.load_pytorch_model()to load the concrete network implementation and sets it to evaluation mode. - Dataset iteration – Iterates over the input
ExperimentDataset(experiment_dataset.py#L23), retrieving raw tensor data and labels from eachDataPoint. - Input reshaping – Reshapes data to match
network.get_input_shape(). - Inference – Performs a forward pass without gradient computation to obtain output logits.
- Top-k extraction – Uses
torch.topkto retrieve the highest-scoring output indices, converting them back to CPU for comparison. - Selection logic – Compares the true label against the top-k predictions, keeping the point based on the
sample_correct_predictionsflag. - Subset construction – Calls
dataset.get_subset(selected_indices)to return a newExperimentDatasetcontaining only the retained samples.
Configuration Parameters
The constructor accepts two parameters that control filtering behavior:
| Parameter | Type | Description |
|---|---|---|
sample_correct_predictions |
bool |
When True, retains data points where the true label appears in the top-k predictions. When False, retains only incorrect predictions. |
top_k |
int |
Number of highest-scoring output entries examined for correctness (default is 1). |
Practical Implementation Examples
Basic Usage with Image Datasets
The following example demonstrates loading an ONNX network and filtering an image dataset to keep only correctly classified samples:
from pathlib import Path
from ada_verona.dataset_sampler.predictions_based_sampler import PredictionsBasedSampler
from ada_verona.database.machine_learning_model.onnx_network import ONNXNetwork
from ada_verona.database.dataset.image_file_dataset import ImageFileDataset
# Load an ONNX network
network = ONNXNetwork("examples/example_experiment/data/networks/mnist-net_256x2.onnx")
# Build dataset from folder and CSV labels
dataset = ImageFileDataset(
image_folder=Path("examples/example_experiment/data/images"),
label_file=Path("examples/example_experiment/data/image_labels.csv"),
)
# Create sampler for correct predictions (default)
sampler = PredictionsBasedSampler(sample_correct_predictions=True, top_k=1)
# Filter the dataset
filtered_dataset = sampler.sample(network, dataset)
print(f"Original size: {len(dataset)}")
print(f"Filtered size: {len(filtered_dataset)}")
Filtering Incorrect Predictions
To analyze failure cases or adversarial susceptibility, invert the selection logic by setting sample_correct_predictions=False:
sampler = PredictionsBasedSampler(sample_correct_predictions=False, top_k=3)
filtered_dataset = sampler.sample(network, dataset)
# filtered_dataset contains only samples where the true label is NOT in the top-3 predictions
Integration in Robustness Workflows
The sampler integrates directly into robustness distribution scripts. In examples/scripts/create_robustness_dist_pgd.py#L22‑L63, the sampler is instantiated and applied as follows:
from ada_verona.dataset_sampler.predictions_based_sampler import PredictionsBasedSampler
dataset_sampler = PredictionsBasedSampler(sample_correct_predictions=True)
sampled_dataset = dataset_sampler.sample(network, dataset)
Error Handling and Testing
The sampler propagates exceptions raised during tensor reshaping or model inference, allowing detection of misconfigured networks. The test suite in tests/test_dataset_sampler/test_prediction_based_sampler.py validates correct sampling behavior, incorrect prediction filtering, and proper error handling through methods like test_sample_network_prediction_failure.
Summary
- PredictionsBasedSampler implements the
DatasetSamplerabstract class to filter datasets based on model inference. - It accepts two key parameters:
sample_correct_predictions(boolean filter direction) andtop_k(prediction ranking depth). - The sampler automatically handles GPU device selection, model loading via
network.load_pytorch_model(), and input reshaping vianetwork.get_input_shape(). - It returns a subsetted
ExperimentDatasetthrough theget_subset()method, preserving only data points matching the specified prediction criteria. - Full source implementation resides in [
ada_verona/dataset_sampler/predictions_based_sampler.py](https://github.com/ada-research/verona/blob/main/ada_verona/dataset_sampler/predictions_based_sampler.py).
Frequently Asked Questions
What is the relationship between PredictionsBasedSampler and DatasetSampler?
DatasetSampler is the abstract base class defining the sample(network, dataset) contract in [ada_verona/dataset_sampler/dataset_sampler.py](https://github.com/ada-research/verona/blob/main/ada_verona/dataset_sampler/dataset_sampler.py). PredictionsBasedSampler is the default concrete implementation that uses model predictions to perform the actual filtering, inheriting from this base class and implementing the required sampling logic.
How does the top_k parameter influence sampling behavior?
The top_k parameter determines how many of the highest-scoring output logits the sampler examines when checking for correctness. With top_k=1, only the argmax prediction is considered. With top_k=3, the true label is considered correct if it appears anywhere in the three highest predictions, expanding the set of retained samples when sample_correct_predictions=True.
Can PredictionsBasedSampler work with custom network implementations?
Yes, provided your custom network implements the Network interface defined in [ada_verona/database/machine_learning_model/network.py](https://github.com/ada-research/verona/blob/main/ada_verona/database/machine_learning_model/network.py). The sampler requires load_pytorch_model() and get_input_shape() methods to function correctly during the sampling pipeline.
What happens if network inference fails during sampling?
The sampler raises any exceptions that occur during the reshaping or forward pass operations. This allows calling scripts or tests to catch configuration errors, such as mismatched input shapes or corrupted model files, rather than silently producing empty or incorrect datasets.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →