如何将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
相关产品推荐
相关产品推荐

