按Transformer架构组件拆分: - utils.py: 激活函数和掩码工具 - layers.py: 基础层(Linear, LayerNorm) - attention.py: 多头注意力机制 - feedforward.py: 前馈网络 - block.py: Transformer编码器块 - model.py: SimpleTransformer模型 - inference.py: 推理引擎 - 删除旧的transformer.py文件
20 lines
531 B
Python
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))) |