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

如何将PyTorch Cosine调度器的lr_lambda函数改写为Lambda表达式?

解决方案

直接将原嵌套函数的逻辑整合进Lambda表达式即可,注意要补充导入math模块(原代码遗漏了该依赖),完整代码如下:

import math
from torch.optim.lr_scheduler import LambdaLR

def cosine_scheduler(optimizer, training_steps, warmup_steps):
    lr_lambda = lambda current_step: (
        current_step / max(1, warmup_steps)
        if current_step < warmup_steps
        else max(0.0, 0.5 * (1.0 + math.cos(math.pi * ((current_step - warmup_steps) / max(1, training_steps - warmup_steps)))))
    )
    return LambdaLR(optimizer, lr_lambda)

关键说明

  • 原嵌套函数中的progress变量无法在Lambda表达式中直接赋值,因此将其计算逻辑直接内联到余弦衰减的公式中
  • 用括号包裹Lambda表达式内容,避免三元运算符的优先级冲突,同时提升代码可读性
  • 保留原逻辑中的max(1, ...)以防止除以0的异常,max(0.0, ...)确保学习率不会出现负值

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 15:57:16