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.VIModuleTransformerEncoder is a stack of N encoder layers.
Equivalent of
nn.TransformerEncoderwith 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.TransformerEncoderfor details.This implementation also currently does not support the torch fastpath.