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

求实现神经网络学习率按特定小数序列衰减的自动化函数

实现自定义学习率衰减策略:基于损失停滞的阶梯式细粒度衰减

这个需求挺实用的——比起固定倍数衰减(比如每次乘0.1),这种细粒度的衰减能在学习率接近关键阈值时更平缓地调整,避免跳过最优解。下面我会一步步拆解实现思路,再给出PyTorch/TensorFlow的代码示例。

核心逻辑拆解

我们需要两个关键模块:

  • 损失停滞监控:跟踪连续多少个epoch损失没有下降,达到阈值n时触发衰减
  • 细粒度学习率衰减函数:根据当前学习率的数量级动态调整步长,实现0.2→0.1→0.09→0.08…→0.01→0.009的模式

1. 损失停滞监控逻辑

简单来说就是维护一个计数器:

  • 初始化best_loss为无穷大,no_improve_epochs为0
  • 每个epoch结束后,比较当前验证损失和best_loss:
    • 如果当前损失更低:更新best_loss,重置no_improve_epochs为0
    • 如果损失没下降:no_improve_epochs +=1
  • 当no_improve_epochs >=n时,调用衰减函数调整学习率,可选择重置计数器或继续计数

2. 学习率衰减函数实现

关键是根据当前学习率的数量级动态计算衰减步长:

  • 先把学习率转换为科学计数法形式lr = a × 10^b(其中1 ≤ a <10)
  • 当a >1时,步长为10^b(比如0.2的步长是0.1)
  • 当a ==1时,步长切换为10^(b-1)(比如0.1的步长是0.01,0.01的步长是0.001)

代码实现这个函数(Python):

import math

def decay_lr(current_lr):
    if current_lr <= 0:
        return current_lr  # 避免负学习率
    
    # 计算当前学习率的数量级b
    b = math.floor(math.log10(current_lr))
    scale = 10 ** b
    a = current_lr / scale
    
    if a > 1:
        new_lr = current_lr - scale
    else:
        # 当a等于1时,切换到更小量级的步长
        new_scale = 10 ** (b - 1)
        new_lr = current_lr - new_scale
    
    # 确保学习率不会过低或为负
    return max(new_lr, 1e-10)

测试一下这个函数的输出:

print(decay_lr(0.2))    # 0.1
print(decay_lr(0.1))    # 0.09
print(decay_lr(0.09))   # 0.08
print(decay_lr(0.02))   # 0.01
print(decay_lr(0.01))   # 0.009
print(decay_lr(0.001))  # 0.0009

完全符合你想要的衰减模式!

结合深度学习框架实现完整流程

PyTorch示例

在训练循环中集成监控和衰减逻辑:

import torch
import torch.nn as nn
from torch.utils.data import DataLoader

# 假设你已经定义了模型、数据集、优化器
model = nn.Linear(10, 1)
optimizer = torch.optim.SGD(model.parameters(), lr=0.2)
criterion = nn.MSELoss()
train_loader = DataLoader(...)  # 替换为你的训练数据集
val_loader = DataLoader(...)    # 替换为你的验证数据集

# 监控参数
n = 3  # 连续3个epoch损失没下降就触发衰减
best_val_loss = float('inf')
no_improve_epochs = 0

for epoch in range(100):
    # 训练阶段
    model.train()
    train_loss = 0.0
    for batch in train_loader:
        x, y = batch
        optimizer.zero_grad()
        outputs = model(x)
        loss = criterion(outputs, y)
        loss.backward()
        optimizer.step()
        train_loss += loss.item() * x.size(0)
    train_loss /= len(train_loader.dataset)
    
    # 验证阶段
    model.eval()
    val_loss = 0.0
    with torch.no_grad():
        for batch in val_loader:
            x, y = batch
            outputs = model(x)
            loss = criterion(outputs, y)
            val_loss += loss.item() * x.size(0)
    val_loss /= len(val_loader.dataset)
    
    print(f"Epoch {epoch+1}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}")
    
    # 监控损失并调整学习率
    if val_loss < best_val_loss:
        best_val_loss = val_loss
        no_improve_epochs = 0
        torch.save(model.state_dict(), 'best_model.pth')  # 可选:保存最佳模型
    else:
        no_improve_epochs += 1
        if no_improve_epochs >= n:
            current_lr = optimizer.param_groups[0]['lr']
            new_lr = decay_lr(current_lr)
            # 更新优化器学习率
            for param_group in optimizer.param_groups:
                param_group['lr'] = new_lr
            print(f"Loss hasn't improved in {n} epochs. Decaying LR from {current_lr:.6f} to {new_lr:.6f}")
            no_improve_epochs = 0  # 重置计数器,也可选择不重置

TensorFlow/Keras示例

自定义Callback实现监控和衰减:

import tensorflow as tf
from tensorflow.keras.callbacks import Callback

class CustomLRDecay(Callback):
    def __init__(self, patience=3):
        super().__init__()
        self.patience = patience
        self.best_val_loss = float('inf')
        self.no_improve_epochs = 0
    
    def on_epoch_end(self, epoch, logs=None):
        val_loss = logs.get('val_loss')
        if val_loss < self.best_val_loss:
            self.best_val_loss = val_loss
            self.no_improve_epochs = 0
        else:
            self.no_improve_epochs += 1
            if self.no_improve_epochs >= self.patience:
                current_lr = self.model.optimizer.lr.numpy()
                new_lr = decay_lr(current_lr)
                self.model.optimizer.lr.assign(new_lr)
                print(f"Loss hasn't improved in {self.patience} epochs. Decaying LR from {current_lr:.6f} to {new_lr:.6f}")
                self.no_improve_epochs = 0

# 使用示例
model = tf.keras.Sequential([tf.keras.layers.Dense(1, input_shape=(10,))])
model.compile(optimizer=tf.keras.optimizers.SGD(learning_rate=0.2), loss='mse')
# 替换为你的训练/验证数据集
model.fit(train_dataset, validation_data=val_dataset, epochs=100, callbacks=[CustomLRDecay(patience=3)])

注意事项

  • 避免学习率为0:在衰减函数中加入max(new_lr, 1e-10),防止学习率过低导致模型停止更新
  • 用验证集监控:务必使用验证集损失判断是否停滞,不要用训练集,避免过拟合
  • 计数器重置逻辑:触发衰减后是否重置计数器,可根据需求调整——如果希望连续停滞就持续衰减,可不重置;如果希望每次衰减后重新计数,就重置

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 09:16:38