TCNAttentionSCA-GA102-SCA / tcn_attention_net.py
bbkdevops's picture
Add model artifact tcn_attention_net.py
3a2687e verified
Raw History Blame Contribute Delete
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
@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