- html - 出于某种原因,IE8 对我的 Sass 文件中继承的 html5 CSS 不友好?
- JMeter 在响应断言中使用 span 标签的问题
- html - 在 :hover and :active? 上具有不同效果的 CSS 动画
- html - 相对于居中的 html 内容固定的 CSS 重复背景?
这是我的 tensorflow keras 模型,(如果它让事情变得困难,你可以忽略 dropout 层)
import tensorflow as tf
optimizers = tf.keras.optimizers
Sequential = tf.keras.models.Sequential
Dense = tf.keras.layers.Dense
Dropout = tf.keras.layers.Dropout
to_categorical = tf.keras.utils.to_categorical
model = Sequential()
model.add(Dense(256, input_shape=(20,), activation="relu"))
model.add(Dropout(0.1))
model.add(Dense(256, activation="relu"))
model.add(Dropout(0.1))
model.add(Dense(256, activation="relu"))
model.add(Dropout(0.1))
model.add(Dense(3, activation="softmax"))
adam = optimizers.Adam(lr=1e-3) # I don't mind rmsprop either
model.compile(optimizer=adam, loss='categorical_crossentropy', metrics=['accuracy'])
model.summary()
我已将模型结构和权重保存为
model.save("sim_score.h5", overwrite=True)
model.save_weights('sim_score_weights.h5', overwrite=True)
在执行 model.predict(X_test)
时,我得到 [0.23, 0.63, 0.14]
这是 3 个输出类的预测概率。
How would I visualize how much weight/importance each of my initial 20 features have in this model w.r.t the 3 output softmax?
例如,我的第二列对最终结果的影响可以忽略不计,而第五列对输出预测的影响可能是第 20 列的 3 倍。第 5 列的绝对效果是什么并不重要,只要计算出相对重要性就足够了,例如 第 5 列 = 0.3,第 20 列 = 0.1
等等,对于 20 x 3 矩阵
。
查看此animation直觉或 Tensorflow playground 。可视化不必显示训练过程中权重如何变化,而只需显示训练结束时权重的快照图像。
In fact, the solution does not even have to be a visualization, it can even be an array of 20 elements x 3 outputs having the relative importance of each feature w.r.t the 3 output softmax and importance relative to the other features.
了解中间层的重要性只是一个额外的好处。
我想要可视化 20 个特征的原因是出于透明度目的(目前模型感觉就像一个黑匣子)。我对 matplotlib、pyplot、seaborn 很熟悉。我也知道 Tensorboard,但找不到任何使用 Softmax 的简单 Dense Relu 网络的示例。
我认为获得 20 x 3 权重的一种耗时方法是从 0 - 1
进行域搜索 20 个特征
,增量为 0.5
通过发送各种输入并尝试据此推断特征的重要性(这将是 3 的 20 次方
~= 34 亿
可能的样本空间,并且我添加的功能越多,情况就会呈指数级恶化),然后应用条件概率对相对权重进行逆向工程,但我不确定是否有通过 TensorBoard 或某些自定义逻辑的更简单/优化的方法。
有人可以通过代码片段帮助可视化模型/计算 20 个特征与 3 个输出的 20 x 3 = 60 相对权重,或者提供有关如何实现此目标的引用吗?
最佳答案
我遇到过类似的问题,但我更关心模型参数(权重和偏差)的可视化,而不是模型特征[因为我也想探索和查看黑匣子]。
例如,以下是具有 2 个隐藏层的浅层神经网络的片段。
model = Sequential()
model.add(Dense(128, input_dim=13, kernel_initializer='uniform', activation='relu'))
model.add(Dropout(0.1))
model.add(Dense(64, kernel_initializer='uniform', activation='relu'))
model.add(Dropout(0.1))
model.add(Dense(64, kernel_initializer='uniform', activation='relu'))
model.add(Dropout(0.1))
model.add(Dense(8, kernel_initializer='uniform', activation='softmax'))
# Compile model
model.compile(loss='sparse_categorical_crossentropy', optimizer='adam', metrics=['accuracy'])
# Using TensorBoard to visualise the Model
ks=TensorBoard(log_dir="/your_full_path/logs/{}".format(time()), histogram_freq=1, write_graph=True, write_grads=True, batch_size=10)
# Fit the model
model.fit(X, Y, epochs = 64, shuffle = True, batch_size=10, verbose = 2, validation_split=0.2, callbacks=[ks])
为了能够可视化参数,需要记住一些重要的事情:
始终确保 model.fit() 函数中有一个validation_split[否则直方图无法可视化]。
确保 histogram_freq > 0 的值始终!![否则不会计算直方图]。
对 TensorBoard 的回调必须在 model.fit() 中以列表形式指定。
一次,这就完成了;转到 cmd 并键入以下命令:
tensorboard --logdir=logs/
这将为您提供一个本地地址,您可以通过该地址在网络浏览器上访问 TensorBoard。所有直方图、分布、损失和准确度函数都将以图表形式提供,并且可以从顶部的菜单栏中进行选择。
希望这个答案给出有关可视化模型参数的过程的提示(我自己也遇到了一些困难,因为上述几点不能同时使用)。
如果有帮助请告诉我。
以下是 keras 文档链接供您引用:
关于python - 计算/可视化 Tensorflow Keras Dense 模型层与输出类的相对连接权重,我们在Stack Overflow上找到一个类似的问题: https://stackoverflow.com/questions/52703096/
我知道这个问题可能已经被问过,但我检查了所有这些,我认为我的情况有所不同(请友善)。所以我有两个数据集,第一个是测试数据集,第二个是我保存在数据框中的预测(预测值,这就是没有数据列的原因)。我想合并两
在 .loc 方法的帮助下,我根据同一数据框中另一列中的值来识别 Panda 数据框中某一列中的值。 下面给出了代码片段供您引用: var1 = output_df['Player'].loc[out
当我在 Windows 中使用 WinSCP 通过 Ubuntu 连接到 VMware 时,它提示: The server rejected SFTP connection, but it lis
我正在开发一个使用 xml web 服务的 android 应用程序。在 wi-fi 网络中连接时工作正常,但在 3G 网络中连接时失败(未找到 http 404)。 这不仅仅发生在设备中。为了进行测
我有一个XIB包含我的控件的文件,加载到 Interface Builder(Snow Leopard 上的 Xcode 4.0.2)中。 文件的所有者被设置为 someClassController
我在本地计算机上管理 MySQL 数据库,并通过运行以下程序通过 C 连接到它: #include #include #include int main(int argc, char** arg
我不知道为什么每次有人访问我网站上的页面时,都会打开一个与数据库的新连接。最终我到达了大约 300 并收到错误并且页面不再加载。我认为它应该工作的方式是,我将 maxIdle 设置为 30,这意味着
希望清理 NMEA GPS 中的 .txt 文件。我当前的代码如下。 deletes = ['$GPGGA', '$GPGSA', '$GPGSV', '$PSRF156', ] searchquer
我有一个 URL、一个用户名和一个密码。我想在 C# .Net WinForms 中建立 VPN 连接。 你能告诉我从哪里开始吗?任何第三方 API? 代码示例将受到高度赞赏... 最佳答案 您可以像
有没有更好的方法将字符串 vector 转换为字符 vector ,字符串之间的终止符为零。 因此,如果我有一个包含以下字符串的 vector "test","my","string",那么我想接收一
我正在编写一个库,它不断检查 android 设备的连接,并在设备连接、断开连接或互联网连接变慢时给出回调。 https://github.com/muddassir235/connection_ch
我的操作系统:Centos 7 + CLOUDLINUX 7.7当我尝试从服务器登录Mysql时 [root@server3 ~]# Mysql -u root -h localhost -P 330
我收到错误:Puma 发现此错误:无法打开到本地主机的 TCP 连接:9200(连接被拒绝 - 连接(2)用于“本地主机”端口 9200)(Faraday::ConnectionFailed)在我的
请给我一些解决以下错误的方法。 这是一个聊天应用....代码和错误如下:: conversations_controller.rb def create if Conversation.bet
我想将两个单元格中的数据连接到一个单元格中。我还想只组合那些具有相同 ID 的单元格。 任务 ID 名称 4355.2 参与者 4355.2 领袖 4462.1 在线 4462.1 快速 4597.1
我经常需要连接 TSQL 中的字段... 使用“+”运算符时 TSQL 强制您处理的两个问题是 Data Type Precedence和 NULL 值。 使用数据类型优先级,问题是转换错误。 1)
有没有在 iPad 或 iPhone 应用程序中使用 Facebook 连接。 这个想法是登录这个应用程序,然后能够看到我的哪些 facebook 用户也在使用该应用程序及其功能。 最佳答案 是的。
我在连接或打印字符串时遇到了一个奇怪的问题。我有一个 char * ,可以将其设置为字符串文字的几个值之一。 char *myStrLiteral = NULL; ... if(blah) myS
对于以下数据 - let $x := "Yahooooo !!!! Select one number - " let $y := 1 2 3 4 5 6 7 我想得到
我正在看 UDEMY for perl 的培训视频,但是视频不清晰,看起来有错误。 培训展示了如何使用以下示例连接 2 个字符串: #!usr/bin/perl print $str = "Hi";
我是一名优秀的程序员,十分优秀!