一、正则化之weight_decay

1、Regularization:减少方差的策略

误差可分解为:偏差、方差与噪声之和。

误差=偏差+方差+噪声

偏差:度量了学习算法的期望预测真实结果的偏离程度,即刻画了学习算法本身的拟合能力;

方差:度量了同样大小的训练集的变动所导致的学习性能的变化,即刻画了数据扰动所造成的影响;

噪声:表达了在当前任务上任何学习算法所能达到的期望泛化误差的下界;

损失函数:衡量模型输出与真实标签的差异;

损失函数(loss function):

代价函数(cost function):

目标函数(objective function)

obj = Cost + Regularization Term

针对Regularization Term常用的有两种:

  1. L1 Regularization Term
  2. L2 Regularization Term

(1) L1 Regularization Term

(2)L2 Regularization Term

左边的图对应L1正则化,右边对应L2正则化;

regularization解决overfitting(L2正则化解决过拟合问题)

regularization可以使得训练曲线变得更加平缓,在训练集上的误差变大,但是在测试集上的误差变小。

最初的loss function只是考虑了prediction的error,而regularization是在原来loss function的基础上加了一个正则化项,就是上面对应的L1和L2;

针对增加的正则化项,其中主要有两个参数, 和  , 因此就期望参数  的值越小甚至接近0;因为参数值接近0的function是比较平滑的,这里为什么没有考虑偏置  呢?因为这个参数值大小与function的平滑程度是没有关系的,偏置的大小只是把function上下移动而已;针对较平滑的function,由于输出对输入是不敏感的,测试的时候,一些噪声对这个平滑的function的影响就会较小。

还有就是   这个值是需要手动去调整以取得最好的值;具体可以参考李宏毅老师的讲解:

我们喜欢比较平滑的function,因为它对noise不那么sensitive;但是我们又不喜欢太平滑的function,因为它就失去了对data拟合的能力;而function的平滑程度,就需要通过调整  来决定;

 L2 Regularization = weight decay(权值衰减)

代码部分:

import torch
import torch.nn as nn
import matplotlib.pyplot as plt
from tools.common_tools import set_seed
from torch.utils.tensorboard import SummaryWriter
set_seed(1)
n_hidden = 200
max_iter = 2000
disp_interval = 200
lr_init = 0.01def gen_data(num_data=10, x_range=(-1, 1)):w = 1.5train_x = torch.linspace(*x_range, num_data).unsqueeze_(1)train_y = w*train_x + torch.normal(0, 0.5, size=train_x.size())test_x = torch.linspace(*x_range, num_data).unsqueeze_(1)test_y = w*test_x + torch.normal(0, 0.3, size=test_x.size())return train_x, train_y, test_x, test_ytrain_x, train_y, test_x, test_y = gen_data(x_range=(-1, 1))class MLP(nn.Module):def __init__(self, neural_num):super(MLP, self).__init__()self.linears = nn.Sequential(nn.Linear(1, neural_num),nn.ReLU(inplace=True),nn.Linear(neural_num, neural_num),nn.ReLU(inplace=True),nn.Linear(neural_num, neural_num),nn.ReLU(inplace=True),nn.Linear(neural_num, 1),)def forward(self, x):return self.linears(x)net_normal = MLP(neural_num=n_hidden)
net_weight_decay = MLP(neural_num=n_hidden)# 优化器
optim_normal = torch.optim.SGD(net_normal.parameters(), lr=lr_init, momentum=0.9)
optim_wdecay = torch.optim.SGD(net_weight_decay.parameters(), lr=lr_init, momentum=0.9, weight_decay=1e-2)# 损失函数
loss_func = torch.nn.MSELoss()writer = SummaryWriter(comment='_test_tensorboard', filename_suffix="12345678")
for epoch in range(max_iter):# forwardpred_normal, pred_wdecay = net_normal(train_x), net_weight_decay(train_x)loss_normal, loss_wdecay = loss_func(pred_normal, train_y), loss_func(pred_wdecay, train_y)optim_normal.zero_grad()optim_wdecay.zero_grad()loss_normal.backward()loss_wdecay.backward()optim_normal.step()optim_wdecay.step()if (epoch+1) % disp_interval == 0:# 可视化for name, layer in net_normal.named_parameters():writer.add_histogram(name + '_grad_normal', layer.grad, epoch)writer.add_histogram(name + '_data_normal', layer, epoch)for name, layer in net_weight_decay.named_parameters():writer.add_histogram(name + '_grad_weight_decay', layer.grad, epoch)writer.add_histogram(name + '_data_weight_decay', layer, epoch)test_pred_normal, test_pred_wdecay = net_normal(test_x), net_weight_decay(test_x)# 绘图plt.scatter(train_x.data.numpy(), train_y.data.numpy(), c='blue', s=50, alpha=0.3, label='train')plt.scatter(test_x.data.numpy(), test_y.data.numpy(), c='red', s=50, alpha=0.3, label='test')plt.plot(test_x.data.numpy(), test_pred_normal.data.numpy(), 'r-', lw=3, label='no weight decay')plt.plot(test_x.data.numpy(), test_pred_wdecay.data.numpy(), 'b--', lw=3, label='weight decay')plt.text(-0.25, -1.5, 'no weight decay loss={:.6f}'.format(loss_normal.item()), fontdict={'size': 15, 'color': 'red'})plt.text(-0.25, -2, 'weight decay loss={:.6f}'.format(loss_wdecay.item()), fontdict={'size': 15, 'color': 'red'})plt.ylim((-2.5, 2.5))plt.legend(loc='upper left')plt.title("Epoch: {}".format(epoch+1))plt.show()plt.close()

虽然no weight decay拟合了所有的点,但是过拟合了,我们需要的是平滑的;

正则化之weight-decay相关推荐

  1. 权值衰减weight decay的理解

    1. 介绍 权值衰减weight decay即L2正则化,目的是通过在Loss函数后加一个正则化项,通过使权重减小的方式,一定减少模型过拟合的问题. L1正则化:即对权重矩阵的每个元素绝对值求和, λ ...

  2. tf.nn.l2_loss() 与 权重衰减(weight decay)

    权重衰减(weight decay)   L2正则化的目的就是为了让权重衰减到更小的值,在一定程度上减少模型过拟合的问题,所以权重衰减也叫L2正则化.   L2正则化就是在代价函数后面再加上一个正则化 ...

  3. 深度学习:权重衰减(weight decay)与学习率衰减(learning rate decay)

    正则化方法:防止过拟合,提高泛化能力 避免过拟合的方法有很多:early stopping.数据集扩增(Data augmentation).正则化(Regularization)包括L1.L2(L2 ...

  4. weight decay 的矩阵描述

    weight decay(权重衰减) 又叫regularization(正则化).下面叙述如何用矩阵简明的描述loss表达式,以及矩阵求导问题. loss表达式 L ( w , b ) = η 2 ∣ ...

  5. 权重衰减(weight decay)在贝叶斯推断(Bayesian inference)下的理解

    权重衰减(weight decay)在贝叶斯推断(Bayesian inference)下的理解 摘要 权重衰减 贝叶斯(Bayes inference) 视角下的权重衰减 似然函数(log like ...

  6. weight decay(权值衰减)、momentum(冲量)和normalization

    一.weight decay(权值衰减)的使用既不是为了提高你所说的收敛精确度也不是为了提高收敛速度,其最终目的是防止过拟合.在损失函数中,weight decay是放在正则项(regularizat ...

  7. weight decay (权值衰减)

    http://blog.sina.com.cn/s/blog_890c6aa30100z7su.html 在机器学习或者模式识别中,会出现overfitting,而当网络逐渐overfitting时网 ...

  8. 深度学习的权重衰减是什么_【深度学习理论】一文搞透Dropout、L1L2正则化/权重衰减...

    前言 本文主要内容--一文搞透深度学习中的正则化概念,常用正则化方法介绍,重点介绍Dropout的概念和代码实现.L1-norm/L2-norm的概念.L1/L2正则化的概念和代码实现- 要是文章看完 ...

  9. 深度学习中的优化算法与实现

    点击上方"3D视觉工坊",选择"星标" 干货第一时间送达 GiantPandaCV导语:这篇文章的内容主要是参考 沐神的mxnet/gluon视频中,Aston ...

  10. 卷积神经网络超详细介绍

    文章目录 1.卷积神经网络的概念 2. 发展过程 3.如何利用CNN实现图像识别的任务 4.CNN的特征 5.CNN的求解 6.卷积神经网络注意事项 7.CNN发展综合介绍 8.LeNet-5结构分析 ...

最新文章

  1. [k8s] 第八章 数据存储
  2. 再谈应用环境下的TIME_WAIT和CLOSE_WAIT
  3. mysql导入工具 行提交_使用命令行工具mysqlimport导入数据
  4. IT外包 OpenEIM 强调CMMI等级
  5. ad导出元件清单_【原创分享】 Altium Designer 一键导出坐标和BOM脚本,V0.6
  6. python 函数参数枚举_Python中的枚举:如何在方法参数中强制执行
  7. python协程详解_彻底搞懂python协程-第一篇(关键词1-4)
  8. python之IO多路复用
  9. 机器学习面试概念重点汇总
  10. React:引入echarts绘制图表
  11. win10电脑找不到xps查看器的详细解决步骤
  12. 小红帽Linux系统命令重启,Linux系统常用命令之一
  13. chariot iperf使用_iperf局域网性能工具
  14. 1%大气密度也能飞?NASA把无人机送上火星,最具野心探测计划启动
  15. 随机效应估算与固定效应估算_面板工具变量法学习手册(固定效应与随机效应方法、过度识别检验、预测等)...
  16. 相机成像时间与曝光时间的关系
  17. H5 --(解决)ios的webview中上/下拉露出黑灰色背景问题
  18. 原来ChatGPT可以充当这么多角色
  19. 三星Galaxy S10可能把加密货币推向数百万名精通高兴技术的手机用户
  20. 如何在AD中导入CAD画的DXF/DWG文件?

热门文章

  1. 解决Vue3的undefined问题
  2. 计算机老师一句话,写给老师的一句话短句 感谢老师的简单一句话
  3. Springboot 整合 druid
  4. mysql端口establish_sqlserver提示The Network Adapter could not establish the con
  5. 更好的Google Glass:棱镜变长、Intel Atom处理器和外置电池组
  6. 小米手机android目录在哪里设置字体,[小米手机]小米手机MIUI自己制作.MTZ字体包方法 无需ROOT权限...
  7. 【ECLIPSE 二】eclipse java web 版本修改问题 3.0-2.5
  8. avi格式如何转换成mp4格式
  9. 三校生计算机教学计划,三校生高考英语教学计划
  10. USRPx310的射频板UBX160