moved_repo_first_update

This commit is contained in:
rpotter6298
2026-02-24 10:39:48 +01:00
commit 9894a23f09
98 changed files with 35387 additions and 0 deletions
+125
View File
@@ -0,0 +1,125 @@
# 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