PyTorch Trainer类中是否可将get_optimizer返回值定义为实例属性?
结论
你将优化器实例化逻辑封装到Trainer类内部的实现是合规的,两种场景下都优先建议你保留将优化器赋值为self.optimizer实例属性的写法。
针对最初的疑问:赋值实例属性 VS 每次调用get_optimizer()
优先选择赋值为实例属性,原因如下:
- 优化器本身是有状态对象,内部会保存动量、自适应学习率缓存等训练过程中动态更新的参数,如果你每次用到优化器都调用
get_optimizer()生成新实例,会直接丢失之前的训练状态,导致训练逻辑完全错误。 - 可读性和可维护性更高,后续调试、打印参数、记录日志时直接访问
self.optimizer即可,不需要重复调用实例化方法。 - 可以兼顾灵活性,你可以保留构造函数的
optimizer入参,加一个简单的判断逻辑:如果用户外部传入自定义优化器就用传入的实例,没有传入再调用内部方法生成,同时满足默认场景的一致性要求和自定义场景的扩展性,示例如下:
def __init__( self, params, model, device=torch.device("cpu"), optimizer=None, scheduler=None, wandb_run=None, early_stopping: callbacks.EarlyStopping = None, ): self.params = params self.model = model self.device = device # 优先用外部传入的优化器,没有则内部生成 self.optimizer = optimizer if optimizer is not None else self.get_optimizer(model, params.optimizer_params) # 其余初始化逻辑保持不变
- 注意你当前代码中的
self.get_optimizer()调用缺少必填的model和optimizer_params参数,运行会报错,需要补全参数。
针对补充的五折训练场景
这种场景仍然适合定义self.optimizer实例属性:
- 你每次调用
fit()方法都重新实例化优化器,刚好符合五折交叉验证每折训练重置优化器状态的需求,赋值为实例属性后,单折训练过程中的梯度反传、参数更新、学习率调度都可以直接访问self.optimizer,不需要将优化器作为参数在fit()内部的各个方法间传递,代码更简洁。 - 复用同一个Trainer实例执行多次
fit()时,每次重赋值self.optimizer会自动覆盖上一折的旧优化器实例,不会出现状态污染的问题。 - 如果不希望优化器被外部代码意外修改,可以将属性名改为
self._optimizer,用单前导下划线标明是类内部属性,符合Python编码规范。
内容的提问来源于stack exchange,提问作者ilovewt
相关产品推荐
相关产品推荐

