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

梯度下降theta值收敛异常问题求助(附Python代码)

梯度下降代码运行异常,无法得到预期收敛图表

我编写了一段梯度下降代码,但运行效果不佳。最终我绘制了一幅包含偏置和权重值的图表,每个点根据循环初始给定的theta(权重、偏置)值的收敛结果进行着色。我尝试自行计算梯度,但效果仍然不佳,希望能得到预期的图表。

import numpy as np
from random import randint,random
import matplotlib . pyplot as plt


def calculh(theta, X):
    h = 0
    h+=theta[0]*X # w*X
    h+= theta[-1] # +b
    return h


def calculY(sigma, h) :
    return sigma(h) # sigma peut-etre tanh, signoide etc.


def erreurJ(theta, sigma):
    somme = 0
    somme = 1/4*(sigma(theta[1])**2+sigma(theta[0]+theta[1])**2)
    return somme


def gradient(X, Y, Ysol, sigmaprime, h):
    return ((Y-Ysol)*sigmaprime(h)*X ,(Y-Ysol)*sigmaprime(h)*1)
def grad(theta):
    w,b = theta[0],theta[1]
    #print(theta)
    return [2*b**3+3*b**2*w+3*b*w**2-2*b+w**3-w,b**3+3*b**2*w+3*b*w**2-b+w**3-w]
# *X correspond a 0 ou 1  : nos 2 entrées ; *1 correspond a derivee de b

def pasfixe(theta, eta, epsilon, X, Y, Ysol, sigma, sigmaprime, h):
    n=0
    while np.linalg.norm(gradient(X, Y, Ysol, sigmaprime, h)) > epsilon and n<10000 :
        for i in range(len(theta)) :
            theta[i] = theta[i] - eta*gradient(X, Y, Ysol, sigmaprime, h)[i]
            h = calculh(theta, X)
            Y = calculY(sigma, h)
            n+=1
            if theta[i]>100 : ### cas de divergence
                return [100,100],Y
    return theta,Y

sigma = lambda z : z**2-1
sigmaprime = lambda z : 2*z
eta = 0.1

X = 1
Ysol = 0
listeY = []
listetheta = []
lst = [[3*random()*(-1)**randint(0,1),3*random()*(-1)**randint(0,1)] for i in range(5000)]
nb = 0
for i in lst:
        nb+=1
        if nb%50 == 0:
            print(nb)
        theta = i[:]
        h = calculh(theta, X)
        Y = calculY(sigma, h)
        CalculTheta = pasfixe( theta, eta, 10**-4, X,Y, Ysol, sigma, sigmaprime, h)
        listetheta.append(CalculTheta[0])
        listeY.append(CalculTheta[1])


for i in range (len(listeY)):
          listeY[i] = round(listeY[i],2)
print (listeY)

for i in range (len(listetheta)):
      for j in range(2):
          listetheta[i][j] = round(listetheta[i][j],2)
print (listetheta)

for i in range(len(lst)):
    if [int(listetheta[i][0]),int(listetheta[i][1])] in [[-2,1]]:
        plt.plot(lst[i][0],lst[i][1],"bo")
    elif [int(listetheta[i][0]),int(listetheta[i][1])] in [[2,-1]]:
        plt.plot(lst[i][0],lst[i][1],"co")
    elif  [int(listetheta[i][0]),int(listetheta[i][1])] in [[0,-1]]:
        plt.plot(lst[i][0],lst[i][1],"go")
    elif  [int(listetheta[i][0]),int(listetheta[i][1])] in [[0,1]]:
        plt.plot(lst[i][0],lst[i][1],"mo")
    elif  int(listetheta[i][0])**2 +int(listetheta[i][1])**2 >= 10:
        plt.plot(lst[i][0],lst[i][1],"ro")

plt.show()

预期图表为:根据初始theta值的收敛结果对散点着色——收敛到[-2,1]的点标记为蓝色,收敛到[2,-1]的点标记为青色,收敛到[0,-1]的点标记为绿色,收敛到[0,1]的点标记为品红色,发散(模长≥10)的点标记为红色。

内容的提问来源于stack exchange,提问作者ismail rachid

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 09:40:55