MXNet动手学深度学习笔记:Gluon实现正则化

1#coding:utf-8 2''' 3正则化 4''' 5import mxnet as mx 6from mxnet import gluon 7from mxnet import ndarray 8from mxnet import autograd 9import numpy as np 10import matplotlib.pyplot as plt 11from mxnet import nd 12import random 13 14num_train = 20 15num_test = 100 16num_inputs = 200 17#模型真实参数 18true_w = nd.ones((num_inputs,1)) 19true_b = 0.05 20 21#生成测试和训练数据 22X = nd.random.normal(shape=(num_train+num_test,num_inputs)) 23y = nd.dot(X,true_w) + true_b 24y += 0.01 * nd.random.normal(shape=y.shape) 25 26X_train,X_test = X[:num_train,:],X[num_train:,:] 27y_train,y_test = y[:num_train],y[num_train:] 28 29batch_size = 1 30dataset_train = gluon.data.ArrayDataset(X_train, y_train) 31data_iter_train = gluon.data.DataLoader(dataset_train, batch_size,shuffle=True) 32 33# 定义损失函数 34square_loss = gluon.loss.L2Loss() 35 36# 定义测试函数 37def test(net,X,y): 38 return square_loss(net(X),y).mean().asscalar() 39 40# 定义训练函数 41def train(weight_decay): 42 epochs = 10 43 learning_rate = 0.005 44 net = gluon.nn.Sequential() 45 with net.name_scope(): 46 net.add(gluon.nn.Dense(1)) 47 48 net.collect_params().initialize() 49 trainer = gluon.Trainer(net.collect_params(),'sgd', 50 {'learning_rate':learning_rate,'wd':weight_decay}) 51 52 train_loss = [] 53 test_loss = [] 54 55 for e in range(epochs): 56 for data,label in data_iter_train: 57 with autograd.record(): 58 output = net(data) 59 loss = square_loss(output,label) 60 loss.backward() 61 trainer.step(batch_size) 62 63 train_loss.append(test(net,X_train,y_train)) 64 test_loss.append(test(net,X_test,y_test)) 65 66 67 plt.plot(train_loss) 68 plt.plot(test_loss) 69 plt.legend(['train','test']) 70 plt.show() 71 72# 未使用正则化 73# train(0) 74 75# 使用正则化 76train(5)
点赞
收藏

评论区

加载中...

相关推荐

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 )