marduk-ra commited on
Commit
302e64f
·
verified ·
1 Parent(s): d562eee

Upload 2 files

Browse files

Package versions fixes

Files changed (2) hide show
  1. app.py +26 -25
  2. 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: list[int] = field(default_factory=lambda: [64, 128, 256])
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) -> dict:
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: dict) -> "FeatureFusionConfig":
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: list[int] = [64, 128, 256],
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 | None = None) -> torch.Tensor:
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 | None = None) -> torch.Tensor:
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: list[int] = [128, 64],
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 | None = None) -> dict[str, torch.Tensor]:
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: dict[str, torch.Tensor],
409
  strict: bool = False,
410
- ) -> tuple[list[str], list[str]]:
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 | None = None
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: list[tuple[np.ndarray, str]] | None,
498
  target_size: int = 256,
499
- ) -> tuple[torch.Tensor, dict[str, Any]] | tuple[None, None]:
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: dict[str, Any],
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: list[tuple[np.ndarray, str]] | None,
588
  image_size: int = 256,
589
  ref_frame: int = 0,
590
- ) -> tuple[Image.Image | None, str]:
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.0.0
4
- huggingface_hub>=0.20.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