Python keras.metrics源代碼分析
前言
metrics用于判斷模型性能。度量函數(shù)類似于損失函數(shù),只是度量的結(jié)果不用于訓(xùn)練模型??梢允褂萌魏螕p失函數(shù)作為度量(如logloss等)。在訓(xùn)練期間監(jiān)控metrics的最佳方式是通過Tensorboard。
官方提供的metrics最重要的概念就是有狀態(tài)(stateful)變量,通過更新狀態(tài)變量,可以不斷累積統(tǒng)計數(shù)據(jù),并可以隨時輸出狀態(tài)變量的計算結(jié)果。這是區(qū)別于losses的重要特性,losses是無狀態(tài)的(stateless)。
本文部分內(nèi)容參考了:
代碼運行環(huán)境為:tf.__version__==2.6.2 。
metrics原理解析(以metrics.Mean為例)
metrics是有狀態(tài)的(stateful),即Metric 實例會存儲、記錄和返回已經(jīng)累積的結(jié)果,有助于未來事務(wù)的信息。下面以tf.keras.metrics.Mean()為例進行解釋:
創(chuàng)建tf.keras.metrics.Mean的實例:
m = tf.keras.metrics.Mean()
通過help(m) 可以看到MRO為:
Mean
Reduce
Metric
keras.engine.base_layer.Layer
...
可見Metric和Mean是 keras.layers.Layer 的子類。相比于類Layer,其子類Mean多出了幾個方法:
- result: 計算并返回標量度量值(tensor形式)或標量字典,即狀態(tài)變量簡單地計算度量值。例如,
m.result(),就是計算均值并返回。 - total: 狀態(tài)變量
m目前累積的數(shù)字總和 - count: 狀態(tài)變量
m目前累積的數(shù)字個數(shù)(m.total/m.count就是m.result()的返回值) - update_state: 累積統(tǒng)計數(shù)字用于計算指標。每次調(diào)用
m.update_state都會更新m.total和m.count; - reset_state: 將狀態(tài)變量重置到初始化狀態(tài);
- reset_states: 等價于reset_state,參見keras源代碼metrics.py L355
- reduction: 目前來看,沒什么用。
這也決定了Mean的特殊性質(zhì)。其使用參見如下代碼:
# 創(chuàng)建狀態(tài)變量m,由于m未剛初始化,
# 所以total,count和result()均為0
m = tf.keras.metrics.Mean()
print("m.total:",m.total)
print("m.count:",m.count)
print("m.result():",m.result())"""
# 輸出:
m.total: <tf.Variable 'total:0' shape=() dtype=float32, numpy=0.0>
m.count: <tf.Variable 'count:0' shape=() dtype=float32, numpy=0.0>
m.result(): tf.Tensor(0.0, shape=(), dtype=float32)
"""
# 更新狀態(tài)變量,可以看到total累加了總和,
# count累積了個數(shù),result()返回total/count
m.update_state([1,2,3])
print("m.total:",m.total)
print("m.count:",m.count)
print("m.result():",m.result())"""
# 輸出:
m.total: <tf.Variable 'total:0' shape=() dtype=float32, numpy=6.0>
m.count: <tf.Variable 'count:0' shape=() dtype=float32, numpy=3.0>
m.result(): tf.Tensor(2.0, shape=(), dtype=float32)
"""
# 重置狀態(tài)變量, 重置到初始化狀態(tài)
m.reset_state()
print("m.total:",m.total)
print("m.count:",m.count)
print("m.result():",m.result())"""
# 輸出:
m.total: <tf.Variable 'total:0' shape=() dtype=float32, numpy=0.0>
m.count: <tf.Variable 'count:0' shape=() dtype=float32, numpy=0.0>
m.result(): tf.Tensor(0.0, shape=(), dtype=float32)
"""
創(chuàng)建自定義metrics
創(chuàng)建無狀態(tài) metrics
與損失函數(shù)類似,任何帶有類似于metric_fn(y_true, y_pred)、返回損失數(shù)組(如輸入一個batch的數(shù)據(jù),會返回一個batch的損失標量)的函數(shù),都可以作為metric傳遞給compile():
import tensorflow as tf
import numpy as np
inputs = tf.keras.Input(shape=(3,))
x = tf.keras.layers.Dense(4, activation=tf.nn.relu)(inputs)
outputs = tf.keras.layers.Dense(1, activation=tf.nn.softmax)(x)
model1 = tf.keras.Model(inputs=inputs, outputs=outputs)
def my_metric_fn(y_true, y_pred):
squared_difference = tf.square(y_true - y_pred)
return tf.reduce_mean(squared_difference, axis=-1) # shape=(None,)
model1.compile(optimizer='adam', loss='mse', metrics=[my_metric_fn])
x = np.random.random((100, 3))
y = np.random.random((100, 1))
model1.fit(x, y, epochs=3)輸出:
Epoch 1/3
4/4 [==============================] - 0s 667us/step - loss: 0.0971 - my_metric_fn: 0.0971
Epoch 2/3
4/4 [==============================] - 0s 667us/step - loss: 0.0958 - my_metric_fn: 0.0958
Epoch 3/3
4/4 [==============================] - 0s 1ms/step - loss: 0.0946 - my_metric_fn: 0.0946
注意,因為本例創(chuàng)建的是無狀態(tài)的度量,所以上面跟蹤的度量值(my_metric_fn后面的值)是每個batch的平均度量值,并不是一個epoch(完整數(shù)據(jù)集)的累積值。(這一點需要理解,這也是為什么要使用有狀態(tài)度量的原因!)
值得一提的是,如果上述代碼使用
model1.compile(optimizer='adam', loss='mse', metrics=["mse"])
進行compile,則輸出的結(jié)果是累積的,在每個epoch結(jié)束時的結(jié)果就是整個數(shù)據(jù)集的結(jié)果,因為metrics=["mse"]是直接調(diào)用了標準庫的有狀態(tài)度量。
通過繼承Metric創(chuàng)建有狀態(tài)metrics
如果想查看整個數(shù)據(jù)集的指標,就需要傳入有狀態(tài)的metrics,這樣就會在一個epoch內(nèi)累加,并在epoch結(jié)束時輸出整個數(shù)據(jù)集的度量值。
創(chuàng)建有狀態(tài)度量指標,需要創(chuàng)建Metric的子類,它可以跨batch維護狀態(tài),步驟如下:
- 在
__init__中創(chuàng)建狀態(tài)變量(state variables) - 更新
update_state()中y_true和y_pred的變量 - 在
result()中返回標量度量結(jié)果 - 在
reset_states()中清除狀態(tài)
class BinaryTruePositives(tf.keras.metrics.Metric):
def __init__(self, name='binary_true_positives', **kwargs):
super(BinaryTruePositives, self).__init__(name=name, **kwargs)
self.true_positives = self.add_weight(name='tp', initializer='zeros')
def update_state(self, y_true, y_pred, sample_weight=None):
y_true = tf.cast(y_true, tf.bool)
y_pred = tf.cast(y_pred, tf.bool)
values = tf.logical_and(tf.equal(y_true, True), tf.equal(y_pred, True))
values = tf.cast(values, self.dtype)
if sample_weight is not None:
sample_weight = tf.cast(sample_weight, self.dtype)
values = tf.multiply(values, sample_weight)
self.true_positives.assign_add(tf.reduce_sum(values))
def result(self):
return self.true_positives
def reset_states(self):
self.true_positives.assign(0)
m = BinaryTruePositives()
m.update_state([0, 1, 1, 1], [0, 1, 0, 0])
print('Intermediate result:', float(m.result()))
m.update_state([1, 1, 1, 1], [0, 1, 1, 0])
print('Final result:', float(m.result()))add_metric()方法
add_metric 方法是 tf.keras.layers.Layer類添加的方法,Layer的父類tf.Module并沒有這個方法,因此在編寫Layer子類如包括自定義層、官方提供的層(Dense)或模型(tf.keras.Model也是Layer的子類)時,可以使用add_metric()來與層相關(guān)的統(tǒng)計量。比如,將類似Dense的自定義層的激活平均值記錄為metric??梢詧?zhí)行以下操作:
class DenseLike(Layer):
"""y = w.x + b"""
...
def call(self, inputs):
output = tf.matmul(inputs, self.w) + self.b
self.add_metric(tf.reduce_mean(output), aggregation='mean', name='activation_mean')
return output將在名稱為activation_mean的度量下跟蹤output,跟蹤的值為每個批次度量值的平均值。
更詳細的信息,參閱官方文檔The base Layer class - add_metric method。
參考
到此這篇關(guān)于Python keras.metrics源代碼分析的文章就介紹到這了,更多相關(guān)Python keras.metrics內(nèi)容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!
相關(guān)文章
Python GUI學(xué)習之登錄系統(tǒng)界面篇
這篇文章主要介紹了Python GUI學(xué)習之登錄系統(tǒng)界面篇,文中通過示例代碼介紹的非常詳細,對大家的學(xué)習或者工作具有一定的參考學(xué)習價值,需要的朋友們下面隨著小編來一起學(xué)習學(xué)習吧2019-08-08
django自定義Field實現(xiàn)一個字段存儲以逗號分隔的字符串
這篇文章主要介紹了django自定義Field實現(xiàn)一個字段存儲以逗號分隔的字符串的示例,需要的朋友可以參考下2014-04-04
在Python中使用MySQL--PyMySQL的基本使用方法
PyMySQL 是在 Python3.x 版本中用于連接 MySQL 服務(wù)器的一個庫,Python2中則使用mysqldb。這篇文章主要介紹了在Python中使用MySQL--PyMySQL的基本使用,需要的朋友可以參考下2019-11-11
詳解使用python crontab設(shè)置linux定時任務(wù)
本篇文章主要介紹了使用python crontab設(shè)置linux定時任務(wù),具有一定的參考價值,有需要的可以了解一下。2016-12-12
Python機器學(xué)習算法庫scikit-learn學(xué)習之決策樹實現(xiàn)方法詳解
這篇文章主要介紹了Python機器學(xué)習算法庫scikit-learn學(xué)習之決策樹實現(xiàn)方法,結(jié)合實例形式分析了決策樹算法的原理及使用sklearn庫實現(xiàn)決策樹的相關(guān)操作技巧,需要的朋友可以參考下2019-07-07
Python爬取門戶論壇評論淺談Python未來發(fā)展方向
這篇文章主要介紹了如何實現(xiàn)Python爬取門戶論壇評論,附含圖片示例代碼,講解了詳細的操作過程,有需要的的朋友可以借鑒參考下,希望可以有所幫助2021-09-09
Python爬蟲中urllib3與urllib的區(qū)別是什么
Urllib3是一個功能強大,條理清晰,用于HTTP客戶端的Python庫。那么Python爬蟲中urllib3與urllib的區(qū)別是什么,本文就詳細的來介紹一下2021-07-07

