ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

如何用 Python 逐个实现 Swish、GELU 与 Mish 等神经网络激活函数并用 doctest 验证输出

如何用 Python 逐个实现 Swish、GELU 与 Mish 等神经网络激活函数并用 doctest 验证输出 如何用 Python 逐个实现 Swish、GELU 与 Mish 等神经网络激活函数并用 doctest 验证输出【免费下载链接】PythonAll Algorithms implemented in Python项目地址: https://gitcode.com/GitHub_Trending/pyt/Python这篇文章对应一个具体任务在 PythonGitHub 推荐项目精选 / pyt / PythonTheAlgorithms 算法仓库中逐个阅读并实现 Swish、GELU 与 Mish 三个神经网络激活函数然后运行每个模块内置的 doctest 测试用例确认输出与文档字符串中声明的期望值一致。仓库中的实现位于 neural_network/activation_functions 目录本文章涉及的四个文件是swish.pySwishSiLU及其带可训练参数的变体gaussian_error_linear_unit.pyGELUmish.pyMishsoftplus.pyMish 依赖的 softplus 函数准备条件三个激活函数文件都只依赖 NumPy每个文件开头都有import numpy as np。项目根目录的 pyproject.toml 在依赖列表中声明了numpy2.1.3因此确保环境中有 NumPy 即可开始。另外注意mish.py 使用了相对导入from .softplus import softplus所以这些文件必须从仓库根目录以模块方式运行python -m不能直接以脚本方式运行 mish.py原因见文末排查一节。逐个实现三个激活函数Swishx · sigmoid(x) 与带参数的变体swish.py 定义了三个函数。文件开头的模块说明指出Swish 是一个平滑、非单调的函数定义为 f(x) x · sigmoid(x)。核心实现只有三行import numpy as np def sigmoid(vector: np.ndarray) - np.ndarray: return 1 / (1 np.exp(-vector)) def sigmoid_linear_unit(vector: np.ndarray) - np.ndarray: return vector * sigmoid(vector) def swish(vector: np.ndarray, trainable_parameter: int) - np.ndarray: return vector * sigmoid(trainable_parameter * vector)sigmoid实现 1/(1 e⁻ˣ)是另外两个函数都会用到的基础组件sigmoid_linear_unit即标准的 Swish/SiLU公式为 x · sigmoid(x)swish多接受一个trainable_parameter参数 α公式变为 x · sigmoid(αx)。文档字符串中说明该参数用于实现不同版本的 Swish 激活函数。docstring 中给出的 doctest 用例也是后续验证时会执行的断言 sigmoid_linear_unit(np.array([-1.0, 1.0, 2.0])) array([-0.26894142, 0.73105858, 1.76159416]) swish(np.array([-1.0, 1.0, 2.0]), 2) array([-0.11920292, 0.88079708, 1.96402758])GELUx · sigmoid(1.702x)gaussian_error_linear_unit.py 中的文件说明写道函数接收一个 K 个实数组成的向量返回x * sigmoid(1.702*x)。文件内同样自带一个局部sigmoid实现与 swish.py 中一致核心只有一行def gaussian_error_linear_unit(vector: np.ndarray) - np.ndarray: return vector * sigmoid(1.702 * vector)docstring 中的 doctest 用例 gaussian_error_linear_unit(np.array([-1.0, 1.0, 2.0])) array([-0.15420423, 0.84579577, 1.93565862])文档字符串注明输入应为形状 (1, n)、由实数组成的 numpy 数组。Mishx · tanh(softplus(x))mish.py 的 docstring 给出了公式f(x) x · tanh(softplus(x)) x · tanh(ln(1 eˣ))。实现中softplus不是本地函数而是从同目录的 softplus.py 导入的from .softplus import softplus def mish(vector: np.ndarray) - np.ndarray: return vector * np.tanh(softplus(vector))softplus.py 中softplus的实现是np.log(1 np.exp(vector))即 ln(1 eˣ)它本身也带 doctest 用例。理解 Mish 时建议先读这个文件因为它就是 Mish 公式中的 softplus 部分。用 doctest 验证输出四个文件的结尾都是同一段代码if __name__ __main__: import doctest doctest.testmod()也就是说以模块方式直接运行某个文件时doctest 会自动执行该文件 docstring 中所有用例并逐条比对实际输出与期望输出。在仓库根目录执行python3 -m neural_network.activation_functions.swish -v-v参数会让 doctest 逐条打印每个用例的Trying:实际执行的调用和Expecting:期望输出末尾给出汇总。以下是实际运行得到的示例输出Trying: swish(np.array([-1.0, 1.0, 2.0]), 2) Expecting: array([-0.11920292, 0.88079708, 1.96402758]) ok ... 5 tests in 4 items. 5 passed and 0 failed. Test passed.GELU 和 Mish 的验证命令同理把模块名换掉即可python3 -m neural_network.activation_functions.gaussian_error_linear_unit -v python3 -m neural_network.activation_functions.mish -v判断验证是否通过看两点每个用例后面跟着ok没有Failed example——Failed example表示某个函数的实际返回值与 docstring 中的期望输出不一致末尾汇总行的形式是N tests in M items./N passed and 0 failed./Test passed.。上面示例中 swish 模块是 5 个用例sigmoid1 个、sigmoid_linear_unit2 个、swish2 个GELU 模块是 3 个用例Mish 模块是 2 个用例。不加-v直接运行时doctest 在全部通过的情况下不打印汇总静默结束只有失败时才会打印失败详情。因此想逐条核对输出时始终加-v。常见问题直接运行 mish.py 报 ImportErrorswish.py和gaussian_error_linear_unit.py没有包内相对导入直接python3 neural_network/activation_functions/swish.py也能跑。但 mish.py 不行直接运行会报错ImportError: attempted relative import with no known parent package原因是第 11 行的from .softplus import softplus是相对导入只有当文件作为包中模块python -m neural_network.activation_functions.mish被导入时才有效。所以 Mish 的验证命令必须使用python3 -m neural_network.activation_functions.mish并保证在仓库根目录下执行使neural_network包能被找到。小结与延伸完成本文的操作后你会得到一条可重复的核对路径读公式 → 看一行核心实现 → 运行python3 -m 模块路径 -v→ 确认0 failed与Test passed.。同一目录下还有 ReLU、Leaky ReLU、ELU、SELU、Softplus、Squareplus、Binary Step 等激活函数的实现见 activation_functions 目录它们都遵循相同的docstring 内嵌 doctest __main__调用doctest.testmod()的结构可以用完全一样的命令逐个验证。【免费下载链接】PythonAll Algorithms implemented in Python项目地址: https://gitcode.com/GitHub_Trending/pyt/Python创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表