Instructions to use nebulette/nano-transformers with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use nebulette/nano-transformers with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("nebulette/nano-transformers", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download model.py from nebulette/nano-transformers: direct link, hf CLI and curl.
- Browser
- Download file 14.9 kB
-
https://huggingface.co/nebulette/nano-transformers/resolve/main/model.py
- Command line
-
hf download hf://nebulette/nano-transformers/model.py
-
curl -L -o model.py https://huggingface.co/nebulette/nano-transformers/resolve/main/model.py
14.9 kB
| import math | |
| from safetensors.torch import load_model | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| SPRINT_NUM_F = 2 | |
| SPRINT_NUM_H = 2 | |
| TEXT_EMBED_DIM = 640 | |
| def adaln( | |
| x: torch.Tensor, | |
| scale: torch.Tensor, | |
| shift: torch.Tensor, | |
| eps: float = 1e-6 | |
| ): | |
| normalized = F.layer_norm( | |
| x, | |
| normalized_shape=(x.shape[-1],), | |
| weight=None, | |
| bias=None, | |
| eps=eps, | |
| ) | |
| return normalized * (1 + scale) + shift | |
| def rope_2d(head_dim, height, width, device): | |
| """Split-half 2D RoPE over centered token coordinates (x frequencies, then y), as (1, N, 1, head_dim // 2, 2, 2) rotations.""" | |
| axis_dim = head_dim // 2 | |
| inv_freq = 1.0 / (10000.0 ** (torch.arange(0, axis_dim, 2, dtype=torch.float32, device=device) / axis_dim)) | |
| y = torch.arange(height, dtype=torch.float32, device=device) - (height - 1) / 2 | |
| x = torch.arange(width, dtype=torch.float32, device=device) - (width - 1) / 2 | |
| y, x = torch.meshgrid(y, x, indexing="ij") | |
| angles = torch.cat([torch.outer(x.flatten(), inv_freq), torch.outer(y.flatten(), inv_freq)], dim=-1) | |
| cos, sin = torch.cos(angles), torch.sin(angles) | |
| return torch.stack([cos, -sin, sin, cos], dim=-1).view(1, height * width, 1, axis_dim, 2, 2) | |
| def modulate(x, shift, scale): | |
| return torch.addcmul(shift, x, 1 + scale) | |
| def rope_split_half(x, pe): | |
| """Differentiable ck.apply_rope_split_half1: pair k is (x[k], x[k + head_dim // 2]).""" | |
| x_ = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(pe.dtype) | |
| return (pe[..., 0] * x_[..., 0] + pe[..., 1] * x_[..., 1]).movedim(-1, -2).reshape(x.shape).type_as(x) | |
| def norm_rope(q, k, q_norm, k_norm, pe): | |
| return rope_split_half(q_norm(q), pe), rope_split_half(k_norm(k), pe) | |
| class Embed(nn.Module): | |
| def __init__(self, in_dim, hidden_size, norm=False, dtype=None, device=None): | |
| super().__init__() | |
| self.proj = nn.Linear(in_dim, hidden_size, bias=True, dtype=dtype, device=device) | |
| self.norm = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) if norm else nn.Identity() | |
| def forward(self, x): | |
| return self.norm(self.proj(x)) | |
| class FeedForward(nn.Module): | |
| def __init__(self, dim, hidden_dim, dtype=None, device=None): | |
| super().__init__() | |
| self.w12 = nn.Linear(dim, hidden_dim * 2, bias=False, dtype=dtype, device=device) | |
| self.w3 = nn.Linear(hidden_dim, dim, bias=False, dtype=dtype, device=device) | |
| def forward(self, x): | |
| x1, x2 = self.w12(x).chunk(2, dim=-1) | |
| return self.w3(F.silu(x1) * x2) | |
| class Attention(nn.Module): | |
| """Self-attention over image tokens with 2D RoPE, optionally joint with (non-rotated) text keys and values.""" | |
| def __init__(self, dim, num_heads, cross_attention, dtype=None, device=None): | |
| super().__init__() | |
| self.num_heads = num_heads | |
| self.head_dim = dim // num_heads | |
| self.qkv_x = nn.Linear(dim, dim * 3, bias=False, dtype=dtype, device=device) | |
| self.kv_y = nn.Linear(dim, dim * 2, bias=False, dtype=dtype, device=device) if cross_attention else None | |
| self.q_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) | |
| self.k_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) | |
| self.proj = nn.Linear(dim, dim, bias=True, dtype=dtype, device=device) | |
| def forward(self, x, txt, pe, txt_bias, transformer_options={}): | |
| b, n, c = x.shape | |
| q, k, v = self.qkv_x(x).view(b, n, 3, self.num_heads, self.head_dim).unbind(2) | |
| q, k = norm_rope(q, k, self.q_norm, self.k_norm, pe) | |
| mask = None | |
| if self.kv_y is not None: | |
| ky, vy = self.kv_y(txt).view(b, -1, 2, self.num_heads, self.head_dim).unbind(2) | |
| k = torch.cat([k, self.k_norm(ky)], dim=1) | |
| v = torch.cat([v, vy], dim=1) | |
| # Image keys are always visible; text keys carry log(emphasis weight), -inf for padding. | |
| mask = F.pad(txt_bias, (n, 0))[:, None, None, :] | |
| # x = optimized_attention(q.reshape(b, n, c), k.reshape(b, -1, c), v.reshape(b, -1, c), self.num_heads, mask=mask, transformer_options=transformer_options) | |
| q = q.transpose(1, 2) | |
| k = k.transpose(1, 2) | |
| v = v.transpose(1, 2) | |
| # Mask from ldm/modules/attention, attention_pytorch() function. | |
| if mask is not None: | |
| # add a batch dimension if there isn't already one | |
| if mask.ndim == 2: | |
| mask = mask.unsqueeze(0) | |
| # add a heads dimension if there isn't already one | |
| if mask.ndim == 3: | |
| mask = mask.unsqueeze(1) | |
| x = F.scaled_dot_product_attention(q, k, v, attn_mask=mask) | |
| x = x.transpose(1, 2).reshape(b, n, c) | |
| return self.proj(x) | |
| class DiTBlock(nn.Module): | |
| """Image block; its adaLN modulation is computed by the caller (adaLN-single).""" | |
| def __init__(self, hidden_size, num_heads, mlp_hidden, cross_attention, dtype=None, device=None): | |
| super().__init__() | |
| self.norm1 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.attn = Attention(hidden_size, num_heads, cross_attention, dtype=dtype, device=device) | |
| self.norm2 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.mlp = FeedForward(hidden_size, mlp_hidden, dtype=dtype, device=device) | |
| def forward(self, x, txt, pe, mod, txt_bias, transformer_options={}): | |
| shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=-1) | |
| x = torch.addcmul(x, gate_msa, self.attn(modulate(self.norm1(x), shift_msa, scale_msa), txt, pe, txt_bias, transformer_options=transformer_options)) | |
| return torch.addcmul(x, gate_mlp, self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))) | |
| class TextRefineAttention(nn.Module): | |
| def __init__(self, dim, num_heads, dtype=None, device=None): | |
| super().__init__() | |
| self.num_heads = num_heads | |
| self.head_dim = dim // num_heads | |
| self.qkv = nn.Linear(dim, dim * 3, bias=False, dtype=dtype, device=device) | |
| self.q_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) | |
| self.k_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) | |
| self.proj = nn.Linear(dim, dim, bias=True, dtype=dtype, device=device) | |
| def forward(self, x, mask, transformer_options={}): | |
| b, n, c = x.shape | |
| q, k, v = self.qkv(x).view(b, n, 3, self.num_heads, self.head_dim).unbind(2) | |
| q, k = self.q_norm(q), self.k_norm(k) | |
| # SDPA expects (B, H, N, D) | |
| q = q.transpose(1, 2) | |
| k = k.transpose(1, 2) | |
| v = v.transpose(1, 2) | |
| # x = optimized_attention(self.q_norm(q).reshape(b, n, c), self.k_norm(k).reshape(b, n, c), v.reshape(b, n, c), self.num_heads, mask=mask, transformer_options=transformer_options) | |
| x = F.scaled_dot_product_attention(q, k, v) | |
| x = x.transpose(1, 2).reshape(b, n, c) | |
| return self.proj(x) | |
| class TextRefineBlock(nn.Module): | |
| def __init__(self, hidden_size, num_heads, mlp_hidden, dtype=None, device=None): | |
| super().__init__() | |
| self.norm1 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.attn = TextRefineAttention(hidden_size, num_heads, dtype=dtype, device=device) | |
| self.norm2 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.mlp = FeedForward(hidden_size, mlp_hidden, dtype=dtype, device=device) | |
| self.adaLN_modulation = nn.Sequential(nn.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)) | |
| def forward(self, x, c, mask, keep, transformer_options={}): | |
| shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=-1) | |
| x = torch.addcmul(x, gate_msa, self.attn(modulate(self.norm1(x), shift_msa, scale_msa), mask, transformer_options=transformer_options)) | |
| x = torch.addcmul(x, gate_mlp, self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))) | |
| return x * keep | |
| class FinalLayer(nn.Module): | |
| def __init__(self, hidden_size, out_channels, dtype=None, device=None): | |
| super().__init__() | |
| self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) | |
| self.adaLN_modulation = nn.Linear(hidden_size, 2 * hidden_size, bias=True, dtype=dtype, device=device) | |
| self.linear = nn.Linear(hidden_size, out_channels, bias=True, dtype=dtype, device=device) | |
| def forward(self, x, c): | |
| shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1) | |
| return self.linear(adaln(x, scale, shift, self.norm_final.eps)) | |
| class TimestepEmbedder(nn.Module): | |
| """ | |
| Embeds scalar timesteps into vector representations. | |
| """ | |
| def __init__(self, hidden_size, frequency_embedding_size=256, output_size=None, dtype=None, device=None, max_period=10000): | |
| super().__init__() | |
| if output_size is None: | |
| output_size = hidden_size | |
| self.mlp = nn.Sequential( | |
| nn.Linear(frequency_embedding_size, hidden_size, bias=True, dtype=dtype, device=device), | |
| nn.SiLU(), | |
| nn.Linear(hidden_size, output_size, bias=True, dtype=dtype, device=device), | |
| ) | |
| self.frequency_embedding_size = frequency_embedding_size | |
| self.max_period = max_period | |
| def forward(self, t, dtype, **kwargs): | |
| t_freq = self.timestep_embedding(t).to(dtype) | |
| t_emb = self.mlp(t_freq) | |
| return t_emb | |
| def timestep_embedding(self, timesteps, repeat_only=False): | |
| """ | |
| Create sinusoidal timestep embeddings. | |
| :param timesteps: a 1-D Tensor of N indices, one per batch element. | |
| These may be fractional. | |
| :return: an [N x dim] Tensor of positional embeddings. | |
| """ | |
| dim = self.frequency_embedding_size | |
| if not repeat_only: | |
| half = dim // 2 | |
| freqs = torch.exp( | |
| -math.log(self.max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=timesteps.device) / half | |
| ) | |
| args = timesteps[:, None].float() * freqs[None] | |
| embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) | |
| if dim % 2: | |
| embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) | |
| else: | |
| # embedding = repeat(timesteps, 'b -> b d', d=dim) | |
| embedding = timesteps[:, None].expand(-1, dim) | |
| return embedding | |
| class Transformer2DModel(nn.Module): | |
| def __init__(self, in_channels=64, hidden_size=1536, num_heads=16, num_blocks=18, num_text_blocks=2, mlp_hidden=4096, dtype=None, device=None): | |
| super().__init__() | |
| self.dtype = dtype | |
| self.in_channels = in_channels | |
| self.num_heads = num_heads | |
| self.head_dim = hidden_size // num_heads | |
| self.s_embedder = Embed(in_channels, hidden_size, dtype=dtype, device=device) | |
| self.t_embedder = TimestepEmbedder(hidden_size, dtype=dtype, device=device) | |
| self.y_embedder = Embed(TEXT_EMBED_DIM, hidden_size, norm=True, dtype=dtype, device=device) | |
| self.shared_encoder_adaLN = nn.Sequential(nn.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)) | |
| self.encoder_adaLN_offsets = nn.ParameterList([nn.Parameter(torch.empty(6 * hidden_size, dtype=dtype, device=device)) for _ in range(num_blocks)]) | |
| self.y_pool_proj = nn.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) | |
| self.sprint_out_proj = nn.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) | |
| self.blocks = nn.ModuleList([DiTBlock(hidden_size, num_heads, mlp_hidden, i % 2 == 0, dtype=dtype, device=device) for i in range(num_blocks)]) | |
| self.text_refine_blocks = nn.ModuleList([TextRefineBlock(hidden_size, num_heads, mlp_hidden, dtype=dtype, device=device) for _ in range(num_text_blocks)]) | |
| self.final_layer = FinalLayer(hidden_size, in_channels, dtype=dtype, device=device) | |
| def run_block(self, i, x, txt, pe, mod, txt_bias, transformer_options): | |
| offset = self.encoder_adaLN_offsets[i].to(dtype=mod.dtype) | |
| return self.blocks[i](x, txt, pe, mod + offset, txt_bias, transformer_options=transformer_options) | |
| def sparse_path(self, x, txt, pe, mod, txt_bias, transformer_options): | |
| g = x | |
| for i in range(SPRINT_NUM_F, len(self.blocks) - SPRINT_NUM_H): | |
| g = self.run_block(i, g, txt, pe, mod, txt_bias, transformer_options) | |
| return x + self.sprint_out_proj(g - x) | |
| def forward(self, x, timestep, context, token_weights, sparse_skip=None, transformer_options={}, **kwargs): | |
| """sparse_skip: optional per-row flags; flagged rows skip the sparse middle blocks.""" | |
| b, _, h, w = x.shape | |
| keep = token_weights > 0 | |
| txt_bias = torch.log(token_weights.clamp(min=1e-4)).masked_fill(~keep, float("-inf")) | |
| # Refine attention is a hard mask; a row with no text key keeps its first one and is zeroed after. | |
| refine_keep = keep.clone() | |
| refine_keep[:, 0] |= ~keep.any(dim=1) | |
| refine_mask = torch.zeros_like(txt_bias).masked_fill(~refine_keep, float("-inf"))[:, None, None, :] | |
| keep = keep.unsqueeze(-1).to(x.dtype) | |
| t = self.t_embedder(timestep * 1000.0, x.dtype).unsqueeze(1) | |
| txt = self.y_embedder(context) | |
| time_condition = F.silu(t) | |
| for block in self.text_refine_blocks: | |
| txt = block(txt, time_condition, refine_mask, keep, transformer_options=transformer_options) | |
| weights = token_weights.unsqueeze(-1) | |
| pooled = (txt * weights).sum(dim=1) / weights.sum(dim=1).clamp(min=1.0) | |
| condition = F.silu(t + self.y_pool_proj(pooled).unsqueeze(1)) | |
| mod = self.shared_encoder_adaLN(condition) | |
| pe = rope_2d(self.head_dim, h, w, x.device) | |
| s = self.s_embedder(x.flatten(2).transpose(1, 2)) | |
| for i in range(SPRINT_NUM_F): | |
| s = self.run_block(i, s, txt, pe, mod, txt_bias, transformer_options) | |
| if sparse_skip is None or not any(sparse_skip): | |
| s = self.sparse_path(s, txt, pe, mod, txt_bias, transformer_options) | |
| elif not all(sparse_skip): | |
| rows = [i for i, skip in enumerate(sparse_skip) if not skip] | |
| s[rows] = self.sparse_path(s[rows], txt[rows], pe, mod[rows], txt_bias[rows], transformer_options) | |
| for i in range(len(self.blocks) - SPRINT_NUM_H, len(self.blocks)): | |
| s = self.run_block(i, s, txt, pe, mod, txt_bias, transformer_options) | |
| x0 = self.final_layer(s, condition).transpose(1, 2).reshape(b, self.in_channels, h, w) | |
| return (x - x0) / timestep.view(-1, 1, 1, 1) | |
| def from_safetensors(path, device=None): | |
| model = Transformer2DModel() | |
| load_model(model, path) | |
| return model.to(device) | |