Python计算函数梯度时出现tuple IndexError索引越界报错该如何解决
问题解答
报错原因
- 你传入的测试参数
x = (1,)是仅含1个元素的元组,仅支持使用索引x[0]访问元素,代码中调用x[1]自然会触发索引越界错误。 - 你的目标函数包含
x1、x2两个自变量,输入的x必须是长度为2的可迭代对象,将测试用例修改为x = (1, 0)这类二元元组即可解决索引报错问题。
代码逻辑修正
你当前编写的代码并没有实现梯度计算功能,存在两处逻辑错误:
- 代码计算的是目标函数的数值,且和你给出的函数定义不符:你定义的函数为
f(x) = w1 * x1^2 + w2 * x2,但代码中额外给x[1]加了平方运算,写成了w2 * (x[1]**2)。 - 梯度是函数对各自变量的偏导组成的向量,对应你给出的函数,梯度计算规则为:
- 对x1的偏导:
2 * w1 * x[0] - 对x2的偏导:
w2
- 对x1的偏导:
修正后的代码示例:
def gradient(w1, w2, x): # 校验输入x的合法性 if len(x) < 2: raise ValueError("自变量x需要包含x1、x2两个元素,长度至少为2") grad_x1 = 2 * w1 * x[0] grad_x2 = w2 return (grad_x1, grad_x2)
正常运行的测试示例
w1 = 5; w2 = 3; x = (1, 2) print(gradient(w1, w2, x)) # 输出结果:(10, 3)
内容的提问来源于stack exchange,提问作者paraglider5
相关产品推荐
相关产品推荐

