关于torch.optim的灵活使用详解(包括重写SGD,加上L1正则)
作者:tsq292978891 发布时间:2023-06-06 07:58:32
torch.optim的灵活使用详解
1. 基本用法:
要构建一个优化器Optimizer,必须给它一个包含参数的迭代器来优化,然后,我们可以指定特定的优化选项,
例如学习速率,重量衰减值等。
注:如果要把model放在GPU中,需要在构建一个Optimizer之前就执行model.cuda(),确保优化器里面的参数也是在GPU中。
例子:
optimizer = optim.SGD(model.parameters(), lr = 0.01, momentum=0.9)
2. 灵活的设置各层的学习率
将model中需要进行BP的层的参数送到torch.optim中,这些层不一定是连续的。
这个时候,Optimizer的参数不是一个可迭代的变量,而是一个可迭代的字典
(字典的key必须包含'params'(查看源码可以得知optimizer通过'params'访问parameters),
其他的key就是optimizer可以接受的,比如说'lr','weight_decay'),可以将这些字典构成一个list,
这样就是一个可迭代的字典了。
注:这个时候,可以在optimizer设置选项作为关键字参数传递,这时它们将被认为是默认值(当字典里面没有这个关键字参数key-value对时,就使用这个默认的参数)
This is useful when you only want to vary a single option, while keeping all others consistent between parameter groups.
例子:
optimizer = SGD([
{'params': model.features12.parameters(), 'lr': 1e-2},
{'params': model.features22.parameters()},
{'params': model.features32.parameters()},
{'params': model.features42.parameters()},
{'params': model.features52.parameters()},
], weight_decay1=5e-4, lr=1e-1, momentum=0.9)
上面创建的optim.SGD类型的Optimizer,lr默认值为1e-1,momentum默认值为0.9。features12的参数学习率为1e-2。
灵活更改各层的学习率
torch.optim.optimizer.Optimizer的初始化函数如下:
__init__(self, params, lr=<object object>, momentum=0, dampening=0, weight_decay=0, nesterov=False)
params (iterable): iterable of parameters to optimize or dicts defining parameter groups (params可以是可迭代的参数,或者一个定义参数组的字典,如上所示,字典的键值包括:params,lr,momentum,dampening,weight_decay,nesterov)
想要改变各层的学习率,可以访问optimizer的param_groups属性。type(optimizer.param_groups) -> list
optimizer.param_groups[0].keys()
Out[21]: ['dampening', 'nesterov', 'params', 'lr', 'weight_decay', 'momentum']
因此,想要更改某层参数的学习率,可以访问optimizer.param_groups,指定某个索引更改'lr'参数就可以。
def adjust_learning_rate(optimizer, decay_rate=0.9):
for para in optimizer.param_groups:
para['lr'] = para['lr']*decay_rate
重写torch.optim,加上L1正则
查看torch.optim.SGD等Optimizer的源码,发现没有L1正则的选项,而L1正则更容易得到稀疏解。
这个时候,可以更改/home/smiles/anaconda2/lib/python2.7/site-packages/torch/optim/sgd.py文件,模拟L2正则化的操作。
L1正则化求导如下:
dw = 1 * sign(w)
更改后的sgd.py如下:
import torch
from torch.optim.optimizer import Optimizer, required
class SGD(Optimizer):
def __init__(self, params, lr=required, momentum=0, dampening=0,
weight_decay1=0, weight_decay2=0, nesterov=False):
defaults = dict(lr=lr, momentum=momentum, dampening=dampening,
weight_decay1=weight_decay1, weight_decay2=weight_decay2, nesterov=nesterov)
if nesterov and (momentum <= 0 or dampening != 0):
raise ValueError("Nesterov momentum requires a momentum and zero dampening")
super(SGD, self).__init__(params, defaults)
def __setstate__(self, state):
super(SGD, self).__setstate__(state)
for group in self.param_groups:
group.setdefault('nesterov', False)
def step(self, closure=None):
"""Performs a single optimization step.
Arguments:
closure (callable, optional): A closure that reevaluates the model
and returns the loss.
"""
loss = None
if closure is not None:
loss = closure()
for group in self.param_groups:
weight_decay1 = group['weight_decay1']
weight_decay2 = group['weight_decay2']
momentum = group['momentum']
dampening = group['dampening']
nesterov = group['nesterov']
for p in group['params']:
if p.grad is None:
continue
d_p = p.grad.data
if weight_decay1 != 0:
d_p.add_(weight_decay1, torch.sign(p.data))
if weight_decay2 != 0:
d_p.add_(weight_decay2, p.data)
if momentum != 0:
param_state = self.state[p]
if 'momentum_buffer' not in param_state:
buf = param_state['momentum_buffer'] = torch.zeros_like(p.data)
buf.mul_(momentum).add_(d_p)
else:
buf = param_state['momentum_buffer']
buf.mul_(momentum).add_(1 - dampening, d_p)
if nesterov:
d_p = d_p.add(momentum, buf)
else:
d_p = buf
p.data.add_(-group['lr'], d_p)
return loss
一个使用的例子:
optimizer = SGD([
{'params': model.features12.parameters()},
{'params': model.features22.parameters()},
{'params': model.features32.parameters()},
{'params': model.features42.parameters()},
{'params': model.features52.parameters()},
], weight_decay1=5e-4, lr=1e-1, momentum=0.9)
来源:https://blog.csdn.net/tsq292978891/article/details/79724781


猜你喜欢
- Python文件操作和异常处理Python作为一门高级编程语言,为我们提供了丰富的文件操作和异常处理机制。在本文中,我们将从以下几个方面讨论
- 前言本文主要介绍属性、事件和插槽这三个vue基础概念、使用方法及其容易被忽略的一些重要细节。如果你阅读别人写的组件,也可以从这三个部分展开,
- var classA = function(){ this.prop1 = 1; } classA.prototype.func1 = fu
- Golang与python线程详解及简单实例在GO中,开启15个线程,每个线程把全局变量遍历增加100000次,因此预测结果是 15*100
- Tkinter是python的GUI模块,内含各种窗口控件,利用其中messagbox可以制作各种信息弹出窗口。以下是制作信息提示框的代码:
- 本文实例讲述了MySQL截取和拆分字符串函数用法。分享给大家供大家参考,具体如下:首先说截取字符串函数:SUBSTRING(commenti
- 泰勒展开与e的求法大家伙儿知道计算机里的 e是怎么求出来的吗?这还要从神奇的泰勒展开讲起……简单
- 从 Google 的一个细节说起:整个虚线框都是“Next”的可点击区域。看似不经意,却直接提升了细节的可用性。其它页码也巧妙地和上面的字母
- 介绍在 Go reflect 包里面对 Type 有一个 Comparable 的定义:package reflecttype Type i
- Python 数字类型Python 中有三种数字类型:intfloatcomplex为变量赋值时,将创建数值类型的变量:实例x = 10 &
- 需求:前端获取到摄像头信息,通过模型来进行判断人像是否在镜头中,镜头是否有被遮挡。实现步骤:1、通过video标签来展示摄像头中的内容2、通
- 系统环境:win10 开发环境:JetBrains PyCharm 2017.1.5 x64 Python版本:2.7假如我们有一个clas
- 如果你的电脑内存较小那么想在本地做一些事情是很有局限性的(哭丧脸),比如想拿一个kaggle上面的竞赛来练练手,你会发现多数训练数据集都是大
- 先上需要用到的全部代码片段(截取) MenuControl.prototype.boxDisplay = false;//是否显示图层选择菜
- 在神经网络入门回顾(感知器、多层感知器)中整理了关于感知器和多层感知器的理论,这里实现关于与门、与非门、或门、异或门的代码,以便对感知器有更
- Oracle的执行计划一句话命令:set autotrace on
- 大多数卷积神经网络都是直接通过写一个Model类来定义的,这样写的代码其实是比较好懂的,特别是在魔改网络的时候也很方便。然后也有一些会通过c
- 【原理介绍】通过NETCONF,网管能够用可视化的界面统一管理网络中的设备,并且安全性高、可靠性强、扩展性强。如下图所示,网管与网络中的所有
- 问题背景在开始正文之前,感谢用户名为怜索的朋友送给了我的博客2021年的第一个赞!import numpy as npimport matp
- 废话不多说,我就直接上代码让大家看看吧!#!/usr/bin/env python# -*- coding: utf-8 -*-# @Fil