"""Gradio web interface for MFIR model inference on Hugging Face Spaces. This module provides a web UI for testing the multi-frame image restoration model. """ from __future__ import annotations from dataclasses import dataclass, field from pathlib import Path from typing import Any, Dict, List, Literal, Optional, Tuple, Union import gradio as gr import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from huggingface_hub import hf_hub_download from PIL import Image from torchvision import transforms from torchvision.ops import deform_conv2d # ============================================================================ # Model Architecture (inline for single-file deployment) # ============================================================================ @dataclass class FeatureFusionConfig: """Configuration for Feature Fusion Model.""" in_channels: int = 3 max_frames: int = 16 encoder_channels: List[int] = field(default_factory=lambda: [64, 128, 256]) encoder_blocks_per_stage: int = 2 offset_channels: int = 64 deform_groups: int = 8 num_deform_layers: int = 3 fusion_type: Literal["attention", "adaptive"] = "attention" fusion_num_heads: int = 4 decoder_blocks_per_stage: int = 2 out_channels: int = 3 def to_dict(self) -> Dict: return { k: list(v) if isinstance(v, list) else v for k, v in self.__dict__.items() } @classmethod def from_dict(cls, d: Dict) -> "FeatureFusionConfig": return cls(**{k: v for k, v in d.items() if k in cls.__dataclass_fields__}) class ResidualBlock(nn.Module): """Basic residual block with two convolutions.""" def __init__(self, channels: int): super().__init__() self.conv1 = nn.Conv2d(channels, channels, 3, 1, 1) self.conv2 = nn.Conv2d(channels, channels, 3, 1, 1) self.lrelu = nn.LeakyReLU(0.1, inplace=True) def forward(self, x: torch.Tensor) -> torch.Tensor: residual = x out = self.lrelu(self.conv1(x)) out = self.conv2(out) out = out + residual out = self.lrelu(out) return out class Encoder(nn.Module): """Multi-scale feature encoder.""" def __init__( self, in_channels: int = 3, channels: Optional[List[int]] = None, blocks_per_stage: int = 2, ): super().__init__() if channels is None: channels = [64, 128, 256] self.channels = channels self.conv_first = nn.Conv2d(in_channels, channels[0], 3, 1, 1) self.stages = nn.ModuleList() self.downsamples = nn.ModuleList() for i, (in_ch, out_ch) in enumerate(zip(channels[:-1], channels[1:])): blocks = nn.Sequential( *[ResidualBlock(in_ch) for _ in range(blocks_per_stage)] ) self.stages.append(blocks) self.downsamples.append( nn.Conv2d(in_ch, out_ch, 3, stride=2, padding=1) ) self.final_blocks = nn.Sequential( *[ResidualBlock(channels[-1]) for _ in range(blocks_per_stage)] ) self.lrelu = nn.LeakyReLU(0.1, inplace=True) def forward(self, x: torch.Tensor) -> torch.Tensor: feat = self.lrelu(self.conv_first(x)) for stage, downsample in zip(self.stages, self.downsamples): feat = stage(feat) feat = self.lrelu(downsample(feat)) feat = self.final_blocks(feat) return feat class DeformableAlignmentModule(nn.Module): """Deformable convolution based alignment module.""" def __init__( self, channels: int = 256, offset_channels: int = 64, deform_groups: int = 8, num_layers: int = 3, ): super().__init__() self.channels = channels self.deform_groups = deform_groups self.num_layers = num_layers self.offset_convs = nn.ModuleList() self.deform_weights = nn.ParameterList() for i in range(num_layers): in_ch = channels * 2 if i == 0 else channels * 2 + 2 * deform_groups * 9 offset_conv = nn.Sequential( nn.Conv2d(in_ch, offset_channels, 3, 1, 1), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(offset_channels, offset_channels, 3, 1, 1), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(offset_channels, 2 * deform_groups * 9, 3, 1, 1), ) nn.init.zeros_(offset_conv[-1].weight) nn.init.zeros_(offset_conv[-1].bias) self.offset_convs.append(offset_conv) deform_weight = nn.Parameter( torch.empty(channels, channels // deform_groups, 3, 3) ) nn.init.kaiming_normal_(deform_weight, a=0.1, mode='fan_out', nonlinearity='leaky_relu') self.deform_weights.append(deform_weight) self.fusion_conv = nn.Conv2d(channels, channels, 3, 1, 1) self.lrelu = nn.LeakyReLU(0.1, inplace=True) def forward(self, ref_feat: torch.Tensor, src_feat: torch.Tensor) -> torch.Tensor: aligned = src_feat prev_offset = None for i in range(self.num_layers): if prev_offset is None: offset_input = torch.cat([ref_feat, aligned], dim=1) else: offset_input = torch.cat([ref_feat, aligned, prev_offset], dim=1) max_offset = 5.0 offset = self.offset_convs[i](offset_input) offset = torch.tanh(offset / 5.0) * max_offset prev_offset = offset aligned = deform_conv2d( aligned, offset, self.deform_weights[i], padding=1, ) aligned = self.lrelu(aligned) aligned = self.fusion_conv(aligned) return aligned class TemporalAttentionFusion(nn.Module): """Temporal attention based fusion module.""" def __init__( self, channels: int = 256, num_heads: int = 4, max_frames: int = 16, ): super().__init__() self.channels = channels self.num_heads = num_heads self.head_dim = channels // num_heads self.scale = self.head_dim ** -0.5 self.max_frames = max_frames self.q_proj = nn.Conv2d(channels, channels, 1) self.k_proj = nn.Conv2d(channels, channels, 1) self.v_proj = nn.Conv2d(channels, channels, 1) self.out_proj = nn.Conv2d(channels, channels, 1) self.temporal_embed = nn.Parameter(torch.randn(1, max_frames, channels, 1, 1) * 0.02) self.refine = nn.Sequential( ResidualBlock(channels), ResidualBlock(channels), ) def _get_temporal_embed(self, num_frames: int) -> torch.Tensor: if num_frames == self.max_frames: return self.temporal_embed embed = self.temporal_embed.squeeze(-1).squeeze(-1) embed = embed.permute(0, 2, 1) embed = F.interpolate(embed, size=num_frames, mode="linear", align_corners=False) embed = embed.permute(0, 2, 1) return embed.unsqueeze(-1).unsqueeze(-1) def forward(self, features: torch.Tensor, ref_idx: Optional[int] = None) -> torch.Tensor: B, N, C, H, W = features.shape if ref_idx is None: ref_idx = 0 temporal_embed = self._get_temporal_embed(N) features = features + temporal_embed ref_feat = features[:, ref_idx] q = self.q_proj(ref_feat) features_flat = features.view(B * N, C, H, W) k = self.k_proj(features_flat).view(B, N, C, H, W) v = self.v_proj(features_flat).view(B, N, C, H, W) q = q.view(B, self.num_heads, self.head_dim, H * W) k = k.view(B, N, self.num_heads, self.head_dim, H * W) v = v.view(B, N, self.num_heads, self.head_dim, H * W) q = q.permute(0, 1, 3, 2) k = k.permute(0, 2, 4, 1, 3) v = v.permute(0, 2, 4, 1, 3) attn = torch.matmul(q.unsqueeze(-2), k.transpose(-2, -1)) * self.scale attn = F.softmax(attn, dim=-1) out = torch.matmul(attn, v).squeeze(-2) out = out.permute(0, 1, 3, 2) out = out.reshape(B, C, H, W) out = self.out_proj(out) out = out + ref_feat out = self.refine(out) return out class AdaptiveFusion(nn.Module): """Adaptive weight based fusion.""" def __init__(self, channels: int = 256): super().__init__() self.channels = channels self.weight_conv1 = nn.Conv2d(channels, channels, 1) self.weight_conv2 = nn.Conv2d(channels, channels, 3, 1, 1) self.weight_conv3 = nn.Conv2d(channels, 1, 3, 1, 1) self.refine = nn.Sequential( ResidualBlock(channels), ResidualBlock(channels), ) def forward(self, features: torch.Tensor, ref_idx: Optional[int] = None) -> torch.Tensor: B, N, C, H, W = features.shape weights_list = [] for i in range(N): feat = features[:, i] w = F.leaky_relu(self.weight_conv1(feat), 0.1) w = F.leaky_relu(self.weight_conv2(w), 0.1) w = self.weight_conv3(w) weights_list.append(w) weights = torch.cat(weights_list, dim=1) weights = F.softmax(weights, dim=1) weights = weights.unsqueeze(2) fused = (features * weights).sum(dim=1) fused = self.refine(fused) return fused class Decoder(nn.Module): """Feature decoder with progressive upsampling.""" def __init__( self, in_channels: int = 256, channels: Optional[List[int]] = None, out_channels: int = 3, blocks_per_stage: int = 2, ): super().__init__() if channels is None: channels = [128, 64] self.upsamples = nn.ModuleList() self.stages = nn.ModuleList() all_channels = [in_channels] + channels for i, (in_ch, out_ch) in enumerate(zip(all_channels[:-1], all_channels[1:])): self.upsamples.append( nn.Sequential( nn.Conv2d(in_ch, out_ch * 4, 3, 1, 1), nn.PixelShuffle(2), nn.LeakyReLU(0.1, inplace=True), ) ) self.stages.append( nn.Sequential( *[ResidualBlock(out_ch) for _ in range(blocks_per_stage)] ) ) self.conv_last = nn.Sequential( nn.Conv2d(channels[-1], channels[-1], 3, 1, 1), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(channels[-1], out_channels, 3, 1, 1), ) nn.init.xavier_uniform_(self.conv_last[-1].weight, gain=0.1) nn.init.constant_(self.conv_last[-1].bias, 0.5) def forward(self, x: torch.Tensor) -> torch.Tensor: for upsample, stage in zip(self.upsamples, self.stages): x = upsample(x) x = stage(x) x = self.conv_last(x) return x class FeatureFusionModel(nn.Module): """Feature-level multi-frame fusion model.""" def __init__(self, config: FeatureFusionConfig): super().__init__() self.config = config self.encoder = Encoder( in_channels=config.in_channels, channels=config.encoder_channels, blocks_per_stage=config.encoder_blocks_per_stage, ) feat_channels = config.encoder_channels[-1] self.alignment = DeformableAlignmentModule( channels=feat_channels, offset_channels=config.offset_channels, deform_groups=config.deform_groups, num_layers=config.num_deform_layers, ) if config.fusion_type == "attention": self.fusion = TemporalAttentionFusion( channels=feat_channels, num_heads=config.fusion_num_heads, max_frames=config.max_frames, ) else: self.fusion = AdaptiveFusion(channels=feat_channels) decoder_channels = config.encoder_channels[-2::-1] self.decoder = Decoder( in_channels=feat_channels, channels=decoder_channels, out_channels=config.out_channels, blocks_per_stage=config.decoder_blocks_per_stage, ) def forward(self, frames: torch.Tensor, ref_idx: Optional[int] = None) -> Dict[str, torch.Tensor]: B, N, C, H, W = frames.shape if ref_idx is None: ref_idx = 0 frames_flat = frames.view(B * N, C, H, W) features_flat = self.encoder(frames_flat) _, feat_ch, feat_h, feat_w = features_flat.shape features = features_flat.view(B, N, feat_ch, feat_h, feat_w) ref_feat = features[:, ref_idx] aligned_features = [] for i in range(N): if i == ref_idx: aligned_features.append(ref_feat) else: aligned = self.alignment(ref_feat, features[:, i]) aligned_features.append(aligned) aligned_features = torch.stack(aligned_features, dim=1) fused_features = self.fusion(aligned_features, ref_idx) decoded = self.decoder(fused_features) output = torch.clamp(decoded, 0.0, 1.0) return {"output": output} def load_state_dict_with_compatibility( self, state_dict: Dict[str, torch.Tensor], strict: bool = False, ) -> Tuple[List[str], List[str]]: if "fusion.temporal_embed" in state_dict: old_embed = state_dict["fusion.temporal_embed"] old_num_frames = old_embed.shape[1] if hasattr(self.fusion, "max_frames") and old_num_frames != self.fusion.max_frames: embed = old_embed.squeeze(-1).squeeze(-1) embed = embed.permute(0, 2, 1) embed = F.interpolate( embed, size=self.fusion.max_frames, mode="linear", align_corners=False, ) embed = embed.permute(0, 2, 1) state_dict["fusion.temporal_embed"] = embed.unsqueeze(-1).unsqueeze(-1) adaptive_mapping = { "fusion.weight_net.0.weight": "fusion.weight_conv1.weight", "fusion.weight_net.0.bias": "fusion.weight_conv1.bias", "fusion.weight_net.2.weight": "fusion.weight_conv2.weight", "fusion.weight_net.2.bias": "fusion.weight_conv2.bias", "fusion.weight_net.4.weight": "fusion.weight_conv3.weight", "fusion.weight_net.4.bias": "fusion.weight_conv3.bias", } for old_key, new_key in adaptive_mapping.items(): if old_key in state_dict and new_key not in state_dict: state_dict[new_key] = state_dict.pop(old_key) return self.load_state_dict(state_dict, strict=strict) # ============================================================================ # Global State and Helper Functions # ============================================================================ class ModelState: """Global state for loaded model.""" model: Optional[FeatureFusionModel] = None device: str = "cpu" state = ModelState() def load_model() -> str: """Load model from Hugging Face Hub.""" if state.model is not None: return "Model already loaded" device = "cuda" if torch.cuda.is_available() else "cpu" try: checkpoint_path = hf_hub_download( repo_id="marduk-ra/MFIR", filename="temporal_fusion_model.pth", ) ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False) config_dict = ckpt["config"] if isinstance(config_dict, dict): config_dict = config_dict.copy() elif hasattr(config_dict, "to_dict"): config_dict = config_dict.to_dict() else: config_dict = vars(config_dict) config_dict.pop("input_frames", None) if "max_frames" not in config_dict: config_dict["max_frames"] = 16 config = FeatureFusionConfig.from_dict(config_dict) model = FeatureFusionModel(config) model.load_state_dict_with_compatibility(ckpt["state_dict"]) model.to(device) model.eval() state.model = model state.device = device return f"Model loaded on {device}" except Exception as e: return f"Error loading model: {e}" def preprocess_images( images: Optional[List[Tuple[np.ndarray, str]]], target_size: int = 256, ) -> Tuple[Optional[torch.Tensor], Optional[Dict[str, Any]]]: """Preprocess input images to tensor with aspect ratio preservation.""" if not images: return None, None frames = [] preprocess_info = None for idx, img_data in enumerate(images): if isinstance(img_data, tuple): img_path = img_data[0] if isinstance(img_path, str): img = Image.open(img_path).convert("RGB") else: img = Image.fromarray(img_path).convert("RGB") elif isinstance(img_data, np.ndarray): img = Image.fromarray(img_data).convert("RGB") else: img = Image.open(img_data).convert("RGB") orig_w, orig_h = img.size if orig_w >= orig_h: new_w = target_size new_h = int(orig_h * target_size / orig_w) else: new_h = target_size new_w = int(orig_w * target_size / orig_h) img_resized = img.resize((new_w, new_h), Image.LANCZOS) tensor = transforms.ToTensor()(img_resized) pad_h = target_size - new_h pad_w = target_size - new_w pad_top = pad_h // 2 pad_bottom = pad_h - pad_top pad_left = pad_w // 2 pad_right = pad_w - pad_left tensor_padded = torch.nn.functional.pad( tensor, (pad_left, pad_right, pad_top, pad_bottom), mode="reflect", ) frames.append(tensor_padded) if idx == 0: preprocess_info = { "original_size": (orig_w, orig_h), "resized_size": (new_w, new_h), "padding": (pad_left, pad_top, pad_right, pad_bottom), } frames_tensor = torch.stack(frames, dim=0) return frames_tensor.unsqueeze(0), preprocess_info def tensor_to_pil(tensor: torch.Tensor) -> Image.Image: """Convert tensor (C, H, W) to PIL Image.""" img_np = tensor.cpu().numpy() img_np = np.transpose(img_np, (1, 2, 0)) img_np = np.clip(img_np * 255, 0, 255).astype(np.uint8) return Image.fromarray(img_np) def postprocess_output( tensor: torch.Tensor, preprocess_info: Dict[str, Any], ) -> Image.Image: """Remove padding from output tensor and convert to PIL Image.""" if tensor.dim() == 4: tensor = tensor[0] pad_left, pad_top, pad_right, pad_bottom = preprocess_info["padding"] _, H, W = tensor.shape cropped = tensor[ :, pad_top : H - pad_bottom if pad_bottom > 0 else H, pad_left : W - pad_right if pad_right > 0 else W, ] return tensor_to_pil(cropped) def process_images( images: Optional[List[Tuple[np.ndarray, str]]], image_size: int = 256, ref_frame: int = 0, ) -> Tuple[Optional[Image.Image], str]: """Process input images through the model.""" if state.model is None: load_result = load_model() if state.model is None: return None, load_result if not images: return None, "Please upload at least 2 images" if len(images) < 2: return None, "Please upload at least 2 images" if len(images) > 16: return None, "Maximum 16 images supported" try: frames, preprocess_info = preprocess_images(images, target_size=image_size) if frames is None: return None, "Failed to preprocess images" frames = frames.to(state.device) ref_idx = min(ref_frame, frames.shape[1] - 1) with torch.no_grad(): result = state.model(frames, ref_idx=ref_idx) output = result["output"] output_pil = postprocess_output(output, preprocess_info) del result, output, frames if state.device == "cuda": torch.cuda.empty_cache() return output_pil, f"Processed {len(images)} frames (reference: frame {ref_idx})" except Exception as e: return None, f"Processing error: {e}" # ============================================================================ # Gradio Interface # ============================================================================ def create_ui() -> gr.Blocks: """Create and return the Gradio interface.""" with gr.Blocks(title="MFIR - Multi-Frame Image Restoration") as demo: gr.Markdown(""" # MFIR - Multi-Frame Image Restoration Upload multiple degraded frames of the same scene, and the model will fuse them into a single high-quality image. **How it works:** The model aligns all frames to a reference frame using deformable convolutions, then fuses them using temporal attention to extract the best information from each frame. """) with gr.Row(): with gr.Column(): input_gallery = gr.Gallery( label="Upload Frames (2-16 images)", show_label=True, columns=4, rows=2, height="auto", object_fit="scale-down", ) with gr.Row(): image_size = gr.Dropdown( choices=[128, 256, 512], value=256, label="Processing Size", info="Larger = better quality, slower", ) ref_frame = gr.Slider( minimum=0, maximum=15, value=0, step=1, label="Reference Frame", info="Which frame to use as reference", ) process_btn = gr.Button("Restore Image", variant="primary", size="lg") with gr.Column(): output_image = gr.Image( label="Restored Output", type="pil", show_label=True, ) status_text = gr.Textbox( label="Status", interactive=False, ) # Examples section gr.Markdown("---\n### Examples") gr.Examples( examples=[ [ [ "images/sample_01/input1.png", "images/sample_01/input2.png", "images/sample_01/input3.png", "images/sample_01/input4.png", "images/sample_01/input5.png", ], 256, 0, ], [ [ "images/sample_02/frame_000.png", "images/sample_02/frame_001.png", "images/sample_02/frame_002.png", "images/sample_02/frame_003.png", "images/sample_02/frame_004.png", "images/sample_02/frame_005.png", "images/sample_02/frame_006.png", "images/sample_02/frame_007.png", "images/sample_02/frame_008.png", "images/sample_02/frame_009.png", "images/sample_02/frame_010.png", "images/sample_02/frame_011.png", ], 256, 0, ], [ [ "images/sample_03/frame_000.png", "images/sample_03/frame_001.png", "images/sample_03/frame_002.png", "images/sample_03/frame_003.png", "images/sample_03/frame_004.png", "images/sample_03/frame_005.png", "images/sample_03/frame_006.png", "images/sample_03/frame_007.png", "images/sample_03/frame_008.png", "images/sample_03/frame_009.png", ], 256, 0, ], [ [ "images/sample_04/frame_000.png", "images/sample_04/frame_001.png", "images/sample_04/frame_002.png", "images/sample_04/frame_003.png", "images/sample_04/frame_004.png", "images/sample_04/frame_005.png", ], 256, 0, ], [ [ "images/sample_05/frame_000.png", "images/sample_05/frame_001.png", "images/sample_05/frame_002.png", "images/sample_05/frame_003.png", ], 256, 0, ], ], inputs=[input_gallery, image_size, ref_frame], label="Click to load sample images", ) gr.Markdown(""" --- **Tips:** - Upload 2-16 frames of the same scene - Frames can have different degradations (blur, noise, etc.) - The model works best when frames have slight variations - Processing size affects quality and speed **Author:** Veli Karaarslan | [GitHub](https://github.com/allcodernet/MFIR) """) process_btn.click( fn=process_images, inputs=[input_gallery, image_size, ref_frame], outputs=[output_image, status_text], ) return demo demo = create_ui() if __name__ == "__main__": demo.launch()