基于神经网络思维的Logistic Regression优化函数成本列表长度不匹配错误
问题排查:Logistic Regression的optimize函数AssertionError修复
你实现的Logistic Regression梯度下降优化函数optimize触发了AssertionError,提示costs列表长度应为2但实际是1。
错误原因
核心问题是代码缩进错误:for循环体内部的关键逻辑(梯度提取、参数更新、cost收集)都被写到了循环外面,导致循环只重复执行了propagate计算,而参数更新和cost记录只在循环结束后执行一次,完全不符合梯度下降的迭代逻辑。
比如原代码中,dw、db的获取,参数更新,costs.append等代码都不在for循环的缩进块内,这些逻辑仅执行一次,而非每次迭代都执行,自然无法收集到预期数量的cost值。
修复后的代码
def optimize(w, b, X, Y, num_iterations=200, learning_rate=0.009, print_cost=False): """ This function optimizes w and b by running a gradient descent algorithm Arguments: w -- weights, a numpy array of size (num_px * num_px * 3, 1) b -- bias, a scalar X -- data of shape (num_px * num_px * 3, number of examples) Y -- true "label" vector (containing 0 if non-cat, 1 if cat), of shape (1, number of examples) num_iterations -- number of iterations of the optimization loop learning_rate -- learning rate of the gradient descent update rule print_cost -- True to print the loss every 100 steps Returns: params -- dictionary containing the weights w and bias b grads -- dictionary containing the gradients of the weights and bias with respect to the cost function costs -- list of all the costs computed during the optimization, this will be used to plot the learning curve. Tips: You basically need to write down two steps and iterate through them: 1) Calculate the cost and the gradient for the current parameters. Use propagate(). 2) Update the parameters using gradient descent rule for w and b. """ import copy w = copy.deepcopy(w) b = copy.deepcopy(b) costs = [] for i in range(num_iterations): # Cost and gradient calculation grads, cost = propagate(w, b, X, Y) # Retrieve derivatives from grads dw = grads["dw"] db = grads["db"] # Update rule w = w - (learning_rate * dw) b = b - (learning_rate * db) # Record the costs if (i % 100 == 0): costs.append(cost) # Print the cost every 100 training iterations if print_cost: print ("Cost after iteration %i: %f" %(i, cost)) params = {"w": w, "b": b} grads = {"dw": dw, "db": db} return params, grads, costs
修复说明
- 将
dw、db的提取、参数更新代码、costs.append逻辑全部缩进,放入for循环的代码块内,确保每次迭代都执行梯度计算、参数更新和cost记录。 - 当
num_iterations=200时,i会在0和100时满足i%100==0,此时costs列表会添加两次cost,长度变为2,符合测试用例的要求。
内容的提问来源于stack exchange,提问作者Hassan Afif
相关产品推荐
相关产品推荐

