AnyMo / modeling_components.py
Breezelled's picture
Release AnyMo model and raw-IMU pipeline
5261696 verified
Raw History Blame Contribute Delete
14.7 kB
"""Frozen AnyMo IMU encoder and tokenizer architectures used at inference."""
from __future__ import annotations
import math
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
SEGMENT_NAMES = (
"Pelvis", "L5", "L3", "T12", "T8", "Neck", "Head",
"R_Shoulder", "R_UpperArm", "R_Forearm", "R_Hand",
"L_Shoulder", "L_UpperArm", "L_Forearm", "L_Hand",
"R_UpperLeg", "R_LowerLeg", "R_Foot", "R_Toe",
"L_UpperLeg", "L_LowerLeg", "L_Foot", "L_Toe",
)
def _hop_distance(num_node: int, edges: list[tuple[int, int]], max_hop: int = 1) -> np.ndarray:
adjacency = np.zeros((num_node, num_node), dtype=np.float32)
for i, j in edges:
adjacency[j, i] = 1.0
adjacency[i, j] = 1.0
distance = np.full((num_node, num_node), np.inf, dtype=np.float32)
arrivals = np.stack([np.linalg.matrix_power(adjacency, d) for d in range(max_hop + 1)]) > 0
for d in range(max_hop, -1, -1):
distance[arrivals[d]] = d
return distance
def _normalize_digraph(adjacency: np.ndarray) -> np.ndarray:
degree = np.sum(adjacency, axis=0)
inverse = np.zeros_like(adjacency)
for index in range(adjacency.shape[0]):
if degree[index] > 0:
inverse[index, index] = degree[index] ** -1
return np.dot(adjacency, inverse)
class Graph:
def __init__(self, strategy: str = "spatial", max_hop: int = 1, dilation: int = 1):
self.max_hop = int(max_hop)
self.dilation = int(dilation)
self.num_node = 23
self.center = 0
self_link = [(i, i) for i in range(self.num_node)]
neighbor_link = [
(1, 0), (2, 1), (3, 2), (4, 3), (5, 4), (6, 5),
(7, 4), (8, 7), (9, 8), (10, 9),
(11, 4), (12, 11), (13, 12), (14, 13),
(15, 0), (16, 15), (17, 16), (18, 17),
(19, 0), (20, 19), (21, 20), (22, 21),
]
self.edge = self_link + neighbor_link
self.hop_dis = _hop_distance(self.num_node, self.edge, max_hop=self.max_hop)
self.A = self._build_adjacency(strategy)
def _build_adjacency(self, strategy: str) -> np.ndarray:
valid_hops = range(0, self.max_hop + 1, self.dilation)
adjacency = np.zeros((self.num_node, self.num_node), dtype=np.float32)
for hop in valid_hops:
adjacency[self.hop_dis == hop] = 1.0
normalized = _normalize_digraph(adjacency)
if strategy == "uniform":
return normalized[None]
if strategy == "distance":
parts = np.zeros((len(valid_hops), self.num_node, self.num_node), dtype=np.float32)
for index, hop in enumerate(valid_hops):
parts[index][self.hop_dis == hop] = normalized[self.hop_dis == hop]
return parts
parts: list[np.ndarray] = []
for hop in valid_hops:
root = np.zeros_like(normalized)
close = np.zeros_like(normalized)
further = np.zeros_like(normalized)
for i in range(self.num_node):
for j in range(self.num_node):
if self.hop_dis[j, i] != hop:
continue
if self.hop_dis[j, self.center] == self.hop_dis[i, self.center]:
root[j, i] = normalized[j, i]
elif self.hop_dis[j, self.center] > self.hop_dis[i, self.center]:
close[j, i] = normalized[j, i]
else:
further[j, i] = normalized[j, i]
parts.append(root if hop == 0 else root + close)
if hop != 0:
parts.append(further)
return np.stack(parts)
class ConvTemporalGraphical(nn.Module):
def __init__(self, in_channels: int, out_channels: int, kernel_size: int):
super().__init__()
self.kernel_size = int(kernel_size)
self.conv = nn.Conv2d(in_channels, out_channels * kernel_size, kernel_size=(1, 1))
def forward(self, x: torch.Tensor, adjacency: torch.Tensor):
x = self.conv(x)
n, kc, t, v = x.size()
x = x.view(n, self.kernel_size, kc // self.kernel_size, t, v)
return torch.einsum("nkctv,kvw->nctw", x, adjacency).contiguous(), adjacency
class STGCNBlock(nn.Module):
def __init__(self, in_channels: int, out_channels: int, kernel_size: tuple[int, int], stride: int = 1,
dropout: float = 0.0, residual: bool = True):
super().__init__()
padding = ((kernel_size[0] - 1) // 2, 0)
self.gcn = ConvTemporalGraphical(in_channels, out_channels, kernel_size[1])
self.tcn = nn.Sequential(
nn.BatchNorm2d(out_channels), nn.ReLU(inplace=True),
nn.Conv2d(out_channels, out_channels, (kernel_size[0], 1), (stride, 1), padding),
nn.BatchNorm2d(out_channels), nn.Dropout(dropout, inplace=True),
)
if not residual:
self.residual = lambda value: 0
elif in_channels == out_channels and stride == 1:
self.residual = lambda value: value
else:
self.residual = nn.Sequential(
nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=(stride, 1)),
nn.BatchNorm2d(out_channels),
)
self.relu = nn.ReLU(inplace=True)
def forward(self, x: torch.Tensor, adjacency: torch.Tensor):
residual = self.residual(x)
x, adjacency = self.gcn(x, adjacency)
return self.relu(self.tcn(x) + residual), adjacency
class STGCNEncoder(nn.Module):
def __init__(self, in_channels: int = 6, latent_dim: int = 256, dropout: float = 0.0):
super().__init__()
self.graph = Graph()
adjacency = torch.tensor(self.graph.A, dtype=torch.float32, requires_grad=False)
self.register_buffer("A", adjacency)
kernel_size = (9, adjacency.size(0))
self.in_channels = int(in_channels)
self.latent_dim = int(latent_dim)
self.data_bn = nn.BatchNorm1d(in_channels * self.graph.num_node)
self.mask_token = nn.Parameter(torch.zeros(in_channels))
self.st_gcn_networks = nn.ModuleList((
STGCNBlock(in_channels, 64, kernel_size, 1, residual=False),
STGCNBlock(64, 64, kernel_size, 1, dropout=dropout),
STGCNBlock(64, 64, kernel_size, 1, dropout=dropout),
STGCNBlock(64, 64, kernel_size, 1, dropout=dropout),
STGCNBlock(64, 128, kernel_size, 2, dropout=dropout),
STGCNBlock(128, 128, kernel_size, 1, dropout=dropout),
STGCNBlock(128, 128, kernel_size, 1, dropout=dropout),
STGCNBlock(128, 256, kernel_size, 2, dropout=dropout),
STGCNBlock(256, 256, kernel_size, 1, dropout=dropout),
STGCNBlock(256, latent_dim, kernel_size, 1, dropout=dropout),
))
self.edge_importance = nn.ParameterList(
[nn.Parameter(torch.ones(self.A.size())) for _ in self.st_gcn_networks]
)
def forward(self, x: torch.Tensor, visible_node_mask: torch.Tensor | None = None,
return_debug: bool = False) -> dict[str, torch.Tensor]:
n, c, t, v, m = x.size()
normalized = x.permute(0, 4, 3, 1, 2).contiguous().view(n * m, v * c, t)
normalized = self.data_bn(normalized).view(n, m, v, c, t)
normalized = normalized.permute(0, 1, 3, 4, 2).contiguous().mean(dim=1)
masked = normalized
if visible_node_mask is not None:
mask = visible_node_mask[:, None, None, :]
token = self.mask_token.view(1, self.in_channels, 1, 1)
masked = torch.where(mask, normalized, token)
out = masked
for gcn, importance in zip(self.st_gcn_networks, self.edge_importance):
out, _ = gcn(out, self.A * importance)
node_seq_latent = out.permute(0, 2, 3, 1).contiguous()
result = {
"node_seq_latent": node_seq_latent,
"global_seq_latent": node_seq_latent.mean(dim=2),
}
if return_debug:
result.update(normalized_input=normalized, masked_input=masked)
return result
class TemporalConvBlock(nn.Module):
def __init__(self, channels: int, kernel_size: int = 3, dilation: int = 1):
super().__init__()
padding = ((kernel_size - 1) // 2) * dilation
self.norm1 = nn.GroupNorm(1, channels)
self.conv1 = nn.Conv1d(channels, channels, kernel_size, padding=padding, dilation=dilation)
self.norm2 = nn.GroupNorm(1, channels)
self.conv2 = nn.Conv1d(channels, channels, kernel_size, padding=padding, dilation=dilation)
self.act = nn.GELU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
residual = x
x = self.conv1(self.act(self.norm1(x)))
x = self.conv2(self.act(self.norm2(x)))
return x + residual
class TemporalConvDecoder(nn.Module):
def __init__(self, input_dim: int, output_dim: int, hidden_dim: int = 256,
kernel_size: int = 3, dilations: tuple[int, ...] = (1, 2, 4)):
super().__init__()
self.in_norm = nn.LayerNorm(input_dim)
self.in_proj = nn.Conv1d(input_dim, hidden_dim, kernel_size=1)
self.blocks = nn.ModuleList(
[TemporalConvBlock(hidden_dim, kernel_size, dilation) for dilation in dilations]
)
self.out_norm = nn.GroupNorm(1, hidden_dim)
self.out_proj = nn.Conv1d(hidden_dim, output_dim, kernel_size=1)
self.act = nn.GELU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.in_proj(self.in_norm(x).transpose(1, 2))
for block in self.blocks:
x = block(x)
return self.out_proj(self.act(self.out_norm(x))).transpose(1, 2).contiguous()
def interleave_codes(codes: torch.Tensor) -> torch.Tensor:
if codes.ndim != 3:
raise ValueError(f"Expected codes [B, L, N], got {tuple(codes.shape)}")
return codes.reshape(codes.shape[0], codes.shape[1] * codes.shape[2])
class EMAProductQuantizer(nn.Module):
def __init__(self, num_codebooks: int, codebook_size: int, codebook_dim: int,
decay: float = 0.99, epsilon: float = 1e-5, dead_code_threshold_ratio: float = 0.2):
super().__init__()
self.num_codebooks = int(num_codebooks)
self.codebook_size = int(codebook_size)
self.codebook_dim = int(codebook_dim)
self.decay = float(decay)
self.epsilon = float(epsilon)
self.dead_code_threshold_ratio = float(dead_code_threshold_ratio)
embedding = torch.randn(self.num_codebooks, self.codebook_size, self.codebook_dim) / math.sqrt(self.codebook_dim)
self.register_buffer("embedding", embedding)
self.register_buffer("cluster_size", torch.zeros(self.num_codebooks, self.codebook_size))
self.register_buffer("embed_avg", embedding.clone())
self.register_buffer("last_dead_code_replacements", torch.zeros(self.num_codebooks, dtype=torch.long))
self.register_buffer("total_dead_code_replacements", torch.zeros(self.num_codebooks, dtype=torch.long))
def forward(self, x: torch.Tensor):
if x.ndim != 3:
raise ValueError(f"Expected x [B, L, C], got {tuple(x.shape)}")
b, length, _ = x.shape
flat = x.reshape(b * length, self.num_codebooks, self.codebook_dim).permute(1, 0, 2).contiguous()
quantized_chunks, code_chunks, commitment = [], [], []
for index in range(self.num_codebooks):
embedding = self.embedding[index]
chunk = flat[index]
distances = chunk.pow(2).sum(1, keepdim=True) - 2 * chunk @ embedding.t() + embedding.pow(2).sum(1)[None]
ids = distances.argmin(dim=1)
quantized = embedding.index_select(0, ids)
quantized_chunks.append(quantized)
code_chunks.append(ids)
commitment.append(F.mse_loss(chunk, quantized.detach()))
quantized_flat = torch.stack(quantized_chunks)
quantized_st = flat + (quantized_flat - flat).detach()
quantized = quantized_st.permute(1, 0, 2).reshape(b, length, -1)
codes = torch.stack(code_chunks).permute(1, 0).reshape(b, length, self.num_codebooks)
return quantized, codes, torch.stack(commitment).sum()
def decode_codes(self, codes: torch.Tensor) -> torch.Tensor:
chunks = []
for index in range(self.num_codebooks):
chunks.append(self.embedding[index].index_select(0, codes[:, :, index].reshape(-1)))
return torch.stack(chunks, dim=1).reshape(codes.shape[0], codes.shape[1], -1)
class MotionPQVAE(nn.Module):
def __init__(self, input_dim: int = 256, bottleneck_dim: int = 128, num_codebooks: int = 2,
codebook_size: int = 2048, codebook_dim: int = 64, commitment_weight: float = 0.25,
ema_decay: float = 0.99, ema_epsilon: float = 1e-5,
dead_code_threshold_ratio: float = 0.2):
super().__init__()
if bottleneck_dim != num_codebooks * codebook_dim:
raise ValueError("bottleneck_dim must equal num_codebooks * codebook_dim")
self.input_dim = int(input_dim)
self.bottleneck_dim = int(bottleneck_dim)
self.num_codebooks = int(num_codebooks)
self.codebook_size = int(codebook_size)
self.codebook_dim = int(codebook_dim)
self.commitment_weight = float(commitment_weight)
self.pre_quant = nn.Sequential(nn.LayerNorm(input_dim), nn.Linear(input_dim, bottleneck_dim))
self.quantizer = EMAProductQuantizer(
num_codebooks, codebook_size, codebook_dim, ema_decay, ema_epsilon, dead_code_threshold_ratio
)
self.decoder = TemporalConvDecoder(bottleneck_dim, input_dim, max(bottleneck_dim, input_dim))
def encode_codes(self, x: torch.Tensor) -> torch.Tensor:
return self.quantizer(self.pre_quant(x))[1]
def encode_tokens(self, x: torch.Tensor) -> torch.Tensor:
return interleave_codes(self.encode_codes(x))
def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
pre_quant = self.pre_quant(x)
quantized, codes, commitment_loss = self.quantizer(pre_quant)
reconstruction = self.decoder(quantized)
recon_loss = F.smooth_l1_loss(reconstruction, x)
return {
"pre_quant": pre_quant,
"quantized": quantized,
"codes": codes,
"tokens": interleave_codes(codes),
"reconstruction": reconstruction,
"recon_loss": recon_loss,
"commitment_loss": commitment_loss,
"dead_code_replacements": self.quantizer.last_dead_code_replacements.clone(),
"loss": recon_loss + self.commitment_weight * commitment_loss,
}