import torch import torch.nn as nn import numpy as np from PIL import Image import gradio as gr from skimage import color from torchvision.models import vit_b_16, ViT_B_16_Weights # ==================================================================== # 1. MODEL ARCHITECTURE CLASSES # ==================================================================== class ReassembleBlock(nn.Module): def __init__(self, in_channels, out_channels=256, spatial_size=14, scale_factor=1): super(ReassembleBlock, self).__init__() self.spatial_size = spatial_size self.project = nn.Conv2d(in_channels, out_channels, kernel_size=1) if scale_factor > 1: self.resample = nn.ConvTranspose2d( out_channels, out_channels, kernel_size=int(scale_factor*2), stride=int(scale_factor), padding=int(scale_factor//2) ) elif scale_factor < 1: stride = int(1 / scale_factor) self.resample = nn.Conv2d( out_channels, out_channels, kernel_size=3, stride=stride, padding=1 ) else: self.resample = nn.Identity() def forward(self, tokens): batch_size = tokens.shape[0] cls_token = tokens[:, 0:1, :] spatial_tokens = tokens[:, 1:, :] spatial_tokens = spatial_tokens + cls_token grid = spatial_tokens.reshape(batch_size, self.spatial_size, self.spatial_size, -1) grid = grid.permute(0, 3, 1, 2) out = self.project(grid) out = self.resample(out) return out class ResidualConvUnit(nn.Module): def __init__(self, features): super(ResidualConvUnit, self).__init__() self.conv1 = nn.Conv2d(features, features, kernel_size=3, padding=1) self.conv2 = nn.Conv2d(features, features, kernel_size=3, padding=1) self.relu = nn.ReLU(inplace=True) def forward(self, x): out = self.relu(x) out = self.conv1(out) out = self.relu(out) out = self.conv2(out) return out + x class FeatureFusionBlock(nn.Module): def __init__(self, features=256): super(FeatureFusionBlock, self).__init__() self.resConfUnit1 = ResidualConvUnit(features) self.resConfUnit2 = ResidualConvUnit(features) self.upsample = nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True) self.project = nn.Conv2d(features, features, kernel_size=3, padding=1) def forward(self, x, previous_stage_output=None): out = self.resConfUnit1(x) if previous_stage_output is not None: out = out + previous_stage_output out = self.resConfUnit2(out) out = self.upsample(out) out = self.project(out) return out class PretrainedViTEncoder(nn.Module): def __init__(self): super(PretrainedViTEncoder, self).__init__() print("Downloading pre-trained ImageNet weights...") self.vit = vit_b_16(weights=ViT_B_16_Weights.IMAGENET1K_V1) self.vit.heads = nn.Identity() self.extracted_features = [] self._register_hooks() def _hook_fn(self, module, input, output): self.extracted_features.append(output) def _register_hooks(self): target_layers = [2, 5, 8, 11] for i, layer in enumerate(self.vit.encoder.layers): if i in target_layers: layer.register_forward_hook(self._hook_fn) def forward(self, x): self.extracted_features = [] self.vit(x) return self.extracted_features class DPT(nn.Module): def __init__(self, vit_encoder, embed_dim=768, features=256, spatial_size=14): super(DPT, self).__init__() self.encoder = vit_encoder self.reassemble_4 = ReassembleBlock(embed_dim, features, spatial_size, scale_factor=4) self.reassemble_8 = ReassembleBlock(embed_dim, features, spatial_size, scale_factor=2) self.reassemble_16 = ReassembleBlock(embed_dim, features, spatial_size, scale_factor=1) self.reassemble_32 = ReassembleBlock(embed_dim, features, spatial_size, scale_factor=0.5) self.fusion_32 = FeatureFusionBlock(features) self.fusion_16 = FeatureFusionBlock(features) self.fusion_8 = FeatureFusionBlock(features) self.fusion_4 = FeatureFusionBlock(features) self.head = nn.Sequential( nn.Conv2d(features, features // 2, kernel_size=3, padding=1), nn.ReLU(True), nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True), nn.Conv2d(features // 2, 2, kernel_size=1), nn.Tanh() ) def forward(self, x): layer_tokens = self.encoder(x) t3, t6, t9, t12 = layer_tokens f4 = self.reassemble_4(t3) f8 = self.reassemble_8(t6) f16 = self.reassemble_16(t9) f32 = self.reassemble_32(t12) out = self.fusion_32(f32) out = self.fusion_16(f16, previous_stage_output=out) out = self.fusion_8(f8, previous_stage_output=out) out = self.fusion_4(f4, previous_stage_output=out) color_map = self.head(out) return color_map # ==================================================================== # 2. INITIALIZATION AND INFERENCE LOGIC # ==================================================================== device = torch.device("cpu") # Web servers usually run on CPU # Load the model structure print("Initializing Model...") encoder = PretrainedViTEncoder() model = DPT(vit_encoder=encoder).to(device) # Load the weights try: model.load_state_dict(torch.load("best_dpt_colorizer.pth", map_location=device)) print("✅ Weights loaded successfully!") except Exception as e: print(f"⚠️ Warning: Could not load weights. {e}") model.eval() # Set to evaluation mode! def colorize_image(input_image): """ Takes a PIL image from Gradio, runs it through the DPT model, and returns a colorized PIL image. """ if input_image is None: return None # 1. Resize the image to match our training (224x224) original_size = input_image.size img_resized = input_image.resize((224, 224)).convert("RGB") # 2. Convert to LAB and extract the L channel img_np = np.array(img_resized) lab_img = color.rgb2lab(img_np).astype(np.float32) # Extract L and normalize exactly like we did in training [-1, 1] l_channel = lab_img[:, :, 0] l_normalized = (l_channel / 50.0) - 1.0 # Convert to Tensor and duplicate 3 times l_tensor = torch.from_numpy(l_normalized).unsqueeze(0).unsqueeze(0) # [1, 1, 224, 224] l_tensor_3ch = l_tensor.repeat(1, 3, 1, 1).to(device) # [1, 3, 224, 224] # 3. Predict the ab channels with torch.no_grad(): ab_predicted = model(l_tensor_3ch) # Shape: [1, 2, 224, 224] # 4. Denormalize the predictions ab_out = ab_predicted.squeeze(0).cpu().numpy() # [2, 224, 224] ab_out = ab_out.transpose(1, 2, 0) # [224, 224, 2] # Denormalize from [-1, 1] back to [-128, 127] ab_denormalized = ab_out * 128.0 # 5. Reconstruct the LAB image l_channel_original = np.expand_dims(l_channel, axis=-1) # [224, 224, 1] reconstructed_lab = np.concatenate((l_channel_original, ab_denormalized), axis=-1) # 6. Convert back to RGB reconstructed_rgb = color.lab2rgb(reconstructed_lab) # The output of lab2rgb is [0.0, 1.0]. Convert to [0, 255] for standard images final_img_np = (reconstructed_rgb * 255).astype(np.uint8) final_img_pil = Image.fromarray(final_img_np) # 7. Optional: Resize back to the user's original image size final_img_pil = final_img_pil.resize(original_size, Image.LANCZOS) return final_img_pil # ==================================================================== # 3. BUILD THE WEB INTERFACE # ==================================================================== # ==================================================================== # 3. BUILD THE WEB INTERFACE # ==================================================================== demo = gr.Interface( fn=colorize_image, inputs=gr.Image(type="pil", label="Upload Black & White Image"), outputs=gr.Image(type="pil", label="Colorized Output"), title="DPT Image Colorization", description="Upload a black and white image, and our Dense Prediction Transformer will add the color back in!", flagging_mode="never" # 🔥 Updated for Gradio v4+ ) if __name__ == "__main__": demo.launch()