Loading PyTorch Networks in the VERONA Verification Pipeline: A Complete Guide
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)
The abstract base class Network in 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)
The PyTorchNetwork class in 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, 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:
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:
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:
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
PyTorchNetworkclass, which implements the abstractNetworkinterface defined innetwork.py. - Lazy loading occurs via
load_pytorch_model(), which constructs aTorchModelWrapperon first access, handling device placement and evaluation mode automatically. - Input normalization is managed by
TorchModelWrapperintorch_model_wrapper.py, which reshapes tensors and converts NumPy arrays to PyTorch tensors before forwarding. - Verification modules such as
AttackEstimationModuleconsume wrapped models through theVerificationContext, 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.
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 →