You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用SymPy求解RMSNorm导数时的求和元素关联问题

使用SymPy求解RMSNorm导数时的求和元素关联问题

问题描述

我在SymPy中定义了RMSNorm的正向传播公式,代码如下:

import sympy as sp

# Define the symbols
x = sp.Symbol('x')  # Input variable
n = sp.Symbol('n')  # Number of elements
gamma = sp.Symbol('gamma')
epsilon = sp.Symbol('epsilon')  # Small constant to avoid division by zero

# Define the RMS normalization equation
mean_square = sp.Sum(x**2, (x, 1, n)) / n
rms = sp.sqrt(mean_square + epsilon)
fwd_out = x * gamma / rms

# Display the equation
sp.pprint(fwd_out)

但在对fwd_out关于x求导时遇到了问题:当我执行sp.diff(fwd_out, x),SymPy并没有把rms当作x的函数,而是将其视为常数。比如直接求rms对x的导数会得到0:

sp.diff(rms, x)  # 返回0

这不符合RMSNorm论文中的定义——rms应该是关于输入x的函数,因为它是对所有输入元素的平方求平均后开方得到的。

我现在使用的是Python 3.12.9和SymPy 1.12.1,请问有没有办法让SymPy将rms正确识别为x的函数?


解决方法

问题的核心在于:你在求和时使用了和输入变量相同的符号x作为求和变量,导致SymPy无法区分单个输入元素与求和迭代变量。正确的做法是使用索引变量来表示输入的每个元素,这样SymPy就能关联起rms与输入元素的依赖关系。

修正后的代码如下:

from sympy import *
from sympy.abc import n, gamma, epsilon

# 用IndexedBase定义输入向量x,每个元素用索引i访问
x = IndexedBase("x")
i = symbols('i', cls=Idx)  # 定义索引变量i

# 重新定义RMSNorm公式
mean_squared = Sum(x[i] ** 2, (i, 1, n)) / n
rms = sqrt(mean_squared + epsilon)
# 这里的x[i]表示单个输入元素,和求和的索引对应
fwd_out = x[i] * gamma / rms

# 对单个输入元素x[i]求导
d_fwd_out = diff(fwd_out, x[i])

# 验证rms对x[i]的导数(不再是0)
d_rms = diff(rms, x[i])

这样修改后,SymPy就能正确识别rms与输入元素x[i]的依赖关系,求导结果会符合RMSNorm论文中的预期。


备注:内容来源于stack exchange,提问作者algoProg

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.13 19:35:26