model.py 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107
  1. """TSPulse: a small time-series pulse Transformer autoencoder.
  2. Encodes an arbitrary cycle waveform (sequence of [signal, second, 1, 1]
  3. channels) into a 128-dim fingerprint vector. Trained with masked
  4. reconstruction: 15% of the patches are replaced by a learnable [MASK] token
  5. and the decoder head must recover the original signal, forcing the encoder to
  6. build a compact representation of the waveform.
  7. """
  8. from __future__ import annotations
  9. import torch
  10. from torch import nn
  11. from .dataset import N_CHANNELS, PATCH_SIZE, SEQ_LEN
  12. NUM_PATCHES = SEQ_LEN // PATCH_SIZE
  13. class TSPulse(nn.Module):
  14. def __init__(
  15. self,
  16. dim: int = 128,
  17. depth: int = 4,
  18. heads: int = 4,
  19. mlp_ratio: float = 4.0,
  20. patch_size: int = PATCH_SIZE,
  21. in_channels: int = N_CHANNELS,
  22. seq_len: int = SEQ_LEN,
  23. ):
  24. super().__init__()
  25. self.dim = dim
  26. self.patch_size = patch_size
  27. self.in_channels = in_channels
  28. self.seq_len = seq_len
  29. self.num_patches = seq_len // patch_size
  30. self.patch_embed = nn.Conv1d(in_channels, dim, kernel_size=patch_size, stride=patch_size)
  31. self.mask_token = nn.Parameter(torch.zeros(1, 1, dim))
  32. self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, dim))
  33. layer = nn.TransformerEncoderLayer(
  34. d_model=dim,
  35. nhead=heads,
  36. dim_feedforward=int(dim * mlp_ratio),
  37. activation="gelu",
  38. batch_first=True,
  39. norm_first=True,
  40. )
  41. self.encoder = nn.TransformerEncoder(layer, num_layers=depth)
  42. self.norm = nn.LayerNorm(dim)
  43. self.head = nn.Linear(dim, patch_size * in_channels)
  44. self._reset_parameters()
  45. def _reset_parameters(self) -> None:
  46. nn.init.trunc_normal_(self.pos_embed, std=0.02)
  47. nn.init.trunc_normal_(self.mask_token, std=0.02)
  48. nn.init.xavier_uniform_(self.head.weight)
  49. nn.init.zeros_(self.head.bias)
  50. def _embed(self, x: torch.Tensor) -> torch.Tensor:
  51. # x: (B, L, C) -> (B, C, L) -> conv -> (B, dim, P) -> (B, P, dim)
  52. return self.patch_embed(x.transpose(1, 2)).transpose(1, 2)
  53. def forward(
  54. self,
  55. x: torch.Tensor,
  56. mask: torch.BoolTensor | None = None,
  57. ) -> tuple[torch.Tensor, torch.Tensor]:
  58. """Return (fingerprint, reconstruction).
  59. fingerprint: (B, dim)
  60. reconstruction: (B, L, C)
  61. """
  62. tokens = self._embed(x) + self.pos_embed
  63. if mask is not None:
  64. tokens = tokens.masked_fill(mask.unsqueeze(-1), 0.0)
  65. tokens = tokens + self.mask_token.masked_fill(~mask.unsqueeze(-1), 0.0)
  66. encoded = self.encoder(tokens)
  67. encoded = self.norm(encoded)
  68. fingerprint = encoded.mean(dim=1)
  69. recon = self.head(encoded) # (B, P, patch*C)
  70. recon = recon.reshape(-1, self.num_patches, self.patch_size, self.in_channels)
  71. recon = recon.permute(0, 2, 1, 3).reshape(-1, self.seq_len, self.in_channels)
  72. return fingerprint, recon
  73. def encode_cycle(self, x: torch.Tensor) -> torch.Tensor:
  74. """Encode a single cycle (L, C) -> 128-dim fingerprint (no masking)."""
  75. self.eval()
  76. with torch.no_grad():
  77. fingerprint, _ = self.forward(x.unsqueeze(0))
  78. return fingerprint.squeeze(0)
  79. def make_random_mask(
  80. batch_size: int,
  81. num_patches: int,
  82. ratio: float,
  83. device: torch.device,
  84. ) -> torch.BoolTensor:
  85. num_masked = max(1, int(num_patches * ratio))
  86. mask = torch.zeros(batch_size, num_patches, dtype=torch.bool, device=device)
  87. for row in range(batch_size):
  88. indices = torch.randperm(num_patches, device=device)[:num_masked]
  89. mask[row, indices] = True
  90. return mask