KTV歌曲推荐

前言

上一篇写了推荐系统最古老的的一种算法叫协同过滤,古老并不是不好用,其实还是很好用的一种算法,随着时代的进步,出现了神经网络和因子分解等更优秀的算法解决不同的问题。 这里主要说一下逻辑回归,逻辑回归主要用于打分的预估。我这里没有打分的数据所以用性别代替。 这里的例子就是用歌曲列表预判用户性别。

什么是逻辑回归

逻辑回归的资料比较多,我比较推荐大家看刷一下bilibili上李宏毅老师的视频,这里我只说一些需要注意的点。

网络结构

逻辑回归可以理解为一种单层神经网络,网络结构如图:

激活函数选择

逻辑回归一般选sigmoid或者softmax

  • 图的上半部分就是二元逻辑回归激活函数是sigmoid
  • 图的下半部分是多元逻辑回归没有激活函数直接接了一个softmax

别问我啥是sigmoid啥是softmax,问就是百度。

损失函数选择

损失函数逻辑回归常用的有三种(其实有很多不止三种,自己查API喽):

  • binary_crossentropy
  • categorical_crossentropy
  • sparse_categorical_crossentrop 这里其实用binary更合适,但是我这里选的categorical_crossentropy,因为我懒得改了,而且我后面会做其他功能

梯度下降选择

梯度下降方式有很多,我这里选择随机梯度下降,sgd其实我觉得adam更合适,看大家心情了。至于为啥

数据准备

这次的数据是1万条KTV唱歌数据,别问我数据哪来的。问就是别人给的。

X是用户唱歌数据的one-hot

Y是用户的性别one-hot

下面是真正的技术

代码实现

  • 数据拆分为 80%训练 20%测试
  • 这里虽然只有两类但是还是用了softmax,不影响
  • 训练工具是keras

数据获取

下面代码都干了些啥呢,主要是两个matrix。

一个是用户唱歌的onehot->song_hot_matrix。

一个是用户性别的onehot->decades_hot_matrix。 代码不重要,主要看字。

1import elasticsearch 2import elasticsearch.helpers 3import re 4import numpy as np 5import operator 6import datetime 7 8 9es_client = elasticsearch.Elasticsearch(hosts=["localhost:9200"]) 10 11def trim_song_name(song_name): 12 """ 13 处理歌名,过滤掉无用内容和空白 14 """ 15 song_name = song_name.strip() 16 song_name = re.sub("【.*?】", "", song_name) 17 song_name = re.sub("(.*?)", "", song_name) 18 return song_name 19 20def trim_address_name(address_name): 21 """ 22 处理地址 23 """ 24 return str(address_name).strip() 25 26def get_data(size=0): 27 """ 28 获取uid=>作品名list的字典 29 """ 30 cur_size=0 31 song_dic = {} 32 user_address_dic = {} 33 user_decades_dic = {} 34 35 search_result = elasticsearch.helpers.scan( 36 es_client, 37 index="ktv_user_info", 38 doc_type="ktv_works", 39 scroll="10m", 40 query={ 41 "query":{ 42 "range": { 43 "birthday": { 44 "gt": 63072662400 45 } 46 } 47 } 48 } 49 ) 50 51 for hit_item in search_result: 52 cur_size += 1 53 if size>0 and cur_size>size: 54 break 55 56 user_info = hit_item["_source"] 57 item = get_work_info(hit_item["_id"]) 58 if item is None: 59 continue 60 61 work_list = item['item_list'] 62 if len(work_list)<2: 63 continue 64 65 if user_info['gender']==0: 66 continue 67 if user_info['gender']==1: 68 user_info['gender']="男" 69 if user_info['gender']==2: 70 user_info['gender']="女" 71 72 song_dic[item['uid']] = [trim_song_name(item['songname']) for item in work_list] 73 74 75 user_decades_dic[item['uid']] = user_info['gender'] 76 user_address_dic[item['uid']] = trim_address_name(user_info['address']) 77 78 return (song_dic, user_address_dic, user_decades_dic) 79 80def get_user_info(uid): 81 """ 82 获取用户信息 83 """ 84 ret = es_client.get( 85 index="ktv_user_info", 86 doc_type="ktv_works", 87 id=uid 88 ) 89 return ret['_source'] 90 91def get_work_info(uid): 92 """ 93 获取用户信息 94 """ 95 try: 96 ret = es_client.get( 97 index="ktv_works", 98 doc_type="ktv_works", 99 id=uid 100 ) 101 return ret['_source'] 102 except Exception as ex: 103 return None 104 105 106def get_uniq_song_sort_list(song_dict): 107 """ 108 合并重复歌曲并按歌曲名排序 109 """ 110 return sorted(list(set(np.concatenate(list(song_dict.values())).tolist()))) 111 112from sklearn import preprocessing 113%run label_encoder.ipynb 114 115user_count = 4000 116song_count = 0 117 118# 获得用户唱歌数据 119song_dict, user_address_dict, user_decades_dict = get_data(user_count) 120 121# 歌曲字典 122song_label_encoder = LabelEncoder() 123song_label_encoder.fit_dict(song_dict, "", True) 124song_hot_matrix = song_label_encoder.encode_hot_dict(song_dict, True) 125 126user_decades_encoder = LabelEncoder() 127user_decades_encoder.fit_dict(user_decades_dict) 128decades_hot_matrix = user_decades_encoder.encode_hot_dict(user_decades_dict, False)

song_hot_matrix

uid

洗刷刷

麻雀

你的答案

0

0

1

0

1

1

1

0

2

1

0

0

3

0

0

0

decades_hot_matrix

uid

0

1

0

1

0

1

2

1

0

3

0

1

模型训练

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

准确率测试集评估

数据训练完了用拆分出来的20%数据测试一下:

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)) 5Y_test = np.argmax(test_y, axis=1) 6y_pred = model.predict_classes(test_X) 7print(classification_report(Y_test, y_pred))

输出:

1accuracy: 78.43% 2 precision recall f1-score support 3 4 0 0.72 0.90 0.80 220 5 1 0.88 0.68 0.77 239 6 7 accuracy 0.78 459 8 macro avg 0.80 0.79 0.78 459 9weighted avg 0.80 0.78 0.78 459

人工测试

然后让小伙伴们一起来玩耍,嗯准确率100%,完美!

1def pred(song_list=[]): 2 blong_hot_matrix = song_label_encoder.encode_hot_dict({"bblong":song_list}, True) 3 y_pred = model.predict_classes(blong_hot_matrix) 4 return user_decades_encoder.decode_list(y_pred) 5 6# # 男A 7# print(pred(["一路向北", "暗香", "菊花台"])) 8# # 男B 9# print(pred(["不要说话", "平凡之路", "李白"])) 10# # 女A 11# print(pred(["知足", "被风吹过的夏天", "龙卷风"])) 12# # 男C 13# print(pred(["情人","再见","无赖","离人","你的样子"])) 14# # 男D 15# print(pred(["小情歌","我好想你","无与伦比的美丽"])) 16# # 男E 17# print(pred(["忐忑","最炫民族风","小苹果"]))
点赞
收藏

评论区

加载中...

相关推荐

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

mysql设置时区

mysql设置时区mysql\_query("SETtime\_zone'8:00'")ordie('时区设置失败,请联系管理员!');中国在东8区所以加8方法二:selectcount(user\_id)asdevice,CONVERT\_TZ(FROM\_UNIXTIME(reg\_time),'08:00','0