"""TSPulse: a small time-series pulse Transformer autoencoder. Encodes an arbitrary cycle waveform (sequence of [signal, second, 1, 1] channels) into a 128-dim fingerprint vector. Trained with masked reconstruction: 15% of the patches are replaced by a learnable [MASK] token and the decoder head must recover the original signal, forcing the encoder to build a compact representation of the waveform. """ from __future__ import annotations import torch from torch import nn from .dataset import N_CHANNELS, PATCH_SIZE, SEQ_LEN NUM_PATCHES = SEQ_LEN // PATCH_SIZE class TSPulse(nn.Module): def __init__( self, dim: int = 128, depth: int = 4, heads: int = 4, mlp_ratio: float = 4.0, patch_size: int = PATCH_SIZE, in_channels: int = N_CHANNELS, seq_len: int = SEQ_LEN, ): super().__init__() self.dim = dim self.patch_size = patch_size self.in_channels = in_channels self.seq_len = seq_len self.num_patches = seq_len // patch_size self.patch_embed = nn.Conv1d(in_channels, dim, kernel_size=patch_size, stride=patch_size) self.mask_token = nn.Parameter(torch.zeros(1, 1, dim)) self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, dim)) layer = nn.TransformerEncoderLayer( d_model=dim, nhead=heads, dim_feedforward=int(dim * mlp_ratio), activation="gelu", batch_first=True, norm_first=True, ) self.encoder = nn.TransformerEncoder(layer, num_layers=depth) self.norm = nn.LayerNorm(dim) self.head = nn.Linear(dim, patch_size * in_channels) self._reset_parameters() def _reset_parameters(self) -> None: nn.init.trunc_normal_(self.pos_embed, std=0.02) nn.init.trunc_normal_(self.mask_token, std=0.02) nn.init.xavier_uniform_(self.head.weight) nn.init.zeros_(self.head.bias) def _embed(self, x: torch.Tensor) -> torch.Tensor: # x: (B, L, C) -> (B, C, L) -> conv -> (B, dim, P) -> (B, P, dim) return self.patch_embed(x.transpose(1, 2)).transpose(1, 2) def forward( self, x: torch.Tensor, mask: torch.BoolTensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """Return (fingerprint, reconstruction). fingerprint: (B, dim) reconstruction: (B, L, C) """ tokens = self._embed(x) + self.pos_embed if mask is not None: tokens = tokens.masked_fill(mask.unsqueeze(-1), 0.0) tokens = tokens + self.mask_token.masked_fill(~mask.unsqueeze(-1), 0.0) encoded = self.encoder(tokens) encoded = self.norm(encoded) fingerprint = encoded.mean(dim=1) recon = self.head(encoded) # (B, P, patch*C) recon = recon.reshape(-1, self.num_patches, self.patch_size, self.in_channels) recon = recon.permute(0, 2, 1, 3).reshape(-1, self.seq_len, self.in_channels) return fingerprint, recon def encode_cycle(self, x: torch.Tensor) -> torch.Tensor: """Encode a single cycle (L, C) -> 128-dim fingerprint (no masking).""" self.eval() with torch.no_grad(): fingerprint, _ = self.forward(x.unsqueeze(0)) return fingerprint.squeeze(0) def make_random_mask( batch_size: int, num_patches: int, ratio: float, device: torch.device, ) -> torch.BoolTensor: num_masked = max(1, int(num_patches * ratio)) mask = torch.zeros(batch_size, num_patches, dtype=torch.bool, device=device) for row in range(batch_size): indices = torch.randperm(num_patches, device=device)[:num_masked] mask[row, indices] = True return mask