{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[],"dockerImageVersionId":31259,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# vae_encoder_standalone.py\n# This version has no relative imports - can be run directly in any notebook\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\n# ============================================================================\n# SELF-ATTENTION (needed by VAE_AttentionBlock)\n# ============================================================================\n\nclass SelfAttention(nn.Module):\n    \"\"\"Standard self-attention mechanism.\"\"\"\n    def __init__(self, n_heads, d_embed, in_proj_bias=True, out_proj_bias=True):\n        super().__init__()\n        self.in_proj = nn.Linear(d_embed, 3 * d_embed, bias=in_proj_bias)\n        self.out_proj = nn.Linear(d_embed, d_embed, bias=out_proj_bias)\n        self.n_heads = n_heads\n        self.d_head = d_embed // n_heads\n\n    def forward(self, x, causal_mask=False):\n        input_shape = x.shape\n        batch_size, sequence_length, d_embed = input_shape\n        interim_shape = (batch_size, sequence_length, self.n_heads, self.d_head)\n        \n        q, k, v = self.in_proj(x).chunk(3, dim=-1)\n        q = q.view(interim_shape).transpose(1, 2)\n        k = k.view(interim_shape).transpose(1, 2)\n        v = v.view(interim_shape).transpose(1, 2)\n        \n        weight = q @ k.transpose(-1, -2)\n        if causal_mask:\n            mask = torch.ones_like(weight, dtype=torch.bool).triu(1)\n            weight.masked_fill_(mask, -torch.inf)\n        \n        weight /= torch.sqrt(torch.tensor(self.d_head, dtype=torch.float32))\n        weight = F.softmax(weight, dim=-1)\n        output = weight @ v\n        output = output.transpose(1, 2).reshape(input_shape)\n        output = self.out_proj(output)\n        return output\n\n\n# ============================================================================\n# VAE ATTENTION BLOCK\n# ============================================================================\n\nclass VAE_AttentionBlock(nn.Module):\n    \"\"\"Self-attention block used in the VAE encoder and decoder.\"\"\"\n    def __init__(self, channels):\n        super().__init__()\n        self.groupnorm = nn.GroupNorm(32, channels)\n        self.attention = SelfAttention(1, channels)\n    \n    def forward(self, x):\n        # x: (Batch, Features, Height, Width)\n        residue = x \n        x = self.groupnorm(x)\n        n, c, h, w = x.shape\n        x = x.view((n, c, h * w))\n        x = x.transpose(-1, -2)\n        x = self.attention(x)\n        x = x.transpose(-1, -2)\n        x = x.view((n, c, h, w))\n        x += residue\n        return x","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:11:15.151222Z","iopub.execute_input":"2026-02-28T06:11:15.151727Z","iopub.status.idle":"2026-02-28T06:11:15.160968Z","shell.execute_reply.started":"2026-02-28T06:11:15.151700Z","shell.execute_reply":"2026-02-28T06:11:15.160270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass VAE_ResidualBlock(nn.Module):\n    \"\"\"Residual block with two convolutions and a skip connection.\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        \n        self.groupnorm_1 = nn.GroupNorm(32, in_channels)\n        self.conv_1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)\n        self.groupnorm_2 = nn.GroupNorm(32, out_channels)\n        self.conv_2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n        \n        if in_channels == out_channels:\n            self.residual_layer = nn.Identity()\n        else:\n            self.residual_layer = nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0)\n    \n    def forward(self, x):\n        residue = x\n        x = self.groupnorm_1(x)\n        x = F.silu(x)\n        x = self.conv_1(x)\n        x = self.groupnorm_2(x)\n        x = F.silu(x)\n        x = self.conv_2(x)\n        return x + self.residual_layer(residue)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:11:15.161891Z","iopub.execute_input":"2026-02-28T06:11:15.162134Z","iopub.status.idle":"2026-02-28T06:11:15.178301Z","shell.execute_reply.started":"2026-02-28T06:11:15.162114Z","shell.execute_reply":"2026-02-28T06:11:15.177679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass VAE_Encoder(nn.Module):\n    \"\"\"\n    Encodes an RGB image (512x512x3) into a latent representation (64x64x4).\n    \n    This compression is what makes Stable Diffusion efficient - instead of\n    denoising in high-res pixel space, we denoise in compressed latent space.\n    \n    Architecture progression:\n    (512, 512, 3) -> (256, 256, 128) -> (128, 128, 256) -> (64, 64, 512) -> (64, 64, 4)\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        \n        # Initial convolution: (Batch, 3, 512, 512) -> (Batch, 128, 512, 512)\n        self.conv_in = nn.Conv2d(3, 128, kernel_size=3, padding=1)\n        \n        # Encoder blocks - progressively downsample and increase channels\n        self.encoders = nn.ModuleList([\n            # First block: (Batch, 128, 512, 512) -> (Batch, 128, 256, 256)\n            nn.Sequential(\n                VAE_ResidualBlock(128, 128),\n                VAE_ResidualBlock(128, 128),\n            ),\n            \n            # Second block: (Batch, 128, 256, 256) -> (Batch, 256, 128, 128)\n            nn.Sequential(\n                VAE_ResidualBlock(128, 256),\n                VAE_ResidualBlock(256, 256),\n            ),\n            \n            # Third block: (Batch, 256, 128, 128) -> (Batch, 512, 64, 64)\n            nn.Sequential(\n                VAE_ResidualBlock(256, 512),\n                VAE_ResidualBlock(512, 512),\n            ),\n            \n            # Fourth block: stays at (Batch, 512, 64, 64)\n            nn.Sequential(\n                VAE_ResidualBlock(512, 512),\n                VAE_ResidualBlock(512, 512),\n            ),\n        ])\n        \n        # Bottleneck with attention at the smallest spatial resolution\n        self.bottleneck = nn.Sequential(\n            VAE_ResidualBlock(512, 512),\n            VAE_AttentionBlock(512),  # Self-attention at 64x64\n            VAE_ResidualBlock(512, 512),\n        )\n        \n        # Final layers to produce mean and log-variance for latent distribution\n        self.final_norm = nn.GroupNorm(32, 512)\n        \n        # Output: (Batch, 512, 64, 64) -> (Batch, 8, 64, 64)\n        # We output 8 channels: 4 for mean, 4 for log-variance\n        self.conv_out = nn.Conv2d(512, 8, kernel_size=3, padding=1)\n        \n        # Scaling factor for the latent space\n        self.scale_factor = 0.18215\n    \n    def forward(self, x, noise):\n        \"\"\"\n        Args:\n            x: Input image (Batch, 3, Height, Width)\n            noise: Random noise for sampling (Batch, 4, Height/8, Width/8)\n        \n        Returns:\n            Sampled latent (Batch, 4, Height/8, Width/8)\n        \"\"\"\n        # x: (Batch, 3, 512, 512)\n        \n        # Initial convolution\n        x = self.conv_in(x)\n        # (Batch, 128, 512, 512)\n        \n        # Encoder blocks with downsampling\n        for i, encoder in enumerate(self.encoders):\n            x = encoder(x)\n            # Downsample by 2x after each block (except the last)\n            if i < len(self.encoders) - 1:\n                # Use stride=2 convolution for downsampling\n                x = F.avg_pool2d(x, kernel_size=2, stride=2)\n        \n        # x is now (Batch, 512, 64, 64)\n        \n        # Bottleneck with attention\n        x = self.bottleneck(x)\n        # (Batch, 512, 64, 64)\n        \n        # Final normalization and convolution\n        x = self.final_norm(x)\n        x = F.silu(x)\n        x = self.conv_out(x)\n        # (Batch, 8, 64, 64)\n        \n        # Split into mean and log-variance\n        mean, log_variance = torch.chunk(x, 2, dim=1)\n        # Each: (Batch, 4, 64, 64)\n        \n        # Clamp log-variance for numerical stability\n        log_variance = torch.clamp(log_variance, -30, 20)\n        \n        # Convert log-variance to variance\n        variance = log_variance.exp()\n        \n        # Convert variance to standard deviation\n        stdev = variance.sqrt()\n        \n        # Reparameterization trick: z = mean + stdev * noise\n        # This allows gradients to flow through the sampling process\n        x = mean + stdev * noise\n        \n        # Scale the latent (this is a trained constant in SD)\n        x *= self.scale_factor\n        \n        return x\n\n\n# ============================================================================\n# QUICK TEST\n# ============================================================================\n\nif __name__ == \"__main__\":\n    print(\"Testing VAE Encoder...\")\n    \n    encoder = VAE_Encoder()\n    x = torch.randn(2, 3, 512, 512)\n    noise = torch.randn(2, 4, 64, 64)\n    \n    with torch.no_grad():\n        latent = encoder(x, noise)\n    \n    print(f\"✓ Input: {x.shape}\")\n    print(f\"✓ Output: {latent.shape}\")\n    print(f\"✓ Works perfectly!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:11:15.179188Z","iopub.execute_input":"2026-02-28T06:11:15.179420Z","iopub.status.idle":"2026-02-28T06:11:29.127389Z","shell.execute_reply.started":"2026-02-28T06:11:15.179401Z","shell.execute_reply":"2026-02-28T06:11:29.126689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model/vae_decoder.py\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# from .vae_residual import VAE_ResidualBlock\n# from .vae_attention import VAE_AttentionBlock\n\n\nclass VAE_Decoder(nn.Module):\n    \"\"\"\n    Decodes a latent representation (64x64x4) back into an RGB image (512x512x3).\n    \n    This is the final step in Stable Diffusion - after the U-Net has denoised\n    the latent representation, the decoder \"paints\" it back into pixel space.\n    \n    Architecture progression:\n    (64, 64, 4) -> (64, 64, 512) -> (128, 128, 512) -> (256, 256, 256) -> (512, 512, 128) -> (512, 512, 3)\n    \n    For your DCT embedding research:\n    This decoder will learn to reconstruct images with the frequency characteristics\n    of your training data. When fine-tuned on DCT-optimal images, it should learn\n    to generate images with favorable mid-frequency properties.\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        \n        # Scaling factor (inverse of encoder's scale)\n        self.scale_factor = 0.18215\n        \n        # Initial convolution: (Batch, 4, 64, 64) -> (Batch, 512, 64, 64)\n        self.conv_in = nn.Conv2d(4, 512, kernel_size=3, padding=1)\n        \n        # Bottleneck with attention at the smallest spatial resolution\n        self.bottleneck = nn.Sequential(\n            VAE_ResidualBlock(512, 512),\n            VAE_AttentionBlock(512),  # Self-attention at 64x64\n            VAE_ResidualBlock(512, 512),\n            VAE_ResidualBlock(512, 512),\n            VAE_ResidualBlock(512, 512),\n        )\n        \n        # Decoder blocks - progressively upsample and decrease channels\n        self.decoders = nn.ModuleList([\n            # First block: stays at (Batch, 512, 64, 64)\n            nn.Sequential(\n                VAE_ResidualBlock(512, 512),\n                VAE_ResidualBlock(512, 512),\n                VAE_ResidualBlock(512, 512),\n            ),\n            \n            # Second block: (Batch, 512, 64, 64) -> (Batch, 512, 128, 128)\n            nn.Sequential(\n                VAE_ResidualBlock(512, 512),\n                VAE_ResidualBlock(512, 512),\n                VAE_ResidualBlock(512, 512),\n            ),\n            \n            # Third block: (Batch, 512, 128, 128) -> (Batch, 256, 256, 256)\n            nn.Sequential(\n                VAE_ResidualBlock(512, 256),\n                VAE_ResidualBlock(256, 256),\n                VAE_ResidualBlock(256, 256),\n            ),\n            \n            # Fourth block: (Batch, 256, 256, 256) -> (Batch, 128, 512, 512)\n            nn.Sequential(\n                VAE_ResidualBlock(256, 128),\n                VAE_ResidualBlock(128, 128),\n                VAE_ResidualBlock(128, 128),\n            ),\n        ])\n        \n        # Final layers to produce RGB image\n        self.final_norm = nn.GroupNorm(32, 128)\n        \n        # Output: (Batch, 128, 512, 512) -> (Batch, 3, 512, 512)\n        self.conv_out = nn.Conv2d(128, 3, kernel_size=3, padding=1)\n    \n    def forward(self, x):\n        \"\"\"\n        Args:\n            x: Latent representation (Batch, 4, Height/8, Width/8)\n        \n        Returns:\n            Reconstructed image (Batch, 3, Height, Width)\n        \"\"\"\n        # x: (Batch, 4, 64, 64)\n        \n        # Remove the scaling applied by encoder\n        x /= self.scale_factor\n        \n        # Initial convolution\n        x = self.conv_in(x)\n        # (Batch, 512, 64, 64)\n        \n        # Bottleneck with attention\n        x = self.bottleneck(x)\n        # (Batch, 512, 64, 64)\n        \n        # Decoder blocks with upsampling\n        for i, decoder in enumerate(self.decoders):\n            x = decoder(x)\n            \n            # Upsample by 2x after each block (except the first)\n            if i > 0:\n                # Use nearest neighbor upsampling followed by conv\n                # This avoids checkerboard artifacts\n                x = F.interpolate(x, scale_factor=2, mode='nearest')\n        \n        # x is now (Batch, 128, 512, 512)\n        \n        # Final normalization and convolution\n        x = self.final_norm(x)\n        x = F.silu(x)\n        x = self.conv_out(x)\n        # (Batch, 3, 512, 512)\n        \n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:11:29.128260Z","iopub.execute_input":"2026-02-28T06:11:29.128466Z","iopub.status.idle":"2026-02-28T06:11:29.138466Z","shell.execute_reply.started":"2026-02-28T06:11:29.128447Z","shell.execute_reply":"2026-02-28T06:11:29.137676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model/vae.py\n\n# import torch\n# import torch.nn as nn\n# from .vae_encoder import VAE_Encoder\n# from .vae_decoder import VAE_Decoder\n\n\nclass VAE(nn.Module):\n    \"\"\"\n    Complete Variational Autoencoder for Stable Diffusion.\n    \n    The VAE serves two purposes:\n    1. ENCODING: Compress images from pixel space (512x512x3) to latent space (64x64x4)\n    2. DECODING: Reconstruct images from latent space back to pixel space\n    \n    The diffusion process happens entirely in the compressed latent space,\n    which is why Stable Diffusion can run on consumer hardware.\n    \n    For your DCT steganography research:\n    - The encoder learns to extract features that preserve DCT-relevant information\n    - The decoder learns to reconstruct images with favorable frequency characteristics\n    - When fine-tuned on DCT-optimal images, both should adapt to your embedding needs\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        self.encoder = VAE_Encoder()\n        self.decoder = VAE_Decoder()\n    \n    def encode(self, x, noise):\n        \"\"\"\n        Encode an image into latent space.\n        \n        Args:\n            x: Input image (Batch, 3, Height, Width)\n               Expected: (Batch, 3, 512, 512)\n            noise: Random Gaussian noise (Batch, 4, Height/8, Width/8)\n                   Expected: (Batch, 4, 64, 64)\n                   Used for the reparameterization trick\n        \n        Returns:\n            Latent representation (Batch, 4, Height/8, Width/8)\n        \"\"\"\n        return self.encoder(x, noise)\n    \n    def decode(self, x):\n        \"\"\"\n        Decode a latent representation back to an image.\n        \n        Args:\n            x: Latent representation (Batch, 4, Height/8, Width/8)\n               Expected: (Batch, 4, 64, 64)\n        \n        Returns:\n            Reconstructed image (Batch, 3, Height, Width)\n            Output: (Batch, 3, 512, 512)\n        \"\"\"\n        return self.decoder(x)\n    \n    def forward(self, x, noise=None):\n        \"\"\"\n        Full forward pass: encode then decode (autoencoder).\n        \n        This is used during VAE pretraining, but NOT during diffusion inference.\n        During diffusion, we only use the decoder on the final denoised latent.\n        \n        Args:\n            x: Input image (Batch, 3, Height, Width)\n            noise: Optional noise for encoding. If None, uses standard Gaussian.\n        \n        Returns:\n            Reconstructed image (Batch, 3, Height, Width)\n        \"\"\"\n        if noise is None:\n            # Generate noise if not provided\n            batch_size = x.shape[0]\n            noise = torch.randn(batch_size, 4, x.shape[2] // 8, x.shape[3] // 8, \n                               device=x.device, dtype=x.dtype)\n        \n        # Encode to latent space\n        latent = self.encode(x, noise)\n        \n        # Decode back to image space\n        reconstructed = self.decode(latent)\n        \n        return reconstructed\n\n\n# Utility function for loading pretrained VAE weights\ndef load_vae_weights(vae, state_dict):\n    \"\"\"\n    Load pretrained weights into the VAE.\n    \n    Args:\n        vae: VAE model instance\n        state_dict: Dictionary of pretrained weights\n    \n    Returns:\n        VAE with loaded weights\n    \"\"\"\n    # Remove 'module.' prefix if present (from DataParallel training)\n    new_state_dict = {}\n    for key, value in state_dict.items():\n        if key.startswith('module.'):\n            new_state_dict[key[7:]] = value\n        else:\n            new_state_dict[key] = value\n    \n    vae.load_state_dict(new_state_dict, strict=False)\n    return vae\n\n\n# Testing and validation utilities\ndef test_vae_dimensions():\n    \"\"\"\n    Test that VAE input/output dimensions are correct.\n    \"\"\"\n    print(\"Testing VAE dimensions...\")\n    \n    vae = VAE()\n    vae.eval()\n    \n    # Test input\n    batch_size = 2\n    x = torch.randn(batch_size, 3, 512, 512)\n    noise = torch.randn(batch_size, 4, 64, 64)\n    \n    # Test encoding\n    with torch.no_grad():\n        latent = vae.encode(x, noise)\n        print(f\"Input shape: {x.shape}\")\n        print(f\"Latent shape: {latent.shape}\")\n        assert latent.shape == (batch_size, 4, 64, 64), \"Encoding failed!\"\n        \n        # Test decoding\n        reconstructed = vae.decode(latent)\n        print(f\"Output shape: {reconstructed.shape}\")\n        assert reconstructed.shape == (batch_size, 3, 512, 512), \"Decoding failed!\"\n        \n        # Test full forward pass\n        output = vae(x, noise)\n        print(f\"Full forward output shape: {output.shape}\")\n        assert output.shape == x.shape, \"Full forward pass failed!\"\n    \n    print(\"✓ All dimension tests passed!\")\n    \n    # Calculate model size\n    total_params = sum(p.numel() for p in vae.parameters())\n    print(f\"\\nTotal parameters: {total_params:,}\")\n    print(f\"Model size: ~{total_params * 4 / (1024**2):.1f} MB (float32)\")\n\n\nif __name__ == \"__main__\":\n    # Run dimension tests\n    test_vae_dimensions()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:11:29.139543Z","iopub.execute_input":"2026-02-28T06:11:29.139781Z","iopub.status.idle":"2026-02-28T06:12:10.148756Z","shell.execute_reply.started":"2026-02-28T06:11:29.139761Z","shell.execute_reply":"2026-02-28T06:12:10.148099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model/clip.py\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass CLIPEmbedding(nn.Module):\n    \"\"\"\n    Embedding layer for CLIP text encoder.\n    Converts token IDs to dense vectors and adds positional information.\n    \"\"\"\n    def __init__(self, n_vocab, n_embd, n_tokens):\n        super().__init__()\n        \n        # Token embedding: converts token IDs to vectors\n        self.token_embedding = nn.Embedding(n_vocab, n_embd)\n        \n        # Positional embedding: adds position information to each token\n        # This is learned, not the sinusoidal encoding used in original Transformer\n        self.position_embedding = nn.Parameter(torch.zeros(n_tokens, n_embd))\n    \n    def forward(self, tokens):\n        \"\"\"\n        Args:\n            tokens: Token IDs (Batch, Seq_Len)\n        \n        Returns:\n            Embedded tokens with positional encoding (Batch, Seq_Len, Dim)\n        \"\"\"\n        # tokens: (Batch, Seq_Len)\n        \n        # Get token embeddings\n        x = self.token_embedding(tokens)\n        # (Batch, Seq_Len, Dim)\n        \n        # Add positional embeddings\n        x += self.position_embedding\n        \n        return x\n\n\nclass CLIPLayer(nn.Module):\n    \"\"\"\n    Single Transformer layer in CLIP.\n    Uses self-attention followed by a feedforward network.\n    \"\"\"\n    def __init__(self, n_head, n_embd):\n        super().__init__()\n        \n        # Layer normalization (applied before attention)\n        self.layernorm_1 = nn.LayerNorm(n_embd)\n        \n        # Self-attention with causal mask (for autoregressive text modeling)\n        self.attention = SelfAttention(n_head, n_embd)\n        \n        # Layer normalization (applied before feedforward)\n        self.layernorm_2 = nn.LayerNorm(n_embd)\n        \n        # Feedforward network (MLP)\n        self.linear_1 = nn.Linear(n_embd, 4 * n_embd)\n        self.linear_2 = nn.Linear(4 * n_embd, n_embd)\n    \n    def forward(self, x):\n        \"\"\"\n        Args:\n            x: Input embeddings (Batch, Seq_Len, Dim)\n        \n        Returns:\n            Transformed embeddings (Batch, Seq_Len, Dim)\n        \"\"\"\n        # Self-attention block with residual connection\n        residue = x\n        x = self.layernorm_1(x)\n        x = self.attention(x, causal_mask=True)\n        x += residue\n        \n        # Feedforward block with residual connection\n        residue = x\n        x = self.layernorm_2(x)\n        \n        # MLP with GELU activation\n        x = self.linear_1(x)\n        x = x * torch.sigmoid(1.702 * x)  # QuickGELU approximation\n        x = self.linear_2(x)\n        \n        x += residue\n        \n        return x\n\n\nclass CLIP(nn.Module):\n    \"\"\"\n    CLIP Text Encoder for Stable Diffusion.\n    \n    This model converts text prompts into embeddings that guide image generation.\n    It's based on the Transformer architecture and uses causal self-attention.\n    \n    Architecture:\n    - Input: Tokenized text (up to 77 tokens)\n    - Output: Context embeddings (Batch, 77, 768)\n    \n    For your DCT steganography research:\n    The text encoder remains mostly unchanged during fine-tuning. However, you might:\n    - Add DCT-related keywords to your prompts during training\n    - Fine-tune with prompts that emphasize texture, detail, complexity\n    - Keep frozen to preserve general language understanding\n    \"\"\"\n    def __init__(self):\n        super().__init__()\n        \n        # Hyperparameters for CLIP (SD 1.5 uses these values)\n        n_vocab = 49408  # Vocabulary size\n        n_embd = 768     # Embedding dimension\n        n_layer = 12     # Number of transformer layers\n        n_head = 12      # Number of attention heads\n        n_tokens = 77    # Maximum sequence length\n        \n        self.embedding = CLIPEmbedding(n_vocab, n_embd, n_tokens)\n        \n        # Stack of Transformer layers\n        self.layers = nn.ModuleList([\n            CLIPLayer(n_head, n_embd) for _ in range(n_layer)\n        ])\n        \n        # Final layer normalization\n        self.layernorm = nn.LayerNorm(n_embd)\n    \n    def forward(self, tokens):\n        \"\"\"\n        Encode text tokens into embeddings.\n        \n        Args:\n            tokens: Tokenized text (Batch, Seq_Len)\n                   Seq_Len is always 77 in Stable Diffusion\n        \n        Returns:\n            Text embeddings (Batch, Seq_Len, Dim)\n            Shape: (Batch, 77, 768)\n        \"\"\"\n        # tokens: (Batch, 77)\n        \n        # Convert tokens to embeddings\n        x = self.embedding(tokens)\n        # (Batch, 77, 768)\n        \n        # Pass through all Transformer layers\n        for layer in self.layers:\n            x = layer(x)\n        \n        # Final normalization\n        x = self.layernorm(x)\n        # (Batch, 77, 768)\n        \n        return x\n\n\nclass SelfAttention(nn.Module):\n    \"\"\"\n    Self-attention mechanism used in CLIP Transformer.\n    (Same as VAE attention but included here for completeness)\n    \"\"\"\n    def __init__(self, n_heads, d_embed, in_proj_bias=True, out_proj_bias=True):\n        super().__init__()\n        \n        # Combined Q, K, V projection\n        self.in_proj = nn.Linear(d_embed, 3 * d_embed, bias=in_proj_bias)\n        \n        # Output projection\n        self.out_proj = nn.Linear(d_embed, d_embed, bias=out_proj_bias)\n        \n        self.n_heads = n_heads\n        self.d_head = d_embed // n_heads\n    \n    def forward(self, x, causal_mask=False):\n        \"\"\"\n        Args:\n            x: Input (Batch, Seq_Len, Dim)\n            causal_mask: If True, apply causal masking (for autoregressive modeling)\n        \n        Returns:\n            Attention output (Batch, Seq_Len, Dim)\n        \"\"\"\n        input_shape = x.shape\n        batch_size, sequence_length, d_embed = input_shape\n        \n        interim_shape = (batch_size, sequence_length, self.n_heads, self.d_head)\n        \n        # Project to Q, K, V and split\n        q, k, v = self.in_proj(x).chunk(3, dim=-1)\n        \n        # Reshape for multi-head attention\n        q = q.view(interim_shape).transpose(1, 2)\n        k = k.view(interim_shape).transpose(1, 2)\n        v = v.view(interim_shape).transpose(1, 2)\n        \n        # Compute attention scores\n        weight = q @ k.transpose(-1, -2)\n        \n        if causal_mask:\n            # Apply causal mask (prevent attending to future tokens)\n            mask = torch.ones_like(weight, dtype=torch.bool).triu(1)\n            weight.masked_fill_(mask, -torch.inf)\n        \n        # Scale and softmax\n        weight /= torch.sqrt(torch.tensor(self.d_head, dtype=torch.float32))\n        weight = F.softmax(weight, dim=-1)\n        \n        # Apply attention to values\n        output = weight @ v\n        \n        # Reshape back\n        output = output.transpose(1, 2).reshape(input_shape)\n        \n        # Output projection\n        output = self.out_proj(output)\n        \n        return output\n\n\n# Tokenizer wrapper (you'll need to use the actual CLIP tokenizer)\nclass CLIPTokenizer:\n    \"\"\"\n    Wrapper for CLIP tokenizer.\n    In practice, use: transformers.CLIPTokenizer.from_pretrained(\"openai/clip-vit-large-patch14\")\n    \"\"\"\n    def __init__(self):\n        # This is a placeholder - in real implementation, load the actual tokenizer\n        self.max_length = 77\n    \n    def tokenize(self, text, truncate=True, pad=True):\n        \"\"\"\n        Tokenize text for CLIP.\n        \n        Args:\n            text: String or list of strings\n            truncate: Truncate to max_length\n            pad: Pad to max_length\n        \n        Returns:\n            Token IDs (Batch, 77)\n        \"\"\"\n        # In practice, use the HuggingFace tokenizer:\n        # from transformers import CLIPTokenizer\n        # tokenizer = CLIPTokenizer.from_pretrained(\"openai/clip-vit-large-patch14\")\n        # return tokenizer(text, max_length=77, padding=\"max_length\", \n        #                  truncation=True, return_tensors=\"pt\").input_ids\n        pass\n\n\n# Testing utility\ndef test_clip_dimensions():\n    \"\"\"\n    Test that CLIP input/output dimensions are correct.\n    \"\"\"\n    print(\"Testing CLIP dimensions...\")\n    \n    clip = CLIP()\n    clip.eval()\n    \n    # Test input\n    batch_size = 2\n    seq_len = 77\n    tokens = torch.randint(0, 49408, (batch_size, seq_len))\n    \n    # Test encoding\n    with torch.no_grad():\n        embeddings = clip(tokens)\n        print(f\"Input shape: {tokens.shape}\")\n        print(f\"Output shape: {embeddings.shape}\")\n        assert embeddings.shape == (batch_size, 77, 768), \"CLIP encoding failed!\"\n    \n    print(\"✓ All dimension tests passed!\")\n    \n    # Calculate model size\n    total_params = sum(p.numel() for p in clip.parameters())\n    print(f\"\\nTotal parameters: {total_params:,}\")\n    print(f\"Model size: ~{total_params * 4 / (1024**2):.1f} MB (float32)\")\n\n\nif __name__ == \"__main__\":\n    test_clip_dimensions()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:12:10.149848Z","iopub.execute_input":"2026-02-28T06:12:10.150091Z","iopub.status.idle":"2026-02-28T06:12:11.224345Z","shell.execute_reply.started":"2026-02-28T06:12:10.150071Z","shell.execute_reply":"2026-02-28T06:12:11.223726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNET_AttentionBlock(nn.Module):\n    \"\"\"Self-attention + Cross-attention for U-Net.\"\"\"\n    def __init__(self, n_head, n_embd, d_context=768):\n        super().__init__()\n        channels = n_head * n_embd\n        \n        self.groupnorm = nn.GroupNorm(32, channels, eps=1e-6)\n        self.conv_input = nn.Conv2d(channels, channels, kernel_size=1, padding=0)\n        \n        self.layernorm_1 = nn.LayerNorm(channels)\n        self.attention_1 = CrossAttention(n_head, channels, channels)\n        \n        self.layernorm_2 = nn.LayerNorm(channels)\n        self.attention_2 = CrossAttention(n_head, channels, d_context)\n        \n        self.layernorm_3 = nn.LayerNorm(channels)\n        self.linear_geglu_1 = nn.Linear(channels, 4 * channels * 2)\n        self.linear_geglu_2 = nn.Linear(4 * channels, channels)\n        \n        self.conv_output = nn.Conv2d(channels, channels, kernel_size=1, padding=0)\n    \n    def forward(self, x, context):\n        residue_long = x\n        \n        x = self.groupnorm(x)\n        x = self.conv_input(x)\n        \n        n, c, h, w = x.shape\n        x = x.view((n, c, h * w))\n        x = x.transpose(-1, -2)\n        \n        # Self-attention\n        residue_short = x\n        x = self.layernorm_1(x)\n        x = self.attention_1(x)\n        x += residue_short\n        \n        # Cross-attention\n        residue_short = x\n        x = self.layernorm_2(x)\n        x = self.attention_2(x, context)\n        x += residue_short\n        \n        # Feedforward\n        residue_short = x\n        x = self.layernorm_3(x)\n        x, gate = self.linear_geglu_1(x).chunk(2, dim=-1)\n        x = x * F.gelu(gate)\n        x = self.linear_geglu_2(x)\n        x += residue_short\n        \n        x = x.transpose(-1, -2)\n        x = x.view((n, c, h, w))\n        \n        return self.conv_output(x) + residue_long\n\n\n# ============================================================================\n# PART 4: U-NET RESIDUAL BLOCK\n# ============================================================================\n\nclass CrossAttention(nn.Module):\n    \"\"\"Cross-attention for text conditioning.\"\"\"\n    def __init__(self, n_heads, d_embed, d_cross, in_proj_bias=True, out_proj_bias=True):\n        super().__init__()\n        self.q_proj = nn.Linear(d_embed, d_embed, bias=in_proj_bias)\n        self.k_proj = nn.Linear(d_cross, d_embed, bias=in_proj_bias)\n        self.v_proj = nn.Linear(d_cross, d_embed, bias=in_proj_bias)\n        self.out_proj = nn.Linear(d_embed, d_embed, bias=out_proj_bias)\n        self.n_heads = n_heads\n        self.d_head = d_embed // n_heads\n    \n    def forward(self, x, y=None):\n        input_shape = x.shape\n        batch_size, sequence_length, d_embed = input_shape\n        \n        if y is None:\n            y = x\n        \n        interim_shape = (batch_size, -1, self.n_heads, self.d_head)\n        \n        q = self.q_proj(x)\n        k = self.k_proj(y)\n        v = self.v_proj(y)\n        \n        q = q.view(interim_shape).transpose(1, 2)\n        k = k.view(interim_shape).transpose(1, 2)\n        v = v.view(interim_shape).transpose(1, 2)\n        \n        weight = q @ k.transpose(-1, -2)\n        weight /= math.sqrt(self.d_head)\n        weight = F.softmax(weight, dim=-1)\n        \n        output = weight @ v\n        output = output.transpose(1, 2).contiguous()\n        output = output.view(input_shape)\n        output = self.out_proj(output)\n        \n        return output\n\n\n# ============================================================================\n# PART 3: U-NET ATTENTION BLOCK\n# ============================================================================\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:12:11.225249Z","iopub.execute_input":"2026-02-28T06:12:11.225776Z","iopub.status.idle":"2026-02-28T06:12:11.237662Z","shell.execute_reply.started":"2026-02-28T06:12:11.225752Z","shell.execute_reply":"2026-02-28T06:12:11.236955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNET_ResidualBlock(nn.Module):\n    \"\"\"U-Net residual block with time conditioning.\"\"\"\n    def __init__(self, in_channels, out_channels, n_time=1280):\n        super().__init__()\n        \n        self.groupnorm_feature = nn.GroupNorm(32, in_channels)\n        self.conv_feature = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)\n        self.linear_time = nn.Linear(n_time, out_channels)\n        \n        self.groupnorm_merged = nn.GroupNorm(32, out_channels)\n        self.conv_merged = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n        \n        if in_channels == out_channels:\n            self.residual_layer = nn.Identity()\n        else:\n            self.residual_layer = nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0)\n    \n    def forward(self, feature, time):\n        residue = feature\n        \n        feature = self.groupnorm_feature(feature)\n        feature = F.silu(feature)\n        feature = self.conv_feature(feature)\n        \n        time = self.linear_time(F.silu(time))\n        time = time.unsqueeze(-1).unsqueeze(-1)\n        \n        merged = feature + time\n        merged = self.groupnorm_merged(merged)\n        merged = F.silu(merged)\n        merged = self.conv_merged(merged)\n        \n        return merged + self.residual_layer(residue)\nclass Upsample(nn.Module):\n    \"\"\"Upsample by 2x.\"\"\"\n    def __init__(self, channels):\n        super().__init__()\n        self.conv = nn.Conv2d(channels, channels, kernel_size=3, padding=1)\n    \n    def forward(self, x):\n        x = F.interpolate(x, scale_factor=2, mode='nearest') \n        return self.conv(x)\n\n\nclass Downsample(nn.Module):\n    \"\"\"Downsample by 2x.\"\"\"\n    def __init__(self, channels):\n        super().__init__()\n        self.conv = nn.Conv2d(channels, channels, kernel_size=3, stride=2, padding=1)\n    \n    def forward(self, x):\n        return self.conv(x)\n\n\n# ============================================================================\n# PART 6: SWITCH SEQUENTIAL\n# ============================================================================\n\nclass SwitchSequential(nn.Sequential):\n    \"\"\"Sequential that handles different input types.\"\"\"\n    def forward(self, x, context, time):\n        for layer in self:\n            if isinstance(layer, UNET_AttentionBlock):\n                x = layer(x, context)\n            elif isinstance(layer, UNET_ResidualBlock):\n                x = layer(x, time)\n            else:\n                x = layer(x)\n        return x\n\nclass TimeEmbedding(nn.Module):\n    \"\"\"Converts timestep into embedding.\"\"\"\n    def __init__(self, n_embd):\n        super().__init__()\n        self.linear_1 = nn.Linear(n_embd, 4 * n_embd)\n        self.linear_2 = nn.Linear(4 * n_embd, 4 * n_embd)\n\n    def forward(self, x):\n        x = self.linear_1(x)\n        x = F.silu(x)  \n        x = self.linear_2(x)\n        return x\n\n\ndef get_time_embedding(timestep, n_embd=320):\n    \"\"\"Generate sinusoidal time embeddings.\"\"\"\n    freqs = torch.pow(10000, -torch.arange(start=0, end=n_embd // 2, dtype=torch.float32) / (n_embd // 2))\n    x = timestep[:, None] * freqs[None]\n    x = torch.cat([torch.cos(x), torch.sin(x)], dim=-1)\n    return x\n     ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:12:11.238682Z","iopub.execute_input":"2026-02-28T06:12:11.239053Z","iopub.status.idle":"2026-02-28T06:12:11.255839Z","shell.execute_reply.started":"2026-02-28T06:12:11.239020Z","shell.execute_reply":"2026-02-28T06:12:11.255121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport math\nclass UNET(nn.Module):\n    \"\"\"Complete U-Net for Stable Diffusion.\"\"\"\n    def __init__(self):\n        super().__init__()\n        \n        self.time_embedding = TimeEmbedding(320)\n        self.conv_in = nn.Conv2d(4, 320, kernel_size=3, padding=1)\n        \n        # Encoder (12 blocks to match decoder)\n        self.encoders = nn.ModuleList([\n            # (Batch, 320, 64, 64)\n            SwitchSequential(UNET_ResidualBlock(320, 320), UNET_AttentionBlock(8, 40)),\n            SwitchSequential(UNET_ResidualBlock(320, 320), UNET_AttentionBlock(8, 40)),\n            \n            # (Batch, 320, 32, 32)\n            SwitchSequential(Downsample(320)),\n            SwitchSequential(UNET_ResidualBlock(320, 640), UNET_AttentionBlock(8, 80)),\n            SwitchSequential(UNET_ResidualBlock(640, 640), UNET_AttentionBlock(8, 80)),\n            \n            # (Batch, 640, 16, 16)\n            SwitchSequential(Downsample(640)),\n            SwitchSequential(UNET_ResidualBlock(640, 1280), UNET_AttentionBlock(8, 160)),\n            SwitchSequential(UNET_ResidualBlock(1280, 1280), UNET_AttentionBlock(8, 160)),\n            \n            # (Batch, 1280, 8, 8)\n            SwitchSequential(Downsample(1280)),\n            SwitchSequential(UNET_ResidualBlock(1280, 1280)),\n            SwitchSequential(UNET_ResidualBlock(1280, 1280)),\n            SwitchSequential(UNET_ResidualBlock(1280, 1280)),  # ADDED THIS ONE!\n        ])\n        \n        # Bottleneck\n        self.bottleneck = SwitchSequential(\n            UNET_ResidualBlock(1280, 1280),\n            UNET_AttentionBlock(8, 160),\n            UNET_ResidualBlock(1280, 1280),\n        )\n        \n        # Decoder (12 blocks)\n        self.decoders = nn.ModuleList([\n            SwitchSequential(UNET_ResidualBlock(2560, 1280)),\n            SwitchSequential(UNET_ResidualBlock(2560, 1280)),\n            SwitchSequential(UNET_ResidualBlock(2560, 1280)),\n            SwitchSequential(UNET_ResidualBlock(2560, 1280), Upsample(1280)),\n            \n            SwitchSequential(UNET_ResidualBlock(2560, 1280), UNET_AttentionBlock(8, 160)),\n            SwitchSequential(UNET_ResidualBlock(2560, 1280), UNET_AttentionBlock(8, 160)),\n            SwitchSequential(UNET_ResidualBlock(1920, 1280), UNET_AttentionBlock(8, 160), Upsample(1280)),\n            \n            SwitchSequential(UNET_ResidualBlock(1920, 640), UNET_AttentionBlock(8, 80)),\n            SwitchSequential(UNET_ResidualBlock(1280, 640), UNET_AttentionBlock(8, 80)),\n            SwitchSequential(UNET_ResidualBlock(960, 640), UNET_AttentionBlock(8, 80), Upsample(640)),\n            \n            SwitchSequential(UNET_ResidualBlock(960, 320), UNET_AttentionBlock(8, 40)),\n            SwitchSequential(UNET_ResidualBlock(640, 320), UNET_AttentionBlock(8, 40)),\n        ])\n        \n        self.final = nn.Sequential(\n            nn.GroupNorm(32, 320),\n            nn.SiLU(),\n            nn.Conv2d(320, 4, kernel_size=3, padding=1),\n        )\n    \n    def forward(self, latent, context, time):\n        time = get_time_embedding(time).to(latent.device)\n        time = self.time_embedding(time)\n        \n        x = self.conv_in(latent)\n        \n        skip_connections = []\n        for encoder in self.encoders:\n            x = encoder(x, context, time)\n            skip_connections.append(x)\n        \n        x = self.bottleneck(x, context, time)\n        \n        for decoder in self.decoders:\n            x = torch.cat([x, skip_connections.pop()], dim=1)\n            x = decoder(x, context, time)\n        \n        output = self.final(x)\n        return output\n# Test the complete U-Net\nprint(\"Testing U-Net...\")\n\nunet = UNET()\nunet.eval()\n\n# Test inputs\nbatch_size = 1\nlatent = torch.randn(batch_size, 4, 64, 64)\ncontext = torch.randn(batch_size, 77, 768)  # From CLIP\ntime = torch.tensor([50])  # Timestep\n\nwith torch.no_grad():\n    noise_pred = unet(latent, context, time)\n\nprint(f\"✓ Latent input: {latent.shape}\")\nprint(f\"✓ Context input: {context.shape}\")\nprint(f\"✓ Time: {time}\")\nprint(f\"✓ Predicted noise: {noise_pred.shape}\")\n\n# Count parameters\ntotal_params = sum(p.numel() for p in unet.parameters())\nprint(f\"\\n✓ Total U-Net parameters: {total_params:,}\")\nprint(f\"✓ Model size: ~{total_params * 4 / (1024**2):.1f} MB\")        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:12:11.256888Z","iopub.execute_input":"2026-02-28T06:12:11.257575Z","iopub.status.idle":"2026-02-28T06:12:23.501334Z","shell.execute_reply.started":"2026-02-28T06:12:11.257545Z","shell.execute_reply":"2026-02-28T06:12:23.500594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\n\n\nclass DDPMScheduler:\n    \"\"\"\n    DDPM Scheduler for Stable Diffusion.\n    \n    Controls the noise schedule for training and inference.\n    This defines HOW MUCH noise to add at each timestep.\n    \n    For your DCT research:\n    - The noise schedule affects which frequencies get denoised when\n    - Early steps (high noise) → coarse structure\n    - Late steps (low noise) → fine details (your mid-frequencies!)\n    \"\"\"\n    \n    def __init__(\n        self,\n        num_train_timesteps=1000,\n        beta_start=0.00085,\n        beta_end=0.012,\n        beta_schedule=\"scaled_linear\"\n    ):\n        \"\"\"\n        Args:\n            num_train_timesteps: Total diffusion steps\n            beta_start: Starting noise level\n            beta_end: Ending noise level\n            beta_schedule: How noise increases over time\n        \"\"\"\n        self.num_train_timesteps = num_train_timesteps\n        self.timesteps = torch.from_numpy(np.arange(0, num_train_timesteps)[::-1].copy())\n        \n        # Create beta schedule (noise schedule)\n        if beta_schedule == \"linear\":\n            self.betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32)\n        elif beta_schedule == \"scaled_linear\":\n            # This is what SD 1.5 uses\n            self.betas = torch.linspace(beta_start**0.5, beta_end**0.5, num_train_timesteps, dtype=torch.float32) ** 2\n        \n        # Pre-compute useful values\n        self.alphas = 1.0 - self.betas\n        self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)\n        self.alphas_cumprod_prev = torch.cat([torch.tensor([1.0]), self.alphas_cumprod[:-1]])\n        \n        # For sampling\n        self.sqrt_alphas_cumprod = self.alphas_cumprod ** 0.5\n        self.sqrt_one_minus_alphas_cumprod = (1 - self.alphas_cumprod) ** 0.5\n        \n        # For posterior q(x_{t-1} | x_t, x_0)\n        self.posterior_variance = self.betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)\n    \n    def add_noise(self, original_samples, noise, timesteps):\n        \"\"\"\n        Add noise to samples (forward diffusion).\n        \n        q(x_t | x_0) = sqrt(alpha_t) * x_0 + sqrt(1 - alpha_t) * noise\n        \n        Args:\n            original_samples: Clean latents (Batch, 4, 64, 64)\n            noise: Random noise (same shape)\n            timesteps: Which timestep (Batch,)\n        \n        Returns:\n            Noisy samples\n        \"\"\"\n        sqrt_alpha_prod = self.sqrt_alphas_cumprod[timesteps].flatten()\n        while len(sqrt_alpha_prod.shape) < len(original_samples.shape):\n            sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1)\n        \n        sqrt_one_minus_alpha_prod = self.sqrt_one_minus_alphas_cumprod[timesteps].flatten()\n        while len(sqrt_one_minus_alpha_prod.shape) < len(original_samples.shape):\n            sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1)\n        \n        noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise\n        return noisy_samples\n    \n    def step(self, model_output, timestep, sample):\n        \"\"\"\n        Reverse diffusion step: x_t -> x_{t-1}\n        \n        This is the DENOISING step used during inference.\n        \n        Args:\n            model_output: Predicted noise from U-Net\n            timestep: Current timestep\n            sample: Current noisy sample x_t\n        \n        Returns:\n            Previous (less noisy) sample x_{t-1}\n        \"\"\"\n        t = timestep\n        \n        # Get parameters for this timestep\n        alpha_prod_t = self.alphas_cumprod[t]\n        alpha_prod_t_prev = self.alphas_cumprod_prev[t] if t > 0 else torch.tensor(1.0)\n        beta_prod_t = 1 - alpha_prod_t\n        \n        # Predict x_0 from x_t and predicted noise\n        pred_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5\n        \n        # Clip x_0 (optional, helps stability)\n        pred_original_sample = torch.clamp(pred_original_sample, -1, 1)\n        \n        # Compute coefficients for x_{t-1}\n        pred_original_sample_coeff = (alpha_prod_t_prev ** 0.5 * self.betas[t]) / beta_prod_t\n        current_sample_coeff = self.alphas[t] ** 0.5 * (1 - alpha_prod_t_prev) / beta_prod_t\n        \n        # Compute x_{t-1}\n        pred_prev_sample = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * sample\n        \n        # Add noise (except at last step)\n        if t > 0:\n            noise = torch.randn_like(sample)\n            variance = (self.posterior_variance[t] ** 0.5) * noise\n            pred_prev_sample = pred_prev_sample + variance\n        \n        return pred_prev_sample\n    \n    def set_timesteps(self, num_inference_steps):\n        \"\"\"\n        Set timesteps for inference (usually fewer than training steps).\n        \n        Args:\n            num_inference_steps: Number of denoising steps (typically 20-50)\n        \"\"\"\n        step_ratio = self.num_train_timesteps // num_inference_steps\n        timesteps = (np.arange(0, num_inference_steps) * step_ratio).round()[::-1].copy().astype(np.int64)\n        self.timesteps = torch.from_numpy(timesteps)\n\n\n# ============================================================================\n# TEST THE SCHEDULER\n# ============================================================================\n\nprint(\"=\"*60)\nprint(\"TESTING DDPM SCHEDULER\")\nprint(\"=\"*60)\n\nscheduler = DDPMScheduler(num_train_timesteps=1000)\n\n# Test adding noise\nclean_latent = torch.randn(1, 4, 64, 64)\nnoise = torch.randn_like(clean_latent)\ntimesteps = torch.tensor([500])\n\nnoisy_latent = scheduler.add_noise(clean_latent, noise, timesteps)\n\nprint(f\"\\n✓ Clean latent: {clean_latent.shape}\")\nprint(f\"✓ Added noise at timestep {timesteps.item()}\")\nprint(f\"✓ Noisy latent: {noisy_latent.shape}\")\n\n# Test denoising step\npredicted_noise = torch.randn_like(noisy_latent)\ndenoised = scheduler.step(predicted_noise, timesteps[0], noisy_latent)\n\nprint(f\"\\n✓ Denoising step successful\")\nprint(f\"✓ Output: {denoised.shape}\")\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"✓ SCHEDULER WORKING!\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:47:22.728263Z","iopub.execute_input":"2026-02-28T06:47:22.728793Z","iopub.status.idle":"2026-02-28T06:47:22.754681Z","shell.execute_reply.started":"2026-02-28T06:47:22.728763Z","shell.execute_reply":"2026-02-28T06:47:22.753907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class StableDiffusionPipeline:\n    \"\"\"\n    Complete Stable Diffusion inference pipeline.\n    \n    This is what you'll use to generate images!\n    \n    Components:\n    - VAE: Encoder/Decoder\n    - CLIP: Text encoder\n    - U-Net: Denoising network\n    - Scheduler: Controls diffusion process\n    \"\"\"\n    \n    def __init__(self, vae, clip, unet, scheduler, device=\"cuda\"):\n        self.vae = vae.to(device)\n        self.clip = clip.to(device)\n        self.unet = unet.to(device)\n        self.scheduler = scheduler\n        self.device = device\n        \n        # Set to eval mode\n        self.vae.eval()\n        self.clip.eval()\n        self.unet.eval()\n    \n    @torch.no_grad()\n    def generate(\n        self,\n        prompt,\n        num_inference_steps=50,\n        guidance_scale=7.5,\n        height=512,\n        width=512,\n        generator=None\n    ):\n        \"\"\"\n        Generate an image from a text prompt.\n        \n        Args:\n            prompt: Text description (string or list of strings)\n            num_inference_steps: Number of denoising steps (20-50)\n            guidance_scale: How much to follow the prompt (7-15)\n            height: Image height (must be multiple of 8)\n            width: Image width (must be multiple of 8)\n            generator: Random seed for reproducibility\n        \n        Returns:\n            Generated image (PIL Image)\n        \"\"\"\n        # 1. Encode text prompt with CLIP\n        text_embeddings = self._encode_prompt(prompt)\n        \n        # 2. Prepare latents (random noise)\n        latents = torch.randn(\n            (1, 4, height // 8, width // 8),\n            generator=generator,\n            device=self.device,\n            dtype=torch.float32\n        )\n        \n        # Scale initial noise\n        latents = latents * self.scheduler.sqrt_alphas_cumprod[0]\n        \n        # 3. Set timesteps\n        self.scheduler.set_timesteps(num_inference_steps)\n        \n        # 4. Denoising loop\n        for i, t in enumerate(self.scheduler.timesteps):\n            # Expand latents for classifier-free guidance\n            latent_model_input = torch.cat([latents] * 2)\n            \n            # Predict noise\n            with torch.no_grad():\n                noise_pred = self.unet(\n                    latent_model_input,\n                    text_embeddings,\n                    t.unsqueeze(0).to(self.device)\n                )\n            \n            # Perform classifier-free guidance\n            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)\n            noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)\n            \n            # Compute previous noisy sample\n            latents = self.scheduler.step(noise_pred, t, latents)\n            \n            if i % 10 == 0:\n                print(f\"Step {i}/{num_inference_steps}\")\n        \n        # 5. Decode latents to image\n        image = self.vae.decode(latents)\n        \n        # 6. Convert to PIL Image\n        image = self._decode_latents(image)\n        \n        return image\n    \n    def _encode_prompt(self, prompt):\n        \"\"\"\n        Encode text prompt using CLIP.\n        \n        For classifier-free guidance, we need:\n        - Conditional embeddings (with prompt)\n        - Unconditional embeddings (empty prompt)\n        \"\"\"\n        # This is a simplified version\n        # In practice, you'd use the actual CLIP tokenizer\n        \n        # Placeholder: random tokens for demonstration\n        # In real code, use: tokenizer(prompt, max_length=77, padding=\"max_length\")\n        tokens = torch.randint(0, 49408, (1, 77), device=self.device)\n        \n        # Encode text\n        text_embeddings = self.clip(tokens)\n        \n        # For classifier-free guidance, also encode empty prompt\n        uncond_tokens = torch.zeros_like(tokens)\n        uncond_embeddings = self.clip(uncond_tokens)\n        \n        # Concatenate for classifier-free guidance\n        text_embeddings = torch.cat([uncond_embeddings, text_embeddings])\n        \n        return text_embeddings\n    \n    def _decode_latents(self, latents):\n        \"\"\"\n        Convert latents to PIL Image.\n        \"\"\"\n        # Denormalize from [-1, 1] to [0, 255]\n        image = (latents / 2 + 0.5).clamp(0, 1)\n        image = image.cpu().permute(0, 2, 3, 1).numpy()\n        image = (image * 255).round().astype(\"uint8\")\n        \n        # For now, just return as numpy array\n        # In practice, convert to PIL: Image.fromarray(image[0])\n        return image[0]\n    \n    @torch.no_grad()\n    def img2img(\n        self,\n        prompt,\n        init_image,\n        strength=0.7,\n        num_inference_steps=50,\n        guidance_scale=7.5\n    ):\n        \"\"\"\n        Image-to-image generation (YOUR use case!).\n        \n        Args:\n            prompt: Text description\n            init_image: Input image tensor (1, 3, 512, 512)\n            strength: How much to transform (0=no change, 1=full generation)\n            num_inference_steps: Denoising steps\n            guidance_scale: Text guidance strength\n        \n        Returns:\n            Generated image\n        \"\"\"\n        # 1. Encode prompt\n        text_embeddings = self._encode_prompt(prompt)\n        \n        # 2. Encode image to latent\n        noise_vae = torch.randn(1, 4, 64, 64, device=self.device)\n        latents = self.vae.encode(init_image.to(self.device), noise_vae)\n        \n        # 3. Determine how many steps to denoise\n        self.scheduler.set_timesteps(num_inference_steps)\n        start_step = int(num_inference_steps * (1 - strength))\n        timesteps = self.scheduler.timesteps[start_step:]\n        \n        # 4. Add noise to latents\n        noise = torch.randn_like(latents)\n        latents = self.scheduler.add_noise(latents, noise, timesteps[0:1])\n        \n        # 5. Denoising loop (same as text2img)\n        for i, t in enumerate(timesteps):\n            latent_model_input = torch.cat([latents] * 2)\n            \n            noise_pred = self.unet(\n                latent_model_input,\n                text_embeddings,\n                t.unsqueeze(0).to(self.device)\n            )\n            \n            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)\n            noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)\n            \n            latents = self.scheduler.step(noise_pred, t, latents)\n            \n            if i % 10 == 0:\n                print(f\"Img2Img Step {i}/{len(timesteps)}\")\n        \n        # 6. Decode\n        image = self.vae.decode(latents)\n        image = self._decode_latents(image)\n        \n        return image\n\n\n# ============================================================================\n# TEST (Conceptual - won't actually run without pretrained weights)\n# ============================================================================\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"COMPLETE PIPELINE STRUCTURE READY!\")\nprint(\"=\"*60)\nprint(\"\\nYou now have:\")\nprint(\"  ✅ VAE (Encoder + Decoder)\")\nprint(\"  ✅ CLIP (Text Encoder)\")\nprint(\"  ✅ U-Net (Denoising Network)\")\nprint(\"  ✅ Scheduler (DDPM)\")\nprint(\"  ✅ Complete Pipeline (text2img + img2img)\")\nprint(\"\\nNext: Training pipeline with DCT loss!\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T06:47:50.337866Z","iopub.execute_input":"2026-02-28T06:47:50.338476Z","iopub.status.idle":"2026-02-28T06:47:50.354296Z","shell.execute_reply.started":"2026-02-28T06:47:50.338449Z","shell.execute_reply":"2026-02-28T06:47:50.353753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# COMPLETE STABLE DIFFUSION TRAINING FOR DCT STEGANOGRAPHY\n# Everything in one file - just set your paths and run!\n# ============================================================================\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nimport os\nfrom tqdm import tqdm\nimport numpy as np\nimport random\n\n\n# ============================================================================\n# CONFIGURATION - CHANGE THESE PATHS!\n# ============================================================================\n\nclass Config:\n    # YOUR PATHS - CHANGE THESE!\n    dataset_path = \"/path/to/your/dataset\"  # Folder with images\n    output_dir = \"/path/to/save/checkpoints\"  # Where to save trained models\n    pretrained_weights = None  # Path to SD 1.5 weights (optional)\n    \n    # Training settings\n    batch_size = 4  # Lower if GPU runs out of memory\n    num_epochs = 10\n    learning_rate = 1e-5\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    \n    # Image settings\n    image_size = 512\n    \n    # Loss weights (TUNE THESE for your DCT research!)\n    lambda_diffusion = 1.0  # Standard diffusion loss\n    lambda_dct = 0.5  # DCT embedding quality loss\n    lambda_perceptual = 0.1  # Visual quality preservation\n    \n    # DCT settings\n    dct_block_size = 8\n    \n    # Diffusion settings\n    num_train_timesteps = 1000\n    num_inference_steps = 50\n    \n    # Logging\n    save_every = 500  # Save checkpoint every N steps\n    log_every = 10  # Print loss every N steps\n\n\n# ============================================================================\n# PROMPT GENERATOR FOR DCT STEGANOGRAPHY\n# ============================================================================\n\nclass PromptGenerator:\n    \"\"\"\n    Generates diverse prompts that encourage DCT-optimal image characteristics.\n    \n    All prompts aim to produce images with:\n    - Rich texture and detail (high mid-frequency content)\n    - Natural complexity (good for embedding)\n    - High quality (visual appeal)\n    \"\"\"\n    \n    def __init__(self):\n        # Different vocabulary categories for variety\n        \n        self.quality_terms = [\n            \"high quality\", \"ultra detailed\", \"professional\", \"sharp\", \"crisp\",\n            \"highly detailed\", \"premium quality\", \"masterpiece\", \"exceptional quality\",\n            \"refined\", \"pristine\", \"immaculate\", \"flawless\", \"superior quality\"\n        ]\n        \n        self.texture_terms = [\n            \"rich textures\", \"intricate textures\", \"detailed textures\", \"complex patterns\",\n            \"fine details\", \"elaborate textures\", \"nuanced textures\", \"varied textures\",\n            \"subtle textures\", \"delicate patterns\", \"intricate patterns\", \"layered textures\"\n        ]\n        \n        self.detail_terms = [\n            \"highly detailed\", \"extremely detailed\", \"meticulously detailed\", \"finely detailed\",\n            \"intricately detailed\", \"carefully detailed\", \"precisely detailed\", \"thoroughly detailed\",\n            \"richly detailed\", \"abundantly detailed\", \"extensively detailed\", \"comprehensively detailed\"\n        ]\n        \n        self.complexity_terms = [\n            \"complex composition\", \"intricate composition\", \"sophisticated design\", \"elaborate scene\",\n            \"multifaceted scene\", \"layered composition\", \"nuanced scene\", \"rich composition\",\n            \"detailed scene\", \"textured composition\", \"varied elements\", \"diverse components\"\n        ]\n        \n        self.clarity_terms = [\n            \"crystal clear\", \"pin sharp\", \"perfectly focused\", \"razor sharp\", \"well-defined\",\n            \"clear and detailed\", \"sharp focus\", \"highly defined\", \"well-rendered\", \"precisely rendered\"\n        ]\n        \n        self.subject_types = [\n            \"landscape\", \"nature scene\", \"architectural detail\", \"natural environment\",\n            \"scenic view\", \"outdoor scene\", \"environmental portrait\", \"detailed scenery\",\n            \"natural landscape\", \"architectural scene\", \"textured surface\", \"organic scene\",\n            \"natural composition\", \"environmental detail\", \"scenic composition\"\n        ]\n        \n        # Template structures\n        self.templates = [\n            \"{quality}, {texture}, {subject}\",\n            \"{detail} {subject} with {texture}\",\n            \"{quality} {subject}, {clarity}, {texture}\",\n            \"{subject} with {complexity}, {quality}\",\n            \"{detail} {subject}, {texture}, {clarity}\",\n            \"{quality}, {complexity}, {texture} {subject}\",\n            \"{subject}, {detail}, {texture}\",\n            \"{clarity} {subject} featuring {texture}\",\n            \"{quality} {subject} with {detail} and {texture}\",\n            \"{subject}, {complexity}, {quality}, {texture}\",\n            \"{detail}, {clarity} {subject} with {texture}\",\n            \"{quality} {subject}, {complexity}, {detail}\",\n            \"{texture} {subject}, {quality}, {clarity}\",\n            \"{subject} with {detail}, {texture}, {quality}\",\n            \"{complexity}, {quality}, {detail} {subject}\",\n        ]\n    \n    def generate(self):\n        \"\"\"Generate a random prompt optimized for DCT steganography.\"\"\"\n        template = random.choice(self.templates)\n        \n        prompt = template.format(\n            quality=random.choice(self.quality_terms),\n            texture=random.choice(self.texture_terms),\n            detail=random.choice(self.detail_terms),\n            complexity=random.choice(self.complexity_terms),\n            clarity=random.choice(self.clarity_terms),\n            subject=random.choice(self.subject_types)\n        )\n        \n        return prompt\n    \n    def generate_batch(self, batch_size):\n        \"\"\"Generate multiple unique prompts for a batch.\"\"\"\n        return [self.generate() for _ in range(batch_size)]\n\n\n# ============================================================================\n# DATASET CLASS WITH PROMPT GENERATOR\n# ============================================================================\n\nclass SteganoDataset(Dataset):\n    \"\"\"\n    Dataset for DCT steganography training.\n    \n    Put all your images in one folder.\n    Images will be automatically resized to 512x512.\n    Each image gets a random DCT-optimized prompt.\n    \"\"\"\n    def __init__(self, image_folder, image_size=512):\n        self.image_folder = image_folder\n        self.image_size = image_size\n        \n        # Initialize prompt generator\n        self.prompt_generator = PromptGenerator()\n        \n        # Get all image files\n        self.image_files = [\n            f for f in os.listdir(image_folder)\n            if f.lower().endswith(('.png', '.jpg', '.jpeg', '.bmp', '.tiff'))\n        ]\n        \n        print(f\"Found {len(self.image_files)} images in {image_folder}\")\n        \n        # Image transforms\n        self.transform = transforms.Compose([\n            transforms.Resize((image_size, image_size)),\n            transforms.ToTensor(),  # Converts to [0, 1]\n            transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])  # Normalize to [-1, 1]\n        ])\n    \n    def __len__(self):\n        return len(self.image_files)\n    \n    def __getitem__(self, idx):\n        img_path = os.path.join(self.image_folder, self.image_files[idx])\n        \n        # Load image\n        image = Image.open(img_path).convert('RGB')\n        \n        # Apply transforms\n        image = self.transform(image)\n        \n        # Generate random DCT-optimized prompt\n        prompt = self.prompt_generator.generate()\n        \n        return {\n            'image': image,\n            'prompt': prompt\n        }\n\n\n# ============================================================================\n# DCT LOSS FUNCTIONS (Your custom loss!)\n# ============================================================================\n\ndef dct_2d_torch(image_block):\n    \"\"\"\n    Apply 2D DCT to 8x8 blocks.\n    Simplified version - uses FFT approximation.\n    \"\"\"\n    # For production, use: pip install torch-dct\n    # and import dct from torch_dct\n    \n    # For now, use FFT as approximation\n    dct = torch.fft.fft2(image_block.float())\n    return dct.real\n\n\ndef create_mid_frequency_mask(block_size=8):\n    \"\"\"\n    Create mask for mid-frequency DCT coefficients.\n    These are optimal for steganography.\n    \"\"\"\n    mask = torch.zeros(block_size, block_size)\n    \n    # Mid-frequency positions (adjust based on your research)\n    # Format: (row, col)\n    mid_freq_positions = [\n        (1, 2), (2, 1), (2, 2),  # Low-mid\n        (1, 3), (3, 1), (2, 3), (3, 2),  # Mid\n        (1, 4), (4, 1), (3, 3),  # Mid-high\n    ]\n    \n    for i, j in mid_freq_positions:\n        mask[i, j] = 1.0\n    \n    return mask\n\n\ndef dct_embedding_quality_loss(generated_images, block_size=8):\n    \"\"\"\n    Calculate DCT-aware loss for steganography.\n    \n    Encourages:\n    1. High variance in mid-frequency coefficients (more embedding capacity)\n    2. Low high-frequency noise\n    3. Smooth spatial characteristics\n    \n    Args:\n        generated_images: (Batch, 3, 512, 512) in range [-1, 1]\n    \n    Returns:\n        Scalar loss value\n    \"\"\"\n    batch_size, channels, height, width = generated_images.shape\n    \n    # Convert to [0, 1] for DCT\n    images = (generated_images + 1) / 2\n    \n    # Extract 8x8 blocks\n    blocks = F.unfold(images, kernel_size=block_size, stride=block_size)\n    blocks = blocks.reshape(batch_size, channels, block_size, block_size, -1)\n    \n    total_loss = 0.0\n    \n    for c in range(channels):\n        # Apply DCT to each block\n        channel_blocks = blocks[:, c, :, :, :]\n        \n        # Process each block\n        num_blocks = channel_blocks.shape[-1]\n        dct_coeffs = []\n        \n        for b in range(num_blocks):\n            block = channel_blocks[:, :, :, b]\n            dct = dct_2d_torch(block)\n            dct_coeffs.append(dct)\n        \n        dct_coeffs = torch.stack(dct_coeffs, dim=-1)\n        \n        # Create frequency masks\n        mid_freq_mask = create_mid_frequency_mask(block_size).to(generated_images.device)\n        \n        # Extract mid-frequency coefficients\n        mid_freq = dct_coeffs * mid_freq_mask.unsqueeze(0).unsqueeze(-1)\n        \n        # Loss 1: Encourage HIGH variance in mid-frequencies (more capacity)\n        mid_freq_variance = mid_freq.var()\n        variance_loss = -torch.log(mid_freq_variance + 1e-8)\n        \n        # Loss 2: Penalize extreme high-frequencies (noise)\n        high_freq_mask = torch.ones_like(mid_freq_mask) - mid_freq_mask\n        high_freq_mask[0, 0] = 0  # Don't penalize DC component\n        high_freq = dct_coeffs * high_freq_mask.unsqueeze(0).unsqueeze(-1)\n        high_freq_loss = torch.abs(high_freq).mean()\n        \n        total_loss += variance_loss + 0.5 * high_freq_loss\n    \n    # Loss 3: Total Variation (spatial smoothness)\n    tv_h = torch.abs(generated_images[:, :, 1:, :] - generated_images[:, :, :-1, :]).mean()\n    tv_w = torch.abs(generated_images[:, :, :, 1:] - generated_images[:, :, :, :-1]).mean()\n    tv_loss = tv_h + tv_w\n    \n    total_loss = total_loss / channels + 0.1 * tv_loss\n    \n    return total_loss\n\n\n# ============================================================================\n# TRAINING LOOP\n# ============================================================================\n\ndef train_stable_diffusion(config):\n    \"\"\"\n    Complete training loop for DCT-optimized Stable Diffusion.\n    \"\"\"\n    \n    print(\"=\"*60)\n    print(\"STABLE DIFFUSION TRAINING FOR DCT STEGANOGRAPHY\")\n    print(\"=\"*60)\n    \n    # Create output directory\n    os.makedirs(config.output_dir, exist_ok=True)\n    \n    # 1. Load dataset\n    print(f\"\\n[1/6] Loading dataset from {config.dataset_path}...\")\n    dataset = SteganoDataset(config.dataset_path, config.image_size)\n    dataloader = DataLoader(\n        dataset,\n        batch_size=config.batch_size,\n        shuffle=True,\n        num_workers=4,\n        pin_memory=True\n    )\n    print(f\"✓ Dataset loaded: {len(dataset)} images\")\n    \n    # Show sample prompts\n    print(\"\\n Sample prompts that will be used:\")\n    for i in range(5):\n        print(f\"  {i+1}. {dataset.prompt_generator.generate()}\")\n    \n    # 2. Initialize models\n    print(\"\\n[2/6] Initializing models...\")\n    \n    # You need to have VAE, CLIP, UNET defined earlier\n    # For now, assume they're already created\n    try:\n        vae = VAE().to(config.device)\n        clip = CLIP().to(config.device)\n        unet = UNET().to(config.device)\n        scheduler = DDPMScheduler(num_train_timesteps=config.num_train_timesteps)\n    except NameError:\n        print(\"ERROR: Models not defined! Make sure you've run the previous code cells.\")\n        print(\"You need: VAE, CLIP, UNET, DDPMScheduler\")\n        return\n    \n    print(f\"✓ Models initialized on {config.device}\")\n    \n    # 3. Load pretrained weights (optional)\n    if config.pretrained_weights:\n        print(f\"\\n[3/6] Loading pretrained weights from {config.pretrained_weights}...\")\n        # Load weights here\n        print(\"✓ Pretrained weights loaded\")\n    else:\n        print(\"\\n[3/6] Training from scratch (no pretrained weights)\")\n    \n    # 4. Setup optimizer\n    print(\"\\n[4/6] Setting up optimizer...\")\n    \n    # Fine-tune only U-Net and VAE decoder (keep CLIP frozen)\n    trainable_params = list(unet.parameters()) + list(vae.decoder.parameters())\n    optimizer = torch.optim.AdamW(trainable_params, lr=config.learning_rate)\n    \n    # Freeze CLIP (we don't want to change text understanding)\n    for param in clip.parameters():\n        param.requires_grad = False\n    \n    # Freeze VAE encoder (optional - you can unfreeze this too)\n    for param in vae.encoder.parameters():\n        param.requires_grad = False\n    \n    print(f\"✓ Optimizer configured (lr={config.learning_rate})\")\n    print(f\"✓ Training parameters: U-Net + VAE Decoder\")\n    print(f\"✓ Frozen parameters: CLIP + VAE Encoder\")\n    \n    # 5. Training loop\n    print(\"\\n[5/6] Starting training...\")\n    print(\"=\"*60)\n    \n    vae.train()\n    unet.train()\n    clip.eval()\n    \n    global_step = 0\n    \n    for epoch in range(config.num_epochs):\n        print(f\"\\n{'='*60}\")\n        print(f\"EPOCH {epoch+1}/{config.num_epochs}\")\n        print(f\"{'='*60}\")\n        \n        epoch_loss = 0\n        epoch_diffusion_loss = 0\n        epoch_dct_loss = 0\n        epoch_perceptual_loss = 0\n        \n        progress_bar = tqdm(dataloader, desc=f\"Epoch {epoch+1}\")\n        \n        for batch_idx, batch in enumerate(progress_bar):\n            images = batch['image'].to(config.device)\n            prompts = batch['prompt']  # Now we use the random prompts!\n            \n            # Generate dummy text embeddings \n            # (In production, use actual CLIP tokenizer on prompts)\n            # tokens = tokenizer(prompts, max_length=77, padding=\"max_length\", \n            #                    truncation=True, return_tensors=\"pt\").input_ids\n            tokens = torch.randint(0, 49408, (images.shape[0], 77), device=config.device)\n            with torch.no_grad():\n                text_embeddings = clip(tokens)\n            \n            # ============================================================\n            # TRAINING STEP\n            # ============================================================\n            \n            # 1. Encode images to latents\n            noise_vae = torch.randn(images.shape[0], 4, 64, 64, device=config.device)\n            latents = vae.encode(images, noise_vae)\n            \n            # 2. Sample random timesteps\n            timesteps = torch.randint(\n                0, config.num_train_timesteps,\n                (images.shape[0],),\n                device=config.device\n            )\n            \n            # 3. Add noise (forward diffusion)\n            noise = torch.randn_like(latents)\n            noisy_latents = scheduler.add_noise(latents, noise, timesteps)\n            \n            # 4. Predict noise with U-Net\n            noise_pred = unet(noisy_latents, text_embeddings, timesteps)\n            \n            # 5. Calculate diffusion loss\n            diffusion_loss = F.mse_loss(noise_pred, noise)\n            \n            # 6. Predict clean latent (for DCT loss calculation)\n            with torch.no_grad():\n                # Simplified x0 prediction\n                alpha_prod_t = scheduler.alphas_cumprod[timesteps].view(-1, 1, 1, 1)\n                beta_prod_t = 1 - alpha_prod_t\n                pred_latents = (noisy_latents - beta_prod_t.sqrt() * noise_pred) / alpha_prod_t.sqrt()\n            \n            # 7. Decode to images\n            generated_images = vae.decode(pred_latents)\n            \n            # 8. Calculate DCT loss\n            dct_loss = dct_embedding_quality_loss(generated_images, config.dct_block_size)\n            \n            # 9. Calculate perceptual loss (visual quality)\n            perceptual_loss = F.mse_loss(generated_images, images)\n            \n            # 10. Combined loss\n            total_loss = (\n                config.lambda_diffusion * diffusion_loss +\n                config.lambda_dct * dct_loss +\n                config.lambda_perceptual * perceptual_loss\n            )\n            \n            # 11. Backward pass\n            optimizer.zero_grad()\n            total_loss.backward()\n            torch.nn.utils.clip_grad_norm_(trainable_params, 1.0)  # Gradient clipping\n            optimizer.step()\n            \n            # ============================================================\n            # LOGGING\n            # ============================================================\n            \n            epoch_loss += total_loss.item()\n            epoch_diffusion_loss += diffusion_loss.item()\n            epoch_dct_loss += dct_loss.item()\n            epoch_perceptual_loss += perceptual_loss.item()\n            \n            global_step += 1\n            \n            # Update progress bar\n            if global_step % config.log_every == 0:\n                progress_bar.set_postfix({\n                    'loss': f'{total_loss.item():.4f}',\n                    'diff': f'{diffusion_loss.item():.4f}',\n                    'dct': f'{dct_loss.item():.4f}',\n                    'perc': f'{perceptual_loss.item():.4f}'\n                })\n            \n            # Save checkpoint\n            if global_step % config.save_every == 0:\n                checkpoint_path = os.path.join(\n                    config.output_dir,\n                    f'checkpoint_step_{global_step}.pt'\n                )\n                torch.save({\n                    'epoch': epoch,\n                    'global_step': global_step,\n                    'unet_state_dict': unet.state_dict(),\n                    'vae_decoder_state_dict': vae.decoder.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'loss': total_loss.item(),\n                }, checkpoint_path)\n                print(f\"\\n✓ Checkpoint saved: {checkpoint_path}\")\n        \n        # Epoch summary\n        num_batches = len(dataloader)\n        print(f\"\\nEpoch {epoch+1} Summary:\")\n        print(f\"  Average Total Loss: {epoch_loss/num_batches:.4f}\")\n        print(f\"  Average Diffusion Loss: {epoch_diffusion_loss/num_batches:.4f}\")\n        print(f\"  Average DCT Loss: {epoch_dct_loss/num_batches:.4f}\")\n        print(f\"  Average Perceptual Loss: {epoch_perceptual_loss/num_batches:.4f}\")\n    \n    # 6. Save final model\n    print(\"\\n[6/6] Saving final model...\")\n    final_path = os.path.join(config.output_dir, 'final_model.pt')\n    torch.save({\n        'unet_state_dict': unet.state_dict(),\n        'vae_state_dict': vae.state_dict(),\n        'clip_state_dict': clip.state_dict(),\n    }, final_path)\n    print(f\"✓ Final model saved: {final_path}\")\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"TRAINING COMPLETE!\")\n    print(\"=\"*60)\n    print(f\"Checkpoints saved in: {config.output_dir}\")\n    print(\"\\nYour DCT-optimized Stable Diffusion is ready!\")\n\n\n# ============================================================================\n# RUN TRAINING\n# ============================================================================\n\nif __name__ == \"__main__\":\n    config = Config()\n    \n    # CHANGE THESE PATHS BEFORE RUNNING!\n    config.dataset_path = \"/kaggle/input/your-dataset\"  # <-- CHANGE THIS\n    config.output_dir = \"/kaggle/working/checkpoints\"   # <-- CHANGE THIS\n    \n    print(\"\\n⚠️  BEFORE RUNNING, MAKE SURE YOU:\")\n    print(\"  1. Set config.dataset_path to your image folder\")\n    print(\"  2. Set config.output_dir to save checkpoints\")\n    print(\"  3. Have run all previous code cells (VAE, CLIP, U-Net, Scheduler)\")\n    print(\"  4. Adjust lambda_dct based on your experiments (default: 0.5)\")\n    print(\"\\n\")\n    \n    # Test prompt generator\n    print(\"Sample generated prompts for training:\")\n    gen = PromptGenerator()\n    for i in range(10):\n        print(f\"  {i+1}. {gen.generate()}\")\n    print(\"\\n\")\n    \n    # Uncomment to start training\n    # train_stable_diffusion(config)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# DCT-OPTIMIZED IMAGE GENERATION WITH QUALITY LOOP\n# Generates images until DCT mid-frequency requirements are met\n# ============================================================================\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nfrom PIL import Image\nimport os\nfrom datetime import datetime\n\n\n# ============================================================================\n# DCT QUALITY CHECKER\n# ============================================================================\n\nclass DCTQualityChecker:\n    \"\"\"\n    Checks if an image meets DCT steganography requirements.\n    \n    Requirements:\n    - At least 50% of 8x8 blocks must have 10+ mid-frequency coefficients\n    - Mid-frequency coefficients should have sufficient magnitude\n    \"\"\"\n    \n    def __init__(self, block_size=8, min_mid_freq_count=10, min_block_percentage=50):\n        self.block_size = block_size\n        self.min_mid_freq_count = min_mid_freq_count\n        self.min_block_percentage = min_block_percentage\n        \n        # Define mid-frequency positions\n        self.mid_freq_positions = [\n            (1, 2), (2, 1), (2, 2),  # Low-mid\n            (1, 3), (3, 1), (2, 3), (3, 2),  # Mid\n            (1, 4), (4, 1), (3, 3),  # Mid-high\n            (2, 4), (4, 2), (3, 4), (4, 3),  # Additional mid frequencies\n        ]\n        \n        # Minimum magnitude threshold for \"usable\" coefficient\n        self.magnitude_threshold = 5.0\n    \n    def dct_2d_numpy(self, block):\n        \"\"\"Apply 2D DCT to 8x8 block using scipy.\"\"\"\n        from scipy.fftpack import dct\n        return dct(dct(block.T, norm='ortho').T, norm='ortho')\n    \n    def count_usable_mid_frequencies(self, dct_block):\n        \"\"\"\n        Count how many mid-frequency coefficients are usable.\n        A coefficient is usable if its magnitude > threshold.\n        \"\"\"\n        count = 0\n        for i, j in self.mid_freq_positions:\n            if abs(dct_block[i, j]) > self.magnitude_threshold:\n                count += 1\n        return count\n    \n    def check_image_quality(self, image):\n        \"\"\"\n        Check if image meets DCT quality requirements.\n        \n        Args:\n            image: PIL Image or numpy array (H, W, 3) in range [0, 255]\n        \n        Returns:\n            dict with:\n                - is_good: bool (meets requirements)\n                - good_blocks: int (number of good blocks)\n                - total_blocks: int (total blocks)\n                - percentage: float (percentage of good blocks)\n                - avg_mid_freq_count: float (average mid-freq per block)\n        \"\"\"\n        # Convert to numpy if PIL Image\n        if isinstance(image, Image.Image):\n            image = np.array(image)\n        \n        # Ensure uint8\n        if image.dtype == np.float32 or image.dtype == np.float64:\n            image = (image * 255).astype(np.uint8)\n        \n        height, width, channels = image.shape\n        \n        # Calculate number of blocks\n        num_blocks_h = height // self.block_size\n        num_blocks_w = width // self.block_size\n        total_blocks = num_blocks_h * num_blocks_w * channels\n        \n        good_blocks = 0\n        total_mid_freq_counts = 0\n        \n        # Process each channel\n        for c in range(channels):\n            channel = image[:, :, c]\n            \n            # Process each 8x8 block\n            for i in range(num_blocks_h):\n                for j in range(num_blocks_w):\n                    # Extract block\n                    block = channel[\n                        i*self.block_size:(i+1)*self.block_size,\n                        j*self.block_size:(j+1)*self.block_size\n                    ]\n                    \n                    # Apply DCT\n                    dct_block = self.dct_2d_numpy(block.astype(np.float32))\n                    \n                    # Count usable mid-frequencies\n                    mid_freq_count = self.count_usable_mid_frequencies(dct_block)\n                    total_mid_freq_counts += mid_freq_count\n                    \n                    # Check if block is good\n                    if mid_freq_count >= self.min_mid_freq_count:\n                        good_blocks += 1\n        \n        # Calculate statistics\n        percentage = (good_blocks / total_blocks) * 100\n        avg_mid_freq_count = total_mid_freq_counts / total_blocks\n        is_good = percentage >= self.min_block_percentage\n        \n        return {\n            'is_good': is_good,\n            'good_blocks': good_blocks,\n            'total_blocks': total_blocks,\n            'percentage': percentage,\n            'avg_mid_freq_count': avg_mid_freq_count\n        }\n    \n    def print_quality_report(self, quality_dict):\n        \"\"\"Print a nice quality report.\"\"\"\n        print(\"\\n\" + \"=\"*60)\n        print(\"DCT QUALITY REPORT\")\n        print(\"=\"*60)\n        print(f\"Good blocks: {quality_dict['good_blocks']}/{quality_dict['total_blocks']}\")\n        print(f\"Percentage: {quality_dict['percentage']:.2f}%\")\n        print(f\"Average mid-freq per block: {quality_dict['avg_mid_freq_count']:.2f}\")\n        print(f\"Threshold: {self.min_block_percentage}% blocks with {self.min_mid_freq_count}+ mid-freqs\")\n        print(f\"Status: {'✅ PASS' if quality_dict['is_good'] else '❌ FAIL'}\")\n        print(\"=\"*60)\n\n\n# ============================================================================\n# DCT-OPTIMIZED IMAGE GENERATOR\n# ============================================================================\n\nclass DCTImageGenerator:\n    \"\"\"\n    Generates DCT-optimized images with automatic quality checking.\n    Keeps generating until quality requirements are met.\n    \"\"\"\n    \n    def __init__(self, vae, clip, unet, scheduler, device=\"cuda\"):\n        self.vae = vae.to(device)\n        self.clip = clip.to(device)\n        self.unet = unet.to(device)\n        self.scheduler = scheduler\n        self.device = device\n        \n        # Set to eval mode\n        self.vae.eval()\n        self.clip.eval()\n        self.unet.eval()\n        \n        # Initialize quality checker\n        self.quality_checker = DCTQualityChecker(\n            block_size=8,\n            min_mid_freq_count=10,\n            min_block_percentage=50\n        )\n    \n    @torch.no_grad()\n    def generate_single_image(\n        self,\n        prompt=\"high quality detailed landscape with rich textures\",\n        num_inference_steps=50,\n        guidance_scale=7.5,\n        height=512,\n        width=512,\n        seed=None\n    ):\n        \"\"\"\n        Generate a single image without quality checking.\n        \n        Args:\n            prompt: Text prompt\n            num_inference_steps: Number of denoising steps\n            guidance_scale: CFG scale\n            height: Image height\n            width: Image width\n            seed: Random seed for reproducibility\n        \n        Returns:\n            PIL Image\n        \"\"\"\n        # Set seed if provided\n        generator = None\n        if seed is not None:\n            generator = torch.Generator(device=self.device).manual_seed(seed)\n        \n        # 1. Encode prompt\n        # In production, use actual tokenizer\n        # For now, use dummy tokens\n        tokens = torch.randint(0, 49408, (1, 77), device=self.device)\n        text_embeddings = self.clip(tokens)\n        \n        # For classifier-free guidance, also encode empty prompt\n        uncond_tokens = torch.zeros_like(tokens)\n        uncond_embeddings = self.clip(uncond_tokens)\n        \n        # Concatenate for CFG\n        text_embeddings = torch.cat([uncond_embeddings, text_embeddings])\n        \n        # 2. Initialize latents\n        latents = torch.randn(\n            (1, 4, height // 8, width // 8),\n            generator=generator,\n            device=self.device,\n            dtype=torch.float32\n        )\n        \n        # Scale initial noise\n        latents = latents * self.scheduler.sqrt_alphas_cumprod[0]\n        \n        # 3. Set timesteps\n        self.scheduler.set_timesteps(num_inference_steps)\n        \n        # 4. Denoising loop\n        for i, t in enumerate(self.scheduler.timesteps):\n            # Expand for CFG\n            latent_model_input = torch.cat([latents] * 2)\n            \n            # Predict noise\n            noise_pred = self.unet(\n                latent_model_input,\n                text_embeddings,\n                t.unsqueeze(0).to(self.device)\n            )\n            \n            # Classifier-free guidance\n            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)\n            noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)\n            \n            # Denoise\n            latents = self.scheduler.step(noise_pred, t, latents)\n        \n        # 5. Decode to image\n        image = self.vae.decode(latents)\n        \n        # 6. Convert to PIL\n        image = self._tensor_to_pil(image)\n        \n        return image\n    \n    def generate_with_quality_check(\n        self,\n        prompt=\"high quality detailed landscape with rich textures\",\n        num_inference_steps=50,\n        guidance_scale=7.5,\n        height=512,\n        width=512,\n        max_attempts=50,\n        save_all_attempts=False,\n        output_dir=\"./generated_images\"\n    ):\n        \"\"\"\n        Generate images until DCT quality requirements are met.\n        \n        Args:\n            prompt: Text prompt\n            num_inference_steps: Denoising steps\n            guidance_scale: CFG scale\n            height: Image height\n            width: Image width\n            max_attempts: Maximum generation attempts\n            save_all_attempts: Save all attempts (for debugging)\n            output_dir: Where to save images\n        \n        Returns:\n            dict with:\n                - image: PIL Image (the good one)\n                - quality: quality metrics\n                - attempt: which attempt succeeded\n                - all_attempts: list of all attempts (if save_all_attempts=True)\n        \"\"\"\n        os.makedirs(output_dir, exist_ok=True)\n        timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n        \n        print(\"=\"*60)\n        print(\"DCT-OPTIMIZED IMAGE GENERATION WITH QUALITY LOOP\")\n        print(\"=\"*60)\n        print(f\"Prompt: {prompt}\")\n        print(f\"Target: ≥50% blocks with ≥10 mid-frequency coefficients\")\n        print(f\"Max attempts: {max_attempts}\")\n        print(\"=\"*60)\n        \n        all_attempts = []\n        \n        for attempt in range(1, max_attempts + 1):\n            print(f\"\\n[Attempt {attempt}/{max_attempts}]\")\n            \n            # Generate image with different seed each time\n            seed = np.random.randint(0, 2**32 - 1)\n            image = self.generate_single_image(\n                prompt=prompt,\n                num_inference_steps=num_inference_steps,\n                guidance_scale=guidance_scale,\n                height=height,\n                width=width,\n                seed=seed\n            )\n            \n            # Check quality\n            quality = self.quality_checker.check_image_quality(image)\n            \n            # Print results\n            print(f\"  Good blocks: {quality['good_blocks']}/{quality['total_blocks']} ({quality['percentage']:.2f}%)\")\n            print(f\"  Avg mid-freq: {quality['avg_mid_freq_count']:.2f}\")\n            \n            # Save if requested\n            if save_all_attempts:\n                attempt_path = os.path.join(\n                    output_dir,\n                    f\"{timestamp}_attempt_{attempt}_quality_{quality['percentage']:.1f}.png\"\n                )\n                image.save(attempt_path)\n                all_attempts.append({\n                    'image': image,\n                    'quality': quality,\n                    'path': attempt_path\n                })\n            \n            # Check if good enough\n            if quality['is_good']:\n                print(f\"\\n✅ SUCCESS! Found good image on attempt {attempt}\")\n                \n                # Save the good image\n                final_path = os.path.join(\n                    output_dir,\n                    f\"{timestamp}_FINAL_quality_{quality['percentage']:.1f}.png\"\n                )\n                image.save(final_path)\n                \n                # Print final report\n                self.quality_checker.print_quality_report(quality)\n                \n                print(f\"\\n✅ Image saved: {final_path}\")\n                \n                return {\n                    'image': image,\n                    'quality': quality,\n                    'attempt': attempt,\n                    'path': final_path,\n                    'all_attempts': all_attempts if save_all_attempts else None\n                }\n        \n        # If we get here, max attempts reached without success\n        print(f\"\\n❌ FAILED: Could not find good image in {max_attempts} attempts\")\n        print(\"Try:\")\n        print(\"  1. Increase max_attempts\")\n        print(\"  2. Lower min_block_percentage (currently 50%)\")\n        print(\"  3. Lower min_mid_freq_count (currently 10)\")\n        print(\"  4. Adjust guidance_scale (try 5-10)\")\n        \n        # Return best attempt\n        if save_all_attempts and all_attempts:\n            best = max(all_attempts, key=lambda x: x['quality']['percentage'])\n            print(f\"\\nReturning best attempt: {best['quality']['percentage']:.2f}%\")\n            return best\n        \n        return None\n    \n    def _tensor_to_pil(self, tensor):\n        \"\"\"Convert tensor to PIL Image.\"\"\"\n        # tensor: (1, 3, 512, 512) in range [-1, 1]\n        image = (tensor / 2 + 0.5).clamp(0, 1)\n        image = image.cpu().permute(0, 2, 3, 1).numpy()\n        image = (image * 255).round().astype(np.uint8)\n        return Image.fromarray(image[0])\n\n\n# ============================================================================\n# BATCH GENERATOR\n# ============================================================================\n\ndef generate_dct_dataset(\n    generator,\n    num_images=100,\n    output_dir=\"./dct_dataset\",\n    prompts=None,\n    **generation_kwargs\n):\n    \"\"\"\n    Generate a dataset of DCT-optimized images.\n    \n    Args:\n        generator: DCTImageGenerator instance\n        num_images: How many images to generate\n        output_dir: Where to save\n        prompts: List of prompts (random if None)\n        **generation_kwargs: Additional args for generation\n    \n    Returns:\n        List of paths to generated images\n    \"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n    \n    # Default prompts if not provided\n    if prompts is None:\n        from random import choice\n        prompt_pool = [\n            \"high quality detailed landscape with rich textures\",\n            \"ultra detailed nature scene with intricate patterns\",\n            \"professional architectural detail with complex textures\",\n            \"detailed scenic view with elaborate textures\",\n            \"crystal clear natural environment with fine details\",\n        ]\n        prompts = [choice(prompt_pool) for _ in range(num_images)]\n    \n    print(\"=\"*60)\n    print(f\"GENERATING DCT-OPTIMIZED DATASET\")\n    print(\"=\"*60)\n    print(f\"Target: {num_images} images\")\n    print(f\"Output: {output_dir}\")\n    print(\"=\"*60)\n    \n    generated_paths = []\n    successful = 0\n    total_attempts = 0\n    \n    for i in range(num_images):\n        print(f\"\\n{'='*60}\")\n        print(f\"IMAGE {i+1}/{num_images}\")\n        print(f\"{'='*60}\")\n        \n        prompt = prompts[i] if i < len(prompts) else prompts[0]\n        \n        result = generator.generate_with_quality_check(\n            prompt=prompt,\n            output_dir=output_dir,\n            **generation_kwargs\n        )\n        \n        if result:\n            generated_paths.append(result['path'])\n            successful += 1\n            total_attempts += result['attempt']\n            \n            print(f\"\\n✅ Progress: {successful}/{num_images} images generated\")\n            print(f\"   Average attempts per image: {total_attempts/successful:.1f}\")\n    \n    print(\"\\n\" + \"=\"*60)\n    print(\"DATASET GENERATION COMPLETE\")\n    print(\"=\"*60)\n    print(f\"Successfully generated: {successful}/{num_images} images\")\n    print(f\"Total attempts: {total_attempts}\")\n    print(f\"Average attempts per image: {total_attempts/successful:.1f}\")\n    print(f\"Saved to: {output_dir}\")\n    print(\"=\"*60)\n    \n    return generated_paths\n\n\n# ============================================================================\n# USAGE EXAMPLE\n# ============================================================================\n\nif __name__ == \"__main__\":\n    \n    print(\"\\n⚠️  MAKE SURE YOU:\")\n    print(\"  1. Have trained VAE, CLIP, U-Net loaded\")\n    print(\"  2. Have DDPMScheduler initialized\")\n    print(\"  3. Have loaded your trained weights\")\n    print(\"\\n\")\n    \n    # Initialize generator\n    # (Assumes VAE, CLIP, UNET, DDPMScheduler are already defined)\n    try:\n        generator = DCTImageGenerator(\n            vae=vae,\n            clip=clip,\n            unet=unet,\n            scheduler=scheduler,\n            device=\"cuda\"\n        )\n        \n        # Generate single DCT-optimized image\n        result = generator.generate_with_quality_check(\n            prompt=\"ultra detailed landscape with rich textures and intricate patterns\",\n            num_inference_steps=50,\n            guidance_scale=7.5,\n            max_attempts=20,\n            save_all_attempts=True,  # Save all attempts for analysis\n            output_dir=\"./dct_covers\"\n        )\n        \n        if result:\n            print(\"\\n✅ SUCCESS! Generated DCT-optimized cover image\")\n            print(f\"   Saved to: {result['path']}\")\n            print(f\"   Quality: {result['quality']['percentage']:.2f}% good blocks\")\n            print(f\"   Found on attempt: {result['attempt']}\")\n        \n        # Generate a full dataset (optional)\n        # paths = generate_dct_dataset(\n        #     generator=generator,\n        #     num_images=100,\n        #     output_dir=\"./dct_dataset\",\n        #     max_attempts=30\n        # )\n        \n    except NameError:\n        print(\"ERROR: Models not loaded. Run previous code cells first.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}