Wiener Attention

class wiener_attention.attention_mechanism.WienerSelfAttention(*args: Any, **kwargs: Any)[source]

Bases: Module

Wiener 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: Module

This 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:

  1. 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.

get_filter_shape(input_shape)[source]
identity(mesh, val=1, **kwargs)[source]
make_delta(shape)[source]
make_doubly_block(X)[source]

Makes Doubly Blocked Toeplitz of a matrix X [r, c]

make_penalty(shape, eta=0.0, penalty_function=None, std=None, device='cpu')[source]
make_toeplitz(a)[source]

Makes toeplitz matrix of a vector A

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]

rms(x)[source]
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

wiener_attention.model.make_bert_tokenizer()[source]

Initializes a BERT tokenizer.

Returns:

The tokenizer associated with the model.

Return type:

tokenizer