torch_blue.vi.VITransformerEncoder

class torch_blue.vi.VITransformerEncoder(encoder_layer: VITransformerEncoderLayer, num_layers: int, norm: torch.nn.Module | None = None, return_log_probs: bool = True)

Bases: torch_blue.vi.base.VIModule

TransformerEncoder is a stack of N encoder layers.

Equivalent of nn.TransformerEncoder with variational inference. See its documentation for usage.

forward(src: torch.Tensor, mask: torch.Tensor | None = None, src_key_padding_mask: torch.Tensor | None = None, is_causal: bool | None = None) torch.Tensor

Pass the input through the encoder layers in turn.

See documentation of nn.TransformerEncoder for details.

This implementation also currently does not support the torch fastpath.