# Loading PyTorch Networks in the VERONA Verification Pipeline: A Complete Guide

> Learn to load PyTorch networks in the VERONA verification pipeline using the PyTorchNetwork class for seamless integration with adversarial verification modules.

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

---

**VERONA loads PyTorch models through the `PyTorchNetwork` class, which wraps raw `torch.nn.Module` instances in a `TorchModelWrapper` to handle input reshaping and device placement before feeding them into adversarial verification modules.**

The `ada-research/verona` repository provides a formal verification framework for neural networks that treats PyTorch architectures as first-class citizens. Loading PyTorch networks in the VERONA verification pipeline requires understanding a small but essential abstraction layer that bridges standard `torch.nn.Module` definitions with the framework's `VerificationContext` and `VerificationModule` interfaces.

## How VERONA Handles PyTorch Model Loading

VERONA decouples raw PyTorch models from the verification logic through an abstract `Network` base class. This design allows the pipeline to treat PyTorch and ONNX models interchangeably while handling framework-specific details like tensor reshaping and device management internally.

### The Network Abstraction Layer ([`network.py`](https://github.com/ada-research/verona/blob/main/network.py))

The abstract base class `Network` 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) defines the common interface that all network types must implement. Key methods include `load_pytorch_model()`, `get_input_shape()`, and `get_name()`. This abstraction ensures that verification modules like `AttackEstimationModule` can request a PyTorch-compatible model without knowing whether the underlying implementation is a native PyTorch network or a converted ONNX graph.

### PyTorchNetwork: The Concrete Implementation ([`pytorch_network.py`](https://github.com/ada-research/verona/blob/main/pytorch_network.py))

The `PyTorchNetwork` class in [`ada_verona/database/machine_learning_model/pytorch_network.py`](https://github.com/ada-research/verona/blob/main/ada_verona/database/machine_learning_model/pytorch_network.py) provides the concrete implementation for PyTorch models. It stores three critical pieces of state: the raw `torch.nn.Module` instance, the expected input shape tuple, and a network identifier.

The `load_pytorch_model()` method implements lazy initialization. On first call, it moves the stored model to the appropriate compute device, switches it to `eval()` mode, and constructs a `TorchModelWrapper` instance. Subsequent calls return the cached wrapper, avoiding redundant initialization overhead during iterative verification processes.

## The TorchModelWrapper: Input Normalization and Device Handling

Located in [`ada_verona/database/machine_learning_model/torch_model_wrapper.py`](https://github.com/ada-research/verona/blob/main/ada_verona/database/machine_learning_model/torch_model_wrapper.py), the `TorchModelWrapper` class solves two common friction points when integrating PyTorch models into verification pipelines: input tensor reshaping and device placement.

The wrapper inherits from `torch.nn.Module`, allowing it to function as a drop-in replacement for the underlying model. Its `forward()` method accepts either PyTorch tensors or NumPy arrays, automatically converting the latter. It then reshapes the input to match the expected `input_shape` defined during `PyTorchNetwork` initialization, and ensures the tensor resides on the same device as the wrapped model before forwarding.

This abstraction frees verification modules from handling shape mismatches or device management logic, allowing them to focus on adversarial attack implementation and property verification.

## Loading Models in the Verification Pipeline

The actual consumption of PyTorch networks occurs within the `VerificationContext` and concrete `VerificationModule` implementations, particularly the `AttackEstimationModule`.

### VerificationContext and AttackEstimationModule

The `VerificationContext` class couples three essential components: a `Network` instance (in this case, a `PyTorchNetwork`), a `DataPoint` containing the input tensor and ground-truth label, and a property generator that defines the verification specification.

When `AttackEstimationModule.verify()` is invoked, it retrieves the wrapped PyTorch model by calling `verification_context.network.load_pytorch_model()`. The module then moves the input data to the correct device, executes the adversarial attack (such as PGD), and evaluates whether the perturbed input causes a misclassification, producing a SAT/UNSAT result.

## Practical Implementation Examples

### Defining and Storing a PyTorch Model

To integrate a custom architecture into VERONA, instantiate `PyTorchNetwork` with the raw model and its expected input dimensions:

```python
import torch.nn as nn
from pathlib import Path
from ada_verona.database.machine_learning_model.pytorch_network import PyTorchNetwork

class SimpleCNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = nn.Conv2d(1, 8, kernel_size=3)
        self.fc = nn.Linear(8 * 26 * 26, 10)

    def forward(self, x):
        x = torch.relu(self.conv(x))
        return self.fc(x.view(x.size(0), -1))

raw_model = SimpleCNN()
input_shape = (1, 28, 28)

network = PyTorchNetwork(
    model=raw_model,
    input_shape=input_shape,
    name="simple_cnn",
    path=Path("models/simple_cnn.pt")
)

```

### Running Verification with Loaded Models

Use the `AttackEstimationModule` to verify robustness properties against the loaded network:

```python
from ada_verona.verification_module.attack_estimation_module import AttackEstimationModule
from ada_verona.verification_module.attacks.pgd_attack import PGDAttack
from ada_verona.database.verification_context import VerificationContext
from ada_verona.database.dataset.data_point import DataPoint
import torch

# Prepare data

data_point = DataPoint(label=3, data=torch.randn(1, 28, 28))

# Build context

verification_context = VerificationContext(
    network=network,
    data_point=data_point,
    property_generator=One2AnyPropertyGenerator()
)

# Configure attack and estimator

attack = PGDAttack(steps=40, step_size=0.01)
estimator = AttackEstimationModule(attack=attack, top_k=1)

# Execute verification

result = estimator.verify(verification_context, epsilon=0.1)
print(result)

```

### Direct TorchModelWrapper Usage

For custom verification scripts outside the standard module framework, instantiate the wrapper directly:

```python
from ada_verona.database.machine_learning_model.torch_model_wrapper import TorchModelWrapper
import numpy as np

wrapper = TorchModelWrapper(
    torch_model=raw_model,
    input_shape=(1, 28, 28)
)

# NumPy arrays are automatically converted

np_input = np.random.rand(1, 28, 28).astype(np.float32)
output = wrapper(np_input)  # Returns torch.Tensor

```

## Summary

- **VERONA abstracts PyTorch models** through the `PyTorchNetwork` class, which implements the abstract `Network` interface defined in [`network.py`](https://github.com/ada-research/verona/blob/main/network.py).
- **Lazy loading** occurs via `load_pytorch_model()`, which constructs a `TorchModelWrapper` on first access, handling device placement and evaluation mode automatically.
- **Input normalization** is managed by `TorchModelWrapper` in [`torch_model_wrapper.py`](https://github.com/ada-research/verona/blob/main/torch_model_wrapper.py), which reshapes tensors and converts NumPy arrays to PyTorch tensors before forwarding.
- **Verification modules** such as `AttackEstimationModule` consume wrapped models through the `VerificationContext`, enabling seamless adversarial robustness checking without manual tensor management.

## Frequently Asked Questions

### What is the difference between PyTorchNetwork and TorchModelWrapper?

`PyTorchNetwork` is a metadata container and factory class that stores the raw `torch.nn.Module`, its expected input shape, and network identifier. It implements the abstract `Network` interface and provides the `load_pytorch_model()` method. `TorchModelWrapper` is a concrete `torch.nn.Module` subclass that handles runtime tensor operations—reshaping inputs, converting NumPy arrays to tensors, and ensuring device consistency—before forwarding data to the underlying model.

### How does VERONA handle device placement for PyTorch models?

Device management is handled automatically within the `load_pytorch_model()` method of `PyTorchNetwork`. When the wrapper is first instantiated, the method detects the available device (CPU or CUDA), moves the stored model to that device, switches it to `eval()` mode, and passes the device context to `TorchModelWrapper`. The wrapper then ensures all incoming tensors are moved to the same device before forward passes, eliminating manual `.to(device)` calls in verification modules.

### Can I use custom PyTorch architectures with VERONA?

Yes, any architecture that inherits from `torch.nn.Module` is compatible. VERONA does not impose architectural constraints—convolutional networks, transformers, or custom layers are all supported. You simply instantiate your model, pass it to `PyTorchNetwork` along with the expected `input_shape` tuple, and the framework handles the rest. The `TorchModelWrapper` will correctly reshape inputs to match your specified dimensions regardless of the internal architecture.

### Where does the model loading happen in the verification pipeline?

Model loading occurs lazily within verification modules when they call `verification_context.network.load_pytorch_model()`. For example, in `AttackEstimationModule.verify()`, the module retrieves the wrapped model from the context's network property. The actual instantiation of `TorchModelWrapper` happens on the first call to `load_pytorch_model()` in `PyTorchNetwork`, after which the cached wrapper is reused for subsequent verification queries to avoid redundant model initialization overhead.