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

    相关文章

    Java中使用Java Mail实现邮件服务功能示例

    《Java中使用JavaMail实现邮件服务功能示例》:本文主要介绍Java中使用JavaMail实现邮件服务功能的相关资料,文章还提供了一个发送邮件的示例代码,包括创建参数类、邮件类和执行结... 目录前言一、历史背景二编程、pom依赖三、API说明(一)Session (会话)(二)Message编程客

    Java中List转Map的几种具体实现方式和特点

    《Java中List转Map的几种具体实现方式和特点》:本文主要介绍几种常用的List转Map的方式,包括使用for循环遍历、Java8StreamAPI、ApacheCommonsCollect... 目录前言1、使用for循环遍历:2、Java8 Stream API:3、Apache Commons

    C#提取PDF表单数据的实现流程

    《C#提取PDF表单数据的实现流程》PDF表单是一种常见的数据收集工具,广泛应用于调查问卷、业务合同等场景,凭借出色的跨平台兼容性和标准化特点,PDF表单在各行各业中得到了广泛应用,本文将探讨如何使用... 目录引言使用工具C# 提取多个PDF表单域的数据C# 提取特定PDF表单域的数据引言PDF表单是一

    使用Python实现高效的端口扫描器

    《使用Python实现高效的端口扫描器》在网络安全领域,端口扫描是一项基本而重要的技能,通过端口扫描,可以发现目标主机上开放的服务和端口,这对于安全评估、渗透测试等有着不可忽视的作用,本文将介绍如何使... 目录1. 端口扫描的基本原理2. 使用python实现端口扫描2.1 安装必要的库2.2 编写端口扫

    PyCharm接入DeepSeek实现AI编程的操作流程

    《PyCharm接入DeepSeek实现AI编程的操作流程》DeepSeek是一家专注于人工智能技术研发的公司,致力于开发高性能、低成本的AI模型,接下来,我们把DeepSeek接入到PyCharm中... 目录引言效果演示创建API key在PyCharm中下载Continue插件配置Continue引言

    MySQL分表自动化创建的实现方案

    《MySQL分表自动化创建的实现方案》在数据库应用场景中,随着数据量的不断增长,单表存储数据可能会面临性能瓶颈,例如查询、插入、更新等操作的效率会逐渐降低,分表是一种有效的优化策略,它将数据分散存储在... 目录一、项目目的二、实现过程(一)mysql 事件调度器结合存储过程方式1. 开启事件调度器2. 创

    使用Python实现操作mongodb详解

    《使用Python实现操作mongodb详解》这篇文章主要为大家详细介绍了使用Python实现操作mongodb的相关知识,文中的示例代码讲解详细,感兴趣的小伙伴可以跟随小编一起学习一下... 目录一、示例二、常用指令三、遇到的问题一、示例from pymongo import MongoClientf

    SQL Server使用SELECT INTO实现表备份的代码示例

    《SQLServer使用SELECTINTO实现表备份的代码示例》在数据库管理过程中,有时我们需要对表进行备份,以防数据丢失或修改错误,在SQLServer中,可以使用SELECTINT... 在数据库管理过程中,有时我们需要对表进行备份,以防数据丢失或修改错误。在 SQL Server 中,可以使用 SE

    基于Go语言实现一个压测工具

    《基于Go语言实现一个压测工具》这篇文章主要为大家详细介绍了基于Go语言实现一个简单的压测工具,文中的示例代码讲解详细,感兴趣的小伙伴可以跟随小编一起学习一下... 目录整体架构通用数据处理模块Http请求响应数据处理Curl参数解析处理客户端模块Http客户端处理Grpc客户端处理Websocket客户端

    Java CompletableFuture如何实现超时功能

    《JavaCompletableFuture如何实现超时功能》:本文主要介绍实现超时功能的基本思路以及CompletableFuture(之后简称CF)是如何通过代码实现超时功能的,需要的... 目录基本思路CompletableFuture 的实现1. 基本实现流程2. 静态条件分析3. 内存泄露 bug