Download tiny_digit_diffusion.py from shibatch/tinydigitdiffusion3m: direct link, hf CLI and curl.
- Browser
- Download file 13.3 kB
-
https://huggingface.co/shibatch/tinydigitdiffusion3m/resolve/main/tiny_digit_diffusion.py
- Command line
-
hf download hf://shibatch/tinydigitdiffusion3m/tiny_digit_diffusion.py
-
curl -L -o tiny_digit_diffusion.py https://huggingface.co/shibatch/tinydigitdiffusion3m/resolve/main/tiny_digit_diffusion.py
13.3 kB
| #!/usr/bin/env python3 | |
| """Tiny conditional diffusion model for 1-8 digit MNIST strings.""" | |
| from __future__ import annotations | |
| import json | |
| import math | |
| from dataclasses import asdict, dataclass | |
| from pathlib import Path | |
| import torch | |
| import torch.nn.functional as F | |
| from PIL import Image, ImageDraw | |
| from safetensors.torch import load_file, save_file | |
| from torch import nn | |
| PAD_TOKEN = 10 | |
| NULL_TOKEN = 11 | |
| class ModelConfig: | |
| image_height: int = 32 | |
| image_width: int = 256 | |
| max_digits: int = 8 | |
| slot_width: int = 32 | |
| base_channels: int = 48 | |
| channel_mults: tuple[int, ...] = (1, 1.5, 2, 8 / 3) | |
| embedding_dim: int = 192 | |
| token_embedding_dim: int = 32 | |
| condition_channels: int = 8 | |
| attention_heads: int = 4 | |
| diffusion_steps: int = 400 | |
| def from_json(cls, path: str | Path) -> "ModelConfig": | |
| values = json.loads(Path(path).read_text(encoding="utf-8")) | |
| if "channel_mults" in values: | |
| values["channel_mults"] = tuple(values["channel_mults"]) | |
| return cls(**values) | |
| def save(self, path: str | Path) -> None: | |
| Path(path).write_text(json.dumps(asdict(self), indent=2), encoding="utf-8") | |
| def prompt_to_tokens(prompt: str, config: ModelConfig) -> tuple[torch.Tensor, int]: | |
| prompt = prompt.strip() | |
| if not prompt or len(prompt) > config.max_digits or not prompt.isascii() or not prompt.isdigit(): | |
| raise ValueError(f"Prompt must contain 1-{config.max_digits} ASCII digits, got {prompt!r}.") | |
| tokens = torch.full((config.max_digits,), PAD_TOKEN, dtype=torch.long) | |
| start = (config.max_digits - len(prompt)) // 2 | |
| tokens[start : start + len(prompt)] = torch.tensor([int(ch) for ch in prompt]) | |
| return tokens, len(prompt) | |
| class SinusoidalTimeEmbedding(nn.Module): | |
| def __init__(self, dim: int): | |
| super().__init__() | |
| self.dim = dim | |
| def forward(self, timesteps: torch.Tensor) -> torch.Tensor: | |
| half = self.dim // 2 | |
| scale = math.log(10_000) / max(half - 1, 1) | |
| frequencies = torch.exp(-scale * torch.arange(half, device=timesteps.device)) | |
| angles = timesteps.float()[:, None] * frequencies[None] | |
| embedding = torch.cat([angles.sin(), angles.cos()], dim=-1) | |
| if self.dim % 2: | |
| embedding = F.pad(embedding, (0, 1)) | |
| return embedding | |
| class ResBlock(nn.Module): | |
| def __init__(self, in_channels: int, out_channels: int, embedding_dim: int): | |
| super().__init__() | |
| self.norm1 = nn.GroupNorm(8, in_channels) | |
| self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1) | |
| self.embedding = nn.Linear(embedding_dim, out_channels) | |
| self.norm2 = nn.GroupNorm(8, out_channels) | |
| self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1) | |
| self.skip = nn.Conv2d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity() | |
| def forward(self, x: torch.Tensor, embedding: torch.Tensor) -> torch.Tensor: | |
| h = self.conv1(F.silu(self.norm1(x))) | |
| h = h + self.embedding(F.silu(embedding))[:, :, None, None] | |
| h = self.conv2(F.silu(self.norm2(h))) | |
| return h + self.skip(x) | |
| class SelfAttention2d(nn.Module): | |
| def __init__(self, channels: int, heads: int): | |
| super().__init__() | |
| if channels % heads: | |
| raise ValueError("Attention channels must be divisible by the number of heads.") | |
| self.heads = heads | |
| self.norm = nn.GroupNorm(8, channels) | |
| self.qkv = nn.Conv2d(channels, channels * 3, 1) | |
| self.proj = nn.Conv2d(channels, channels, 1) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| batch, channels, height, width = x.shape | |
| head_dim = channels // self.heads | |
| qkv = self.qkv(self.norm(x)).view(batch, 3, self.heads, head_dim, height * width) | |
| q, k, v = qkv.unbind(dim=1) | |
| q = q.transpose(-2, -1) | |
| k = k.transpose(-2, -1) | |
| v = v.transpose(-2, -1) | |
| h = F.scaled_dot_product_attention(q, k, v) | |
| h = h.transpose(-2, -1).reshape(batch, channels, height, width) | |
| return x + self.proj(h) | |
| class TinyDigitDiffusion(nn.Module): | |
| def __init__(self, config: ModelConfig): | |
| super().__init__() | |
| self.config = config | |
| channels = [int(config.base_channels * mult) for mult in config.channel_mults] | |
| if any(c % 8 for c in channels): | |
| raise ValueError(f"All channels must be divisible by 8, got {channels}.") | |
| time_dim = 64 | |
| self.time_embedding = nn.Sequential( | |
| SinusoidalTimeEmbedding(time_dim), | |
| nn.Linear(time_dim, config.embedding_dim), | |
| nn.SiLU(), | |
| nn.Linear(config.embedding_dim, config.embedding_dim), | |
| ) | |
| self.token_embedding = nn.Embedding(12, config.token_embedding_dim) | |
| self.position_embedding = nn.Parameter( | |
| torch.randn(config.max_digits, config.token_embedding_dim) * 0.02 | |
| ) | |
| self.length_embedding = nn.Embedding(config.max_digits + 1, config.embedding_dim) | |
| self.condition_mlp = nn.Sequential( | |
| nn.Linear(config.max_digits * config.token_embedding_dim, config.embedding_dim), | |
| nn.SiLU(), | |
| nn.Linear(config.embedding_dim, config.embedding_dim), | |
| ) | |
| self.spatial_condition = nn.Embedding(12, config.condition_channels) | |
| self.input_conv = nn.Conv2d(1 + config.condition_channels, channels[0], 3, padding=1) | |
| self.down_blocks = nn.ModuleList() | |
| self.downsamples = nn.ModuleList() | |
| for i, channel in enumerate(channels): | |
| self.down_blocks.append(ResBlock(channel, channel, config.embedding_dim)) | |
| if i < len(channels) - 1: | |
| self.downsamples.append(nn.Conv2d(channel, channels[i + 1], 4, stride=2, padding=1)) | |
| self.mid1 = ResBlock(channels[-1], channels[-1], config.embedding_dim) | |
| self.mid_attention = SelfAttention2d(channels[-1], config.attention_heads) | |
| self.mid2 = ResBlock(channels[-1], channels[-1], config.embedding_dim) | |
| self.upsamples = nn.ModuleList() | |
| self.up_blocks = nn.ModuleList() | |
| for i in range(len(channels) - 1, 0, -1): | |
| self.upsamples.append(nn.ConvTranspose2d(channels[i], channels[i - 1], 4, stride=2, padding=1)) | |
| self.up_blocks.append(ResBlock(channels[i - 1] * 2, channels[i - 1], config.embedding_dim)) | |
| self.output_norm = nn.GroupNorm(8, channels[0]) | |
| self.output_conv = nn.Conv2d(channels[0], 1, 3, padding=1) | |
| def _condition( | |
| self, tokens: torch.Tensor, lengths: torch.Tensor, height: int, width: int | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| token_features = self.token_embedding(tokens) + self.position_embedding[None] | |
| global_condition = self.condition_mlp(token_features.flatten(1)) | |
| global_condition = global_condition + self.length_embedding(lengths) | |
| spatial = self.spatial_condition(tokens).transpose(1, 2) | |
| spatial = spatial.repeat_interleave(self.config.slot_width, dim=-1) | |
| if spatial.shape[-1] != width: | |
| spatial = F.interpolate(spatial, size=width, mode="nearest") | |
| spatial = spatial[:, :, None, :].expand(-1, -1, height, -1) | |
| return global_condition, spatial | |
| def forward( | |
| self, | |
| noisy_images: torch.Tensor, | |
| timesteps: torch.Tensor, | |
| tokens: torch.Tensor, | |
| lengths: torch.Tensor, | |
| ) -> torch.Tensor: | |
| condition, spatial = self._condition(tokens, lengths, noisy_images.shape[-2], noisy_images.shape[-1]) | |
| embedding = self.time_embedding(timesteps) + condition | |
| h = self.input_conv(torch.cat([noisy_images, spatial], dim=1)) | |
| skips = [] | |
| for i, block in enumerate(self.down_blocks): | |
| h = block(h, embedding) | |
| skips.append(h) | |
| if i < len(self.downsamples): | |
| h = self.downsamples[i](h) | |
| h = self.mid2(self.mid_attention(self.mid1(h, embedding)), embedding) | |
| for upsample, block, skip in zip(self.upsamples, self.up_blocks, reversed(skips[:-1])): | |
| h = upsample(h) | |
| h = block(torch.cat([h, skip], dim=1), embedding) | |
| return self.output_conv(F.silu(self.output_norm(h))) | |
| def cosine_beta_schedule(steps: int, s: float = 0.008) -> torch.Tensor: | |
| x = torch.linspace(0, steps, steps + 1, dtype=torch.float64) | |
| cumulative = torch.cos(((x / steps + s) / (1 + s)) * math.pi * 0.5).square() | |
| cumulative = cumulative / cumulative[0] | |
| betas = 1 - cumulative[1:] / cumulative[:-1] | |
| return betas.clamp(1e-5, 0.999).float() | |
| class DiffusionSchedule: | |
| def __init__(self, steps: int, device: torch.device): | |
| self.steps = steps | |
| self.betas = cosine_beta_schedule(steps).to(device) | |
| self.alphas = 1 - self.betas | |
| self.alpha_bars = self.alphas.cumprod(dim=0) | |
| def add_noise( | |
| self, clean: torch.Tensor, noise: torch.Tensor, timesteps: torch.Tensor | |
| ) -> torch.Tensor: | |
| alpha_bar = self.alpha_bars[timesteps][:, None, None, None] | |
| return alpha_bar.sqrt() * clean + (1 - alpha_bar).sqrt() * noise | |
| def ddim_sample( | |
| model: TinyDigitDiffusion, | |
| prompts: list[str], | |
| device: torch.device, | |
| sampling_steps: int = 50, | |
| guidance_scale: float = 3.0, | |
| seed: int = 0, | |
| ) -> torch.Tensor: | |
| config = model.config | |
| if not 2 <= sampling_steps <= config.diffusion_steps: | |
| raise ValueError(f"sampling_steps must be in [2, {config.diffusion_steps}].") | |
| encoded = [prompt_to_tokens(prompt, config) for prompt in prompts] | |
| tokens = torch.stack([item[0] for item in encoded]).to(device) | |
| lengths = torch.tensor([item[1] for item in encoded], device=device) | |
| null_tokens = torch.full_like(tokens, NULL_TOKEN) | |
| null_lengths = torch.zeros_like(lengths) | |
| generator = torch.Generator(device=device).manual_seed(seed) | |
| images = torch.randn( | |
| len(prompts), 1, config.image_height, config.image_width, | |
| generator=generator, device=device, | |
| ) | |
| schedule = DiffusionSchedule(config.diffusion_steps, device) | |
| timesteps = torch.linspace( | |
| config.diffusion_steps - 1, 0, sampling_steps, device=device | |
| ).round().long() | |
| was_training = model.training | |
| model.eval() | |
| for i, timestep in enumerate(timesteps): | |
| t = torch.full((len(prompts),), int(timestep), device=device, dtype=torch.long) | |
| if guidance_scale == 1.0: | |
| # At exactly 1.0, CFG reduces algebraically to the conditional | |
| # prediction. Avoiding the unconditional half nearly halves | |
| # inference memory and compute and also makes the recommended path | |
| # straightforward. | |
| predicted_noise = model(images, t, tokens, lengths) | |
| else: | |
| model_input = torch.cat([images, images], dim=0) | |
| time_input = torch.cat([t, t], dim=0) | |
| token_input = torch.cat([null_tokens, tokens], dim=0) | |
| length_input = torch.cat([null_lengths, lengths], dim=0) | |
| eps_unconditional, eps_conditional = model( | |
| model_input, time_input, token_input, length_input | |
| ).chunk(2) | |
| predicted_noise = eps_unconditional + guidance_scale * ( | |
| eps_conditional - eps_unconditional | |
| ) | |
| alpha_bar = schedule.alpha_bars[timestep] | |
| if i + 1 < len(timesteps): | |
| previous_alpha_bar = schedule.alpha_bars[timesteps[i + 1]] | |
| else: | |
| previous_alpha_bar = torch.ones((), device=device) | |
| predicted_clean = ( | |
| images - (1 - alpha_bar).sqrt() * predicted_noise | |
| ) / alpha_bar.sqrt() | |
| predicted_clean = predicted_clean.clamp(-1, 1) | |
| images = previous_alpha_bar.sqrt() * predicted_clean + ( | |
| 1 - previous_alpha_bar | |
| ).sqrt() * predicted_noise | |
| if was_training: | |
| model.train() | |
| return images.clamp(-1, 1) | |
| def save_prompt_sheet(images: torch.Tensor, prompts: list[str], path: str | Path) -> None: | |
| images = ((images.detach().cpu() + 1) * 127.5).round().clamp(0, 255).byte() | |
| row_height = images.shape[-2] + 10 | |
| label_width = 90 | |
| sheet = Image.new("L", (label_width + images.shape[-1], row_height * len(prompts)), 255) | |
| draw = ImageDraw.Draw(sheet) | |
| for i, (image, prompt) in enumerate(zip(images, prompts)): | |
| y = i * row_height | |
| draw.text((4, y + 10), prompt, fill=0) | |
| digit_image = Image.fromarray(image[0].numpy(), mode="L") | |
| sheet.paste(digit_image, (label_width, y + 5)) | |
| Path(path).parent.mkdir(parents=True, exist_ok=True) | |
| sheet.save(path) | |
| def save_weights(model: TinyDigitDiffusion, path: str | Path) -> None: | |
| tensors = {name: value.detach().cpu().contiguous() for name, value in model.state_dict().items()} | |
| save_file(tensors, str(path)) | |
| def load_model(model_dir: str | Path, device: torch.device) -> TinyDigitDiffusion: | |
| model_dir = Path(model_dir) | |
| config = ModelConfig.from_json(model_dir / "config.json") | |
| model = TinyDigitDiffusion(config) | |
| model.load_state_dict(load_file(model_dir / "model.safetensors", device=str(device))) | |
| return model.to(device).eval() | |
| def count_parameters(model: nn.Module) -> int: | |
| return sum(parameter.numel() for parameter in model.parameters()) | |