File size: 1,906 Bytes
9313a90 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 | 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)
|