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