DGFF / models.py
abbaab's picture
Upload folder using huggingface_hub (part 5)
d610c9f verified
Raw History Blame Contribute Delete
6 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from ultralytics import YOLO
class ResBlock(nn.Module):
def __init__(self, ch):
super().__init__()
self.block = nn.Sequential(
nn.Conv2d(ch, ch, 3, padding=1, bias=False), nn.BatchNorm2d(ch), nn.ReLU(True),
nn.Conv2d(ch, ch, 3, padding=1, bias=False), nn.BatchNorm2d(ch),
)
self.relu = nn.ReLU(True)
def forward(self, x):
return self.relu(x + self.block(x))
class EncoderBlock(nn.Module):
def __init__(self, in_ch, out_ch):
super().__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),
nn.BatchNorm2d(out_ch), nn.ReLU(True),
)
self.res = ResBlock(out_ch)
self.down = nn.MaxPool2d(2)
def forward(self, x):
x = self.res(self.conv(x))
skip = x
return self.down(x), skip
class DecoderBlock(nn.Module):
def __init__(self, in_ch, skip_ch, out_ch, feedback_ch=32):
super().__init__()
self.merge = nn.Sequential(
nn.Conv2d(in_ch + skip_ch, out_ch, 1, bias=False),
nn.BatchNorm2d(out_ch), nn.ReLU(True),
)
self.res = ResBlock(out_ch)
if feedback_ch is not None:
self.feedback_proj = nn.Conv2d(feedback_ch, out_ch, 1, bias=False)
else:
self.feedback_proj = None
def forward(self, x, skip, feedback=None):
x = F.interpolate(x, size=skip.shape[-2:], mode='bilinear', align_corners=False)
x = self.merge(torch.cat([x, skip], dim=1))
if feedback is not None and self.feedback_proj is not None:
if feedback.shape[-2:] != x.shape[-2:]:
feedback = F.interpolate(feedback, size=x.shape[-2:],
mode='bilinear', align_corners=False)
x = x + self.feedback_proj(feedback)
return self.res(x)
class LLEN(nn.Module):
def __init__(self, base_ch=32):
super().__init__()
c = base_ch
self.enc1 = EncoderBlock(3, c)
self.enc2 = EncoderBlock(c, c * 2)
self.enc3 = EncoderBlock(c * 2, c * 4)
self.enc4 = EncoderBlock(c * 4, c * 8)
self.bottleneck = nn.Sequential(ResBlock(c * 8), ResBlock(c * 8))
self.dec4 = DecoderBlock(c * 8, c * 8, c * 4, feedback_ch=c * 4)
self.dec3 = DecoderBlock(c * 4, c * 4, c * 2, feedback_ch=c * 2)
self.dec2 = DecoderBlock(c * 2, c * 2, c, feedback_ch=c)
self.dec1 = DecoderBlock(c, c, c, feedback_ch=None)
self.head = nn.Sequential(
nn.Conv2d(c, c, 3, padding=1, bias=False), nn.ReLU(True),
nn.Conv2d(c, 3, 1), nn.Sigmoid(),
)
def forward(self, x, dgff_p5=None, dgff_p4=None, dgff_p3=None):
x, s1 = self.enc1(x)
x, s2 = self.enc2(x)
x, s3 = self.enc3(x)
x, s4 = self.enc4(x)
x = self.bottleneck(x)
x = self.dec4(x, s4, feedback=dgff_p5)
x = self.dec3(x, s3, feedback=dgff_p4)
x = self.dec2(x, s2, feedback=dgff_p3)
x = self.dec1(x, s1)
return self.head(x)
def count_parameters(self):
return sum(p.numel() for p in self.parameters() if p.requires_grad)
class GatedAdapter(nn.Module):
def __init__(self, in_ch: int, out_ch: int):
super().__init__()
self.proj = nn.Sequential(
nn.Conv2d(in_ch, out_ch, kernel_size=1, bias=False),
nn.BatchNorm2d(out_ch),
nn.ReLU(inplace=True),
)
self.gate = nn.Sequential(
nn.Conv2d(in_ch, out_ch, kernel_size=1, bias=False),
nn.Sigmoid(),
)
def forward(self, f_det, target_size):
f_det = nn.functional.interpolate(
f_det, size=target_size, mode='bilinear', align_corners=False
)
return self.gate(f_det) * self.proj(f_det)
class DGFFModule(nn.Module):
def __init__(self, yolo_weights='yolov8n.pt', llen_base_ch=32, freeze_backbone=True):
super().__init__()
yolo = YOLO(yolo_weights)
self.backbone = yolo.model.model[:10]
if freeze_backbone:
for p in self.backbone.parameters():
p.requires_grad_(False)
with torch.no_grad():
probe = torch.zeros(1, 3, 64, 64)
out = probe
for i, layer in enumerate(self.backbone):
out = layer(out)
if i == 4:
actual_p3 = out.shape[1]
if i == 6:
actual_p4 = out.shape[1]
if i == 9:
actual_p5 = out.shape[1]
self.p3_ch = actual_p3
self.p4_ch = actual_p4
self.p5_ch = actual_p5
c = llen_base_ch
self._llen_base_ch = c
self.adapter_p3 = GatedAdapter(self.p3_ch, c)
self.adapter_p5 = GatedAdapter(self.p5_ch, c * 4)
self.adapter_p4 = GatedAdapter(self.p4_ch, c * 2)
def forward(self, enhanced_img):
x = enhanced_img
p4_raw = None
p5_raw = None
p3_raw = None
for i, layer in enumerate(self.backbone):
x = layer(x)
if i == 4:
p3_raw = x
if i == 6:
p4_raw = x
if i == 9:
p5_raw = x
c = self._llen_base_ch
H_input, W_input = enhanced_img.shape[-2:]
p5_size = (H_input // 8, W_input // 8)
p4_size = (H_input // 4, W_input // 4)
p3_size = (H_input // 2, W_input // 2)
p5_feedback = self.adapter_p5(p5_raw, p5_size)
p4_feedback = self.adapter_p4(p4_raw, p4_size)
p3_feedback = self.adapter_p3(p3_raw, p3_size)
return p4_feedback, p5_feedback, p3_feedback
def count_parameters(self):
total = sum(p.numel() for p in self.parameters())
trainable = sum(p.numel() for p in self.parameters() if p.requires_grad)
return total, trainable