Files
transformer_study/transformer/utils.py
T
XingfenD 7856f5bf4c 拆分transformer.py为独立模块
按Transformer架构组件拆分:
- utils.py: 激活函数和掩码工具
- layers.py: 基础层(Linear, LayerNorm)
- attention.py: 多头注意力机制
- feedforward.py: 前馈网络
- block.py: Transformer编码器块
- model.py: SimpleTransformer模型
- inference.py: 推理引擎
- 删除旧的transformer.py文件
2026-07-17 11:23:11 +08:00

40 lines
1.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
工具函数模块
包含激活函数和掩码创建函数
"""
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)