| """A compact causal log-mel TCN trained from scratch.""" |
|
|
| from __future__ import annotations |
|
|
| from collections.abc import Mapping, Sequence |
| from dataclasses import asdict, dataclass |
|
|
| import torch |
| import torch.nn.functional as F |
| from torch import Tensor, nn |
|
|
| from .common import TurnDetectionOutput |
| from .pooling import MaskedAttentiveStatisticsPooling |
|
|
|
|
| @dataclass(frozen=True) |
| class TinyTCNConfig: |
| n_mels: int = 80 |
| channels: int = 192 |
| num_blocks: int = 8 |
| kernel_size: int = 5 |
| dilation_cycle: Sequence[int] = (1, 2, 4, 8, 16, 32, 64, 128) |
| attention_channels: int = 128 |
| head_hidden: int = 128 |
| dropout: float = 0.1 |
| auxiliary_fillers: bool = True |
|
|
| @classmethod |
| def from_mapping(cls, values: Mapping[str, object]) -> TinyTCNConfig: |
| known = {field.name for field in cls.__dataclass_fields__.values()} |
| clean = {key: value for key, value in values.items() if key in known} |
| if "dilation_cycle" in clean: |
| clean["dilation_cycle"] = tuple(int(x) for x in clean["dilation_cycle"]) |
| return cls(**clean) |
|
|
| def to_dict(self) -> dict: |
| result = asdict(self) |
| result["dilation_cycle"] = list(self.dilation_cycle) |
| return result |
|
|
|
|
| class FrameLayerNorm(nn.Module): |
| """Layer-normalize channels independently at every frame. |
| |
| Unlike GroupNorm on a ``[B, C, T]`` tensor, this does not let the amount of |
| right padding change statistics for valid frames. |
| """ |
|
|
| def __init__(self, channels: int) -> None: |
| super().__init__() |
| self.norm = nn.LayerNorm(channels) |
|
|
| def forward(self, x: Tensor) -> Tensor: |
| return self.norm(x.transpose(1, 2)).transpose(1, 2) |
|
|
|
|
| class CausalDepthwiseSeparableBlock(nn.Module): |
| def __init__( |
| self, |
| channels: int, |
| kernel_size: int, |
| dilation: int, |
| dropout: float, |
| ) -> None: |
| super().__init__() |
| self.left_padding = dilation * (kernel_size - 1) |
| self.depthwise = nn.Conv1d( |
| channels, |
| channels, |
| kernel_size=kernel_size, |
| dilation=dilation, |
| groups=channels, |
| bias=False, |
| ) |
| self.pointwise = nn.Conv1d(channels, channels, kernel_size=1, bias=False) |
| self.norm = FrameLayerNorm(channels) |
| self.dropout = nn.Dropout(dropout) |
|
|
| def forward(self, x: Tensor, mask: Tensor | None = None) -> Tensor: |
| residual = x |
| x = F.pad(x, (self.left_padding, 0)) |
| x = self.depthwise(x) |
| x = self.pointwise(x) |
| x = self.norm(x) |
| x = F.silu(x) |
| x = self.dropout(x) |
| x = x + residual |
| if mask is not None: |
| x = x * mask.unsqueeze(1).to(x.dtype) |
| return x |
|
|
|
|
| class TinyTurnDetector(nn.Module): |
| """Approximately 0.4M parameter endpoint detector. |
| |
| The model consumes normalized log-mel features shaped ``[B, 80, T]``. It |
| is causal up to the final pooling operation, so prefixes can be cached or |
| re-evaluated at VAD checkpoints without future-frame leakage. |
| """ |
|
|
| model_type = "tiny_tcn" |
|
|
| def __init__(self, config: TinyTCNConfig | None = None) -> None: |
| super().__init__() |
| config = config or TinyTCNConfig() |
| if config.num_blocks < 1: |
| raise ValueError("num_blocks must be positive") |
| if config.kernel_size < 2: |
| raise ValueError("kernel_size must be at least 2") |
| if not config.dilation_cycle: |
| raise ValueError("dilation_cycle cannot be empty") |
| self.config = config |
| self.input_projection = nn.Sequential( |
| nn.Conv1d(config.n_mels, config.channels, kernel_size=1, bias=False), |
| FrameLayerNorm(config.channels), |
| nn.SiLU(), |
| ) |
| self.blocks = nn.ModuleList( |
| CausalDepthwiseSeparableBlock( |
| config.channels, |
| config.kernel_size, |
| int(config.dilation_cycle[index % len(config.dilation_cycle)]), |
| config.dropout, |
| ) |
| for index in range(config.num_blocks) |
| ) |
| self.pool = MaskedAttentiveStatisticsPooling(config.channels, config.attention_channels) |
| self.embedding = nn.Sequential( |
| nn.Linear(2 * config.channels, config.head_hidden), |
| nn.LayerNorm(config.head_hidden), |
| nn.SiLU(), |
| nn.Dropout(config.dropout), |
| ) |
| self.endpoint_head = nn.Linear(config.head_hidden, 1) |
| if config.auxiliary_fillers: |
| self.midfiller_head: nn.Linear | None = nn.Linear(config.head_hidden, 1) |
| self.endfiller_head: nn.Linear | None = nn.Linear(config.head_hidden, 1) |
| else: |
| self.midfiller_head = None |
| self.endfiller_head = None |
|
|
| def forward(self, log_mel: Tensor, attention_mask: Tensor | None = None) -> TurnDetectionOutput: |
| if log_mel.ndim != 3: |
| raise ValueError("log_mel must have shape [batch, n_mels, frames]") |
| if log_mel.shape[1] != self.config.n_mels: |
| raise ValueError(f"expected {self.config.n_mels} mel bins, got {log_mel.shape[1]}") |
| if attention_mask is None: |
| attention_mask = torch.ones( |
| (log_mel.shape[0], log_mel.shape[-1]), |
| dtype=torch.bool, |
| device=log_mel.device, |
| ) |
| else: |
| attention_mask = attention_mask.to(device=log_mel.device, dtype=torch.bool) |
| if attention_mask.shape != (log_mel.shape[0], log_mel.shape[-1]): |
| raise ValueError("attention_mask must have shape [batch, frames]") |
|
|
| x = log_mel * attention_mask.unsqueeze(1).to(log_mel.dtype) |
| x = self.input_projection(x) |
| x = x * attention_mask.unsqueeze(1).to(x.dtype) |
| for block in self.blocks: |
| x = block(x, attention_mask) |
| pooled = self.pool(x, attention_mask) |
| embedding = self.embedding(pooled) |
| endpoint = self.endpoint_head(embedding).squeeze(-1) |
| mid = ( |
| self.midfiller_head(embedding).squeeze(-1) if self.midfiller_head is not None else None |
| ) |
| end = ( |
| self.endfiller_head(embedding).squeeze(-1) if self.endfiller_head is not None else None |
| ) |
| return TurnDetectionOutput(endpoint, mid, end, embedding) |
|
|
| @torch.no_grad() |
| def predict_proba(self, log_mel: Tensor, attention_mask: Tensor | None = None) -> Tensor: |
| """Endpoint probability convenience method for inference only.""" |
|
|
| return torch.sigmoid(self(log_mel, attention_mask).endpoint_logits) |
|
|
| def model_config(self) -> dict: |
| return {"type": self.model_type, **self.config.to_dict()} |
|
|