最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

keras打印loss對權(quán)重的導(dǎo)數(shù)方式

 更新時(shí)間:2020年06月10日 09:25:49   作者:HackerTom  
這篇文章主要介紹了keras打印loss對權(quán)重的導(dǎo)數(shù)方式,具有很好的參考價(jià)值,希望對大家有所幫助。一起跟隨小編過來看看吧

Notes

懷疑模型梯度爆炸,想打印模型 loss 對各權(quán)重的導(dǎo)數(shù)看看。如果如果fit來訓(xùn)練的話,可以用keras.callbacks.TensorBoard實(shí)現(xiàn)。

但此次使用train_on_batch來訓(xùn)練的,用K.gradients和K.function實(shí)現(xiàn)。

Codes

以一份 VAE 代碼為例

# -*- coding: utf8 -*-
import keras
from keras.models import Model
from keras.layers import Input, Lambda, Conv2D, MaxPooling2D, Flatten, Dense, Reshape
from keras.losses import binary_crossentropy
from keras.datasets import mnist, fashion_mnist
import keras.backend as K
from scipy.stats import norm
import numpy as np
import matplotlib.pyplot as plt

BATCH = 128
N_CLASS = 10
EPOCH = 5
IN_DIM = 28 * 28
H_DIM = 128
Z_DIM = 2

(x_train, y_train), (x_test, y_test) = fashion_mnist.load_data()
x_train = x_train.reshape(len(x_train), -1).astype('float32') / 255.
x_test = x_test.reshape(len(x_test), -1).astype('float32') / 255.

def sampleing(args):
  """reparameterize"""
  mu, logvar = args
  eps = K.random_normal([K.shape(mu)[0], Z_DIM], mean=0.0, stddev=1.0)
  return mu + eps * K.exp(logvar / 2.)

# encode
x_in = Input([IN_DIM])
h = Dense(H_DIM, activation='relu')(x_in)
z_mu = Dense(Z_DIM)(h) # mean,不用激活
z_logvar = Dense(Z_DIM)(h) # log variance,不用激活
z = Lambda(sampleing, output_shape=[Z_DIM])([z_mu, z_logvar]) # 只能有一個(gè)參數(shù)
encoder = Model(x_in, [z_mu, z_logvar, z], name='encoder')

# decode
z_in = Input([Z_DIM])
h_hat = Dense(H_DIM, activation='relu')(z_in)
x_hat = Dense(IN_DIM, activation='sigmoid')(h_hat)
decoder = Model(z_in, x_hat, name='decoder')

# VAE
x_in = Input([IN_DIM])
x = x_in
z_mu, z_logvar, z = encoder(x)
x = decoder(z)
out = x
vae = Model(x_in, [out, out], name='vae')

# loss_kl = 0.5 * K.sum(K.square(z_mu) + K.exp(z_logvar) - 1. - z_logvar, axis=1)
# loss_recon = binary_crossentropy(K.reshape(vae_in, [-1, IN_DIM]), vae_out) * IN_DIM
# loss_vae = K.mean(loss_kl + loss_recon)

def loss_kl(y_true, y_pred):
  return 0.5 * K.sum(K.square(z_mu) + K.exp(z_logvar) - 1. - z_logvar, axis=1)


# vae.add_loss(loss_vae)
vae.compile(optimizer='rmsprop',
      loss=[loss_kl, 'binary_crossentropy'],
      loss_weights=[1, IN_DIM])
vae.summary()

# 獲取模型權(quán)重 variable
w = vae.trainable_weights
print(w)

# 打印 KL 對權(quán)重的導(dǎo)數(shù)
# KL 要是 Tensor,不能是上面的函數(shù) `loss_kl`
grad = K.gradients(0.5 * K.sum(K.square(z_mu) + K.exp(z_logvar) - 1. - z_logvar, axis=1),
          w)
print(grad) # 有些是 None 的
grad = grad[grad is not None] # 去掉 None,不然報(bào)錯(cuò)

# 打印梯度的函數(shù)
# K.function 的輸入和輸出必要是 list!就算只有一個(gè)
show_grad = K.function([vae.input], [grad])

# vae.fit(x_train, # y_train, # 不能傳 y_train
#     batch_size=BATCH,
#     epochs=EPOCH,
#     verbose=1,
#     validation_data=(x_test, None))

''' 以 train_on_batch 方式訓(xùn)練 '''
for epoch in range(EPOCH):
  for b in range(x_train.shape[0] // BATCH):
    idx = np.random.choice(x_train.shape[0], BATCH)
    x = x_train[idx]
    l = vae.train_on_batch([x], [x, x])

  # 計(jì)算梯度
  gd = show_grad([x])
  # 打印梯度
  print(gd)

# show manifold
PIXEL = 28
N_PICT = 30
grid_x = norm.ppf(np.linspace(0.05, 0.95, N_PICT))
grid_y = grid_x

figure = np.zeros([N_PICT * PIXEL, N_PICT * PIXEL])
for i, xi in enumerate(grid_x):
  for j, yj in enumerate(grid_y):
    noise = np.array([[xi, yj]]) # 必須秩為 2,兩層中括號
    x_gen = decoder.predict(noise)
    # print('x_gen shape:', x_gen.shape)
    x_gen = x_gen[0].reshape([PIXEL, PIXEL])
    figure[i * PIXEL: (i+1) * PIXEL,
        j * PIXEL: (j+1) * PIXEL] = x_gen

fig = plt.figure(figsize=(10, 10))
plt.imshow(figure, cmap='Greys_r')
fig.savefig('./variational_autoencoder.png')
plt.show()

補(bǔ)充知識:keras 自定義損失 自動求導(dǎo)時(shí)出現(xiàn)None

問題記錄,keras 自定義損失 自動求導(dǎo)時(shí)出現(xiàn)None,后來想到是因?yàn)閭魅氲淖兞繘]有使用,所以keras無法求出偏導(dǎo),修改后問題解決。就是不愿使用的變量×0,求導(dǎo)后還是0就可以了。

def my_complex_loss_graph(y_label, emb_uid, lstm_out,y_true_1,y_true_2,y_true_3,out_1,out_2,out_3):
 
  mse_out_1 = mean_squared_error(y_true_1, out_1)
  mse_out_2 = mean_squared_error(y_true_2, out_2)
  mse_out_3 = mean_squared_error(y_true_3, out_3)
  # emb_uid= K.reshape(emb_uid, [-1, 32])
  cosine_sim = tf.reduce_sum(0.5*tf.square(emb_uid-lstm_out))
 
  cost=0*cosine_sim+K.sum([0.5*mse_out_1 , 0.25*mse_out_2,0.25*mse_out_3],axis=1,keepdims=True)
  # print(mse_out_1)
  final_loss = cost
 
  return K.mean(final_loss)

以上這篇keras打印loss對權(quán)重的導(dǎo)數(shù)方式就是小編分享給大家的全部內(nèi)容了,希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。

相關(guān)文章

  • Python+Pygame繪制小球的實(shí)例詳解

    Python+Pygame繪制小球的實(shí)例詳解

    這篇文章主要為大家詳細(xì)介紹了如何利用Python?Pygame繪制小球(漸變大的小球、自由下落的小球、循環(huán)上下反彈的小球),感興趣的小伙伴可以了解一下
    2022-10-10
  • 使用python采集腳本之家電子書資源并自動下載到本地的實(shí)例腳本

    使用python采集腳本之家電子書資源并自動下載到本地的實(shí)例腳本

    這篇文章主要介紹了python采集jb51電子書資源并自動下載到本地實(shí)例教程,非常不錯(cuò),具有一定的參考借鑒價(jià)值,需要的朋友可以參考下
    2018-10-10
  • python使用range函數(shù)計(jì)算一組數(shù)和的方法

    python使用range函數(shù)計(jì)算一組數(shù)和的方法

    這篇文章主要介紹了python使用range函數(shù)計(jì)算一組數(shù)和的方法,涉及Python中range函數(shù)的使用技巧,具有一定參考借鑒價(jià)值,需要的朋友可以參考下
    2015-05-05
  • TensorFlow命名空間和TensorBoard圖節(jié)點(diǎn)實(shí)例

    TensorFlow命名空間和TensorBoard圖節(jié)點(diǎn)實(shí)例

    今天小編就為大家分享一篇TensorFlow命名空間和TensorBoard圖節(jié)點(diǎn)實(shí)例,具有很好的參考價(jià)值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2020-01-01
  • Python實(shí)現(xiàn)定時(shí)執(zhí)行任務(wù)的三種方式簡單示例

    Python實(shí)現(xiàn)定時(shí)執(zhí)行任務(wù)的三種方式簡單示例

    這篇文章主要介紹了Python實(shí)現(xiàn)定時(shí)執(zhí)行任務(wù)的三種方式,結(jié)合簡單實(shí)例形式分析了Python使用time,os,sched等模塊定時(shí)執(zhí)行任務(wù)的相關(guān)操作技巧,需要的朋友可以參考下
    2019-03-03
  • Django瀑布流的實(shí)現(xiàn)示例

    Django瀑布流的實(shí)現(xiàn)示例

    在瀏覽一些網(wǎng)站時(shí),經(jīng)常會看到類似于這種滿屏都是圖片,本文主要介紹了Django瀑布流的實(shí)現(xiàn)示例,具有一定的參考價(jià)值,感興趣的可以了解一下
    2023-03-03
  • Python中l(wèi)en()函數(shù)用法使用示例

    Python中l(wèi)en()函數(shù)用法使用示例

    這篇文章主要介紹了Python中的len()函數(shù),包括其基礎(chǔ)用法、適用范圍、常見使用場景以及在第三方庫(如NumPy和pandas)中的應(yīng)用,文中通過代碼介紹的非常詳細(xì),需要的朋友可以參考下
    2025-03-03
  • Python強(qiáng)大郵件處理庫Imbox安裝及用法示例

    Python強(qiáng)大郵件處理庫Imbox安裝及用法示例

    這篇文章主要給大家介紹了關(guān)于Python強(qiáng)大郵件處理庫Imbox安裝及用法的相關(guān)資料,Imbox是一個(gè)Python 庫,用于從IMAP郵箱中讀取郵件,它提供了簡單易用的接口,幫助開發(fā)者處理郵件,需要的朋友可以參考下
    2024-03-03
  • Python3中的指針你了解嗎

    Python3中的指針你了解嗎

    Python這個(gè)編程語言雖然沒有指針類型,但是Python中的可變參量也可以像指針一樣,改變一個(gè)數(shù)值之后,所有指向該數(shù)值的可變參量都會隨之而改變,這篇文章主要介紹了Python3中的“指針”,需要的朋友可以參考下
    2024-02-02
  • 詳解python讀取matlab數(shù)據(jù)(.mat文件)

    詳解python讀取matlab數(shù)據(jù)(.mat文件)

    本文主要介紹了python讀取matlab數(shù)據(jù),文中通過示例代碼介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2021-12-12

最新評論

大安市| 堆龙德庆县| 伊金霍洛旗| 离岛区| 南郑县| 巨野县| 萍乡市| 平陆县| 博爱县| 姚安县| 绍兴市| 墨江| 梧州市| 郯城县| 伽师县| 佳木斯市| 乐陵市| 郴州市| 武定县| 大城县| 平原县| 锡林郭勒盟| 威信县| 绥德县| 贺州市| 阳西县| 中山市| 十堰市| 大石桥市| 广德县| 建德市| 盱眙县| 穆棱市| 新建县| 历史| 吉隆县| 丹棱县| 高密市| 屯昌县| 葫芦岛市| 武定县|