You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用Keras-vis可视化多输入神经网络的注意力机制?

Can Keras-vis Visualize Attention for Multi-Input Dense Networks?

Absolutely! You can definitely use Keras-vis to visualize attention for both input_a and input_b in your multi-input dense neural network. While Keras-vis is often linked to image models, it works seamlessly with dense feature inputs—you just need to shift focus to feature importance scores (the core of attention here) instead of spatial heatmaps.

Step-by-Step Approach

1. Understand the Core Idea

For dense inputs, "attention visualization" translates to measuring how much each feature in input_a or input_b contributes to the final classification decision. Keras-vis computes gradients of the target class output with respect to the input features; higher gradient magnitudes mean the feature has a bigger impact on the prediction.

2. Adapt Your Model for Gradient Calculations

Keras-vis works best when the output layer uses a linear activation (instead of softmax) for gradient clarity. We’ll temporarily modify the output layer’s activation function to make gradient calculations more straightforward.

3. Example Code Implementation

Let’s walk through a concrete example using your network structure:

First, Build Your Multi-Input Model

from keras.layers import Input, Dense, concatenate
from keras.models import Model

# Define inputs
input_a = Input(shape=(10,), name="input_a")
input_b = Input(shape=(15,), name="input_b")

# Individual dense layers
x_a = Dense(20, activation="relu")(input_a)
x_b = Dense(20, activation="relu")(input_b)

# Merge and final classification
merged = concatenate([x_a, x_b])
output = Dense(5, activation="softmax")(merged)

model = Model(inputs=[input_a, input_b], outputs=output)
model.compile(optimizer="adam", loss="categorical_crossentropy")

Visualize Attention for Each Input

from vis.visualization import visualize_saliency
from vis.utils import utils
from keras import activations
import matplotlib.pyplot as plt

# Replace output layer activation with linear (better for gradient calculations)
output_layer_idx = utils.find_layer_idx(model, "dense_2")  # Adjust to your output layer name
model.layers[output_layer_idx].activation = activations.linear
model = utils.apply_modifications(model)

# Pick a target class to visualize (e.g., class 0)
target_class = 0

# Prepare a sample input pair from your dataset
sample_input_a = ...  # Shape: (1, 10) – one sample from input_a
sample_input_b = ...  # Shape: (1, 15) – matching sample from input_b

# Generate attention scores for input_a (fix input_b, compute gradients for input_a)
attention_a = visualize_saliency(
    model,
    layer_idx=output_layer_idx,
    filter_indices=target_class,
    seed_input=[sample_input_a, sample_input_b],
    input_indices=0  # Focus on input_a (index 0 in the model's input list)
)

# Generate attention scores for input_b (fix input_a, compute gradients for input_b)
attention_b = visualize_saliency(
    model,
    layer_idx=output_layer_idx,
    filter_indices=target_class,
    seed_input=[sample_input_a, sample_input_b],
    input_indices=1  # Focus on input_b (index 1 in the model's input list)
)

# Plot the results as bar charts (intuitive for dense features)
plt.figure(figsize=(14, 5))

plt.subplot(1, 2, 1)
plt.bar(range(10), attention_a[0])
plt.title("Attention Scores for Input_A Features")
plt.xlabel("Feature Index")
plt.ylabel("Gradient Magnitude (Importance)")

plt.subplot(1, 2, 2)
plt.bar(range(15), attention_b[0])
plt.title("Attention Scores for Input_B Features")
plt.xlabel("Feature Index")
plt.ylabel("Gradient Magnitude (Importance)")

plt.tight_layout()
plt.show()

Key Notes

  • input_indices Parameter: This is critical for multi-input models—it tells Keras-vis which input to compute gradients for. Use 0 for input_a and 1 for input_b.
  • Visualization Type: For dense features, bar charts are more intuitive than heatmaps, as each bar corresponds to a specific input feature’s importance.
  • Gradient Magnitude: Higher values mean the feature strongly influences the target class prediction (either positively or negatively; you can adjust the code to show signed gradients if needed).
  • Compatibility: Ensure your Keras-vis version matches your Keras/TensorFlow version to avoid compatibility issues.

内容的提问来源于stack exchange,提问作者AlexGuevara

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.21 04:08:46