如何跟踪scikit-learn逻辑回归的训练损失以优化参数?
解决方法:通过Callback追踪LogisticRegression每轮迭代的损失与预测结果
由于你使用的是saga求解器(支持迭代过程的回调),可以通过自定义训练回调类来记录每轮迭代的损失、模型参数,进而推导各轮的预测结果。以下是具体实现步骤:
1. 自定义训练追踪回调类
这个类会在每轮迭代结束时自动被调用,负责保存损失值、模型系数和截距:
from sklearn.pipeline import make_pipeline from sklearn.linear_model import LogisticRegression import numpy as np class TrainingTracker: def __init__(self, X_train, y_train, X_val=None, y_val=None): self.losses = [] # 存储每轮损失值 self.coefs = [] # 存储每轮模型系数 self.intercepts = [] # 存储每轮模型截距 self.train_preds = [] # 可选:存储每轮训练集预测结果 self.val_preds = [] # 可选:存储每轮验证集预测结果 self.X_train = X_train self.y_train = y_train self.X_val = X_val self.y_val = y_val def __call__(self, iteration, model): # 记录当前轮次的损失 self.losses.append(model.loss_) # 记录当前模型参数(必须用copy,避免后续迭代覆盖) self.coefs.append(model.coef_.copy()) self.intercepts.append(model.intercept_.copy()) # 可选:计算并保存训练集预测结果 logit_train = np.dot(self.X_train, model.coef_.T) + model.intercept_ prob_train = 1 / (1 + np.exp(-logit_train)) self.train_preds.append((prob_train >= 0.5).astype(int)) # 可选:计算并保存验证集预测结果(如果传入验证集) if self.X_val is not None: logit_val = np.dot(self.X_val, model.coef_.T) + model.intercept_ prob_val = 1 / (1 + np.exp(-logit_val)) self.val_preds.append((prob_val >= 0.5).astype(int))
2. 在Pipeline中集成回调
创建Pipeline时,将回调实例传入LogisticRegression的callback参数(注意:sklearn版本需≥0.24,该参数才被支持):
# 初始化追踪器,可传入验证集以记录验证结果 tracker = TrainingTracker(X_train, y_train, X_val=X_val, y_val=y_val) # 构建带回调的Pipeline logreg = make_pipeline( LogisticRegression( random_state=42, C=.06, penalty="l1", max_iter=5000, solver="saga", callback=tracker # 绑定自定义回调 ) ) # 启动训练 logreg.fit(X_train, y_train)
3. 分析追踪到的数据
训练完成后,你可以直接从tracker中提取所需信息:
- 查看每轮损失变化:
tracker.losses(是一个列表,索引对应迭代轮次) - 查看第N轮的训练集预测结果:
tracker.train_preds[N] - 查看第N轮的模型参数:
tracker.coefs[N]、tracker.intercepts[N]
如果需要单独验证某一轮的预测逻辑,也可以手动用参数计算:
def get_prediction(X, coef, intercept): logit = np.dot(X, coef.T) + intercept prob = 1 / (1 + np.exp(-logit)) return (prob >= 0.5).astype(int) # 例如获取第400轮的验证集预测 pred_400 = get_prediction(X_val, tracker.coefs[400], tracker.intercepts[400])
关键注意事项
- 版本要求:
callback参数仅在sklearn 0.24及以上版本支持,若版本过低需先升级。 - 参数拷贝:保存系数时必须用
copy(),否则所有轮次的系数会指向同一个内存对象,最终只保留最后一轮的值。 - 性能权衡:如果数据集极大,保存每轮的预测结果会占用较多内存,可根据需求选择性记录(比如只记录每10轮的结果)。
内容的提问来源于stack exchange,提问作者neverreally
相关产品推荐
相关产品推荐

