用螺旋数重写 Transformer Attention:让大模型自带“相位记忆“的 PyTorch 实现
发布时间:2026/9/28 18:05:50来源:尧图网络
摘要标准 Transformer 的 Softmax Attention 本质是无尺度纯旋转——每个 token 等权参与注意力长序列时信息被稀释推理链缺乏几何约束。本文基于螺旋生成论的I² -N给出一种Spiral Attention螺旋注意力 的 PyTorch 实现把 Query/Key/Value 从复数域扩展到螺旋数域让大模型天然具备相位连续性和尺度记忆。附完整可运行代码。一、标准 Attention 的先天缺陷先回顾 Scaled Dot-Product Attention$Attention(Q, K, V) softmax\left(\frac{QK^T}{\sqrt{d_k}}\right)V$拆开看它的问题问题数学本质长序列信息稀释Softmax 归一化强制所有 token 的注意力权重和为 1远处 token 权重趋零位置信息依赖外挂必须额外加 Positional Encoding / RoPE否则模型不知道 token 顺序推理链无方向约束每个 head 独立旋转无相位连续性概念无置信度传播中间层的注意力权重不携带这一步有多确定的信息无法自然处理非平稳数据对频率变化的信号chirp、语音、金融时序建模能力弱根因QK^T产生的是实数相似度丢失了相位和尺度两个维度。螺旋生成论说如果让 Q/K/V 在螺旋数域运算Attention 就同时携带相位方向/语义转向尺度置信度/重要性衰减二、螺旋注意力从i² -1到I² -N2.1 螺旋数的 PyTorch 表示螺旋数z a I·b其中I² -N可以映射为二维实向量$z \leftrightarrow \begin{pmatrix} a \\ \sqrt{N} \cdot b \end{pmatrix}$乘法规则$(a_1 I b_1)(a_2 I b_2) (a_1 a_2 - N b_1 b_2) I(a_1 b_2 a_2 b_1)$当N 1时退化为标准复数乘法。2.2 螺旋 Attention 公式把 Q、K、V 从ℝ^{d}映射到螺旋数域^{d/2}每两个实数组成一个螺旋数$S_{ij} \frac{Q_i \star K_j}{\sqrt{d_k \cdot N}}$其中⋆是螺旋内积$Q_i \star K_j \sum_{k1}^{d/2} \left( q_{i,2k} k_{j,2k} N \cdot q_{i,2k-1} k_{j,2k-1} \right)$然后过螺旋 Softmax保留相位信息$A_{ij} \frac{\exp(S_{ij} / \tau)}{Z_j}$输出$O_i \sum_j A_{ij} \star V_j$关键差异N参数控制注意力的聚焦程度N → 0接近标准实数 AttentionN 1标准复数 AttentionN 1强聚焦远距离 token 衰减更快适合长序列推理N 1弱聚焦保留更多全局信息适合创意生成三、PyTorch 完整实现3.1 螺旋数线性层import torch import torch.nn as nn import torch.nn.functional as F class SpiralLinear(nn.Module): 螺旋数线性变换层 输入: (batch, seq_len, d_model) 其中 d_model 必须是偶数 每两个相邻维度组成一个螺旋数: (a, b) - a I*b def __init__(self, d_in, d_out, N1.0): super().__init__() assert d_in % 2 0 and d_out % 2 0 self.N N self.d_in_half d_in // 2 self.d_out_half d_out // 2 # 权重矩阵: 实部和虚部分开 self.W_real nn.Parameter(torch.randn(d_out_half, d_in_half) * 0.02) self.W_imag nn.Parameter(torch.randn(d_out_half, d_in_half) * 0.02) def forward(self, x): x: (batch, seq_len, d_in) - (batch, seq_len, d_out) batch, seq_len, _ x.shape # 重塑为螺旋数: (batch, seq_len, d/2, 2) x x.reshape(batch, seq_len, -1, 2) a x[..., 0] # 实部 b x[..., 1] # 虚部系数 # 螺旋乘法: (a I*b) * (W_real I*W_imag) # (a*W_real - N*b*W_imag) I*(a*W_imag b*W_real) real_out F.linear(a, self.W_real) - self.N * F.linear(b, self.W_imag) imag_out F.linear(a, self.W_imag) F.linear(b, self.W_real) # 拼接回 (batch, seq_len, d_out) output torch.stack([real_out, imag_out], dim-1) return output.reshape(batch, seq_len, -1)3.2 螺旋注意力层class SpiralAttention(nn.Module): 螺旋注意力机制 N 参数控制聚焦程度: - N 1: 强聚焦远距离衰减快 - N 1: 弱聚焦保留全局信息 - N 1: 退化为标准复数注意力 def __init__(self, d_model, n_heads, N1.0, dropout0.1): super().__init__() assert d_model % (2 * n_heads) 0 self.d_model d_model self.n_heads n_heads self.d_head d_model // n_heads self.N N self.scale (self.d_head // 2) * N # 螺旋缩放因子 # Q/K/V 投影螺旋线性层 self.W_q SpiralLinear(d_model, d_model, N) self.W_k SpiralLinear(d_model, d_model, N) self.W_v SpiralLinear(d_model, d_model, N) self.W_o SpiralLinear(d_model, d_model, N) self.dropout nn.Dropout(dropout) def forward(self, x, maskNone): x: (batch, seq_len, d_model) batch, seq_len, _ x.shape # 投影并分头 Q self.W_q(x).reshape(batch, seq_len, self.n_heads, self.d_head) K self.W_k(x).reshape(batch, seq_len, self.n_heads, self.d_head) V self.W_v(x).reshape(batch, seq_len, self.n_heads, self.d_head) # 转置为 (batch, n_heads, seq_len, d_head) Q Q.transpose(1, 2) K K.transpose(1, 2) V V.transpose(1, 2) # 螺旋内积: (batch, n_heads, seq_len, seq_len) # 每两个维度组成一个螺旋数 Q_reshape Q.reshape(*Q.shape[:-1], -1, 2) K_reshape K.reshape(*K.shape[:-1], -1, 2) a_q, b_q Q_reshape[..., 0], Q_reshape[..., 1] a_k, b_k K_reshape[..., 0], K_reshape[..., 1] # 螺旋内积: sum(a_q * a_k N * b_q * b_k) scores_real torch.sum(a_q.unsqueeze(-2) * a_k.unsqueeze(-3), dim-1) scores_imag self.N * torch.sum(b_q.unsqueeze(-2) * b_k.unsqueeze(-3), dim-1) scores (scores_real scores_imag) / (self.scale ** 0.5) # Causal mask (如果提供) if mask is not None: scores scores.masked_fill(mask 0, -1e9) # 螺旋 Softmax沿 key 维度 attn F.softmax(scores, dim-1) attn self.dropout(attn) # 加权求和: (batch, n_heads, seq_len, d_head) # 螺旋数加权: sum(attn * V) V_reshape V.reshape(*V.shape[:-1], -1, 2) a_v, b_v V_reshape[..., 0], V_reshape[..., 1] out_a torch.sum(attn.unsqueeze(-1) * a_v.unsqueeze(-3), dim-2) out_b torch.sum(attn.unsqueeze(-1) * b_v.unsqueeze(-3), dim-2) output torch.stack([out_a, out_b], dim-1) output output.reshape(*output.shape[:-2], self.d_model) # 输出投影 output output.transpose(1, 2).reshape(batch, seq_len, self.d_model) output self.W_o(output) return output, attn3.3 螺旋 Transformer 块class SpiralTransformerBlock(nn.Module): def __init__(self, d_model, n_heads, N1.0, d_ff2048, dropout0.1): super().__init__() self.attention SpiralAttention(d_model, n_heads, N, dropout) self.norm1 nn.LayerNorm(d_model) self.norm2 nn.LayerNorm(d_model) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model), nn.Dropout(dropout) ) def forward(self, x, maskNone): # Pre-LN 残差 attn_out, attn_weights self.attention(self.norm1(x), mask) x x attn_out x x self.ffn(self.norm2(x)) return x, attn_weights3.4 测试对比标准 Attention vs 螺旋 Attentiondef test_comparison(): batch, seq_len, d_model, n_heads 2, 64, 128, 4 x torch.randn(batch, seq_len, d_model) # 标准 Transformer Block from torch.nn import TransformerEncoderLayer standard_block TransformerEncoderLayer( d_modeld_model, nheadn_heads, dim_feedforward512, batch_firstTrue ) # 螺旋 Transformer Block spiral_block SpiralTransformerBlock( d_modeld_model, n_headsn_heads, N1.2 ) with torch.no_grad(): y_std standard_block(x) y_spiral, attn spiral_block(x) print(fStandard output shape: {y_std.shape}) print(fSpiral output shape: {y_spiral.shape}) print(fSpiral attention shape: {attn.shape}) print(fSpiral attention sum per query: {attn.sum(dim-1)[0, 0, :5]}) print(✅ 螺旋 Attention 运行成功) test_comparison()四、螺旋 Attention 的物理直觉4.1 为什么 N 能控制聚焦N 值物理意义适合场景N 0.5弱螺旋旋转主导伸缩弱创意写作、头脑风暴N 1.0标准复数纯旋转通用任务退化到基线N 1.2中等聚焦代码生成、推理N 2.0强聚焦伸缩主导数学证明、长链逻辑N 动态每层不同 N混合任务4.2 与 RoPE 的对比特性RoPE螺旋 Attention位置编码方式旋转矩阵乘 Q/K内建在螺旋内积中外推能力依赖 base 参数N 参数天然控制衰减实现复杂度需修改 attention 计算替换线性层即可相位连续性间接通过旋转直接螺旋结构保证计算开销5~10%15~25%可优化五、训练策略如何让 N 自己学class LearnableSpiralAttention(SpiralAttention): def __init__(self, d_model, n_heads, N_init1.0, dropout0.1): super().__init__(d_model, n_heads, NN_init, dropoutdropout) # 把 N 变成可学习参数 self.log_N nn.Parameter(torch.log(torch.tensor(N_init))) property def N(self): return torch.exp(self.log_N).item() def forward(self, x, maskNone): # 动态更新 scale self.scale (self.d_head // 2) * self.N return super().forward(x, mask)训练时N会自动调整如果模型发现需要聚焦→N增大如果模型需要发散→N减小不同层可以学出不同N→ 形成螺旋深度六、实验设想螺旋 Transformer 能做什么任务预期优势长文本推理32KN 自动增大远距离注意力衰减缓解 lost-in-the-middle数学证明链相位连续性减少逻辑矛盾代码生成置信度传播帮助发现潜在 bug多轮对话螺旋记忆让上下文更连贯时间序列预测天然建模相位变化chirp 类信号多模态融合不同模态用不同 N自动对齐七、 螺旋生成论系列作品必藏网址作者张智明平台ZenodoCERN 运营开放获取 核心数学与计算《螺旋数原理公理系统与各向异性复数理论》https://doi.org/10.5281/zenodo.20602099《螺旋生成元一个跨学科统一数学框架的探索》https://doi.org/10.5281/zenodo.21555082《螺旋计算量子计算的新基础——从几何原理到可扩展量子计算架构》https://doi.org/10.5281/zenodo.21356615《螺旋元逻辑从 i²-1 到万物理论的统一框架假说》https://doi.org/10.5281/zenodo.21806751 物理与信号《螺旋波物理与数学基础 (HGO)》https://doi.org/10.5281/zenodo.21416056《螺旋统计力学从因果闭环到可检验预言》https://doi.org/10.5281/zenodo.21416056 AI / 工程《生成式 AI 与提示词工程原理、方法与实战》https://doi.org/10.5281/zenodo.20839550《螺旋工程学从生成论到可控构造》https://doi.org/10.5281/zenodo.21254457 全集索引Spiral-Generation Theory: A Comprehensive Compendium of Workshttps://doi.org/10.5281/zenodo.21211001螺旋生成论全集索引、术语表与开放问题汇编https://doi.org/10.5281/zenodo.21320146 作者主页ORCIDhttps://orcid.org/0009-0003-7777-7694八、CSDN 式总结标准 Transformer 的 Attention 是i² -1的产物——纯旋转、无尺度、靠 Softmax 强行归一化。螺旋 Attention 用I² -N把相位和尺度编码进注意力机制本身相位 → 语义方向尺度 → 置信度/重要性N 参数 → 聚焦程度可学习螺旋内积 → 天然的位置感知你不需要推翻 Transformer 架构只需要把nn.Linear换成SpiralLinear把scaled_dot_product_attention换成SpiralAttention——一行代码不改模型结构底层数学直接升级。好框架不一定颠覆一切但能让你在现有架构上多一个从数学结构上优化的旋钮。
网站建设高端定制企业官网