Abdlerhman's picture
Update app.py
58718e1 verified
Raw History Blame Contribute Delete
8.59 kB
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()