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:

  1. Device selection – Automatically selects CUDA if available via torch.device.
  2. Model initialization – Calls network.load_pytorch_model() to load the concrete network implementation and sets it to evaluation mode.
  3. Dataset iteration – Iterates over the input ExperimentDataset (experiment_dataset.py#L23), retrieving raw tensor data and labels from each DataPoint.
  4. Input reshaping – Reshapes data to match network.get_input_shape().
  5. Inference – Performs a forward pass without gradient computation to obtain output logits.
  6. Top-k extraction – Uses torch.topk to retrieve the highest-scoring output indices, converting them back to CPU for comparison.
  7. Selection logic – Compares the true label against the top-k predictions, keeping the point based on the sample_correct_predictions flag.
  8. Subset construction – Calls dataset.get_subset(selected_indices) to return a new ExperimentDataset containing 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

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →