Download clean/video/mintime/models/utils.py from deepsafe/model-code: direct link, hf CLI and curl.
- Browser
- Download file 2.29 kB
-
https://huggingface.co/deepsafe/model-code/resolve/main/clean/video/mintime/models/utils.py
- Command line
-
hf download hf://deepsafe/model-code/clean/video/mintime/models/utils.py
-
curl -L -o utils.py https://huggingface.co/deepsafe/model-code/resolve/main/clean/video/mintime/models/utils.py
2.29 kB
| from math import log, pi | |
| import torch | |
| from torch import nn, einsum | |
| import torch.nn.functional as F | |
| from einops import rearrange, repeat | |
| def rotate_every_two(x): | |
| x = rearrange(x, '... (d j) -> ... d j', j = 2) | |
| x1, x2 = x.unbind(dim = -1) | |
| x = torch.stack((-x2, x1), dim = -1) | |
| return rearrange(x, '... d j -> ... (d j)') | |
| def apply_rot_emb(q, k, rot_emb): | |
| sin, cos = rot_emb | |
| rot_dim = sin.shape[-1] | |
| (q, q_pass), (k, k_pass) = map(lambda t: (t[..., :rot_dim], t[..., rot_dim:]), (q, k)) | |
| q, k = map(lambda t: t * cos + rotate_every_two(t) * sin, (q, k)) | |
| q, k = map(lambda t: torch.cat(t, dim = -1), ((q, q_pass), (k, k_pass))) | |
| return q, k | |
| class AxialRotaryEmbedding(nn.Module): | |
| def __init__(self, dim, max_freq = 10): | |
| super().__init__() | |
| self.dim = dim | |
| scales = torch.logspace(0., log(max_freq / 2) / log(2), self.dim // 4, base = 2) | |
| self.register_buffer('scales', scales) | |
| def forward(self, h, w, device): | |
| scales = rearrange(self.scales, '... -> () ...') | |
| scales = scales.to(device) | |
| h_seq = torch.linspace(-1., 1., steps = h, device = device) | |
| h_seq = h_seq.unsqueeze(-1) | |
| w_seq = torch.linspace(-1., 1., steps = w, device = device) | |
| w_seq = w_seq.unsqueeze(-1) | |
| h_seq = h_seq * scales * pi | |
| w_seq = w_seq * scales * pi | |
| x_sinu = repeat(h_seq, 'i d -> i j d', j = w) | |
| y_sinu = repeat(w_seq, 'j d -> i j d', i = h) | |
| sin = torch.cat((x_sinu.sin(), y_sinu.sin()), dim = -1) | |
| cos = torch.cat((x_sinu.cos(), y_sinu.cos()), dim = -1) | |
| sin, cos = map(lambda t: rearrange(t, 'i j d -> (i j) d'), (sin, cos)) | |
| sin, cos = map(lambda t: repeat(t, 'n d -> () n (d j)', j = 2), (sin, cos)) | |
| return sin, cos | |
| class RotaryEmbedding(nn.Module): | |
| def __init__(self, dim): | |
| super().__init__() | |
| inv_freqs = 1. / (10000 ** (torch.arange(0, dim, 2).float() / dim)) | |
| self.register_buffer('inv_freqs', inv_freqs) | |
| def forward(self, n, device): | |
| seq = torch.arange(n, device = device) | |
| freqs = einsum('i, j -> i j', seq, self.inv_freqs) | |
| freqs = torch.cat((freqs, freqs), dim = -1) | |
| freqs = rearrange(freqs, 'n d -> () n d') | |
| return freqs.sin(), freqs.cos() |