""" 工具函数模块 包含激活函数和掩码创建函数 """ 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)