MXNET:权重衰减

构建数据集

1# -*- coding: utf-8 -*- 2from mxnet import init 3from mxnet import ndarray as nd 4from mxnet.gluon import loss as gloss 5import gb 6 7n_train = 20 8n_test = 100 9 10num_inputs = 200 11true_w = nd.ones((num_inputs, 1)) * 0.01 12true_b = 0.05 13features = nd.random.normal(shape=(n_train+n_test, num_inputs)) 14labels = nd.dot(features, true_w) + true_b 15labels += nd.random.normal(scale=0.01, shape=labels.shape) 16train_features, test_features = features[:n_train, :], features[n_train:, :] 17train_labels, test_labels = labels[:n_train], labels[n_train:]

数据迭代器

1from mxnet import autograd 2from mxnet.gluon import data as gdata 3 4batch_size = 1 5num_epochs = 10 6learning_rate = 0.003 7 8train_iter = gdata.DataLoader(gdata.ArrayDataset( 9 train_features, train_labels), batch_size, shuffle=True) 10loss = gloss.L2Loss()

训练并展示结果

gb.semilogy函数:绘制训练和测试数据的loss

1from mxnet import gluon 2from mxnet.gluon import nn 3 4def fit_and_plot(weight_decay): 5 net = nn.Sequential() 6 net.add(nn.Dense(1)) 7 net.initialize(init.Normal(sigma=1)) 8 # 对权重参数做 L2 范数正则化,即权重衰减。 9 trainer_w = gluon.Trainer(net.collect_params('.*weight'), 'sgd', { 10 'learning_rate': learning_rate, 'wd': weight_decay}) 11 # 不对偏差参数做 L2 范数正则化。 12 trainer_b = gluon.Trainer(net.collect_params('.*bias'), 'sgd', { 13 'learning_rate': learning_rate}) 14 train_ls = [] 15 test_ls = [] 16 for _ in range(num_epochs): 17 for X, y in train_iter: 18 with autograd.record(): 19 l = loss(net(X), y) 20 l.backward() 21 # 对两个 Trainer 实例分别调用 step 函数。 22 trainer_w.step(batch_size) 23 trainer_b.step(batch_size) 24 train_ls.append(loss(net(train_features), 25 train_labels).mean().asscalar()) 26 test_ls.append(loss(net(test_features), 27 test_labels).mean().asscalar()) 28 gb.semilogy(range(1, num_epochs + 1), train_ls, 'epochs', 'loss', 29 range(1, num_epochs + 1), test_ls, ['train', 'test']) 30 return 'w[:10]:', net[0].weight.data()[:, :10], 'b:', net[0].bias.data() 31print fit_and_plot(5)
  • 使用 Gluon 的 wd 超参数可以使用权重衰减来应对过拟合问题。
  • 我们可以定义多个 Trainer 实例对不同的模型参数使用不同的迭代方法。
点赞
收藏

评论区

加载中...

相关推荐

MySQL:[Err] 1292 - Incorrect datetime value: ‘0000-00-00 00:00:00‘ for column ‘CREATE_TIME‘ at row 1

文章目录问题用navicat导入数据时,报错:原因这是因为当前的MySQL不支持datetime为0的情况。解决修改sql\mode:sql\mode:SQLMode定义了MySQL应支持的SQL语法、数据校验等,这样可以更容易地在不同的环境中使用MySQL。全局s

Oracle 分组与拼接字符串同时使用

SELECTT.,ROWNUMIDFROM(SELECTT.EMPLID,T.NAME,T.BU,T.REALDEPART,T.FORMATDATE,SUM(T.S0)S0,MAX(UPDATETIME)CREATETIME,LISTAGG(TOCHAR(

MySQL部分从库上面因为大量的临时表tmp_table造成慢查询

背景描述Time:20190124T00:08:14.70572408:00User@Host:@Id:Schema:sentrymetaLast_errno:0Killed:0Query_time:0.315758Lock_

皕杰报表之UUID

​在我们用皕杰报表工具设计填报报表时,如何在新增行里自动增加id呢?能新增整数排序id吗?目前可以在新增行里自动增加id,但只能用uuid函数增加UUID编码,不能新增整数排序id。uuid函数说明:获取一个UUID,可以在填报表中用来创建数据ID语法:uuid()或uuid(sep)参数说明:sep布尔值,生成的uuid中是否包含分隔符'',缺省为

手写Java HashMap源码

HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程22

2020年前端实用代码段,为你的工作保驾护航

有空的时候,自己总结了几个代码段,在开发中也经常使用,谢谢。1、使用解构获取json数据let jsonData  id: 1,status: "OK",data: 'a', 'b';let  id, status, data: number   jsonData;console.log(id, status, number )