TensorFlow(2) 使用TF构建多层感知机预测MNIST数据集
import tensorflow as tf import tensorflow.examples.tutorials.mnist.input_data as input_data import matplotlib.pyplot as plt #读入数据---------------------------------------------------------------------- mnist input_data.read_data_sets(MNIST_data/,one_hotTrue)#label为one-hot-encoding #查看读入的数据格式------------------------------------------------------------- print(train,mnist.train.num_examples, ,validation,mnist.validation.num_examples, ,test,mnist.test.num_examples) print(train images:,mnist.train.images.shape, labels:,mnist.train.labels.shape) #查看label的值 mnist.train.labels[0] #函数将读入的数据图像显示 def plot_image(image): #上面读入的图像是一维的用reshape转成矩阵输出 plt.imshow(image.reshape(28,28),cmapbinary) plt.show() plot_image(mnist.train.images[0]) #函数输出10张图像带标签label和预测值prediction def plot_images_labels_prediction(images, labels, prediction, idx, num10): figplt.gcf() fig.set_size_inches(12,14) if num25: num25 for i in range(0, num): ax plt.subplot(5,5,1i) ax.imshow(np.reshape(images[idx],(28,28)),cmapbinary) title labelstr(np.argmax(labels[idx])) if len(prediction)0: title,predictionstr(prediction[idx]) ax.set_title(title, fontsize10) ax.set_xticks([]);ax.set_yticks([]) idx1 plt.show() plot_images_labels_prediction(mnist.validation.images, mnist.validation.labels,[],0) #构建MLP模型--------------------------------------------------------------------- #自定义layer def layer(output_dim, input_dim, inputs, activationNone): W tf.Variable(tf.random_normal([input_dim, output_dim])) b tf.Variable(tf.random_normal([1, output_dim])) WXb tf.matmul(inputs,W)b if activation is None: outputs WXb else: outputs activation(WXb) return outputs #两个隐藏层h1、h2一个输入x一个输出y_predict x tf.placeholder(float,[None, 784]) h1 layer(output_dim1000,input_dim784,inputsx,activationtf.nn.relu) h2 layer(output_dim1000,input_dim1000,inputsh1,activationtf.nn.relu) y_predict layer(output_dim10,input_dim1000,inputsh2,activationNone) #定义标签值 y_label tf.placeholder(float,[None,10]) #损失函数 loss_function tf.reduce_mean( tf.nn.softmax_cross_entropy_with_logits(logitsy_predict,labelsy_label)) #优化器使loss最小化学习率为0.001 optimizer tf.train.AdamOptimizer(learning_rate0.001).minimize(loss_function) #预测结果的正确性 correct_prediction tf.equal(tf.argmax(y_label,1),tf.argmax(y_predict,1)) #计算精度cast():转换值的类型,reduce_mean():计算平均值 accuracy tf.reduce_mean(tf.cast(correct_prediction,float)) #开始训练--------------------------------------------------------------------- trainEpochs 15 #epoch batchSize 100 totalBatchs int(len(mnist.train.images)/batchSize) #一个周期总的批次 loss_list[];epoch_list[];accuracy_list[] #记录训练过程的loss,epoch,accuracy from time import time startTime time() sess tf.Session() sess.run(tf.global_variables_initializer()) for epoch in range(trainEpochs): for i in range(totalBatchs): batch_x, batch_y mnist.train.next_batch(batchSize) #读取下一个批次的数据循环读取 sess.run(optimizer, feed_dict{x:batch_x,y_label:batch_y}) loss, acc sess.run([loss_function,accuracy],feed_dict\ {x:mnist.validation.images,y_label:mnist.validation.labels}) epoch_list.append(epoch) loss_list.append(loss) accuracy_list.append(acc) print(Train Epoch:, %02d%(epoch1), Loss,\ {:.9f}.format(loss), Accuary,acc) duration time()-startTime print(Train Finished takes:, duration) #训练结果显示-------------------------------------------------------------------- %matplotlib inline fig plt.gcf()#获取当前的figure图 fig.set_size_inches(4,2) plt.plot(epoch_list, loss_list, labelloss) plt.ylabel(loss) plt.xlabel(epoch) plt.legend([loss], locupper right) plt.plot(epoch_list,accuracy_list,labelaccuracy) fig plt.gcf() fig.set_size_inches(4,2) plt.ylim(0.8,1)#设置y轴范围 plt.ylabel(accuracy) plt.xlabel(epoch) plt.legend([accuarcy], locupper right) #评估------------------------------------------------------------------------------ #测试集test准确率 print(Accuracy:, sess.run(accuracy,feed_dict\ {x:mnist.test.images,y_label:mnist.test.labels})) #预测test prediction_result sess.run(tf.argmax(y_predict,1),feed_dict{x:mnist.test.images}) #显示真实值和预测值及图像 plot_images_labels_prediction(mnist.test.images,mnist.test.labels,prediction_result,0)

相关新闻

3个维度重塑你的英雄联盟体验:League Akari如何成为你的智能游戏伙伴

3个维度重塑你的英雄联盟体验:League Akari如何成为你的智能游戏伙伴

3个维度重塑你的英雄联盟体验:League Akari如何成为你的智能游戏伙伴 【免费下载链接】League-Toolkit An all-in-one toolkit for LeagueClient. Gathering power 🚀. 项目地址: https://gitcode.com/gh_mirrors/le/League-Toolkit 当英雄联盟的…

2026/7/28 15:24:33 阅读更多 →
中欧 PHP 开发者大会因多元化争议而取消

中欧 PHP 开发者大会因多元化争议而取消

没有女性的演讲者名单导致提倡多元化的男性演讲者退出会议,最终使得计划于德国举办的中欧 PHP 开发者大会宣布取消。 上周末,原定于 10 月 4 日至 6 日在德国德累斯顿举行的 PHP 会议 PHP Central Europe developer conference (PHP.CE) 因多元化争议宣布…

2026/7/28 15:24:33 阅读更多 →
SpringBoot在线投稿系统开发与优化实践

SpringBoot在线投稿系统开发与优化实践

1. 项目概述 在线投稿系统是学术期刊、会议和内容平台的核心基础设施,它直接关系到内容生产的效率和质量管控。基于SpringBoot框架开发的投稿系统,能够为编辑部、审稿人和作者提供全流程的数字化解决方案。这个毕业设计项目采用Java技术栈实现&#xff0…

2026/7/28 15:24:33 阅读更多 →

最新新闻

市面上测试稳定可靠的内存颗粒NAND芯片测试座生产商测试精度高

市面上测试稳定可靠的内存颗粒NAND芯片测试座生产商测试精度高

在当前高度竞争的半导体行业中,内存颗粒NAND芯片的测试是确保产品质量和性能的关键环节。一个稳定可靠的测试座对于提高测试精度、降低误测率以及延长设备寿命至关重要。本文将通过具体数据和案例,分析深圳市谷易电子有限公司(以下简称“谷易…

2026/7/28 15:34:37 阅读更多 →
无人机体系化竞争:从单机对抗到系统集成的技术壁垒与工业逻辑

无人机体系化竞争:从单机对抗到系统集成的技术壁垒与工业逻辑

最近和几位做硬件和嵌入式开发的朋友聊天,话题不知怎么就拐到了无人机上。一位朋友提到,他最近看了一些关于无人机在军事和民用领域应用的讨论,发现一个挺有意思的现象:很多人一提到“无人机对抗”,脑子里浮现的还是那种单机对单机、比拼飞行速度和挂载能力的画面,就像电…

2026/7/28 15:34:37 阅读更多 →
字节 Agent 岗二面:RAG 的 Top-K 是不是越大越好?

字节 Agent 岗二面:RAG 的 Top-K 是不是越大越好?

👔 面试官:你们 RAG 系统里,检索 Top-K 一般设多少? 🙋‍♂️ 我:最开始是 Top5,后来发现有些复杂问题信息不够,就调到了 Top10,再后来有些跨文档比较的问题还是答不全&…

2026/7/28 15:34:37 阅读更多 →
面试官:“什么是大模型项目的分词器?”,我:“文本切成一个个词,让模型能处理”,他:“只会表面功夫?”

面试官:“什么是大模型项目的分词器?”,我:“文本切成一个个词,让模型能处理”,他:“只会表面功夫?”

👔面试官:来讲讲什么是大模型项目的分词器?原理是什么? 🙋‍♂️我:分词器就是把文本切成一个个词,让模型能处理。 👔面试官:……「切成词」是表面理解。模型为什么需要…

2026/7/28 15:34:37 阅读更多 →
暗黑破坏神2存档修改器Diablo Edit2:重新定义你的游戏体验

暗黑破坏神2存档修改器Diablo Edit2:重新定义你的游戏体验

暗黑破坏神2存档修改器Diablo Edit2:重新定义你的游戏体验 【免费下载链接】diablo_edit Diablo II Character editor. 项目地址: https://gitcode.com/gh_mirrors/di/diablo_edit 你是否厌倦了暗黑破坏神2中无尽的刷装备过程?是否因为技能点分配…

2026/7/28 15:34:36 阅读更多 →
编写程序梳理自己人生最想解决的一个生活痛点,围绕痛点长期迭代方案,完成持续性创新。

编写程序梳理自己人生最想解决的一个生活痛点,围绕痛点长期迭代方案,完成持续性创新。

终身痛点追踪器:用 Python 锁定人生最想解决的一个问题,持续迭代创新方案。说明:本文为纯技术实践分享,不涉及任何课程推广、营销引流或商业产品。所有代码可在本地离线运行。一、实际应用场景描述在《心理健康与创新能力》课程中…

2026/7/28 15:33:36 阅读更多 →

日新闻

告别臃肿!3步让你的暗影精灵笔记本重获新生

告别臃肿!3步让你的暗影精灵笔记本重获新生

告别臃肿!3步让你的暗影精灵笔记本重获新生 【免费下载链接】OmenSuperHub Control Omen laptop performance, fan speeds, and keyboard lighting, and unlock power limits. 项目地址: https://gitcode.com/gh_mirrors/om/OmenSuperHub 你是否也曾为官方Om…

2026/7/28 0:00:43 阅读更多 →
RAG必踩坑!财报法规检索不准?这款开源工具让答案浮出水面,准确率飙升98.7%!

RAG必踩坑!财报法规检索不准?这款开源工具让答案浮出水面,准确率飙升98.7%!

做 RAG 的人应该都踩过这个致命的坑:把几百页的财报、法规、技术手册扔给向量库,问一个具体问题,搜出来的全是沾边但没用的内容 —— 关键信息要么被硬切块拆碎了,要么藏在几十条结果的最下面。语义相似≠真正相关,这个…

2026/7/28 0:00:43 阅读更多 →
抖音视频文案提取工具全指南:免费2026版、手机App、在线工具一网打尽

抖音视频文案提取工具全指南:免费2026版、手机App、在线工具一网打尽

2026年做短视频运营,从抖音上扒文案早就不是偷偷抄笔记的事了。我刚开始做内容的时候,每天刷半小时抖音,手动把爆款视频的口播敲进备忘录,一条2分钟的视频得花十来分钟,碰到语速快的还要反复回听。后来试了一圈工具&am…

2026/7/28 0:00:43 阅读更多 →

周新闻

深度学习道路桥梁裂缝检测系统 道路桥梁裂缝检测数据集 道路桥梁病害识别检测数据集

深度学习道路桥梁裂缝检测系统 道路桥梁裂缝检测数据集 道路桥梁病害识别检测数据集

深度学习道路桥梁裂缝检测系统 数据集6000张 完整源码已标注数据集训练好的模型环境配置教程程序运行说明文档,可以直接使用!系统支持图片、视频、摄像头等多种方式检测裂缝,功能强大实用。 1数据集6000张 8各类别

2026/7/28 12:04:22 阅读更多 →
深度学习YOLO模型如何训练 PUBG 绝地求生目标检测数据集

深度学习YOLO模型如何训练 PUBG 绝地求生目标检测数据集

pubg数据集 精选原图1.42万数据 1.49万标签 无任何重复、算法增强或冗余图像! pubg绝地求生目标检测数据集 1分类:e_body,14905个标签,txt格式 共计14244张图,99%为640*640尺寸图像 适合yolo目标检测、AI训练关键词&am…

2026/7/28 8:29:16 阅读更多 →
Apex英雄目标检测数据集 深度学习框架YOLO如何训练APEX数据集

Apex英雄目标检测数据集 深度学习框架YOLO如何训练APEX数据集

Apex检测数据集数据集详情检测类别: allies enemy tag图片总量:7247张训练集:5139张验证集:1425张测试集:683张标注状态:全部已标注,即拿即用数据格式:支持YOLO格式及其他格式&#…

2026/7/28 5:03:42 阅读更多 →

月新闻