Upload 2 files
Browse filesPackage versions fixes
- app.py +26 -25
- requirements.txt +2 -2
app.py
CHANGED
|
@@ -5,28 +5,25 @@ This module provides a web UI for testing the multi-frame image restoration mode
|
|
| 5 |
|
| 6 |
from __future__ import annotations
|
| 7 |
|
|
|
|
| 8 |
from pathlib import Path
|
| 9 |
-
from typing import Any
|
| 10 |
|
| 11 |
import gradio as gr
|
| 12 |
import numpy as np
|
| 13 |
import torch
|
|
|
|
|
|
|
| 14 |
from huggingface_hub import hf_hub_download
|
| 15 |
from PIL import Image
|
| 16 |
from torchvision import transforms
|
|
|
|
| 17 |
|
| 18 |
|
| 19 |
# ============================================================================
|
| 20 |
# Model Architecture (inline for single-file deployment)
|
| 21 |
# ============================================================================
|
| 22 |
|
| 23 |
-
from dataclasses import dataclass, field
|
| 24 |
-
from typing import Literal
|
| 25 |
-
|
| 26 |
-
import torch.nn as nn
|
| 27 |
-
import torch.nn.functional as F
|
| 28 |
-
from torchvision.ops import deform_conv2d
|
| 29 |
-
|
| 30 |
|
| 31 |
@dataclass
|
| 32 |
class FeatureFusionConfig:
|
|
@@ -34,7 +31,7 @@ class FeatureFusionConfig:
|
|
| 34 |
|
| 35 |
in_channels: int = 3
|
| 36 |
max_frames: int = 16
|
| 37 |
-
encoder_channels:
|
| 38 |
encoder_blocks_per_stage: int = 2
|
| 39 |
offset_channels: int = 64
|
| 40 |
deform_groups: int = 8
|
|
@@ -44,14 +41,14 @@ class FeatureFusionConfig:
|
|
| 44 |
decoder_blocks_per_stage: int = 2
|
| 45 |
out_channels: int = 3
|
| 46 |
|
| 47 |
-
def to_dict(self) ->
|
| 48 |
return {
|
| 49 |
k: list(v) if isinstance(v, list) else v
|
| 50 |
for k, v in self.__dict__.items()
|
| 51 |
}
|
| 52 |
|
| 53 |
@classmethod
|
| 54 |
-
def from_dict(cls, d:
|
| 55 |
return cls(**{k: v for k, v in d.items() if k in cls.__dataclass_fields__})
|
| 56 |
|
| 57 |
|
|
@@ -79,10 +76,12 @@ class Encoder(nn.Module):
|
|
| 79 |
def __init__(
|
| 80 |
self,
|
| 81 |
in_channels: int = 3,
|
| 82 |
-
channels:
|
| 83 |
blocks_per_stage: int = 2,
|
| 84 |
):
|
| 85 |
super().__init__()
|
|
|
|
|
|
|
| 86 |
self.channels = channels
|
| 87 |
self.conv_first = nn.Conv2d(in_channels, channels[0], 3, 1, 1)
|
| 88 |
self.stages = nn.ModuleList()
|
|
@@ -213,7 +212,7 @@ class TemporalAttentionFusion(nn.Module):
|
|
| 213 |
embed = embed.permute(0, 2, 1)
|
| 214 |
return embed.unsqueeze(-1).unsqueeze(-1)
|
| 215 |
|
| 216 |
-
def forward(self, features: torch.Tensor, ref_idx: int
|
| 217 |
B, N, C, H, W = features.shape
|
| 218 |
if ref_idx is None:
|
| 219 |
ref_idx = 0
|
|
@@ -265,7 +264,7 @@ class AdaptiveFusion(nn.Module):
|
|
| 265 |
ResidualBlock(channels),
|
| 266 |
)
|
| 267 |
|
| 268 |
-
def forward(self, features: torch.Tensor, ref_idx: int
|
| 269 |
B, N, C, H, W = features.shape
|
| 270 |
|
| 271 |
weights_list = []
|
|
@@ -292,11 +291,13 @@ class Decoder(nn.Module):
|
|
| 292 |
def __init__(
|
| 293 |
self,
|
| 294 |
in_channels: int = 256,
|
| 295 |
-
channels:
|
| 296 |
out_channels: int = 3,
|
| 297 |
blocks_per_stage: int = 2,
|
| 298 |
):
|
| 299 |
super().__init__()
|
|
|
|
|
|
|
| 300 |
|
| 301 |
self.upsamples = nn.ModuleList()
|
| 302 |
self.stages = nn.ModuleList()
|
|
@@ -372,7 +373,7 @@ class FeatureFusionModel(nn.Module):
|
|
| 372 |
blocks_per_stage=config.decoder_blocks_per_stage,
|
| 373 |
)
|
| 374 |
|
| 375 |
-
def forward(self, frames: torch.Tensor, ref_idx: int
|
| 376 |
B, N, C, H, W = frames.shape
|
| 377 |
|
| 378 |
if ref_idx is None:
|
|
@@ -405,9 +406,9 @@ class FeatureFusionModel(nn.Module):
|
|
| 405 |
|
| 406 |
def load_state_dict_with_compatibility(
|
| 407 |
self,
|
| 408 |
-
state_dict:
|
| 409 |
strict: bool = False,
|
| 410 |
-
) ->
|
| 411 |
if "fusion.temporal_embed" in state_dict:
|
| 412 |
old_embed = state_dict["fusion.temporal_embed"]
|
| 413 |
old_num_frames = old_embed.shape[1]
|
|
@@ -444,14 +445,14 @@ class FeatureFusionModel(nn.Module):
|
|
| 444 |
class ModelState:
|
| 445 |
"""Global state for loaded model."""
|
| 446 |
|
| 447 |
-
model: FeatureFusionModel
|
| 448 |
device: str = "cpu"
|
| 449 |
|
| 450 |
|
| 451 |
state = ModelState()
|
| 452 |
|
| 453 |
|
| 454 |
-
def load_model():
|
| 455 |
"""Load model from Hugging Face Hub."""
|
| 456 |
if state.model is not None:
|
| 457 |
return "Model already loaded"
|
|
@@ -494,9 +495,9 @@ def load_model():
|
|
| 494 |
|
| 495 |
|
| 496 |
def preprocess_images(
|
| 497 |
-
images:
|
| 498 |
target_size: int = 256,
|
| 499 |
-
) ->
|
| 500 |
"""Preprocess input images to tensor with aspect ratio preservation."""
|
| 501 |
if not images:
|
| 502 |
return None, None
|
|
@@ -565,7 +566,7 @@ def tensor_to_pil(tensor: torch.Tensor) -> Image.Image:
|
|
| 565 |
|
| 566 |
def postprocess_output(
|
| 567 |
tensor: torch.Tensor,
|
| 568 |
-
preprocess_info:
|
| 569 |
) -> Image.Image:
|
| 570 |
"""Remove padding from output tensor and convert to PIL Image."""
|
| 571 |
if tensor.dim() == 4:
|
|
@@ -584,10 +585,10 @@ def postprocess_output(
|
|
| 584 |
|
| 585 |
|
| 586 |
def process_images(
|
| 587 |
-
images:
|
| 588 |
image_size: int = 256,
|
| 589 |
ref_frame: int = 0,
|
| 590 |
-
) ->
|
| 591 |
"""Process input images through the model."""
|
| 592 |
if state.model is None:
|
| 593 |
load_result = load_model()
|
|
|
|
| 5 |
|
| 6 |
from __future__ import annotations
|
| 7 |
|
| 8 |
+
from dataclasses import dataclass, field
|
| 9 |
from pathlib import Path
|
| 10 |
+
from typing import Any, Dict, List, Literal, Optional, Tuple, Union
|
| 11 |
|
| 12 |
import gradio as gr
|
| 13 |
import numpy as np
|
| 14 |
import torch
|
| 15 |
+
import torch.nn as nn
|
| 16 |
+
import torch.nn.functional as F
|
| 17 |
from huggingface_hub import hf_hub_download
|
| 18 |
from PIL import Image
|
| 19 |
from torchvision import transforms
|
| 20 |
+
from torchvision.ops import deform_conv2d
|
| 21 |
|
| 22 |
|
| 23 |
# ============================================================================
|
| 24 |
# Model Architecture (inline for single-file deployment)
|
| 25 |
# ============================================================================
|
| 26 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
|
| 28 |
@dataclass
|
| 29 |
class FeatureFusionConfig:
|
|
|
|
| 31 |
|
| 32 |
in_channels: int = 3
|
| 33 |
max_frames: int = 16
|
| 34 |
+
encoder_channels: List[int] = field(default_factory=lambda: [64, 128, 256])
|
| 35 |
encoder_blocks_per_stage: int = 2
|
| 36 |
offset_channels: int = 64
|
| 37 |
deform_groups: int = 8
|
|
|
|
| 41 |
decoder_blocks_per_stage: int = 2
|
| 42 |
out_channels: int = 3
|
| 43 |
|
| 44 |
+
def to_dict(self) -> Dict:
|
| 45 |
return {
|
| 46 |
k: list(v) if isinstance(v, list) else v
|
| 47 |
for k, v in self.__dict__.items()
|
| 48 |
}
|
| 49 |
|
| 50 |
@classmethod
|
| 51 |
+
def from_dict(cls, d: Dict) -> "FeatureFusionConfig":
|
| 52 |
return cls(**{k: v for k, v in d.items() if k in cls.__dataclass_fields__})
|
| 53 |
|
| 54 |
|
|
|
|
| 76 |
def __init__(
|
| 77 |
self,
|
| 78 |
in_channels: int = 3,
|
| 79 |
+
channels: Optional[List[int]] = None,
|
| 80 |
blocks_per_stage: int = 2,
|
| 81 |
):
|
| 82 |
super().__init__()
|
| 83 |
+
if channels is None:
|
| 84 |
+
channels = [64, 128, 256]
|
| 85 |
self.channels = channels
|
| 86 |
self.conv_first = nn.Conv2d(in_channels, channels[0], 3, 1, 1)
|
| 87 |
self.stages = nn.ModuleList()
|
|
|
|
| 212 |
embed = embed.permute(0, 2, 1)
|
| 213 |
return embed.unsqueeze(-1).unsqueeze(-1)
|
| 214 |
|
| 215 |
+
def forward(self, features: torch.Tensor, ref_idx: Optional[int] = None) -> torch.Tensor:
|
| 216 |
B, N, C, H, W = features.shape
|
| 217 |
if ref_idx is None:
|
| 218 |
ref_idx = 0
|
|
|
|
| 264 |
ResidualBlock(channels),
|
| 265 |
)
|
| 266 |
|
| 267 |
+
def forward(self, features: torch.Tensor, ref_idx: Optional[int] = None) -> torch.Tensor:
|
| 268 |
B, N, C, H, W = features.shape
|
| 269 |
|
| 270 |
weights_list = []
|
|
|
|
| 291 |
def __init__(
|
| 292 |
self,
|
| 293 |
in_channels: int = 256,
|
| 294 |
+
channels: Optional[List[int]] = None,
|
| 295 |
out_channels: int = 3,
|
| 296 |
blocks_per_stage: int = 2,
|
| 297 |
):
|
| 298 |
super().__init__()
|
| 299 |
+
if channels is None:
|
| 300 |
+
channels = [128, 64]
|
| 301 |
|
| 302 |
self.upsamples = nn.ModuleList()
|
| 303 |
self.stages = nn.ModuleList()
|
|
|
|
| 373 |
blocks_per_stage=config.decoder_blocks_per_stage,
|
| 374 |
)
|
| 375 |
|
| 376 |
+
def forward(self, frames: torch.Tensor, ref_idx: Optional[int] = None) -> Dict[str, torch.Tensor]:
|
| 377 |
B, N, C, H, W = frames.shape
|
| 378 |
|
| 379 |
if ref_idx is None:
|
|
|
|
| 406 |
|
| 407 |
def load_state_dict_with_compatibility(
|
| 408 |
self,
|
| 409 |
+
state_dict: Dict[str, torch.Tensor],
|
| 410 |
strict: bool = False,
|
| 411 |
+
) -> Tuple[List[str], List[str]]:
|
| 412 |
if "fusion.temporal_embed" in state_dict:
|
| 413 |
old_embed = state_dict["fusion.temporal_embed"]
|
| 414 |
old_num_frames = old_embed.shape[1]
|
|
|
|
| 445 |
class ModelState:
|
| 446 |
"""Global state for loaded model."""
|
| 447 |
|
| 448 |
+
model: Optional[FeatureFusionModel] = None
|
| 449 |
device: str = "cpu"
|
| 450 |
|
| 451 |
|
| 452 |
state = ModelState()
|
| 453 |
|
| 454 |
|
| 455 |
+
def load_model() -> str:
|
| 456 |
"""Load model from Hugging Face Hub."""
|
| 457 |
if state.model is not None:
|
| 458 |
return "Model already loaded"
|
|
|
|
| 495 |
|
| 496 |
|
| 497 |
def preprocess_images(
|
| 498 |
+
images: Optional[List[Tuple[np.ndarray, str]]],
|
| 499 |
target_size: int = 256,
|
| 500 |
+
) -> Tuple[Optional[torch.Tensor], Optional[Dict[str, Any]]]:
|
| 501 |
"""Preprocess input images to tensor with aspect ratio preservation."""
|
| 502 |
if not images:
|
| 503 |
return None, None
|
|
|
|
| 566 |
|
| 567 |
def postprocess_output(
|
| 568 |
tensor: torch.Tensor,
|
| 569 |
+
preprocess_info: Dict[str, Any],
|
| 570 |
) -> Image.Image:
|
| 571 |
"""Remove padding from output tensor and convert to PIL Image."""
|
| 572 |
if tensor.dim() == 4:
|
|
|
|
| 585 |
|
| 586 |
|
| 587 |
def process_images(
|
| 588 |
+
images: Optional[List[Tuple[np.ndarray, str]]],
|
| 589 |
image_size: int = 256,
|
| 590 |
ref_frame: int = 0,
|
| 591 |
+
) -> Tuple[Optional[Image.Image], str]:
|
| 592 |
"""Process input images through the model."""
|
| 593 |
if state.model is None:
|
| 594 |
load_result = load_model()
|
requirements.txt
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
torch>=2.1.0
|
| 2 |
torchvision>=0.16.0
|
| 3 |
-
gradio>=4.
|
| 4 |
-
huggingface_hub>=0.
|
| 5 |
numpy>=1.26.0
|
| 6 |
Pillow>=10.0.0
|
|
|
|
| 1 |
torch>=2.1.0
|
| 2 |
torchvision>=0.16.0
|
| 3 |
+
gradio>=4.44.0
|
| 4 |
+
huggingface_hub>=0.24.0
|
| 5 |
numpy>=1.26.0
|
| 6 |
Pillow>=10.0.0
|