回答多选项问题,使用softmax函数,对数几率回归在多个可能不同值上的推广。函数返回值是C个分量的概率向量,每个分量对应一个输出类别概率。分量为概率,C个分量和始终为1。每个样本必须属于某个输出类别,所有可能样本均被覆盖。分量和小于1,存在隐藏类别;分量和大于1,每个样本可能同时属于多个类别。类别数量为2,输出概率与对数几率回归模型输出相同。

变量初始化,需要C个不同权值组,每个组对应一个可能输出,使用权值矩阵。每行与输入特征对应,每列与输出类别对应。

鸢尾花数据集Iris,包含4个数据特征、3类可能输出,权值矩阵4X3。

训练样本每个输出类别损失相加。训练样本期望类别为1,其他为0。只有一个损失值被计入,度量模型为真实类别预测的概率可信度。每个训练样本损失相加,得到训练集总损失值。TensorFlow的softmax交叉熵函数,sparse_softmax_cross_entropy_with_logits版本针对训练集每个样本只对应单个类别优化,softmax_cross_entropy_with_logits版本可使用包含每个样本属于每个类别的概率信息的训练集。模型最终输出是单个类别值。

不需要每个类别都转换独立变量,需要把值转换为0~2整数(总类别数3)。tf.stack创建张量,tf.equal把文件输入与每个可能值比较。tf.argmax找到张量值为真的位置。

推断过程计算测试样本属于每个类别概率。tf. argmax函数选择预测输出值最大概率类别。tf.equal与期望类别比较。tf.reduce_meen计算准确率。

  1. import tensorflow as tf#导入TensorFlow库
  2. import os#导入OS库
  3. W = tf.Variable(tf.zeros([4, 3]), name="weights")#变量权值,矩阵,每个特征权值列对应一个输出类别
  4. b = tf.Variable(tf.zeros([3], name="bias"))#模型偏置,每个偏置对应一个输出类别
  5. def combine_inputs(X):#输入值合并
  6. print "function: combine_inputs"
  7. return tf.matmul(X, W) + b
  8. def inference(X):#计算返回推断模型输出(数据X)
  9. print "function: inference"
  10. return tf.nn.softmax(combine_inputs(X))#调用softmax分类函数
  11. def loss(X, Y):#计算损失(训练数据X及期望输出Y)
  12. print "function: loss"
  13. return tf.reduce_mean(tf.nn.sparse_softmax_cross_entropy_with_logits(logits=combine_inputs(X), labels=Y))#求平均值,针对每个样本只对应单个类别优化
  14. def read_csv(batch_size, file_name, record_defaults):#从csv文件读取数据,加载解析,创建批次读取张量多行数据
  15. filename_queue = tf.train.string_input_producer([os.path.join(os.getcwd(), file_name)])
  16. reader = tf.TextLineReader(skip_header_lines=1)
  17. key, value = reader.read(filename_queue)
  18. decoded = tf.decode_csv(value, record_defaults=record_defaults)#字符串(文本行)转换到指定默认值张量列元组,为每列设置数据类型
  19. return tf.train.shuffle_batch(decoded, batch_size=batch_size, capacity=batch_size * 50, min_after_dequeue=batch_size)#读取文件,加载张量batch_size
  20. def inputs():#读取或生成训练数据X及期望输出Y
  21. print "function: inputs"
  22. #数据来源:https://archive.ics.uci.edu/ml/datasets/Iris
  23. #iris.data改为iris.csv,增加sepal_length, sepal_width, petal_length, petal_width, label字段行首行
  24. sepal_length, sepal_width, petal_length, petal_width, label =\
  25. read_csv(100, "iris.csv", [[0.0], [0.0], [0.0], [0.0], [""]])
  26. #转换属性数据
  27. label_number = tf.to_int32(tf.argmax(tf.to_int32(tf.stack([
  28. tf.equal(label, ["Iris-setosa"]),
  29. tf.equal(label, ["Iris-versicolor"]),
  30. tf.equal(label, ["Iris-virginica"])
  31. ])), 0))#将类名称转抽象为从0开始的类别索引
  32. features = tf.transpose(tf.stack([sepal_length, sepal_width, petal_length, petal_width]))#特征装入矩阵,转置,每行一样本,每列一特征
  33. return features, label_number
  34. def train(total_loss):#训练或调整模型参数(计算总损失)
  35. print "function: train"
  36. learning_rate = 0.01
  37. return tf.train.GradientDescentOptimizer(learning_rate).minimize(total_loss)
  38. def evaluate(sess, X, Y):#评估训练模型
  39. print "function: evaluate"
  40. predicted = tf.cast(tf.arg_max(inference(X), 1), tf.int32)#选择预测输出值最大概率类别
  41. print sess.run(tf.reduce_mean(tf.cast(tf.equal(predicted, Y), tf.float32)))#统计所有正确预测样本数,除以批次样本总数,得到正确预测百分比
  42. with tf.Session() as sess:#会话对象启动数据流图,搭建流程
  43. print "Session: start"
  44. tf.global_variables_initializer().run()
  45. X, Y = inputs()
  46. total_loss = loss(X, Y)
  47. train_op = train(total_loss)
  48. coord = tf.train.Coordinator()
  49. threads = tf.train.start_queue_runners(sess=sess, coord=coord)
  50. training_steps = 1000#实际训练迭代次数
  51. for step in range(training_steps):#实际训练闭环
  52. sess.run([train_op])
  53. if step % 10 == 0:#查看训练过程损失递减
  54. print str(step)+ " loss: ", sess.run([total_loss])
  55. print str(training_steps) + " final loss: ", sess.run([total_loss])
  56. evaluate(sess, X, Y)#模型评估
  57. coord.request_stop()
  58. coord.join(threads)
  59. sess.close()

参考资料:
《面向机器智能的TensorFlow实践》

欢迎加我微信交流:qingxingfengzi
我的微信公众号:qingxingfengzigz
我老婆张幸清的微信公众号:qingqingfeifangz

学习笔记TF010:softmax分类的更多相关文章

  1. 学习笔记TF019:序列分类、IMDB影评分类

    序列分类,预测整个输入序列的类别标签.情绪分析,预测用户撰写文字话题态度.预测选举结果或产品.电影评分. 国际电影数据库(International Movie Database)影评数据集.目标值二 ...

  2. CSS权威指南学习笔记 —— HTML元素分类

    HTML文档由各种元素组成.比如,p.table.span等等.每个元素都会对文档的表现有所影响.CSS中,每个元素都会生成一个框(传说中的盒子),其中包含元素内容. 元素可以根据它的创建方式分为两种 ...

  3. CSS学习笔记之元素分类

    在讲解CSS布局之前,我们需要提前知道一些知识,在CSS中,html中的标签元素大体被分为三种不同的类型:块状元素.内联元素(又叫行内元素)和内联块状元素. 常用的块状元素有: <div> ...

  4. SAS学习笔记36 二分类logistic回归

    这里所拟合模型的AIC和SC统计量的值均小于只有截距的模型的相应统计量的值,说明含有自变量的模型较仅含有常数项的要好 但模型的最大重新换算 R 方为0.0993,说明模型拟合效果并不好,可能有其他危险 ...

  5. CNN学习笔记:目标函数

    CNN学习笔记:目标函数 分类任务中的目标函数 目标函数,亦称损失函数或代价函数,是整个网络模型的指挥棒,通过样本的预测结果与真实标记产生的误差来反向传播指导网络参数学习和表示学习. 假设某分类任务共 ...

  6. 学习笔记分享之汇编---3. 堆栈&标志寄存器

    前言:   此文章收录在本人的<学习笔记分享>分类中,此分类记录本人的学习心得体会,现全部分享出来希望和大家共同交流学习成长.附上分类链接:   https://www.cnblogs.c ...

  7. UFLDL深度学习笔记 (二)SoftMax 回归(矩阵化推导)

    UFLDL深度学习笔记 (二)Softmax 回归 本文为学习"UFLDL Softmax回归"的笔记与代码实现,文中略过了对代价函数求偏导的过程,本篇笔记主要补充求偏导步骤的详细 ...

  8. UFLDL深度学习笔记 (四)用于分类的深度网络

    UFLDL深度学习笔记 (四)用于分类的深度网络 1. 主要思路 本文要讨论的"UFLDL 建立分类用深度网络"基本原理基于前2节的softmax回归和 无监督特征学习,区别在于使 ...

  9. 卷积神经网络用语句子分类---Convolutional Neural Networks for Sentence Classification 学习笔记

    读了一篇文章,用到卷积神经网络的方法来进行文本分类,故写下一点自己的学习笔记: 本文在事先进行单词向量的学习的基础上,利用卷积神经网络(CNN)进行句子分类,然后通过微调学习任务特定的向量,提高性能. ...

随机推荐

  1. (转)Java线程面试题 Top 50

    原文链接:http://www.importnew.com/12773.html   本文由 ImportNew - 李 广 翻译自 javarevisited.欢迎加入Java小组.转载请参见文章末 ...

  2. 《ECMAScript标准入门》第二版读书笔记

    title: <ECMAScript标准入门>第二版 date: 2017-04-10 tags: JavaScript categories: Reading-note 2015年6月, ...

  3. Netd学习笔记

    service netd /system/bin/netd     class main     socket netd stream 0660 root system     socket dnsp ...

  4. JVM年轻代、年老代、永久代

    年轻代: HotSpot JVM把年轻代分为了三部分:1个Eden区和2个Survivor区(分别叫From和To),每次新创建对象时,都会分配到Eden区,当Eden区没有足够的空间进行分配时,虚拟 ...

  5. 极化SAR图像基础知识(1)

    从今天开始学习极化SAR图像,记录于此. 极化散射矩阵S是用来表示单个像素散射特性的一种简便办法,它包含了目标的全部极化信息.

  6. javascript核心概念——new

    如果完全没有编程经验的朋友看到这个词会想到什么? 上过幼儿园的都知道new表示 "新的" 的意思. var a = new Date() 按照字面的意思表示什么? 把一个新的dat ...

  7. C++中的类继承(4)继承种类之单继承&多继承&菱形继承

    单继承是一般的单一继承,一个子类只 有一个直接父类时称这个继承关系为单继承.这种关系比较简单是一对一的关系: 多继承是指 一个子类有两个或以上直接父类时称这个继承关系为多继承.这种继承方式使一个子类可 ...

  8. lua 模块

    lua 模块 概述 lua 模块类似于封装库 将相应功能封装为一个模块, 可以按照面向对象中的类定义去理解和使用 使用 模块文件示例程序 mod = {} mod.constant = "模 ...

  9. 百度UEditor图片上传或文件上传路径自定义

    最近在项目中使用到百度UEditor的图片以及文件上传功能,但在上传的时候路径总是按照预设规则来自动生成,不方便一些特殊文件的维护.于是开始查看文档和源代码,其实操作还是比较简单的,具体如下: 1.百 ...

  10. elasticsearch基础概念

    接近实时(NRT)        Elasticsearch是一个接近实时的搜索平台.这意味着,从索引一个文档直到这个文档能够被搜索到有一个轻微的延迟(通常是1秒).           集群(clu ...