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

为何循环内定义的函数生成的MSE序列长度与外部循环不符?

问题排查与解决方案

你的mse_by_iter长度仅为100,大概率是以下几个核心原因导致的,逐个排查即可解决:

1. 循环次数未达到1000

先确认fit方法的循环范围:

  • 检查循环是否写为for _ in range(self.n_iter),而非range(100)
  • 核对类初始化的n_iter参数:如果__init__默认值设为100,或者实例化时未显式传入n_iter=1000,循环只会执行100次。比如:
    # 错误示例:默认n_iter=100
    def __init__(self, n_iter=100):
        self.n_iter = n_iter
    
    改成n_iter=1000,或者实例化时指定CustomSGDRegressor(n_iter=1000)

2. MSE计算未在每次迭代执行

如果内部mse函数加了不必要的条件判断,导致仅部分迭代会存入MSE值:

def mse():
    if some_condition:  # 比如每10次迭代才计算一次
        current_mse = ...
        mse_by_iter.append(current_mse)

直接移除条件判断,确保每次迭代都计算并追加MSE值。

3. 列表被意外重置

检查mse_by_iter的初始化位置:

  • 若将mse_by_iter = []写在循环内部,每次迭代都会清空列表(最终长度为1),但你的情况是100,这个可能性较低,但仍需确认:列表初始化应放在循环之前,比如fit方法开头或类的__init__中。

修正后的代码示例

以下是标准实现模板,确保MSE序列长度与迭代次数完全匹配:

from sklearn.base import BaseEstimator
import numpy as np

class CustomSGDRegressor(BaseEstimator):
    def __init__(self, n_iter=1000, learning_rate=0.01):
        self.n_iter = n_iter
        self.learning_rate = learning_rate
        # 初始化存储历史的列表
        self.w0_history = []
        self.w1_history = []
        self.mse_history = []
    
    def fit(self, X, y):
        # 初始化权重
        w0, w1 = 0.0, 0.0
        n_samples = X.shape[0]

        # MSE计算函数放在循环外更高效
        def calculate_mse(y_true, y_pred):
            return np.mean((y_true - y_pred) ** 2)

        for _ in range(self.n_iter):
            # 计算预测值
            y_pred = w0 + w1 * X
            # 计算梯度
            grad_w0 = -2 * np.mean(y - y_pred)
            grad_w1 = -2 * np.mean((y - y_pred) * X)
            # 更新权重
            w0 -= self.learning_rate * grad_w0
            w1 -= self.learning_rate * grad_w1
            # 保存权重与MSE
            self.w0_history.append(w0)
            self.w1_history.append(w1)
            current_mse = calculate_mse(y, y_pred)
            self.mse_history.append(current_mse)
        
        return self

快速验证方法

在fit方法的循环内添加一行打印,实时确认列表长度:

for i in range(self.n_iter):
    # ... 其他逻辑代码 ...
    self.mse_history.append(current_mse)
    print(f"迭代{i+1}次,MSE列表长度:{len(self.mse_history)}")

运行后查看最后一次打印的数值,即可快速定位问题所在。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 07:05:32