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