拆分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,40 @@
|
||||
"""
|
||||
工具函数模块
|
||||
包含激活函数和掩码创建函数
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def softmax(x, axis=-1):
|
||||
"""Softmax激活函数,用于将 logits 转换为概率分布"""
|
||||
# 减去最大值防止数值溢出
|
||||
e_x = np.exp(x - np.max(x, axis=axis, keepdims=True))
|
||||
return e_x / np.sum(e_x, axis=axis, keepdims=True)
|
||||
|
||||
|
||||
def relu(x):
|
||||
"""ReLU激活函数"""
|
||||
return np.maximum(0, x)
|
||||
|
||||
|
||||
def create_padding_mask(seq, pad_idx=0):
|
||||
"""
|
||||
创建填充掩码
|
||||
用于忽略序列中的填充位置(padding tokens)
|
||||
seq: 输入序列 [batch_size, seq_len]
|
||||
pad_idx: 填充token的索引
|
||||
返回: 掩码 [batch_size, 1, 1, seq_len]
|
||||
"""
|
||||
return (seq != pad_idx).astype(float)[:, np.newaxis, np.newaxis, :]
|
||||
|
||||
|
||||
def create_look_ahead_mask(size):
|
||||
"""
|
||||
创建前瞻掩码(下三角掩码)
|
||||
用于解码器中,防止位置i看到i之后的信息
|
||||
size: 序列长度
|
||||
返回: 掩码 [size, size]
|
||||
"""
|
||||
mask = np.triu(np.ones((size, size)), k=1).astype(bool)
|
||||
return (~mask).astype(float)
|
||||
Reference in New Issue
Block a user