# sca_core/tcn_attention_net.py """ Sovereign TCNAttentionSCA - Advanced Neural Side-Channel Analysis Architecture. Combines Dilated Temporal Convolutional Networks with Multi-Head Self-Attention, Multi-Channel Waveform Support, Temporal Point-of-Interest (POI) Localization, and Calibrated Cryptanalytic Posterior Estimation. """ import math import os import warnings import torch import torch.nn as nn import torch.nn.functional as F from typing import Optional, List, Tuple, Union, Dict, Any class Chomp1d(nn.Module): """Maintains sequence length in causal convolutions by trimming future padding.""" def __init__(self, chomp_size: int): super(Chomp1d, self).__init__() self.chomp = chomp_size def forward(self, x: torch.Tensor) -> torch.Tensor: if self.chomp == 0: return x return x[:, :, :-self.chomp].contiguous() class TemporalBlock(nn.Module): """ Dilated Residual Temporal Convolutional Block. Supports both Causal (streaming) and Symmetric (offline forensic) receptive fields. """ def __init__( self, n_inputs: int, n_outputs: int, kernel_size: int, stride: int, dilation: int, padding: int, dropout: float = 0.2, causal: bool = True ): super(TemporalBlock, self).__init__() self.causal = causal self.conv1 = nn.Conv1d(n_inputs, n_outputs, kernel_size, stride=stride, padding=padding, dilation=dilation) self.chomp1 = Chomp1d(padding) if causal else nn.Identity() self.relu1 = nn.ReLU() self.dropout1 = nn.Dropout(dropout) self.conv2 = nn.Conv1d(n_outputs, n_outputs, kernel_size, stride=stride, padding=padding, dilation=dilation) self.chomp2 = Chomp1d(padding) if causal else nn.Identity() self.relu2 = nn.ReLU() self.dropout2 = nn.Dropout(dropout) self.net = nn.Sequential( self.conv1, self.chomp1, self.relu1, self.dropout1, self.conv2, self.chomp2, self.relu2, self.dropout2 ) self.downsample = nn.Conv1d(n_inputs, n_outputs, 1) if n_inputs != n_outputs else None self.relu = nn.ReLU() self.init_weights() def init_weights(self): self.conv1.weight.data.normal_(0, 0.01) self.conv2.weight.data.normal_(0, 0.01) if self.downsample is not None: self.downsample.weight.data.normal_(0, 0.01) def forward(self, x: torch.Tensor) -> torch.Tensor: out = self.net(x) res = x if self.downsample is None else self.downsample(x) return self.relu(out + res) class TCNAttentionSCA(nn.Module): """ Production-grade Sovereign Neural Side-Channel Analyzer (SOTA Enhanced). Key Innovations & Enhancements: 1. Multi-Channel Waveform Support (Power + Clock + PMU + EM sensors). 2. Multi-Head Self-Attention for Point-of-Interest (POI) Leakage Extraction. 3. Bidirectional/Symmetric or Causal Temporal Dilation Modes. 4. Optional Residual Attention & Layer Normalization. 5. POI Attention Heatmap Extraction for Explainable Cryptanalysis. 6. Temperature-Calibrated Posterior & Entropy Estimation. 7. 100% Backwards-Compatible with Legacy Checkpoints. """ def __init__( self, trace_length: int, num_classes: int = 256, in_channels: int = 1, num_channels: Optional[List[int]] = None, kernel_size: int = 5, num_heads: int = 4, dropout: float = 0.2, causal: bool = True, residual_attention: bool = False, use_layer_norm: bool = False ): super(TCNAttentionSCA, self).__init__() self.trace_length = trace_length self.num_classes = num_classes self.in_channels = in_channels self.causal = causal self.residual_attention = residual_attention self.use_layer_norm = use_layer_norm if num_channels is None: num_channels = [64, 128, 64] self.num_channels = num_channels layers = [] num_levels = len(num_channels) for i in range(num_levels): dilation_size = 2 ** i in_ch = in_channels if i == 0 else num_channels[i-1] out_ch = num_channels[i] if causal: pad = (kernel_size - 1) * dilation_size else: pad = ((kernel_size - 1) * dilation_size) // 2 layers += [TemporalBlock( n_inputs=in_ch, n_outputs=out_ch, kernel_size=kernel_size, stride=1, dilation=dilation_size, padding=pad, dropout=dropout, causal=causal )] self.network = nn.Sequential(*layers) # Multi-Head Self Attention over the temporal dimension self.attention = nn.MultiheadAttention( embed_dim=num_channels[-1], num_heads=num_heads, batch_first=True ) # Optional LayerNorm (only instantiated if requested to preserve legacy state_dict) if use_layer_norm: self.norm = nn.LayerNorm(num_channels[-1]) else: self.norm = None # Linear Classifier Head self.fc = nn.Linear(num_channels[-1], num_classes) def forward( self, x: torch.Tensor, return_attention: bool = False, return_features: bool = False ) -> Union[torch.Tensor, Tuple[torch.Tensor, ...]]: """ Forward propagation. Args: x: Input tensor. Accepts: - 2D: (batch_size, trace_length) - 3D: (batch_size, in_channels, trace_length) return_attention: If True, returns attention weights matrix. return_features: If True, returns pooled temporal embeddings. Returns: Logits (batch_size, num_classes) by default, or tuple if extra returns requested. """ # Format input to (batch_size, in_channels, trace_length) if not torch.jit.is_tracing() and not torch.jit.is_scripting(): if x.dim() == 2: x = x.unsqueeze(1) elif x.dim() == 3 and x.shape[1] != self.in_channels and x.shape[2] == self.in_channels: x = x.permute(0, 2, 1) else: if x.dim() == 2: x = x.unsqueeze(1) # 1. TCN Hierarchical Feature Extraction: (batch_size, channels, seq_len) tcn_out = self.network(x) # 2. Permute for Temporal Multi-Head Attention: (batch_size, seq_len, channels) tcn_perm = tcn_out.permute(0, 2, 1) # 3. Multi-Head Self-Attention attn_out, attn_weights = self.attention( tcn_perm, tcn_perm, tcn_perm, need_weights=return_attention, average_attn_weights=True ) # 4. Optional Residual Connection and LayerNorm if self.residual_attention: attended = tcn_perm + attn_out else: attended = attn_out if self.norm is not None: attended = self.norm(attended) # 5. Global Temporal Pooling pooled = torch.mean(attended, dim=1) # 6. Classification Logits logits = self.fc(pooled) if return_attention and return_features: return logits, pooled, attn_weights elif return_attention: return logits, attn_weights elif return_features: return logits, pooled return logits @torch.no_grad() def predict_posteriors(self, x: torch.Tensor, temperature: float = 1.0) -> Dict[str, Any]: """ Computes calibrated Bayesian class posteriors, top candidate, and Shannon entropy. """ self.eval() logits = self.forward(x) scaled_logits = logits / max(temperature, 1e-4) probs = F.softmax(scaled_logits, dim=-1) top1_idx = torch.argmax(probs, dim=-1) top1_prob = torch.gather(probs, 1, top1_idx.unsqueeze(-1)).squeeze(-1) # Shannon Entropy H(P) = -sum(p * log2(p)) entropy = -torch.sum(probs * torch.log2(probs + 1e-12), dim=-1) return { "probabilities": probs.cpu().numpy(), "predicted_classes": top1_idx.cpu().numpy(), "confidence": top1_prob.cpu().numpy(), "entropy_bits": entropy.cpu().numpy() } @torch.no_grad() def extract_poi(self, x: torch.Tensor, top_k: int = 10) -> Dict[str, Any]: """ Extracts Points-of-Interest (POI) temporal leakage heatmap from the attention matrix. Returns top-k sample indices with highest attention energy. """ self.eval() _, attn_weights = self.forward(x, return_attention=True) # attn_weights: (batch_size, seq_len, seq_len) temporal_energy = torch.mean(attn_weights, dim=1) # (batch_size, seq_len) mean_poi_curve = torch.mean(temporal_energy, dim=0).cpu().numpy() top_indices = torch.topk(torch.tensor(mean_poi_curve), k=min(top_k, len(mean_poi_curve))).indices.tolist() return { "mean_temporal_energy": mean_poi_curve.tolist(), "top_poi_indices": top_indices, "peak_leakage_cycle": int(top_indices[0]) if top_indices else 0 } def export_torchscript(self, filepath: str, example_input: Optional[torch.Tensor] = None, check_trace: bool = False): """Exports the model to standalone TorchScript format.""" self.eval() if example_input is None: example_input = torch.randn(1, self.in_channels, self.trace_length) os.makedirs(os.path.dirname(filepath), exist_ok=True) with warnings.catch_warnings(): warnings.simplefilter("ignore", category=FutureWarning) try: warnings.simplefilter("ignore", category=torch.jit.TracerWarning) except Exception: pass traced = torch.jit.trace(self, example_input, check_trace=check_trace) traced.save(filepath) return filepath def export_onnx(self, filepath: str, example_input: Optional[torch.Tensor] = None): """Exports the model to ONNX format for hardware accelerator compilation.""" self.eval() if example_input is None: example_input = torch.randn(1, self.in_channels, self.trace_length) os.makedirs(os.path.dirname(filepath), exist_ok=True) with warnings.catch_warnings(): warnings.simplefilter("ignore") torch.onnx.export( self, example_input, filepath, input_names=["trace_waveform"], output_names=["logits"], dynamic_axes={"trace_waveform": {0: "batch_size"}, "logits": {0: "batch_size"}}, opset_version=18 ) return filepath