该片文章是频域先验 + 空间拓扑 + Transformer在 EEG 情绪识别的里程碑,针对现有模型频域特征挖掘不足、Transformer 缺乏归纳偏置、跨被试泛化差三大核心痛点,提出傅里叶邻接 Transformer(FAT),在 SEED、DEAP 数据集上刷新 SOTA。

一、当前 EEG 情绪识别领域的核心问题

1.信号特性痛点

EEG 信号具有非平稳性、低信噪比、被试间差异大的天然缺陷,传统模型难以学习稳定的判别性特征。

2.特征提取缺陷

手工特征(DE/PSD)鲁棒但丢失信息;原生 CNN/RNN 无法捕捉全局依赖;标准 Transformer缺乏归纳偏置,数据低效、泛化差。

3.频域特性挖掘缺失

EEG 脑电信号具有内在周期性,但现有方法未有效解耦周期 / 非周期成分,无法利用频域先验提升性能。

4.空间拓扑建模不足

图神经网络依赖预定义通道连接,跨数据集泛化差;标准自注意力未融合 EEG 电极的物理邻接关系。

5.泛化能力不足

被试内精度高,但跨被试(LOSO)性能暴跌,无法满足实际 BCI 系统落地需求。

二、论文提出的解决方案:FAT 框架

针对上述痛点,提出傅里叶邻接 Transformer(FAT),核心是傅里叶解析线性层 + 傅里叶邻接注意力,实现频域周期特性与空间通道关联的联合建模。

1. 核心模块 1:FAL(傅里叶解析线性层)

创新解耦机制:将输入特征拆分为周期分量(cos/sin 傅里叶基建模)+非周期分量(纯线性变换,无非线性激活)。

作用:替代 Transformer 的 QKV 线性投影,显式捕捉 EEG 周期结构,同时保留非周期成分的线性映射能力,适配脑电信号频域特性。

2. 核心模块 2:FAA(傅里叶邻接注意力)

基于 FAL 将 Q/K/V 分解为周期 / 非周期两路,分别计算注意力分数。

引入可学习的周期 / 非周期通道邻接矩阵,将物理空间关联注入注意力计算,弥补标准 Transformer 无归纳偏置的缺陷。

融合数据驱动的注意力与电极拓扑先验,同时建模全局依赖 + 空间结构。

3. 整体框架

输入:EEG 差分熵(DE)频域特征。

结构:Patch 嵌入 + 位置编码→多层 FAA 替代 MHSA→前馈网络→分类头。

训练:频域数据增强 + Mixup,进一步提升泛化。

三、实验效果与核心结论

实验在SEED 系列(3/4/5/7 分类)、DEAP(效价 / 唤醒度)数据集开展,包含被试内与 ** 跨被试(LOSO)** 验证:

被试内精度

全面超越 SOTA,SEED-V 提升 6.5%、SEED-VII 提升 6.6%;DEAP 效价90.10%、唤醒度89.18%。

跨被试泛化

泛化能力显著提升,SEED-VII 上比 Conformer 高约 5%,解决 Transformer 跨被试差的痛点。

消融验证

FAL 解耦 + 邻接矩阵协同增益,缺一不可;全频带特征输入效果最优。

可解释性

周期 / 非周期邻接矩阵学到差异化的通道关联,与脑科学先验一致。

四、领域见解与创新点

(一)对 EEG 情绪识别领域的核心见解

核心突破方向

单纯堆叠深度学习结构无效,** 融合神经物理先验(频域周期、空间拓扑)** 是提升精度与泛化的唯一路径。

Transformer 适配改造

原生 Transformer 不适合 EEG,必须加入归纳偏置(通道邻接、频域分解、时序约束)。

信号处理与深度学习融合

传统信号处理(傅里叶分析)与深度学习并非对立,先验嵌入比纯数据驱动更优。

落地核心指标

跨被试(LOSO)性能远重要于被试内精度,是情感 BCI 实用化的关键。

(二)论文原生创新点

频域解耦创新

首次在 Transformer 中用 FAL显式解耦 EEG 周期 / 非周期成分,充分利用脑电频域特性。

注意力机制创新

提出 FAA,融合傅里叶分解 + 可学习通道邻接矩阵,同时建模全局依赖与电极拓扑。

框架创新

构建 FAT 端到端框架,兼顾频域、空间、时序三维信息,精度与泛化双 SOTA。

(三)未来可落地创新方向

混合架构融合

FAT + 图卷积 + CNN,结合全局注意力、空间拓扑、局部特征。

跨被试自适应

基于 FAT 的周期特征做域自适应,解决被试差异问题。

轻量化部署

模型剪枝 / 量化,适配穿戴式脑电设备实时推理。

多模态协同

用 FAT 提取 EEG 特征,融合 GSR、眼动等外周信号,构建多模态情感模型。

动态拓扑学习

邻接矩阵随情绪状态动态更新,更贴合脑功能连接的变化特性。

总结

FAT 框架开创了傅里叶频域先验 + 空间邻接拓扑 + Transformer的新范式,完美解决 EEG 情绪识别中频域挖掘不足、Transformer 无归纳偏置、跨被试泛化差三大痛点,是后续 EEG 情绪识别、脑机接口研究的重要参考。

实现代码

from torch import nn
from utils import normalize_A
from einops import rearrange
from FANLayer import FANLayer
import torch

# 定义 EEG Conformer 模型
class ModifiedPatchEmbedding2D(nn.Module):
    def __init__(self, emb_size=40, num_channels=62, num_freq_bands=5):
        super(ModifiedPatchEmbedding2D, self).__init__()
        self.emb_size = emb_size
        self.num_channels = num_channels
        self.num_freq_bands = num_freq_bands
        self.batch_norm_stage1 = nn.BatchNorm2d(emb_size // 2)
        self.batch_norm_stage2 = nn.BatchNorm2d(emb_size)

        # 位置编码
        self.position_encodings = nn.Parameter(torch.randn(1, 1, num_channels, num_freq_bands))

      

    def forward(self, x):
        # 添加位置编码
        x = x + self.position_encodings  # [B, 1, C, num_freq_bands]

        x = x.squeeze(1).permute(0, 2, 1).unsqueeze(-1)  # [B, num_freq_bands, C, 1]

       

        # 调整回 Transformer 格式
        x = x.squeeze(-1).permute(0, 2, 1)  # [B, C, emb_size]
        return x


class PositionalEncoding(nn.Module):
    def __init__(self, emb_size, dropout=0.1, max_len=100):
        super().__init__()
        self.dropout = nn.Dropout(p=dropout)

        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, emb_size, 2).float() * (-torch.log(torch.tensor(10000.0)) / emb_size))
        pe = torch.zeros(max_len, emb_size)
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        pe = pe.unsqueeze(0)  # [1, max_len, emb_size]
        self.register_buffer("pe", pe)

    def forward(self, x):
        seq_len = x.size(1)
        x = x + self.pe[:, :seq_len, :].to(x.device)  # 加位置编码
        return self.dropout(x)




    def forward(self, x, mask=None, dynamic_graph1=None, dynamic_graph2=None):
        """
        x.shape = (B, n, emb_size=40)
        前提:
          - num_heads = 8
          - 每个 head_dim = 5
          - 前 4 个 head 和后 4 个 head 分别融合到 DG1 / DG2
          - DG1.shape = (B,4,n,n), DG2.shape = (B,4,n,n)
        """
        B, n, _ = x.shape

        queries = rearrange(self.queries(x), "b n (h d) -> b h n d", h=8)
        keys = rearrange(self.keys(x), "b n (h d) -> b h n d", h=8)
        values = rearrange(self.values(x), "b n (h d) -> b h n d", h=8)

        energy = torch.einsum("bhqd, bhkd -> bhqk", queries, keys)
        if mask is not None:
            energy = energy.masked_fill(~mask, float("-inf"))

        queries_4_4 = queries.permute(0, 2, 1, 3)

        q_front4 = queries_4_4[:, :, :4, :]
        q_back4 = queries_4_4[:, :, 4:, :]

        q_front4_20 = q_front4.reshape(B, n, 20)

       
        if dynamic_graph1 is not None:
            energy[:, :4, :, :] = energy[:, :4, :, :] + w_front4 * dynamic_graph1
        if dynamic_graph2 is not None:
            energy[:, 4:, :, :] = energy[:, 4:, :, :] + w_back4 * dynamic_graph2

        scaling = (self.emb_size ** 0.5)
        att = torch.softmax(energy / scaling, dim=-1)
        out = torch.einsum("bhqk, bhkd -> bhqd", att, values)
        out = rearrange(out, "b h n d -> b n (h d)")
        out = self.projection(out)
        return out


class ResidualAdd(nn.Module):
    def __init__(self, fn):
        super().__init__()
        self.fn = fn

    def forward(self, x, **kwargs):
        if isinstance(self.fn, FAA):
            return x + self.fn(x, **kwargs)
        else:
            return x + self.fn(x)


class TransformerEncoderBlock(nn.Module):
    def __init__(self, emb_size, num_heads=4, drop_p=0.5, forward_expansion=4, use_dynamic_graph=False):
        super().__init__()
        # 注意力部分用自定义的 AttentionBlock 来替代原先的 nn.Sequential
        self.attention = AttentionBlock(emb_size, num_heads, drop_p, use_dynamic_graph)

        # FeedForward 部分可以继续用 ResidualAdd + nn.Sequential,因为它不需要额外参数
        self.feed_forward = ResidualAdd(
            nn.Sequential(
                nn.LayerNorm(emb_size),
                FeedForwardBlock(emb_size, expansion=forward_expansion, drop_p=drop_p),
                nn.Dropout(drop_p),
            )
        )

    def forward(self, x, mask=None, dynamic_graph1=None, dynamic_graph2=None):
        # 只在多头注意力这里需要传入 dynamic_graph1 / dynamic_graph2
        x = self.attention(x, mask=mask, dynamic_graph1=dynamic_graph1, dynamic_graph2=dynamic_graph2)
        x = self.feed_forward(x)  # FFN 部分无需额外参数
        return x


class TransformerEncoder(nn.Module):
    def __init__(self, depth, emb_size, num_heads=4, drop_p=0.1, num_dim=63, use_dynamic_graph=False):
        super().__init__()

        self.use_dynamic_graph = use_dynamic_graph
        self.dynamic_graph_learner1 = DynamicGraphLearner(n=num_dim, num_heads=num_heads // 2) if use_dynamic_graph else None
        self.dynamic_graph_learner2 = DynamicGraphLearner(n=num_dim, num_heads=num_heads // 2) if use_dynamic_graph else None

        self.layers = nn.ModuleList([
            TransformerEncoderBlock(emb_size, num_heads, drop_p, use_dynamic_graph=use_dynamic_graph)
            for _ in range(depth)
        ])

    def forward(self, x, mask=None):
        dynamic_graph1 = self.dynamic_graph_learner1(x) if self.use_dynamic_graph else None
        dynamic_graph2 = self.dynamic_graph_learner2(x) if self.use_dynamic_graph else None

        for layer in self.layers:
            x = layer(x, mask=mask, dynamic_graph1=dynamic_graph1, dynamic_graph2=dynamic_graph2)

        return x


class ClassificationHead(nn.Sequential):
    def __init__(self, emb_size, n_classes, drop_p=0.3):
        super().__init__(
            nn.LayerNorm(emb_size),  # Normalize features
            nn.Linear(emb_size, n_classes)  # Map to class logits
        )


class FAT(nn.Module):
    def __init__(self, emb_size=40, depth=6, n_classes=4,
                 num_channels=62, num_freq_bands=5, num_heads=8,
                 use_dynamic_graph=True):
        super(FAT, self).__init__()
        self.patch_embedding = ModifiedPatchEmbedding2D(emb_size, num_channels, num_freq_bands)
        self.cls_token = nn.Parameter(torch.randn(1, 1, emb_size))
        self.positional_encoding = PositionalEncoding(emb_size, dropout=0.1,
                                                      max_len=num_channels + 1)
        self.transformer_encoder = TransformerEncoder(
            depth, emb_size, num_heads, drop_p=0.2, num_dim=num_channels+1, use_dynamic_graph=use_dynamic_graph
        )
        self.classification_head = ClassificationHead(emb_size, n_classes)

    def forward(self, x):
        x = self.patch_embedding(x)  # [B, C, emb_size]

        B, C, E = x.shape
        cls_token = self.cls_token.expand(B, -1, -1)
        x = torch.cat((cls_token, x), dim=1)  # [B, C+1, emb_size]

        x = self.positional_encoding(x)       # [B, C+1, emb_size]
        x = self.transformer_encoder(x)       # [B, C+1, emb_size]

        cls_output = x[:, 0, :]              # [B, emb_size]
        logits = self.classification_head(cls_output)
        return logits

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐