tensorflow 固定部分參數(shù)訓(xùn)練,只訓(xùn)練部分參數(shù)的實例
在使用tensorflow來訓(xùn)練一個模型的時候,有時候需要依靠驗證集來判斷模型是否已經(jīng)過擬合,是否需要停止訓(xùn)練。
1.首先想到的是用tf.placeholder()載入不同的數(shù)據(jù)來進(jìn)行計算,比如
def inference(input_):
"""
this is where you put your graph.
the following is just an example.
"""
conv1 = tf.layers.conv2d(input_)
conv2 = tf.layers.conv2d(conv1)
return conv2
input_ = tf.placeholder()
output = inference(input_)
...
calculate_loss_op = ...
train_op = ...
...
with tf.Session() as sess:
sess.run([loss, train_op], feed_dict={input_: train_data})
if validation == True:
sess.run([loss], feed_dict={input_: validate_date})
這種方式很簡單,也很直接了然。
2.但是,如果處理的數(shù)據(jù)量很大的時候,使用 tf.placeholder() 來載入數(shù)據(jù)會嚴(yán)重地拖慢訓(xùn)練的進(jìn)度,因此,常用tfrecords文件來讀取數(shù)據(jù)。
此時,很容易想到,將不同的值傳入inference()函數(shù)中進(jìn)行計算。
train_batch, label_batch = decode_train()
val_train_batch, val_label_batch = decode_validation()
train_result = inference(train_batch)
...
loss = ..
train_op = ...
...
if validation == True:
val_result = inference(val_train_batch)
val_loss = ..
with tf.Session() as sess:
sess.run([loss, train_op])
if validation == True:
sess.run([val_result, val_loss])
這種方式看似能夠直接調(diào)用inference()來對驗證數(shù)據(jù)進(jìn)行前向傳播計算,但是,實則會在原圖上添加上許多新的結(jié)點,這些結(jié)點的參數(shù)都是需要重新初始化的,也是就是說,驗證的時候并不是使用訓(xùn)練的權(quán)重。
3.用一個tf.placeholder來控制是否訓(xùn)練、驗證。
def inference(input_):
...
...
...
return inference_result
train_batch, label_batch = decode_train()
val_batch, val_label = decode_validation()
is_training = tf.placeholder(tf.bool, shape=())
x = tf.cond(is_training, lambda: train_batch, lambda: val_batch)
y = tf.cond(is_training, lambda: train_label, lambda: val_label)
logits = inference(x)
loss = cal_loss(logits, y)
train_op = optimize(loss)
with tf.Session() as sess:
loss, _ = sess.run([loss, train_op], feed_dict={is_training: True})
if validation == True:
loss = sess.run(loss, feed_dict={is_training: False})
使用這種方式就可以在一個大圖里創(chuàng)建一個分支條件,從而通過控制placeholder來控制是否進(jìn)行驗證。
以上這篇tensorflow 固定部分參數(shù)訓(xùn)練,只訓(xùn)練部分參數(shù)的實例就是小編分享給大家的全部內(nèi)容了,希望能給大家一個參考,也希望大家多多支持腳本之家。
相關(guān)文章
Python基礎(chǔ)知識+結(jié)構(gòu)+數(shù)據(jù)類型
這篇文章主要介紹了Python基礎(chǔ)知識+結(jié)構(gòu)+數(shù)據(jù)類型,文章基于python基礎(chǔ)知識圍繞主題展開詳細(xì)內(nèi)容介紹,具有一定的參考價值,需要的小伙伴可以參考一下2022-05-05
PyCharm License Activation激活碼失效問題的解決方法(圖文詳解)
這篇文章主要介紹了PyCharm License Activation激活碼失效問題的解決方法,本文通過圖文并茂的形式給大家介紹的非常詳細(xì),對大家的學(xué)習(xí)或工作具有一定的參考借鑒價值,需要的朋友可以參考下2020-03-03
Pandas使用stack和pivot實現(xiàn)數(shù)據(jù)透視的方法
筆者最近正在學(xué)習(xí)Pandas數(shù)據(jù)分析,將自己的學(xué)習(xí)筆記做成一套系列文章。本節(jié)主要記錄Pandas中使用stack和pivot實現(xiàn)數(shù)據(jù)透視。感興趣的小伙伴們可以參考一下2021-09-09
基于Python實現(xiàn)面向?qū)ο蟀鎸W(xué)生管理系統(tǒng)
這篇文章主要為大家詳細(xì)介紹了如何利用python實現(xiàn)學(xué)生管理系統(tǒng)(面向?qū)ο蟀妫?,文中示例代碼介紹的非常詳細(xì),具有一定的參考價值,感興趣的小伙伴們可以參考一下2022-07-07
python給指定csv表格中的聯(lián)系人群發(fā)郵件(帶附件的郵件)
這篇文章主要介紹了python給指定csv表格中的聯(lián)系人群發(fā)郵件,本文通過代碼講解的非常詳細(xì),具有一定的參考借鑒價值,需要的朋友可以參考下2019-12-12
如何利用pandas工具輸出每行的索引值、及其對應(yīng)的行數(shù)據(jù)
這篇文章主要介紹了如何利用pandas工具輸出每行的索引值、及其對應(yīng)的行數(shù)據(jù),本文給大家介紹的非常詳細(xì),對大家的學(xué)習(xí)或工作具有一定的參考借鑒價值,需要的朋友可以參考下2021-03-03

