TobiasLogic's picture
Upload ObjectModel-v1: code, README, assets (v1 checkpoint)
9b92c75 verified
Raw
History Blame Contribute Delete
16.4 kB
from __future__ import annotations
from copy import deepcopy
from dataclasses import dataclass
from typing import Any
import torch
from torch import Tensor, nn
import torch.nn.functional as F
from .boxes import inverse_sigmoid
class ConvNormAct(nn.Sequential):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int = 1,
stride: int = 1,
groups: int = 1,
activation: bool = True,
) -> None:
padding = kernel_size // 2
layers: list[nn.Module] = [
nn.Conv2d(
in_channels,
out_channels,
kernel_size,
stride,
padding,
groups=groups,
bias=False,
),
nn.BatchNorm2d(out_channels),
]
if activation:
layers.append(nn.SiLU(inplace=True))
super().__init__(*layers)
class GatedConvBlock(nn.Module):
"""Inverted residual block with a cheap learned residual gate."""
def __init__(self, channels: int, expansion: float = 2.0) -> None:
super().__init__()
hidden = int(channels * expansion)
self.expand = ConvNormAct(channels, hidden)
self.depthwise = ConvNormAct(hidden, hidden, 3, groups=hidden)
self.project = ConvNormAct(hidden, channels, activation=False)
self.gate = nn.Parameter(torch.zeros(1))
def forward(self, inputs: Tensor) -> Tensor:
return inputs + torch.tanh(self.gate) * self.project(self.depthwise(self.expand(inputs)))
class BackboneStage(nn.Sequential):
def __init__(self, in_channels: int, out_channels: int, depth: int, stride: int) -> None:
super().__init__(
ConvNormAct(in_channels, out_channels, 3, stride=stride),
*(GatedConvBlock(out_channels) for _ in range(depth)),
)
class CompactBackbone(nn.Module):
def __init__(
self, stem_channels: int, channels: list[int], depths: list[int]
) -> None:
super().__init__()
if len(channels) != 4 or len(depths) != 4:
raise ValueError("Backbone requires four channel and depth values")
self.stem = nn.Sequential(
ConvNormAct(3, stem_channels, 3, stride=2),
ConvNormAct(stem_channels, stem_channels, 3, stride=2),
)
stages: list[nn.Module] = []
in_channels = stem_channels
for index, (out_channels, depth) in enumerate(zip(channels, depths, strict=True)):
stages.append(
BackboneStage(in_channels, out_channels, depth, stride=1 if index == 0 else 2)
)
in_channels = out_channels
self.stages = nn.ModuleList(stages)
self.out_channels = channels[1:]
def forward(self, images: Tensor) -> list[Tensor]:
features = self.stem(images)
outputs = []
for index, stage in enumerate(self.stages):
features = stage(features)
if index > 0:
outputs.append(features)
return outputs
class PyramidFusion(nn.Module):
def __init__(self, in_channels: list[int], hidden_dim: int, depth: int) -> None:
super().__init__()
self.lateral = nn.ModuleList(ConvNormAct(c, hidden_dim) for c in in_channels)
self.refine = nn.ModuleList(
nn.Sequential(*(GatedConvBlock(hidden_dim, expansion=1.5) for _ in range(depth)))
for _ in in_channels
)
def forward(self, inputs: list[Tensor]) -> list[Tensor]:
projected = [layer(x) for layer, x in zip(self.lateral, inputs, strict=True)]
outputs = list(projected)
for index in range(len(outputs) - 2, -1, -1):
outputs[index] = outputs[index] + F.interpolate(
outputs[index + 1], size=outputs[index].shape[-2:], mode="nearest"
)
return [block(x) for block, x in zip(self.refine, outputs, strict=True)]
def sine_position_encoding(
height: int, width: int, dim: int, device: torch.device, dtype: torch.dtype
) -> Tensor:
if dim % 4 != 0:
raise ValueError("Position encoding dimension must be divisible by four")
y, x = torch.meshgrid(
torch.linspace(0, 1, height, device=device, dtype=dtype),
torch.linspace(0, 1, width, device=device, dtype=dtype),
indexing="ij",
)
frequencies = torch.arange(dim // 4, device=device, dtype=dtype)
frequencies = 2.0 * torch.pi * (10000.0 ** (-frequencies / max(dim // 4, 1)))
x = x.flatten()[:, None] * frequencies[None]
y = y.flatten()[:, None] * frequencies[None]
return torch.cat((x.sin(), x.cos(), y.sin(), y.cos()), dim=-1)
class FeedForward(nn.Sequential):
def __init__(self, dim: int, expansion: int = 4, dropout: float = 0.0) -> None:
super().__init__(
nn.Linear(dim, dim * expansion),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(dim * expansion, dim),
nn.Dropout(dropout),
)
class LatentLayer(nn.Module):
def __init__(self, dim: int, num_heads: int, dropout: float) -> None:
super().__init__()
self.norm1 = nn.LayerNorm(dim)
self.attention = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True)
self.norm2 = nn.LayerNorm(dim)
self.ffn = FeedForward(dim, dropout=dropout)
def forward(self, inputs: Tensor) -> Tensor:
normalized = self.norm1(inputs)
inputs = inputs + self.attention(normalized, normalized, normalized, need_weights=False)[0]
return inputs + self.ffn(self.norm2(inputs))
class LatentMemory(nn.Module):
"""Compresses multi-scale maps into a fixed-size global reasoning memory."""
def __init__(
self,
dim: int,
latent_count: int,
pool_sizes: list[int],
layers: int,
num_heads: int,
dropout: float,
) -> None:
super().__init__()
if len(pool_sizes) != 3:
raise ValueError("One latent pool size is required for each pyramid level")
self.pool_sizes = pool_sizes
self.latents = nn.Parameter(torch.empty(latent_count, dim))
self.level_embedding = nn.Parameter(torch.empty(len(pool_sizes), dim))
self.query_norm = nn.LayerNorm(dim)
self.token_norm = nn.LayerNorm(dim)
self.compress = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True)
self.layers = nn.ModuleList(LatentLayer(dim, num_heads, dropout) for _ in range(layers))
nn.init.normal_(self.latents, std=0.02)
nn.init.normal_(self.level_embedding, std=0.02)
def forward(self, features: list[Tensor]) -> Tensor:
tokens = []
for level, (feature, size) in enumerate(zip(features, self.pool_sizes, strict=True)):
pooled = F.adaptive_avg_pool2d(feature, (size, size)).flatten(2).transpose(1, 2)
position = sine_position_encoding(
size, size, feature.shape[1], feature.device, feature.dtype
)
tokens.append(pooled + position[None] + self.level_embedding[level][None, None])
token_memory = self.token_norm(torch.cat(tokens, dim=1))
latents = self.latents[None].expand(features[0].shape[0], -1, -1)
latents = latents + self.compress(
self.query_norm(latents), token_memory, token_memory, need_weights=False
)[0]
for layer in self.layers:
latents = layer(latents)
return latents
class QueryLocalSampler(nn.Module):
"""Samples high-resolution pyramid evidence around each evolving query box."""
def __init__(self, dim: int, num_levels: int, points: int) -> None:
super().__init__()
self.num_levels = num_levels
self.points = points
self.offsets = nn.Linear(dim, num_levels * points * 2)
self.weights = nn.Linear(dim, num_levels * points)
self.output = nn.Linear(dim, dim)
nn.init.zeros_(self.offsets.weight)
nn.init.zeros_(self.offsets.bias)
nn.init.zeros_(self.weights.weight)
nn.init.zeros_(self.weights.bias)
def forward(self, queries: Tensor, boxes: Tensor, features: list[Tensor]) -> Tensor:
batch, query_count, _ = queries.shape
offsets = self.offsets(queries).view(
batch, query_count, self.num_levels, self.points, 2
)
offsets = offsets.tanh() * boxes[..., None, None, 2:] * 0.5
centers = boxes[..., None, None, :2]
sample_points = (centers + offsets).clamp(0.0, 1.0)
weights = self.weights(queries).view(
batch, query_count, self.num_levels * self.points
)
weights = weights.softmax(dim=-1).view(
batch, query_count, self.num_levels, self.points
)
sampled_levels = []
for level, feature in enumerate(features):
grid = sample_points[:, :, level] * 2.0 - 1.0
sampled = F.grid_sample(
feature,
grid,
mode="bilinear",
padding_mode="zeros",
align_corners=False,
)
sampled = sampled.permute(0, 2, 3, 1)
sampled_levels.append(sampled)
sampled_features = torch.stack(sampled_levels, dim=2)
fused = (sampled_features * weights[..., None]).sum(dim=(2, 3))
return self.output(fused)
class DecoderLayer(nn.Module):
def __init__(
self, dim: int, num_heads: int, num_levels: int, local_points: int, dropout: float
) -> None:
super().__init__()
self.norm1 = nn.LayerNorm(dim)
self.self_attention = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True)
self.norm2 = nn.LayerNorm(dim)
self.global_attention = nn.MultiheadAttention(dim, num_heads, dropout, batch_first=True)
self.norm3 = nn.LayerNorm(dim)
self.local_sampler = QueryLocalSampler(dim, num_levels, local_points)
self.norm4 = nn.LayerNorm(dim)
self.ffn = FeedForward(dim, dropout=dropout)
def forward(
self, queries: Tensor, memory: Tensor, boxes: Tensor, features: list[Tensor]
) -> Tensor:
normalized = self.norm1(queries)
queries = queries + self.self_attention(
normalized, normalized, normalized, need_weights=False
)[0]
queries = queries + self.global_attention(
self.norm2(queries), memory, memory, need_weights=False
)[0]
queries = queries + self.local_sampler(self.norm3(queries), boxes, features)
return queries + self.ffn(self.norm4(queries))
class MLP(nn.Sequential):
def __init__(self, input_dim: int, hidden_dim: int, output_dim: int, layers: int) -> None:
modules: list[nn.Module] = []
for index in range(layers):
in_dim = input_dim if index == 0 else hidden_dim
out_dim = output_dim if index == layers - 1 else hidden_dim
modules.append(nn.Linear(in_dim, out_dim))
if index < layers - 1:
modules.append(nn.ReLU(inplace=True))
super().__init__(*modules)
class DenseAuxiliaryHead(nn.Module):
def __init__(self, dim: int, num_classes: int) -> None:
super().__init__()
self.shared = nn.ModuleList(
nn.Sequential(ConvNormAct(dim, dim, 3, groups=dim), ConvNormAct(dim, dim))
for _ in range(3)
)
self.classification = nn.Conv2d(dim, num_classes, 1)
self.regression = nn.Conv2d(dim, 4, 1)
def forward(self, features: list[Tensor]) -> list[dict[str, Tensor]]:
outputs = []
for feature, tower in zip(features, self.shared, strict=True):
hidden = tower(feature)
outputs.append(
{
"logits": self.classification(hidden),
"distances": F.softplus(self.regression(hidden)),
}
)
return outputs
@dataclass(frozen=True)
class ObjectModelV1Spec:
num_classes: int = 80
input_size: int = 640
stem_channels: int = 48
backbone_channels: tuple[int, int, int, int] = (64, 128, 256, 384)
backbone_depths: tuple[int, int, int, int] = (2, 3, 6, 3)
hidden_dim: int = 256
fpn_depth: int = 2
latent_count: int = 64
latent_pool_sizes: tuple[int, int, int] = (12, 6, 3)
latent_layers: int = 2
decoder_layers: int = 6
num_queries: int = 300
num_heads: int = 8
local_points: int = 4
dropout: float = 0.0
dense_aux: bool = True
class ObjectModelV1(nn.Module):
"""NMS-free detector with compressed global memory and local geometric sampling."""
def __init__(self, spec: ObjectModelV1Spec) -> None:
super().__init__()
self.spec = spec
self.backbone = CompactBackbone(
spec.stem_channels, list(spec.backbone_channels), list(spec.backbone_depths)
)
self.neck = PyramidFusion(self.backbone.out_channels, spec.hidden_dim, spec.fpn_depth)
self.memory = LatentMemory(
spec.hidden_dim,
spec.latent_count,
list(spec.latent_pool_sizes),
spec.latent_layers,
spec.num_heads,
spec.dropout,
)
decoder_template = DecoderLayer(
spec.hidden_dim, spec.num_heads, 3, spec.local_points, spec.dropout
)
self.decoder = nn.ModuleList(deepcopy(decoder_template) for _ in range(spec.decoder_layers))
self.query_embedding = nn.Embedding(spec.num_queries, spec.hidden_dim)
self.reference_points = nn.Embedding(spec.num_queries, 4)
self.class_heads = nn.ModuleList(
nn.Linear(spec.hidden_dim, spec.num_classes) for _ in range(spec.decoder_layers)
)
self.box_heads = nn.ModuleList(
MLP(spec.hidden_dim, spec.hidden_dim, 4, 3) for _ in range(spec.decoder_layers)
)
self.dense_head = (
DenseAuxiliaryHead(spec.hidden_dim, spec.num_classes) if spec.dense_aux else None
)
self._reset_parameters()
def _reset_parameters(self) -> None:
prior_probability = 0.01
class_bias = -torch.log(torch.tensor((1.0 - prior_probability) / prior_probability))
for head in self.class_heads:
nn.init.constant_(head.bias, class_bias)
nn.init.zeros_(self.reference_points.weight)
with torch.no_grad():
self.reference_points.weight[:, 2:] = -2.0
for head in self.box_heads:
nn.init.zeros_(head[-1].weight)
nn.init.zeros_(head[-1].bias)
if self.dense_head is not None:
nn.init.constant_(self.dense_head.classification.bias, class_bias)
nn.init.zeros_(self.dense_head.regression.weight)
nn.init.constant_(self.dense_head.regression.bias, 1.0)
def forward(self, images: Tensor) -> dict[str, Any]:
features = self.neck(self.backbone(images))
memory = self.memory(features)
batch = images.shape[0]
queries = self.query_embedding.weight[None].expand(batch, -1, -1)
boxes = self.reference_points.weight.sigmoid()[None].expand(batch, -1, -1)
layer_outputs: list[dict[str, Tensor]] = []
for layer, class_head, box_head in zip(
self.decoder, self.class_heads, self.box_heads, strict=True
):
queries = layer(queries, memory, boxes, features)
boxes = (inverse_sigmoid(boxes) + box_head(queries)).sigmoid()
layer_outputs.append({"pred_logits": class_head(queries), "pred_boxes": boxes})
boxes = boxes.detach() if self.training else boxes
output: dict[str, Any] = dict(layer_outputs[-1])
output["aux_outputs"] = layer_outputs[:-1]
if self.training and self.dense_head is not None:
output["dense_outputs"] = self.dense_head(features)
return output
def build_model(config: dict[str, Any]) -> ObjectModelV1:
model_config = config.get("model", config)
fields = ObjectModelV1Spec.__dataclass_fields__
unknown = set(model_config) - set(fields)
if unknown:
raise ValueError(f"Unknown model configuration keys: {sorted(unknown)}")
values = dict(model_config)
for key in ("backbone_channels", "backbone_depths", "latent_pool_sizes"):
if key in values:
values[key] = tuple(values[key])
return ObjectModelV1(ObjectModelV1Spec(**values))