拆分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:
@@ -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)))
|
||||
Reference in New Issue
Block a user