bioacoustic
audio
birds
perch_v2 / perch_v2.py
rambip's picture
Upload perch_v2.py with huggingface_hub
e300024 verified
Raw
History Blame Contribute Delete
6.95 kB
"""Native PyTorch implementation of the Perch v2 EfficientNet-B3 model.
Note: this file has been entirely generated by chatGPT 5.6 with codex, and I did not review every line.
I take the entire responsability for any error.
This file is the pytorch translation of the onnx artefact published here: https://www.kaggle.com/datasets/nikitababich/perchv2-onnx
You can download it locally with:
```
path = kagglehub.dataset_download("nikitababich/perchv2-onnx")
```
To generate the `.pth` weight file from the onnx, use `tools/extract_perch_backbone.py`
"""
from __future__ import annotations
import math
import torch
import torch.nn.functional as F
from torch import Tensor, nn
_WIDTH = 1.2
_DEPTH = 1.4
_NUM_PERCH_CLASSES = 14795
_NUM_PROTOTYPES = 4
_STAGES = (
(1, 16, 3, 1, 1),
(2, 24, 3, 2, 6),
(2, 40, 5, 2, 6),
(3, 80, 3, 2, 6),
(3, 112, 5, 1, 6),
(4, 192, 5, 2, 6),
(1, 320, 3, 1, 6),
)
def _channels(channels: int) -> int:
scaled = channels * _WIDTH
rounded = max(8, int(scaled + 4) // 8 * 8)
return int(rounded + 8 if rounded < 0.9 * scaled else rounded)
def _blocks(count: int) -> int:
return math.ceil(count * _DEPTH)
def _same_padding(x: Tensor, kernel_size: int, stride: int) -> Tensor:
height, width = x.shape[-2:]
def padding(size: int) -> tuple[int, int]:
total = max((math.ceil(size / stride) - 1) * stride + kernel_size - size, 0)
return total // 2, total - total // 2
top, bottom = padding(height)
left, right = padding(width)
return F.pad(x, (left, right, top, bottom))
torch.fx.wrap("_same_padding")
class SqueezeExcitation(nn.Module):
"""Per-channel squeeze-and-excitation gate."""
def __init__(self, channels: int, reduced_channels: int) -> None:
super().__init__()
self.reduce = nn.Linear(channels, reduced_channels)
self.expand = nn.Linear(reduced_channels, channels)
def forward(self, x: Tensor) -> Tensor:
scale = torch.sigmoid(self.expand(F.silu(self.reduce(x.mean(dim=(-2, -1))))))
return x * scale[:, :, None, None]
class MBConv(nn.Module):
"""Mobile inverted bottleneck block used by Perch."""
def __init__(self, in_channels: int, out_channels: int, kernel_size: int, stride: int, expansion: int) -> None:
super().__init__()
expanded_channels = in_channels * expansion
self.kernel_size = kernel_size
self.stride = stride
self.has_expand = expansion != 1
if self.has_expand:
self.expand_conv = nn.Conv2d(in_channels, expanded_channels, 1, bias=False)
self.expand_bn = nn.BatchNorm2d(expanded_channels)
self.depthwise_conv = nn.Conv2d(
expanded_channels,
expanded_channels,
kernel_size,
stride=stride,
groups=expanded_channels,
bias=False,
)
self.depthwise_bn = nn.BatchNorm2d(expanded_channels)
self.se = SqueezeExcitation(expanded_channels, expanded_channels // (4 * expansion))
self.project_conv = nn.Conv2d(expanded_channels, out_channels, 1, bias=False)
self.project_bn = nn.BatchNorm2d(out_channels)
def forward(self, x: Tensor) -> Tensor:
if self.has_expand:
x = F.silu(self.expand_bn(self.expand_conv(x)))
x = _same_padding(x, self.kernel_size, self.stride)
x = F.silu(self.depthwise_bn(self.depthwise_conv(x)))
return self.project_bn(self.project_conv(self.se(x)))
class ResidualMBConv(nn.Module):
"""MBConv with its eligible residual connection."""
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int,
stride: int,
expansion: int,
residual: bool,
) -> None:
super().__init__()
self.mbconv = MBConv(in_channels, out_channels, kernel_size, stride, expansion)
self.residual = residual
def forward(self, x: Tensor) -> Tensor:
result = self.mbconv(x)
return result + x if self.residual else result
class PerchBackbone(nn.Module):
"""Perch v2 feature extractor with selectable ``blocks.<index>`` layers."""
def __init__(self, input_layout: str = "frequency_time") -> None:
super().__init__()
if input_layout not in {"frequency_time", "time_frequency"}:
raise ValueError("input_layout must be 'frequency_time' or 'time_frequency'")
self.input_layout = input_layout
stem_channels = _channels(32)
self.stem_conv = nn.Conv2d(1, stem_channels, 3, stride=2, bias=False)
self.stem_bn = nn.BatchNorm2d(stem_channels)
blocks: list[nn.Module] = []
in_channels = stem_channels
for stage_blocks, stage_channels, kernel_size, stage_stride, expansion in _STAGES:
out_channels = _channels(stage_channels)
for block_index in range(_blocks(stage_blocks)):
blocks.append(
ResidualMBConv(
in_channels,
out_channels,
kernel_size,
stage_stride if block_index == 0 else 1,
expansion,
residual=block_index > 0,
)
)
in_channels = out_channels
self.blocks = nn.ModuleList(blocks)
self.head_conv = nn.Conv2d(in_channels, _channels(1280), 1, bias=False)
self.head_bn = nn.BatchNorm2d(_channels(1280))
def forward(self, x: Tensor) -> Tensor:
if self.input_layout == "time_frequency":
x = x.transpose(-1, -2).contiguous()
x = F.silu(self.stem_bn(self.stem_conv(x)))
for block in self.blocks:
x = block(x)
return F.silu(self.head_bn(self.head_conv(x)))
class ProtoPNetHead(nn.Module):
"""Perch v2 ProtoPNet classifier."""
def __init__(self) -> None:
super().__init__()
self.prototypes = nn.Parameter(torch.empty(_NUM_PERCH_CLASSES, _channels(1280), _NUM_PROTOTYPES))
self.kernel = nn.Parameter(torch.empty(_NUM_PERCH_CLASSES, _NUM_PROTOTYPES))
self.bias = nn.Parameter(torch.empty(_NUM_PERCH_CLASSES))
def forward(self, spatial_embedding: Tensor) -> Tensor:
normalized_embedding = spatial_embedding / (spatial_embedding.norm(dim=1, keepdim=True) + 1e-5)
similarities = torch.einsum("bdhw,cdp->bhwcp", normalized_embedding, self.prototypes).amax(dim=(1, 2))
return (similarities * self.kernel.clamp_min(0)).sum(dim=-1) + self.bias
class PerchModel(nn.Module):
"""Perch v2 backbone and ProtoPNet classifier."""
def __init__(self, input_layout: str = "frequency_time") -> None:
super().__init__()
self.backbone = PerchBackbone(input_layout)
self.classifier = ProtoPNetHead()
def forward(self, x: Tensor) -> Tensor:
return self.classifier(self.backbone(x))