Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
76 changes: 76 additions & 0 deletions lessonFive/to_jpg.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
#coding=utf-8
import cv2
import numpy as np
import os
import pickle


#文件夹名
str_2 = 'train_data'
str_1 = 'test_data'

#判断文件夹是否存在,不存在的话创建文件夹
if os.path.exists(str_1) == False:
os.mkdir(str_1)
if os.path.exists(str_2) == False:
os.mkdir(str_2)

# 解压缩,返回解压后的字典,f,encoding='bytes'
def unpickle(file):

fo = open(file, 'rb')
dict = pickle.load(fo, encoding='bytes')
fo.close()
return dict

def cifar_jpg(dir_file):
# 生成训练集图片,如果需要png格式,只需要改图片后缀名即可。
for j in range(1, 6):
dataName = dir_file + '/' + "data_batch_" + str(j) # 读取当前目录下的data_batch12345文件,dataName其实也是data_batch文件的路径,本文和脚本文件在同一目录下。
Xtr = unpickle(dataName)
print(Xtr)
print(dataName + " is loading...")

for i in range(0, 10000):
img = np.reshape(Xtr[b'data'][i], (3, 32, 32))
img = img.transpose(1, 2, 0) # 读取image
picName = 'train_data/' + str(Xtr[b'labels'][i]) + '_' + str(i + (j - 1) * 10000) + '.jpg' # Xtr['labels']为图片的标签,值范围0-9,本文中,train文件夹需要存在,并与脚本文件在同一目录下。
cv2.imwrite(picName, img)
print(dataName + " loaded.")

print("test_batch is loading...")

# 生成测试集图片
testName = dir_file + '/' + 'data_batch_6'
testXtr = unpickle(testName)
for i in range(0, 10000):
img = np.reshape(testXtr[b'data'][i], (3, 32, 32))
img = img.transpose(1, 2, 0)
picName = 'test_data/' + str(testXtr[b'labels'][i]) + '_' + str(i) + '.jpg'
cv2.imwrite(picName, img)
print("test_batch loaded.")
return

#标签与名字的对应关系
def label_name():
label_name_dict ={
'airplane': "0",
'automobile': "1",
'bird': "2",
'cat': "3",
'deer': "4",
'dog': "5",
'frog': "6",
'horse': "7",
'ship': "8",
'truck': "9"
}
return label_name_dict

if __name__ == '__main__':
dir_file = 'train_data'
#cifar_jpg(dir_file)
try:
cifar_jpg(dir_file)
except:
print('函数报错')
2 changes: 1 addition & 1 deletion lessonOne/imgClassifierWeb/cnnModel.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,7 +106,7 @@ def fc_layer(flattened_layer, num_inputs, num_outputs):
fc_resultl = tf.matmul(flattened_layer, fc_weights)
return fc_resultl
batch_size=gConfig['percent']*gConfig['dataset_size']/100
self.data_tensor=tf.placeholder(tf.float32,shape=tf.placeholder(tf.float32,shape=[batch_size,gConfig['im_dim'], gConfig['im_dim'],gConfig['num_channels']],name='data_tensor'))
self.data_tensor=tf.placeholder(tf.float32,shape=[batch_size,gConfig['im_dim'], gConfig['im_dim'],gConfig['num_channels']],name='data_tensor')
self.label_tensor=tf.placeholder(tf.int32,shape=[batch_size],name='label_tensor')
keep_prop=tf.Variable(initial_value=0.5,name="keep_prop")
self.fc_result=create_CNN(input_data=self.data_tensor,num_classes=gConfig['num_dataset_classes'],keep_prop=gConfig['keeps'])
Expand Down
130 changes: 130 additions & 0 deletions lessonSix/vgg16/app.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
import flask
import werkzeug
import os
import scipy.misc
import tensorflow as tf
import getConfig
import execute

gConfig = {}
gConfig = getConfig.get_config(config_file='config.ini')

# Creating a new Flask Web application. It accepts the package name.
app = flask.Flask("imgClassifierWeb")


def CNN_predict():
global sess
global model
global graph

"""
global:
"""
global secure_filename
# 从本地目录读取需要分类的图片
img = scipy.misc.imread(os.path.join(app.root_path, secure_filename))

"""
校验图片格式
"""
if (img.ndim) == 3:
"""
是否为32*32
"""
if img.shape[0] == img.shape[1] and img.shape[0] == 32:
"""
是否为3通道,GRB
"""
if img.shape[-1] == 3:

predicted_class = execute.predict_line(sess, model, img, graph)
"""
将返回的结果用页面模板给渲染出来
"""
return flask.render_template(template_name_or_list="prediction_result.html",
predicted_class=predicted_class)
else:
""" 如果检测出图片格式不符合要求,则返回错误并返回上传图片的格式"""
return flask.render_template(template_name_or_list="error.html", img_shape=img.shape)
else:
""" 如果检测出图片格式不符合要求,则返回错误并返回上传图片的格式"""
return flask.render_template(template_name_or_list="error.html", img_shape=img.shape)
return "遇到非图片格式的未知错误,请联系技术人员解决"


"""
flask路由系统:
1、使用flask.Flask.route() 修饰器。
2、使用flask.Flask.add_url_rule()函数。
3、直接访问基于werkzeug路由系统的flask.Flask.url_map.
参考知识链接:https://www.jianshu.com/p/e69016bd8f08
1、@app.route('/index.html')
def index():
return "Hello World!"
2、def index():
return "Hello World!"
index = app.route('/index.html')(index)
app.add_url_rule:app.add_url_rule(rule,endpoint,view_func)
关于rule、ednpoint、view_func以及函数注册路由的原理可以参考:https://www.cnblogs.com/eric-nirnava/p/endpoint.html
"""
app.add_url_rule(rule="/predict/", endpoint="predict", view_func=CNN_predict)
"""
知识点:
flask.request属性
form:
一个从POST和PUT请求解析的 MultiDict(一键多值字典)。
args:
MultiDict,要操作 URL (如 ?key=value )中提交的参数可以使用 args 属性:
searchword = request.args.get('key', '')
values:
CombinedMultiDict,内容是form和args。
可以使用values替代form和args。
cookies:
顾名思义,请求的cookies,类型是dict。
stream:
在可知的mimetype下,如果进来的表单数据无法解码,会没有任何改动的保存到这个·stream·以供使用。很多时候,当请求的数据转换为string时,使用data是最好的方式。这个stream只返回数据一次。
headers:
请求头,字典类型。
data:
包含了请求的数据,并转换为字符串,除非是一个Flask无法处理的mimetype。
files:
MultiDict,带有通过POST或PUT请求上传的文件。
method:
请求方法,比如POST、GET
知识点参考链接:https://blog.csdn.net/yannanxiu/article/details/53116652
werkzeug
"""


def upload_image():
global secure_filename
if flask.request.method == "POST": # 设置request的模式为POST
img_file = flask.request.files["image_file"] # 获取需要分类的图片
secure_filename = werkzeug.secure_filename(img_file.filename) # 生成一个没有乱码的文件名
img_path = os.path.join(app.root_path, secure_filename) # 获取图片的保存路径
img_file.save(img_path) # 将图片保存在应用的根目录下
print("图片上传成功.")
"""

"""
return flask.redirect(flask.url_for(endpoint="predict"))
return "图片上传失败"


"""
"""
app.add_url_rule(rule="/upload/", endpoint="upload", view_func=upload_image, methods=["POST"])


def redirect_upload():
return flask.render_template(template_name_or_list="upload_image.html")


"""
"""
app.add_url_rule(rule="/", endpoint="homepage", view_func=redirect_upload)
sess = tf.Session()
sess, model, graph = execute.init_session(sess, conf='config.ini')
if __name__ == "__main__":
app.run(host="localhost", port=7777, debug=False)
24 changes: 24 additions & 0 deletions lessonSix/vgg16/config.ini
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
[strings]
# Mode : train, test, serve
mode = train
working_directory = model/
dataset_path=/Users/zhaoyingjun/Learning/TensorFlow_code/lessonOne/imgClassifierWeb/train_data/
dataset_test=/Users/zhaoyingjun/Learning/TensorFlow_code/lessonOne/imgClassifierWeb/test_data/

[ints]
steps_per_checkpoint = 10
num_dataset_classes=10
dataset_size=50000
im_dim=32
num_channels = 3
num_files=5
images_per_file=10000
max_gradient_norm=5

[floats]
learning_rate = 0.01
keeps=0.5
learning_rate_decay_factor = 0.9
percent=1
end_learning_rate=0.0

Loading