from __future__ import annotations import torch from torch import nn class PoseGRU(nn.Module): """Small bidirectional GRU for binary sequence classification.""" def __init__( self, input_size: int, hidden_size: int = 64, num_layers: int = 1, dropout: float = 0.2, ) -> None: super().__init__() recurrent_dropout = dropout if num_layers > 1 else 0.0 self.gru = nn.GRU( input_size=input_size, hidden_size=hidden_size, num_layers=num_layers, dropout=recurrent_dropout, batch_first=True, bidirectional=True, ) self.classifier = nn.Sequential( nn.LayerNorm(hidden_size * 2), nn.Dropout(dropout), nn.Linear(hidden_size * 2, 1), ) def forward(self, x: torch.Tensor) -> torch.Tensor: output, _ = self.gru(x) pooled = output.mean(dim=1) return self.classifier(pooled).squeeze(1) class PoseTCN(nn.Module): """Compact temporal CNN with mean/max pooling for short pose sequences.""" def __init__(self, input_size: int, channels: int = 32, dropout: float = 0.2) -> None: super().__init__() self.temporal = nn.Sequential( nn.Conv1d(input_size, channels, kernel_size=5, padding=2), nn.BatchNorm1d(channels), nn.ReLU(), nn.Dropout(dropout), nn.Conv1d(channels, channels, kernel_size=3, padding=1), nn.BatchNorm1d(channels), nn.ReLU(), nn.Dropout(dropout), ) self.classifier = nn.Linear(channels * 2, 1) def forward(self, x: torch.Tensor) -> torch.Tensor: temporal = self.temporal(x.transpose(1, 2)) pooled = torch.cat([temporal.mean(dim=2), temporal.amax(dim=2)], dim=1) return self.classifier(pooled).squeeze(1)