"""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, }