KTV歌曲推荐

前言

上一篇使用逻辑回归预测了用户性别,由于矩阵比较稀疏所以会影响训练速度。所以考虑降维,降维方案有很多,本次只考虑PCA和SVD。

PCA和SVD原理

有兴趣的可以自己去研究一下 https://medium.com/@jonathan_hui/machine-learning-singular-value-decomposition-svd-principal-component-analysis-pca-1d45e885e491

我简述一下:

  • PCA是将高维数据映射到低维坐标系中,让数据尽量稀疏
  • SVD就是非方阵的PCA
  • 实际使用中SVD和PCA并无太大区别
  • 如果特征大于数据记录数,并不能有好的效果,具体原因自己可以去看。

代码

数据获取和处理

以前文章写过很多次,这里略过 原数据shape为:2000*1900

PCA和矩阵转换

查看最佳维度数

1%matplotlib inline 2import numpy as np 3import matplotlib.pyplot as plt 4from sklearn.decomposition import PCA 5pca = PCA().fit(song_hot_matrix) 6plt.plot(np.cumsum(pca.explained_variance_ratio_)) 7plt.xlabel('number of components') 8plt.ylabel('cumulative explained variance');

从图中可以看出大概1500维度已经可以达到90+解释性

保留99%矩阵解释性

1pca = PCA(n_components=0.99, whiten=True) 2song_hot_matrix_pca = pca.fit_transform(song_hot_matrix)

得到压缩后特征为: 2000*1565 并没有压缩多少

模型训练

1import os 2os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" # see issue #152 3os.environ["CUDA_VISIBLE_DEVICES"] = "" 4 5import numpy as np 6from keras.models import Sequential 7from keras.layers import Dense, Activation, Embedding,Flatten,Dropout 8import matplotlib.pyplot as plt 9from keras.utils import np_utils 10from sklearn import datasets 11from sklearn.model_selection import train_test_split 12 13n_class=user_decades_encoder.get_class_count() 14song_count=song_label_encoder.get_class_count() 15print(n_class) 16print(song_count) 17 18train_X,test_X, train_y, test_y = train_test_split(song_hot_matrix_pca, 19 decades_hot_matrix, 20 test_size = 0.2, 21 random_state = 0) 22train_count = np.shape(train_X)[0] 23# 构建神经网络模型 24model = Sequential() 25model.add(Dense(input_dim=song_hot_matrix_pca.shape[1], units=n_class)) 26model.add(Activation('softmax')) 27 28# 选定loss函数和优化器 29model.compile(loss='categorical_crossentropy', optimizer='sgd', metrics=['accuracy']) 30 31# 训练过程 32print('Training -----------') 33for step in range(train_count): 34 scores = model.train_on_batch(train_X, train_y) 35 if step % 50 == 0: 36 print("训练样本 %d 个, 损失: %f, 准确率: %f" % (step, scores[0], scores[1]*100)) 37print('finish!')

训练结果:

1训练样本 4750, 损失: 0.371499, 准确率: 83.207470 2训练样本 4800, 损失: 0.381518, 准确率: 82.193959 3训练样本 4850, 损失: 0.364363, 准确率: 83.763909 4训练样本 4900, 损失: 0.378466, 准确率: 82.551670 5训练样本 4950, 损失: 0.391976, 准确率: 81.756759 6训练样本 5000, 损失: 0.378810, 准确率: 83.505565

测试集验证:

1# 准确率评估 2from sklearn.metrics import classification_report 3scores = model.evaluate(test_X, test_y, verbose=0) 4print("%s: %.2f%%" % (model.metrics_names[1], scores[1]*100)) 5 6 7Y_test = np.argmax(test_y, axis=1) 8y_pred = model.predict_classes(song_hot_matrix_pca.transform(test_X)) 9print(classification_report(Y_test, y_pred))

accuracy: 50.20%

很明显已经过拟合

处理过拟合-增加Dropout

这里使用加Dropout,随机丢弃特征的方式处理过拟合,代码:

1# 构建神经网络模型 2model = Sequential() 3model.add(Dropout(0.5)) 4model.add(Dense(input_dim=song_hot_matrix_pca.shape[1], units=n_class)) 5model.add(Activation('softmax'))

accuracy:70%

处理过拟合-L1L2正则

这里给权重增加正则

1# 构建神经网络模型 2model = Sequential() 3model.add(Dense(input_dim=song_hot_matrix_pca.shape[1], units=n_class, kernel_regularizer=regularizers.l2(0.01))) 4model.add(Activation('softmax'))

accuracy:62%

Well Done

其实SVD的做法与PCA类似,这里不再演示。经过我测试发现,在我的数据集上,PCA虽然加快了训练速度,但是丢弃了太多特征,导致数据很容易过拟合。加入Dropout或者增加正则相可以改善过拟合的情况,下一篇会分享自编码降维。

点赞
收藏

评论区

加载中...

相关推荐

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 )