126 lines
5.1 KiB
Python
Executable File
126 lines
5.1 KiB
Python
Executable File
# classes/image_tower.py
|
||
from __future__ import annotations
|
||
import math
|
||
from typing import Optional
|
||
|
||
import torch
|
||
from torch import nn
|
||
from torchvision import transforms
|
||
from classes.backbones import BACKBONES, list_names, load_backbone_weights
|
||
from classes.SE_attention import SEBlock
|
||
|
||
def build_backbone(name: str, freeze_ratio: float = 0.0, augment: bool = True):
|
||
"""
|
||
Operational builder:
|
||
- instantiate with DEFAULT weights
|
||
- strip classifier → features
|
||
- apply ratio-based freezing over coarse blocks
|
||
- return (model, out_dim, transform)
|
||
"""
|
||
key = (name or "").lower()
|
||
if key not in BACKBONES:
|
||
raise ValueError(f"Unsupported backbone '{name}'. Valid options: {list_names()}")
|
||
|
||
spec = BACKBONES[key]
|
||
m = spec.ctor(weights=spec.weights_default)
|
||
out_dim, m = spec.strip(m)
|
||
load_backbone_weights(key, m)
|
||
|
||
# transforms: use the weights’ mean/std, but keep your augmentation pipeline
|
||
mean = getattr(spec.weights_default, "meta", {}).get("mean", (0.485, 0.456, 0.406))
|
||
std = getattr(spec.weights_default, "meta", {}).get("std", (0.229, 0.224, 0.225))
|
||
crop = 299 if key == "inception_v3" else 224
|
||
|
||
if augment:
|
||
transform = transforms.Compose([
|
||
transforms.Resize(256),
|
||
transforms.CenterCrop(crop),
|
||
transforms.RandomHorizontalFlip(),
|
||
transforms.RandomVerticalFlip(),
|
||
transforms.RandomRotation(15),
|
||
transforms.ColorJitter(0.1, 0.1, 0.1, 0.05),
|
||
transforms.ToTensor(),
|
||
transforms.Normalize(mean=mean, std=std),
|
||
])
|
||
else:
|
||
transform = transforms.Compose([
|
||
transforms.Resize(256),
|
||
transforms.CenterCrop(crop),
|
||
transforms.ToTensor(),
|
||
transforms.Normalize(mean=mean, std=std),
|
||
])
|
||
|
||
# ratio-based freezing: freeze earliest floor(N * freeze_ratio) blocks
|
||
fr = max(0.0, min(1.0, float(freeze_ratio)))
|
||
blocks = spec.blocks(m)
|
||
n = len(blocks)
|
||
freeze_n = int(math.floor(n * fr))
|
||
for b in blocks[:freeze_n]:
|
||
for p in b.parameters():
|
||
p.requires_grad = False
|
||
|
||
return m, out_dim, transform
|
||
|
||
class ImageTower(nn.Module):
|
||
"""
|
||
Vision backbone → pooled features.
|
||
- backbone: one of list_names() (default 'efficientnet_b0')
|
||
- always DEFAULT torchvision weights
|
||
- freeze_ratio ∈ [0,1] freezes earliest floor(N*freeze_ratio) blocks
|
||
- returns [N, out_dim] features from backbone forward
|
||
"""
|
||
def __init__(self, backbone: str = "efficientnet_b0", freeze_ratio: float = 0.0,
|
||
use_se: bool = False, se_reduction: int = 16, se_pre_norm: bool = True,
|
||
augment: bool = True, geometry_dim: int = 0):
|
||
super().__init__()
|
||
self.backbone, base_dim, self.transform = build_backbone(backbone, freeze_ratio, augment=augment)
|
||
self._name = backbone
|
||
# Keep ordered blocks for dynamic freezing/thawing
|
||
key = (self._name or "").lower()
|
||
self._spec = BACKBONES[key]
|
||
self._blocks = self._spec.blocks(self.backbone)
|
||
# Optional tower-level SE over the final feature vector
|
||
self.base_dim = base_dim
|
||
self.geometry_dim = max(0, int(geometry_dim))
|
||
self.out_dim = self.base_dim + self.geometry_dim
|
||
self.tower_ln = nn.LayerNorm(self.base_dim) if se_pre_norm else nn.Identity()
|
||
self.tower_se = SEBlock(self.base_dim, reduction=se_reduction, residual=True) if use_se else None
|
||
|
||
def forward(self, x: torch.Tensor, geometry: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||
y = self.backbone(x)
|
||
# sanity: pooled features, not logits
|
||
assert y.dim() == 2 and y.size(1) == self.base_dim, \
|
||
f"Expected features [N,{self.base_dim}], got {tuple(y.shape)}"
|
||
if self.tower_se is not None:
|
||
y, _ = self.tower_se(self.tower_ln(y))
|
||
if self.geometry_dim > 0:
|
||
if geometry is None or geometry.numel() == 0:
|
||
geom = torch.zeros(y.size(0), self.geometry_dim, device=y.device, dtype=y.dtype)
|
||
else:
|
||
if geometry.dim() == 1:
|
||
geom = geometry.unsqueeze(0)
|
||
else:
|
||
geom = geometry
|
||
geom = geom.to(device=y.device, dtype=y.dtype)
|
||
if geom.size(0) != y.size(0):
|
||
raise ValueError(f"Geometry batch size mismatch: {geom.size(0)} vs {y.size(0)}")
|
||
if geom.size(1) != self.geometry_dim:
|
||
raise ValueError(f"Expected geometry dim {self.geometry_dim}, got {geom.size(1)}")
|
||
y = torch.cat([y, geom], dim=1)
|
||
return y
|
||
|
||
def set_freeze_ratio(self, ratio: float):
|
||
"""Dynamically freeze earliest floor(N*ratio) backbone blocks.
|
||
ratio in [0,1]."""
|
||
r = max(0.0, min(1.0, float(ratio)))
|
||
n = len(self._blocks)
|
||
freeze_n = int(math.floor(n * r))
|
||
# Unfreeze all first
|
||
for b in self._blocks:
|
||
for p in b.parameters():
|
||
p.requires_grad = True
|
||
# Freeze earliest blocks
|
||
for b in self._blocks[:freeze_n]:
|
||
for p in b.parameters():
|
||
p.requires_grad = False
|