Files
transformer_study/transformer/feedforward.py
T
XingfenD 7856f5bf4c 拆分transformer.py为独立模块
按Transformer架构组件拆分:
- utils.py: 激活函数和掩码工具
- layers.py: 基础层(Linear, LayerNorm)
- attention.py: 多头注意力机制
- feedforward.py: 前馈网络
- block.py: Transformer编码器块
- model.py: SimpleTransformer模型
- inference.py: 推理引擎
- 删除旧的transformer.py文件
2026-07-17 11:23:11 +08:00

20 lines
531 B
Python

"""
前馈网络模块
包含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)))