Implementing a CNN for Text Classification in TensorFlow(用tensorflow实现CNN文本分类) 阅读笔记

本文主要是介绍Implementing a CNN for Text Classification in TensorFlow(用tensorflow实现CNN文本分类) 阅读笔记,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

    目前正在学习把深度学习应用到NLP,主要是看些论文和博客,同时做些笔记方便理解,还没入门很多东西还不懂,一知半解。贴出来的原因,一是方便自己查看,二是希望大家指点一下,尽快入门。

    原paper:Convolutional Neural Networks for Sentence Classification

    源代码:https://github.com/dennybritz/cnn-text-classification-tf

    原博客:http://www.wildml.com/2015/12/implementing-a-cnn-for-text-classification-in-tensorflow/


    1. 数据和预处理

      1. 数据集:电影评论数据——Movie Review data from Rotten Tomatoes,包含5331个积极的评论和5331个消极评论,同时包含一个20k的词表

      2. 注意:数据集过小容易过拟合,可以进行10交叉验证

      3. 步骤:

        1. 加载两类数据

        2. 文本数据清洗

        3. 把每个句子填充到最大的句子长度,填充字符是<PAD>,使得每个句子都包含59个单词。相同的长度有利于进行高效的批处理

        4. 根据所有单词的词表,建立一个索引,用一个整数代表一个词,则每个句子由一个整数向量表示

    2. 模型

      1. 第一层把词嵌入到低纬向量;第二层用多个不同大小的filter进行卷积;第三层用max-pool把第二层多个filter的结果转换成一个长的特征向量并加入dropout正规化;第四层用softmax进行分类。

      2. 简化模型,方便理解:

        1. 不适用预训练的word2vec的词向量,而是学习如何嵌入

        2. 不对权重向量强制执行L2正规化

        3. 原paper使用静态词向量和非静态词向量两个同道作为输入,这里只使用一种同道作为输入

    3. 实现

      1. TextCNN类,参数如下:

        1. sequence_length:句子长度,把每个句子统一填充到59个单词

        2. num_classes:输出的类型个数,这里是积极和消极两类

        3. vocab_size:词典长度,需要在嵌入层定义

        4. embeding_size :嵌入的维度

        5. filter_sizes:卷积核的高度

        6. num_filters:每种不同大小的卷积核的个数,这里每种有3个

      2. 输入占位符(定义我们要传给网络的数据)

        1. 如输入占位符,输出占位符和dropout占位符

        2. tf.placeholder创建一个占位符,在训练和测试时才会传入相应的数据。第一个参数是数据类型;第二个参数是tensor的格式,none表示是任何大小;第三个参数是名称

        3. dropout_keep_prob是保留一个神经元的概率,这个概率只在训练的时候用到

      3. 第一层(嵌入层)

        1. tf.device("/cpu:0")使用cpu进行操作,因为tensorflow当gpu可用时默认使用gpu,但是embedding不支持gpu实现,所以使用CPU操作

        2. tf.name_scope,把所有操作加到命名为embedding的顶层节点,用于可视化网络视图

        3. W是我们在训练时得到的嵌入矩阵,通过随机均匀分布进行初始化

        4. tf.nn.embedding_lookup 是真正的embedding操作,结果是一个三维的tensor,[None, sequence_length, embedding_size]

        5. 因为卷积操作conv2d需要4个维度的tensor所以需要给embedding结果增加一个维度,得到[None, sequence_length, embedding_size, 1]

      4. 卷积和max-pooling

        1. 对不同大小的filter建立不同的卷积层,W是卷积的输入矩阵,h是使用relu进行卷积的结果。

        2. “VALID”表示使用narrow卷积,得到的结果大小为[1, sequence_length - filter_size + 1, 1, 1]

        3. 为了更容易理解,需要计算输入输出的大小:"VALID" padding means that we slide the filter over our sentence without padding the edges, performing a narrow convolution that gives us an output of shape[1, sequence_length - filter_size + 1, 1, 1]. Performing max-pooling over the output of a specific filter size leaves us with a tensor of shape[batch_size, 1, 1, num_filters]. This is essentially a feature vector, where the last dimension corresponds to our features. Once we have all the pooled output tensors from each filter size we combine them into one long feature vector of shape[batch_size, num_filters_total]. Using-1 intf.reshape tells TensorFlow to flatten the dimension when possible.

      5. Dropout层

        1. dropout是正规化卷积神经网络最流行的方法,即随机禁用一些神经元

      6. 分数和预测

        1. 用max-pooling得到的向量作为x作为输入,与随机产生的W权重矩阵进行计算得到分数,选择分数高的作为预测类型结果

      7. 交叉熵损失和正确率

      8. 网络可视化

      9. 训练过程

        1. Session是执行graph操作(表示计算任务)的上下文环境,包含变量和序列的状态。每个session执行一个graph。tensorflow包含了默认session,也可以自定义session然后通过session.as_default() 设置为默认视图

        2. graph包含操作和tensors(表示数据),可以在程序中建立多个图,但是通常只需一个图。同一个图可以在多个session中使用,但是不能多个图在一个session中使用。

        3. allow_soft_placement可以在不存在预设运行设备时可以在其他设备运行,例如设置在gpu上运行的操作,当没有gpu时allow_soft_placement使得可以在cpu操作

        4. log_device_placement用于设备的log,方便debugging

        5. FLAGS是程序的命令行输入

      10. CNN初始化和最小化loss

        1. 按照TextCNN的参数进行初始化

        2. tensorflow提供了几种自带的优化器,我们使用Adam优化器求loss的最小值

        3. train_op就是训练步骤,每次更新我们的参数,global_step用于记录训练的次数,在tensorflow中自增

      11. summaries汇总

        1. tensorflow提供了各方面的汇总信息,方便跟踪和可视化训练和预测的过程。summaries是一个序列化的对象,通过SummaryWriter写入到光盘

      12. checkpointing检查点

        1. 用于保存训练参数,方便选择最优的参数,使用tf.train.saver()进行保存

      13. 变量初始化

        1. sess.run(tf.initialize_all_variables()),用于初始化所有我们定义的变量,也可以对特定的变量手动调用初始化,如预训练好的词向量

      14. 定义单一的训练步骤

        1. 定义一个函数用于模型评价、更新批量数据和更新模型参数

        2. feed_dict中包含了我们在网络中定义的占位符的数据,必须要对所有的占位符进行赋值,否则会报错

        3. train_op不返回结果,只是更新网络的参数

      15. 训练循环

        1. 遍历数据并对每次遍历数据调用train_step函数,并定期打印模型评价和检查点

      16. 用tensorboard进行结果可视化

        1. python tensorflow/tensorboard/tensorboard.py --logdir=path/to/log-directory
        2. 问题是没找到tensorboard.py文件,找了半天发现在/home/pyx/.local/lib/python3.5/site-package/tensorflow中,但是报warming,可以忽略
      17. 本实验的几个问题
        1. 训练的指标不是平滑的,原因是我们每个批处理的数据过少
        2. 训练集正确率过高,测试集正确率过低,过拟合。避免过拟合:更多的数据;更强的正规化;更少的模型参数。例如对最后一层的权重进行L2惩罚,使得正确率提升到76%,接近原始paper


                        
















    这篇关于Implementing a CNN for Text Classification in TensorFlow(用tensorflow实现CNN文本分类) 阅读笔记的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



    http://www.chinasem.cn/article/456214

    相关文章

    使用animation.css库快速实现CSS3旋转动画效果

    《使用animation.css库快速实现CSS3旋转动画效果》随着Web技术的不断发展,动画效果已经成为了网页设计中不可或缺的一部分,本文将深入探讨animation.css的工作原理,如何使用以及... 目录1. css3动画技术简介2. animation.css库介绍2.1 animation.cs

    Java进行日期解析与格式化的实现代码

    《Java进行日期解析与格式化的实现代码》使用Java搭配ApacheCommonsLang3和Natty库,可以实现灵活高效的日期解析与格式化,本文将通过相关示例为大家讲讲具体的实践操作,需要的可以... 目录一、背景二、依赖介绍1. Apache Commons Lang32. Natty三、核心实现代

    SpringBoot实现接口数据加解密的三种实战方案

    《SpringBoot实现接口数据加解密的三种实战方案》在金融支付、用户隐私信息传输等场景中,接口数据若以明文传输,极易被中间人攻击窃取,SpringBoot提供了多种优雅的加解密实现方案,本文将从原... 目录一、为什么需要接口数据加解密?二、核心加解密算法选择1. 对称加密(AES)2. 非对称加密(R

    基于Go语言实现Base62编码的三种方式以及对比分析

    《基于Go语言实现Base62编码的三种方式以及对比分析》Base62编码是一种在字符编码中使用62个字符的编码方式,在计算机科学中,,Go语言是一种静态类型、编译型语言,它由Google开发并开源,... 目录一、标准库现状与解决方案1. 标准库对比表2. 解决方案完整实现代码(含边界处理)二、关键实现细

    python通过curl实现访问deepseek的API

    《python通过curl实现访问deepseek的API》这篇文章主要为大家详细介绍了python如何通过curl实现访问deepseek的API,文中的示例代码讲解详细,感兴趣的小伙伴可以跟随小编... API申请和充值下面是deepeek的API网站https://platform.deepsee

    SpringBoot实现二维码生成的详细步骤与完整代码

    《SpringBoot实现二维码生成的详细步骤与完整代码》如今,二维码的应用场景非常广泛,从支付到信息分享,二维码都扮演着重要角色,SpringBoot是一个非常流行的Java基于Spring框架的微... 目录一、环境搭建二、创建 Spring Boot 项目三、引入二维码生成依赖四、编写二维码生成代码五

    MyBatisX逆向工程的实现示例

    《MyBatisX逆向工程的实现示例》本文主要介绍了MyBatisX逆向工程的实现示例,文中通过示例代码介绍的非常详细,对大家的学习或者工作具有一定的参考学习价值,需要的朋友们下面随着小编来一起学习学... 目录逆向工程准备好数据库、表安装MyBATisX插件项目连接数据库引入依赖pom.XML生成实体类、

    C#实现查找并删除PDF中的空白页面

    《C#实现查找并删除PDF中的空白页面》PDF文件中的空白页并不少见,因为它们有可能是作者有意留下的,也有可能是在处理文档时不小心添加的,下面我们来看看如何使用Spire.PDFfor.NET通过C#... 目录安装 Spire.PDF for .NETC# 查找并删除 PDF 文档中的空白页C# 添加与删

    Java实现MinIO文件上传的加解密操作

    《Java实现MinIO文件上传的加解密操作》在云存储场景中,数据安全是核心需求之一,MinIO作为高性能对象存储服务,支持通过客户端加密(CSE)在数据上传前完成加密,下面我们来看看如何通过Java... 目录一、背景与需求二、技术选型与原理1. 加密方案对比2. 核心算法选择三、完整代码实现1. 加密上

    Java使用WebView实现桌面程序的技术指南

    《Java使用WebView实现桌面程序的技术指南》在现代软件开发中,许多应用需要在桌面程序中嵌入Web页面,例如,你可能需要在Java桌面应用中嵌入一部分Web前端,或者加载一个HTML5界面以增强... 目录1、简述2、WebView 特点3、搭建 WebView 示例3.1 添加 JavaFX 依赖3