# Using PredictionsBasedSampler for Dataset Sampling in VERONA

> Filter VERONA datasets with PredictionsBasedSampler. Optimize sampling by selecting correct or incorrect model predictions for your analysis.

- Repository: [ADA research/verona](https://github.com/ada-research/verona)
- Tags: how-to-guide
- Published: 2026-02-23

---

**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`](https://github.com/ada-research/verona/blob/main/ada_verona/dataset_sampler/dataset_sampler.py#L22‑L35). It inherits from `DatasetSampler` ([`predictions_based_sampler.py#L23`](https://github.com/ada-research/verona/blob/main/ada_verona/dataset_sampler/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`](https://github.com/ada-research/verona/blob/main/ada_verona/dataset_sampler/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`](https://github.com/ada-research/verona/blob/main/ada_verona/database/dataset/experiment_dataset.py#L23‑L35)), 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:

```python
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`:

```python
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`](https://github.com/ada-research/verona/blob/main/examples/scripts/create_robustness_dist_pgd.py#L22‑L63), the sampler is instantiated and applied as follows:

```python
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`](https://github.com/ada-research/verona/blob/main/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 `DatasetSampler` abstract class to filter datasets based on model inference.
- It accepts two key parameters: `sample_correct_predictions` (boolean filter direction) and `top_k` (prediction ranking depth).
- The sampler automatically handles GPU device selection, model loading via `network.load_pytorch_model()`, and input reshaping via `network.get_input_shape()`.
- It returns a subsetted `ExperimentDataset` through the `get_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)](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)](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)](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.