拆分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
+20
View File
@@ -0,0 +1,20 @@
"""
前馈网络模块
包含FeedForward类
"""
from .layers import Linear
from .utils import relu
class FeedForward:
"""前馈神经网络(两层全连接网络)"""
def __init__(self, d_model, d_ff):
# 第一层:扩展维度
self.linear1 = Linear(d_model, d_ff)
# 第二层:恢复维度
self.linear2 = Linear(d_ff, d_model)
def forward(self, x):
"""前向传播:线性 -> ReLU -> 线性"""
return self.linear2.forward(relu(self.linear1.forward(x)))