Spaces:
Sleeping
Sleeping
Download app.py from Abdlerhman/DPT-Image-Colorization: direct link, hf CLI and curl.
- Browser
- Download file 8.59 kB
-
https://huggingface.co/spaces/Abdlerhman/DPT-Image-Colorization/resolve/main/app.py
- Command line
-
hf download hf://spaces/Abdlerhman/DPT-Image-Colorization/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Abdlerhman/DPT-Image-Colorization/resolve/main/app.py
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() |