FallKLTN / src /fall_detection /models.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw History Blame Contribute Delete
1.91 kB
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)