Download tcn_attention_net.py from bbkdevops/TCNAttentionSCA-GA102-SCA: direct link, hf CLI and curl.
- Browser
- Download file 11 kB
-
https://huggingface.co/bbkdevops/TCNAttentionSCA-GA102-SCA/resolve/main/tcn_attention_net.py
- Command line
-
hf download hf://bbkdevops/TCNAttentionSCA-GA102-SCA/tcn_attention_net.py
-
curl -L -o tcn_attention_net.py https://huggingface.co/bbkdevops/TCNAttentionSCA-GA102-SCA/resolve/main/tcn_attention_net.py
11 kB
| # 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 | |
| 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() | |
| } | |
| 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 | |