Wiener Attention¶
- class wiener_attention.attention_mechanism.WienerSelfAttention(*args: Any, **kwargs: Any)[source]¶
Bases:
ModuleWiener Self-Attention mechanism for transformer models.
This class implements a custom self-attention mechanism using Wiener filters for similarity computation between query and key vectors.
- Parameters:
config (object) – Configuration object containing model parameters.
similarity_function (callable) – Function to compute similarity between query and key vectors.
gamma (float, optional) – Parameter for the similarity function. Defaults to 0.1.
- num_attention_heads¶
Number of attention heads.
- Type:
int
- attention_head_size¶
Size of each attention head.
- Type:
int
- all_head_size¶
Total size of all attention heads.
- Type:
int
- query¶
Linear layer for query transformation.
- Type:
nn.Linear
- key¶
Linear layer for key transformation.
- Type:
nn.Linear
- value¶
Linear layer for value transformation.
- Type:
nn.Linear
- dropout¶
Dropout layer for regularization.
- Type:
nn.Dropout
- similarity_function¶
Function to compute similarity.
- Type:
callable
- gamma¶
Parameter for the similarity function.
- Type:
float
- forward(hidden_states, attention_mask=None, head_mask=None, encoder_hidden_states=None, encoder_attention_mask=None, past_key_value=None, output_attentions=True)[source]¶
Forward pass of the Wiener Self-Attention mechanism.
- Parameters:
hidden_states (torch.Tensor) – Input hidden states.
attention_mask (torch.Tensor, optional) – Attention mask. Defaults to None.
head_mask (torch.Tensor, optional) – Mask for attention heads. Defaults to None.
encoder_hidden_states (torch.Tensor, optional) – Hidden states from encoder. Defaults to None.
encoder_attention_mask (torch.Tensor, optional) – Attention mask for encoder. Defaults to None.
past_key_value (tuple, optional) – Cached key and value projection states. Defaults to None.
output_attentions (bool, optional) – Whether to output attention weights. Defaults to True.
- Returns:
- A tuple containing:
context_layer (torch.Tensor): Output context layer.
attention_probs (torch.Tensor): Attention probabilities if output_attentions is True.
- Return type:
tuple
- transpose_for_scores(x)[source]¶
Transpose and reshape the input tensor for attention score calculation.
- Parameters:
x (torch.Tensor) – Input tensor of shape (batch_size, seq_length, all_head_size).
- Returns:
Reshaped tensor of shape (batch_size, num_attention_heads, seq_length, attention_head_size).
- Return type:
torch.Tensor
- class wiener_attention.wiener_metric.WienerSimilarityMetric(method='fft', filter_dim=2, filter_scale=2, reduction='mean', mode='reverse', penalty_function=None, store_filters=False, epsilon=0.0001, std=0.0001, clamp_min=None, corr_norm=None, rel_epsilon=False)[source]¶
Bases:
ModuleThis loss was implemented by Cruz et al in the paper Convolve and Conquer: Data Comparison with Wiener Filters. I have modified the implementation in two ways:
Ensured that the identity function can work with any number of dimensions
2. Made the function not return an average, but rather a value for each sample pair in the batch.
Original implementation: https://github.com/dpelacani/WienerLoss
The AWLoss class implements the adaptive Wiener criterion, which aims to compare two data samples through a convolutional filter. The methodology is inspired by the paper `Adaptive Waveform Inversion: Theory`_ (Warner and Guasch, 2014).
A matching filter w can be computed such that it transforms a targetsignal p into the data d under an L2 norm principle:
g = || p*w - d ||^2
Let ‘Z’ be the Toeplitz matrix formulation of p such that the equation above is equivalent to: g = || Zw - d ||^2
Minimizing this functional: dgdw = Z^T (Zw - d) dgdw –> 0 : w = (Z^T @ Z)^(-1) @ Z^T @ d
To stabilize the matrix inversion, an amount is added to the diagonal of (Z^T @ Z) based on a value epsilon such that the inverted matrix is (Z^T @ Z) + max(diagonal(Z^T @ Z)) * epsilon
In 2D, convolving p with w (or w with p) is equivalent to the matrix vector multiplication Zd where Z is the doubly block Toeplitz of the reconstructed image P and w is the flattened array of the 2D kernel W.
Therefore, the system is equivalent to solving || Zw - d ||^2, and the solution to w is given by w = (Z^T @ Z + max(diagonal(Z^T @ Z)) * epsilon)^(-1) @ Z^T @ d
This composes the direct method.
Alternatively, convolution can be performed in the frequency domain with multiplication and division operations. This tends to be much more computationally efficient.
The criterion is evaluated through a symmetrical monotonically decreasing function T and a dirac delta function that rewards when the filter kernel w is close to the identity kernel, and penalizes otherwise.
f = 1/2 ||T * (v - delta)||^2
- Parameters:
method – “fft” for Fast Fourier Transform or “direct” for the Levinson-Durbin recurssion algorithm. Defaults to “fft”
optional – “fft” for Fast Fourier Transform or “direct” for the Levinson-Durbin recurssion algorithm. Defaults to “fft”
filter_dim – the dimensionality of the filter. This parameter should be upper-bounded by the dimensionality of the data. If data is 3-dimensional and filter_dim is set to 2, one filter is computed per channel dimension assuming format [B, NC, H , W]. Current implementation only supports filter dimensions for 1D, 2D and 3D. Defaults to 2
optional – the dimensionality of the filter. This parameter should be upper-bounded by the dimensionality of the data. If data is 3-dimensional and filter_dim is set to 2, one filter is computed per channel dimension assuming format [B, NC, H , W]. Current implementation only supports filter dimensions for 1D, 2D and 3D. Defaults to 2
filter_scale – the scale of the filters compared to the size of the data. Defaults to 2
optional – the scale of the filters compared to the size of the data. Defaults to 2
reduction – specifies the reduction to apply to the output, “mean” or “sum”. Defaults to mean
optional – specifies the reduction to apply to the output, “mean” or “sum”. Defaults to mean
mode – “forward” or “reverse” computation of the filter. For details of the difference, refer to the original paper. Default “reverse”
optional – “forward” or “reverse” computation of the filter. For details of the difference, refer to the original paper. Default “reverse”
penalty_function – the penalty function to apply to the filter. If None, the penalty function is the identity. Takes “identity”, “gaussian” or custom penalty function. Default None
optional – the penalty function to apply to the filter. If None, the penalty function is the identity. Takes “identity”, “gaussian” or custom penalty function. Default None
std – the standard deviation of the gaussian when penalty_function=”gaussian”. Mean is always zero. Default None
optional – the standard deviation of the gaussian when penalty_function=”gaussian”. Mean is always zero. Default None
store_filters – whether to store the filters in memory, useful for debugging. Option to store the filers before or after normalisation with “norm” and “unorm”. Default False.
optional – whether to store the filters in memory, useful for debugging. Option to store the filers before or after normalisation with “norm” and “unorm”. Default False.
epsilon – the stabilization value to compute the filter. Default 1e-4.
optional – the stabilization value to compute the filter. Default 1e-4.
clamp_min – filters are clipped to this minimum value after computation. If None, operation is disabled. Default none
optional – filters are clipped to this minimum value after computation. If None, operation is disabled. Default none
- forward(recon, target, epsilon=None, gamma=0.0, eta=0.0)[source]¶
> The function takes in a reconstructed signal, a target signal, and a few other parameters, and returns the loss
- Args
- recon
the reconstructed signal
- target
the target signal
- epsilon, optional
the stabilization value to compute the filter. If passed, overwrites the class attribute of same name. Default None.
- gamma, optional
noise to add to both target and reconstructed signals for training stabilization. Default 0.
- eta, optional
noise to add to penalty function. Default 0.
- multigauss(mesh, mean, covmatrix)[source]¶
Multivariate gaussian of N dimensions on evenly spaced hypercubed grid. Mesh should be stacked along the last axis E.g. for a 3D gaussian of 20 grid points in each axis mesh should be of shape (20, 20, 20, 3)
- pad_signal(x, shape, val=0)[source]¶
x must be a multichannel signal of shape [batch_size, nchannels, width, height]
- wiener(x, y, fs, epsilon=1e-09)[source]¶
calculates the optimal least squares convolutional Wiener filter that transforms signal x into signal y using the direct Toeplitz matrix implementation
- wienerfft(x, y, fs, prwh=1e-09)[source]¶
George Strong (geowstrong@gmail.com) calculates the optimal least squares convolutional Wiener filter that transforms signal x into signal y using FFT
- wiener_attention.model.make_bert_model(wiener_attention=None, wiener_similarity=None, num_labels=2, gamma=None)[source]¶
Initializes a single-layer BERT model for sequence classification with a custom attention module.
- Parameters:
num_labels (int) – Number of labels for the classification task.
- Returns:
The tokenizer associated with the model. model: A BERT model instance for sequence classification with custom attention.
- Return type:
tokenizer