按Transformer架构组件拆分: - utils.py: 激活函数和掩码工具 - layers.py: 基础层(Linear, LayerNorm) - attention.py: 多头注意力机制 - feedforward.py: 前馈网络 - block.py: Transformer编码器块 - model.py: SimpleTransformer模型 - inference.py: 推理引擎 - 删除旧的transformer.py文件
33 lines
839 B
Python
33 lines
839 B
Python
"""
|
||
Transformer模型实现包
|
||
按架构组件拆分:
|
||
- utils: 激活函数和掩码工具
|
||
- layers: 基础层(Linear, LayerNorm)
|
||
- attention: 多头注意力机制
|
||
- feedforward: 前馈网络
|
||
- block: Transformer块
|
||
- model: SimpleTransformer模型
|
||
- inference: 推理引擎
|
||
"""
|
||
|
||
from .utils import softmax, relu, create_padding_mask, create_look_ahead_mask
|
||
from .layers import Linear, LayerNorm
|
||
from .attention import MultiHeadAttention
|
||
from .feedforward import FeedForward
|
||
from .block import TransformerBlock
|
||
from .model import SimpleTransformer
|
||
from .inference import InferenceEngine
|
||
|
||
__all__ = [
|
||
'softmax',
|
||
'relu',
|
||
'create_padding_mask',
|
||
'create_look_ahead_mask',
|
||
'Linear',
|
||
'LayerNorm',
|
||
'MultiHeadAttention',
|
||
'FeedForward',
|
||
'TransformerBlock',
|
||
'SimpleTransformer',
|
||
'InferenceEngine'
|
||
] |