Download models.py from abbaab/DGFF: direct link, hf CLI and curl.
- Browser
- Download file 6 kB
-
https://huggingface.co/abbaab/DGFF/resolve/main/models.py
- Command line
-
hf download hf://abbaab/DGFF/models.py
-
curl -L -o models.py https://huggingface.co/abbaab/DGFF/resolve/main/models.py
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 | |