dYang1's picture
Release SignBridge-TSL-STGCN with compliance Model Card
3b18ebd verified
Raw History Blame Contribute Delete
4.32 kB
# -*- coding: utf-8 -*-
"""ST-GCN model used by the word-level sign recognizer.
Input shape: (B, T, V, Ch), where V=55 contains upper-body and hand
landmarks. Spatial graph convolution follows the skeleton topology, while
temporal convolution learns how each sign evolves across frames.
"""
from __future__ import annotations
import torch
import torch.nn as nn
from graph import build_adjacency
class SpatialGraphConv(nn.Module):
"""圖卷積:每組鄰接矩陣各學一套權重,再沿著骨架邊聚合。"""
def __init__(self, in_ch, out_ch, num_subsets):
super().__init__()
self.k = num_subsets
self.conv = nn.Conv2d(in_ch, out_ch * num_subsets, kernel_size=1)
def forward(self, x, A): # x: (B, C, T, V), A: (K, V, V)
x = self.conv(x)
B, KC, T, V = x.shape
x = x.view(B, self.k, KC // self.k, T, V)
# 沿著邊把鄰居的特徵聚合過來
x = torch.einsum("bkctv,kvw->bctw", x, A)
return x.contiguous()
class STGCNBlock(nn.Module):
"""空間圖卷積 → 時間卷積 → 殘差。"""
def __init__(self, in_ch, out_ch, num_subsets, stride=1, dropout=0.3, t_kernel=9):
super().__init__()
self.gcn = SpatialGraphConv(in_ch, out_ch, num_subsets)
self.bn_g = nn.BatchNorm2d(out_ch)
self.tcn = nn.Sequential(
nn.Conv2d(out_ch, out_ch, (t_kernel, 1), (stride, 1), ((t_kernel - 1) // 2, 0)),
nn.BatchNorm2d(out_ch),
nn.Dropout(dropout),
)
if in_ch == out_ch and stride == 1:
self.residual = nn.Identity()
else:
self.residual = nn.Sequential(
nn.Conv2d(in_ch, out_ch, 1, (stride, 1)), nn.BatchNorm2d(out_ch)
)
self.relu = nn.ReLU(inplace=True)
def forward(self, x, A):
res = self.residual(x)
h = self.relu(self.bn_g(self.gcn(x, A)))
return self.relu(self.tcn(h) + res)
class SignSTGCN(nn.Module):
def __init__(self, num_nodes, in_channels, num_classes,
channels=(64, 64, 128, 128), strides=(1, 1, 2, 1),
dropout=0.3, strategy="spatial", edge_importance=True):
super().__init__()
A = build_adjacency(num_nodes, strategy)
self.register_buffer("A", torch.tensor(A, dtype=torch.float32))
k = self.A.shape[0]
self.data_bn = nn.BatchNorm1d(in_channels * num_nodes)
blocks, prev = [], in_channels
for ch, st in zip(channels, strides):
blocks.append(STGCNBlock(prev, ch, k, stride=st, dropout=dropout))
prev = ch
self.blocks = nn.ModuleList(blocks)
# 讓模型自己學「哪幾條骨架邊比較重要」,ST-GCN 原論文的做法
if edge_importance:
self.edge_weight = nn.ParameterList(
[nn.Parameter(torch.ones_like(self.A)) for _ in blocks]
)
else:
self.edge_weight = [1.0] * len(blocks)
self.head = nn.Sequential(nn.Dropout(dropout), nn.Linear(prev, num_classes))
def forward(self, x): # x: (B, T, V, Ch)
B, T, V, Ch = x.shape
h = x.permute(0, 3, 1, 2).contiguous() # (B, Ch, T, V)
h = self.data_bn(h.permute(0, 1, 3, 2).reshape(B, Ch * V, T))
h = h.view(B, Ch, V, T).permute(0, 1, 3, 2).contiguous()
for block, w in zip(self.blocks, self.edge_weight):
h = block(h, self.A * w)
h = h.mean(dim=(2, 3)) # 時間與節點都做平均池化
return self.head(h)
def build_from_checkpoint(ckpt: dict):
"""給即時推論用:照 checkpoint 裡存的設定還原模型。"""
import types
if ckpt.get("arch") != "stgcn":
raise ValueError("This release only supports ST-GCN checkpoints")
cfg = types.SimpleNamespace(**ckpt["model_cfg"])
return SignSTGCN(
ckpt["num_nodes"],
ckpt["in_channels"],
len(ckpt["classes"]),
channels=cfg.STGCN_CHANNELS,
strides=cfg.STGCN_STRIDES,
dropout=cfg.DROPOUT,
strategy=cfg.GRAPH_STRATEGY,
edge_importance=cfg.STGCN_EDGE_IMPORTANCE,
)