拆分transformer.py为独立模块

按Transformer架构组件拆分:
- utils.py: 激活函数和掩码工具
- layers.py: 基础层(Linear, LayerNorm)
- attention.py: 多头注意力机制
- feedforward.py: 前馈网络
- block.py: Transformer编码器块
- model.py: SimpleTransformer模型
- inference.py: 推理引擎
- 删除旧的transformer.py文件
This commit is contained in:
2026-07-17 11:23:11 +08:00
parent c57a9bf9e6
commit 7856f5bf4c
10 changed files with 509 additions and 439 deletions
+42
View File
@@ -0,0 +1,42 @@
"""
Transformer块模块
包含TransformerBlock类
"""
from .attention import MultiHeadAttention
from .feedforward import FeedForward
from .layers import LayerNorm
class TransformerBlock:
"""Transformer编码器块"""
def __init__(self, d_model, num_heads, d_ff):
# 多头自注意力层
self.attention = MultiHeadAttention(d_model, num_heads)
# 前馈网络层
self.feed_forward = FeedForward(d_model, d_ff)
# 两个层归一化
self.norm1 = LayerNorm(d_model)
self.norm2 = LayerNorm(d_model)
def forward(self, x, mask=None, use_cache=False):
"""
前向传播(残差连接 + 层归一化)
1. 自注意力 -> 残差连接 -> 层归一化
2. 前馈网络 -> 残差连接 -> 层归一化
x: 输入张量 [batch_size, seq_len, d_model]
mask: 注意力掩码
use_cache: 是否使用KV cache
"""
# 自注意力子层(Q=K=V,自注意力)
attn_output = self.attention.forward(x, x, x, mask, use_cache=use_cache)
x = self.norm1.forward(x + attn_output)
# 前馈网络子层
ff_output = self.feed_forward.forward(x)
x = self.norm2.forward(x + ff_output)
return x
def clear_cache(self):
"""清除该层的KV cache"""
self.attention.clear_cache()