CoolFace
Modelpublic

ControlNet/marlin_vit_base_ytf

sourceHugging Faceccupdated 1y agoView on Hugging Face
1likes28downloads
positional_embedding.py50 linesDownload Raw Back to root
1import torch2from torch import Tensor, nn3 4from .modules import Shape5 6 7class PositionalEmbedding(nn.Module):8 9    def __init__(self, input_shape: Shape, dropout_rate: float = 0.5, trainable: bool = True):10        super().__init__()11        self.input_shape = input_shape12        self.emb = nn.Parameter(torch.zeros(1, *input_shape), requires_grad=trainable)13        self.use_dropout = dropout_rate is not None and dropout_rate != 0.14        if self.use_dropout:15            self.dropout = nn.Dropout(dropout_rate)16 17    def forward(self, x: Tensor) -> Tensor:18        x = x + self.emb19        if self.use_dropout:20            x = self.dropout(x)21        return x22 23    @property24    def trainable(self):25        return self.emb.requires_grad26 27    @trainable.setter28    def trainable(self, value: bool):29        self.emb.requires_grad = value30 31 32class SinCosPositionalEmbedding(PositionalEmbedding):33 34    def __init__(self, input_shape: Shape, dropout_rate: float = 0.5):35        super().__init__(input_shape, dropout_rate, trainable=False)36        self.emb.data = self.make_embedding().unsqueeze(0)37 38    def make_embedding(self) -> Tensor:39        n_position, d_hid = self.input_shape40 41        def get_position_angle_vec(position):42            return position / torch.tensor(10000).pow(43                2 * torch.div(torch.arange(d_hid), 2, rounding_mode='trunc') / d_hid)44 45        sinusoid_table = torch.stack([get_position_angle_vec(pos_i) for pos_i in range(n_position)], 0)46        sinusoid_table[:, 0::2] = torch.sin(sinusoid_table[:, 0::2])  # dim 2i47        sinusoid_table[:, 1::2] = torch.cos(sinusoid_table[:, 1::2])  # dim 2i+148 49        return sinusoid_table.float()50