File size: 1,199 Bytes
adcc0ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36

import torch
import matplotlib.pyplot as plt

def visualize_results(input_tensor, loss_history, neuron_id=None, layer_name=None):
  """
  Visualize the optimized input and the loss history side by side.

  Args:
    input_tensor (torch.Tensor): The optimized imput tensor.
    loss_history (list): List of loss values recorded during optimization.
    neuron_id (int, optional):The target neoron ID being optimized.
    layer_name (str): Name of the layer being optimized.
  """

  input_image = input_tensor.detach().squeeze().permute(1,2,0).cpu().numpy()
  input_image_normalized = (input_image - input_image.min())/(input_image.max()-input_image.min())

  neuron_text = f"{neuron_id}" if isinstance(neuron_id, int) else ", ".join(map(str, neuron_id))
  title = f"Optimized Input\nNeuron(s): {neuron_id} in {layer_name}" if neuron_id and layer_name else "Optimized Input"

  fig, ax = plt.subplots(nrows=1, ncols=2, figsize=(12,5))


  ax[0].imshow(input_image_normalized)
  ax[0].set_title(title)
  ax[0].axis("Off")

  ax[1].plot(loss_history, marker = 'o')
  ax[1].set_title("Loss During Optimization")
  ax[1].set_xlabel("Steps")
  ax[1].set_ylabel("Loss")

  plt.tight_layout()
  return fig