| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107 |
- """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
|