polejowska commited on
Commit
9b3f85e
·
0 Parent(s):

Duplicate from polejowska/vicellst-att

Browse files
.gitattributes ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tflite filter=lfs diff=lfs merge=lfs -text
29
+ *.tgz filter=lfs diff=lfs merge=lfs -text
30
+ *.wasm filter=lfs diff=lfs merge=lfs -text
31
+ *.xz filter=lfs diff=lfs merge=lfs -text
32
+ *.zip filter=lfs diff=lfs merge=lfs -text
33
+ *.zst filter=lfs diff=lfs merge=lfs -text
34
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
35
+ cd45rb_test_imgs/CD45RB_Leukocyte_036_053248_048128_HE.png filter=lfs diff=lfs merge=lfs -text
36
+ cd45rb_test_imgs/CD45RB_Leukocyte_096_008192_025600_HE.png filter=lfs diff=lfs merge=lfs -text
37
+ cd45rb_test_imgs/CD45RB_Leukocyte_096_009216_022528_HE.png filter=lfs diff=lfs merge=lfs -text
38
+ cd45rb_test_imgs/CD45RB_Leukocyte_096_009216_025600_HE.png filter=lfs diff=lfs merge=lfs -text
39
+ cd45rb_test_imgs/CD45RB_Leukocyte_096_013312_024576_HE.png filter=lfs diff=lfs merge=lfs -text
40
+ cd45rb_test_imgs/CD45RB_Leukocyte_096_013312_025600_HE.png filter=lfs diff=lfs merge=lfs -text
41
+ cd45rb_test_imgs/CD45RB_Leukocyte_096_053248_075776_HE.png filter=lfs diff=lfs merge=lfs -text
42
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_046080_050176_HE.png filter=lfs diff=lfs merge=lfs -text
43
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_047104_048128_HE.png filter=lfs diff=lfs merge=lfs -text
44
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_047104_049152_HE.png filter=lfs diff=lfs merge=lfs -text
45
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_048128_050176_HE.png filter=lfs diff=lfs merge=lfs -text
46
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_048128_051200_HE.png filter=lfs diff=lfs merge=lfs -text
47
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_048128_052224_HE.png filter=lfs diff=lfs merge=lfs -text
48
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_050176_048128_HE.png filter=lfs diff=lfs merge=lfs -text
49
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_050176_049152_HE.png filter=lfs diff=lfs merge=lfs -text
50
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_050176_050176_HE.png filter=lfs diff=lfs merge=lfs -text
51
+ cd45rb_test_imgs/CD45RB_Leukocyte_228_090112_066560_HE.png filter=lfs diff=lfs merge=lfs -text
52
+ cd45rb_test_imgs/CD45RB_Leukocyte_228_091136_006144_HE.png filter=lfs diff=lfs merge=lfs -text
53
+ cd45rb_test_imgs/CD45RB_Leukocyte_228_091136_062464_HE.png filter=lfs diff=lfs merge=lfs -text
54
+ cd45rb_test_imgs/CD45RB_Leukocyte_228_092160_008192_HE.png filter=lfs diff=lfs merge=lfs -text
55
+ cd45rb_test_imgs/CD45RB_Leukocyte_228_092160_062464_HE.png filter=lfs diff=lfs merge=lfs -text
56
+ cd45rb_test_imgs/CD45RB_Leukocyte_281_096256_027648_HE.png filter=lfs diff=lfs merge=lfs -text
57
+ cd45rb_test_imgs/CD45RB_Leukocyte_281_097280_022528_HE.png filter=lfs diff=lfs merge=lfs -text
58
+ cd45rb_test_imgs/CD45RB_Leukocyte_281_097280_024576_HE.png filter=lfs diff=lfs merge=lfs -text
59
+ cd45rb_test_imgs/CD45RB_Leukocyte_281_097280_025600_HE.png filter=lfs diff=lfs merge=lfs -text
60
+ cd45rb_test_imgs/CD45RB_Leukocyte_281_097280_026624_HE.png filter=lfs diff=lfs merge=lfs -text
61
+ cd45rb_test_imgs/CD45RB_Leukocyte_281_097280_027648_HE.png filter=lfs diff=lfs merge=lfs -text
62
+ cd45rb_test_imgs/CD45RB_Leukocyte_059_092160_016384_HE.png filter=lfs diff=lfs merge=lfs -text
63
+ cd45rb_test_imgs/CD45RB_Leukocyte_080_044032_053248_HE.png filter=lfs diff=lfs merge=lfs -text
64
+ cd45rb_test_imgs/CD45RB_Leukocyte_080_101376_038912_HE.png filter=lfs diff=lfs merge=lfs -text
65
+ cd45rb_test_imgs/CD45RB_Leukocyte_101_057344_037888_HE.png filter=lfs diff=lfs merge=lfs -text
66
+ cd45rb_test_imgs/CD45RB_Leukocyte_124_019456_009216_HE.png filter=lfs diff=lfs merge=lfs -text
67
+ cd45rb_test_imgs/CD45RB_Leukocyte_205_010240_018432_HE.png filter=lfs diff=lfs merge=lfs -text
68
+ cd45rb_test_imgs/CD45RB_Leukocyte_228_089088_062464_HE.png filter=lfs diff=lfs merge=lfs -text
69
+ cd45rb_test_imgs/CD45RB_Leukocyte_253_054272_057344_HE.png filter=lfs diff=lfs merge=lfs -text
70
+ cd45rb_test_imgs/CD45RB_Leukocyte_253_055296_057344_HE.png filter=lfs diff=lfs merge=lfs -text
README.md ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ title: Vicellst
3
+ emoji: 🦀
4
+ colorFrom: blue
5
+ colorTo: gray
6
+ sdk: gradio
7
+ sdk_version: 3.18.0
8
+ app_file: app.py
9
+ pinned: false
10
+ duplicated_from: polejowska/vicellst-att
11
+ ---
12
+
13
+ Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
app.py ADDED
@@ -0,0 +1,146 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import pathlib
2
+ from constants import MODELS_REPO, MODELS_NAMES
3
+
4
+ import gradio as gr
5
+ import torch
6
+
7
+ from transformers import AutoFeatureExtractor, DetrForObjectDetection
8
+ from visualization import visualize_attention_map, visualize_prediction
9
+ from style import css, description, title
10
+
11
+ from PIL import Image
12
+
13
+
14
+
15
+ def make_prediction(img, feature_extractor, model):
16
+ inputs = feature_extractor(img, return_tensors="pt")
17
+ outputs = model(**inputs)
18
+ img_size = torch.tensor([tuple(reversed(img.size))])
19
+ processed_outputs = feature_extractor.post_process(outputs, img_size)
20
+ print(outputs.keys())
21
+ return (
22
+ processed_outputs[0],
23
+ outputs["decoder_attentions"],
24
+ outputs["encoder_attentions"],
25
+ )
26
+
27
+
28
+ def detect_objects(model_name, image_input, threshold, display_mask=False, img_input_mask=None):
29
+ feature_extractor = AutoFeatureExtractor.from_pretrained(MODELS_REPO[model_name])
30
+
31
+ if "DETR" in model_name:
32
+ model = DetrForObjectDetection.from_pretrained(MODELS_REPO[model_name])
33
+ model_details = "DETR details"
34
+
35
+ (
36
+ processed_outputs,
37
+ decoder_attention_map,
38
+ encoder_attention_map,
39
+ ) = make_prediction(image_input, feature_extractor, model)
40
+
41
+ viz_img = visualize_prediction(
42
+ pil_img=image_input,
43
+ output_dict=processed_outputs,
44
+ threshold=threshold,
45
+ id2label=model.config.id2label,
46
+ display_mask=display_mask,
47
+ mask=img_input_mask
48
+ )
49
+ decoder_attention_map_img = visualize_attention_map(
50
+ image_input, decoder_attention_map
51
+ )
52
+ encoder_attention_map_img = visualize_attention_map(
53
+ image_input, encoder_attention_map
54
+ )
55
+
56
+ return (
57
+ viz_img,
58
+ decoder_attention_map_img,
59
+ encoder_attention_map_img,
60
+ model_details
61
+ )
62
+
63
+
64
+ def set_example_image(example: list):
65
+ print(f"Set example image to: {example[0]}")
66
+ print(f"Set example image mask to: {example[1]}")
67
+ return gr.Image.update(value=example[0]), gr.Image.update(value=example[1])
68
+
69
+
70
+ with gr.Blocks(css=css) as app:
71
+ gr.Markdown(title)
72
+
73
+ with gr.Tabs():
74
+ with gr.TabItem("Image upload and detections visualization"):
75
+ with gr.Row():
76
+ with gr.Column():
77
+ with gr.Row():
78
+ img_input = gr.Image(type="pil")
79
+ img_input_mask = gr.Image(type="pil", visible=False)
80
+ with gr.Row():
81
+ example_images = gr.Dataset(
82
+ components=[img_input, img_input_mask],
83
+ samples=[
84
+ [path.as_posix(), path.as_posix().replace("_HE", "_mask")]
85
+ for path in sorted(
86
+ pathlib.Path("cd45rb_test_imgs").rglob("*_HE.png")
87
+ )
88
+ ],
89
+ samples_per_page=2,
90
+ )
91
+ with gr.Column():
92
+ with gr.Row():
93
+ options = gr.Dropdown(
94
+ value=MODELS_NAMES[0],
95
+ choices=MODELS_NAMES,
96
+ label="Select an object detection model",
97
+ show_label=True,
98
+ )
99
+ with gr.Row():
100
+ slider_input = gr.Slider(
101
+ minimum=0.2, maximum=1, value=0.7, label="Prediction threshold"
102
+ )
103
+ with gr.Row():
104
+ display_mask = gr.Checkbox(
105
+ label="Display masks", default=False
106
+ )
107
+ with gr.Row():
108
+ detect_button = gr.Button("Detect leukocytes")
109
+ with gr.Row():
110
+ with gr.Column():
111
+ gr.Markdown(
112
+ """The selected image with detected bounding boxes by the model"""
113
+ )
114
+ img_output_from_upload = gr.Image(shape=(800, 800))
115
+ with gr.TabItem("Attentions visualization"):
116
+ gr.Markdown("""Encoder attentions""")
117
+ with gr.Row():
118
+ encoder_att_map_output = gr.Image(shape=(850, 850))
119
+ gr.Markdown("""Decoder attentions""")
120
+ with gr.Row():
121
+ decoder_att_map_output = gr.Image(shape=(850, 850))
122
+ with gr.TabItem("Model details"):
123
+ with gr.Row():
124
+ model_details = gr.Markdown(""" """)
125
+ with gr.TabItem("Dataset details"):
126
+ with gr.Row():
127
+ gr.Markdown(description)
128
+
129
+ detect_button.click(
130
+ detect_objects,
131
+ inputs=[options, img_input, slider_input, display_mask, img_input_mask],
132
+ outputs=[
133
+ img_output_from_upload,
134
+ decoder_att_map_output,
135
+ encoder_att_map_output,
136
+ # cross_att_map_output,
137
+ model_details,
138
+ ],
139
+ queue=True,
140
+ )
141
+ example_images.click(
142
+ fn=set_example_image, inputs=[example_images], outputs=[img_input, img_input_mask],
143
+ show_progress=True
144
+ )
145
+
146
+ app.launch(enable_queue=True)
cd45rb_test_imgs/CD45RB_Leukocyte_059_092160_016384_HE.png ADDED

Git LFS Details

  • SHA256: 68daf280db7680ee8d3df9acb7d35a27c953aea73335686ea4f8a152d1dfa4cd
  • Pointer size: 132 Bytes
  • Size of remote file: 2 MB
cd45rb_test_imgs/CD45RB_Leukocyte_059_092160_016384_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_080_044032_053248_HE.png ADDED

Git LFS Details

  • SHA256: 0cfcb269cb6e4e179a884e9c0ee6ea1d70fffe66569d9751d56830b314370497
  • Pointer size: 132 Bytes
  • Size of remote file: 1.53 MB
cd45rb_test_imgs/CD45RB_Leukocyte_080_044032_053248_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_080_101376_038912_HE.png ADDED

Git LFS Details

  • SHA256: 035cbe46da8752d82e1c67748af0f1585ea6aecf52a8b19487942a39990e436c
  • Pointer size: 132 Bytes
  • Size of remote file: 1.98 MB
cd45rb_test_imgs/CD45RB_Leukocyte_080_101376_038912_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_101_057344_037888_HE.png ADDED

Git LFS Details

  • SHA256: 9320674148dd255e09ff54cfbf6a1eac854473ae0989a615b833a2097fdacba2
  • Pointer size: 132 Bytes
  • Size of remote file: 2.05 MB
cd45rb_test_imgs/CD45RB_Leukocyte_101_057344_037888_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_124_019456_009216_HE.png ADDED

Git LFS Details

  • SHA256: f490aaa89ef2b5458d03eb5277795aaf1e161a13f8459a8cdf11d0cd54b21f2f
  • Pointer size: 132 Bytes
  • Size of remote file: 2.1 MB
cd45rb_test_imgs/CD45RB_Leukocyte_124_019456_009216_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_205_010240_018432_HE.png ADDED

Git LFS Details

  • SHA256: a0e6721547fdc24102e34a9ae4b88b3fa539d73cfdf41d2e2b7d6d4314f99b99
  • Pointer size: 132 Bytes
  • Size of remote file: 1.91 MB
cd45rb_test_imgs/CD45RB_Leukocyte_205_010240_018432_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_228_089088_062464_HE.png ADDED

Git LFS Details

  • SHA256: 050d85f373d787597d687afe92fa2bd68f2647ce5a46b18dcb89c49cce47d4fb
  • Pointer size: 132 Bytes
  • Size of remote file: 2.02 MB
cd45rb_test_imgs/CD45RB_Leukocyte_228_089088_062464_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_253_054272_057344_HE.png ADDED

Git LFS Details

  • SHA256: ccab192497326655944efc21642da2c861888d8b148b656784bc183b752cc0c0
  • Pointer size: 132 Bytes
  • Size of remote file: 2.13 MB
cd45rb_test_imgs/CD45RB_Leukocyte_253_054272_057344_mask.png ADDED
cd45rb_test_imgs/CD45RB_Leukocyte_253_055296_057344_HE.png ADDED

Git LFS Details

  • SHA256: 60d7da0539ce7b04035d4ff254309156d38b93c4827f38034d0c745cd38723aa
  • Pointer size: 132 Bytes
  • Size of remote file: 2.06 MB
cd45rb_test_imgs/CD45RB_Leukocyte_253_055296_057344_mask.png ADDED
constants.py ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ COLORS = [
2
+ [0.000, 0.447, 0.741],
3
+ [0.850, 0.325, 0.098],
4
+ [0.929, 0.694, 0.125],
5
+ [0.494, 0.184, 0.556],
6
+ [0.466, 0.674, 0.188],
7
+ [0.301, 0.745, 0.933],
8
+ [0.351, 0.760, 0.903],
9
+ ]
10
+
11
+ MODELS_REPO = {
12
+ "DETR-RESNET-50": "polejowska/detr-resnet-50-CD45RB-1000-att",
13
+ "DEFORMABLE-DETR": "polejowska/deformable-detr-resnet50-leuk",
14
+ "CONDITIONAL-DETR": "polejowska/cdetr-cd45rb-s",
15
+ "DETR-RESNET-101": "polejowska/detr-resnet-101-CD45RB-1000-att",
16
+ "DETR-RESNET-50-TEST": "polejowska/detr-resnet50-leuk",
17
+ }
18
+
19
+ MODELS_NAMES = ["DETR-RESNET-50", "DEFORMABLE-DETR", "CONDITIONAL-DETR", "DETR-RESNET-101", "DETR-RESNET-50-TEST"]
requirements.txt ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ beautifulsoup4==4.9.3
2
+ bs4==0.0.1
3
+ requests-file==1.5.1
4
+ torch==1.10.1
5
+ git+https://github.com/huggingface/transformers.git
6
+ validators==0.18.2
7
+ timm==0.5.4
8
+ opencv-python
style.py ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ title = """<h2 id="title">vicellst</h2>"""
2
+
3
+ description = """
4
+ Citation for the dataset used to train the model and provide examples in this interface:
5
+ > @dataset{komura_daisuke_2022_7412739,
6
+ > author = {Komura, Daisuke},
7
+ > title = {{Large-scale annotation dataset for cell/tissue
8
+ > segmentation in H\&E-stained images : anti-CD45RB
9
+ > (leukocytes)}},
10
+ > month = apr,
11
+ > year = 2022,
12
+ > publisher = {Zenodo},
13
+ > version = {0.3},
14
+ > doi = {10.5281/zenodo.7412739},
15
+ > url = {https://doi.org/10.5281/zenodo.7412739}
16
+ }
17
+ """
18
+
19
+ css = """
20
+ h2#title {
21
+ text-align: left;
22
+ }
23
+ """
utils.py ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ import io
2
+ from PIL import Image
3
+
4
+
5
+ def fig2img(fig):
6
+ buf = io.BytesIO()
7
+ fig.savefig(buf)
8
+ buf.seek(0)
9
+ return Image.open(buf)
visualization.py ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from matplotlib import pyplot as plt
2
+ from PIL import Image
3
+ import numpy as np
4
+ import torch
5
+ import torch.nn.functional as F
6
+
7
+ from constants import COLORS
8
+ from utils import fig2img
9
+
10
+
11
+ def visualize_prediction(
12
+ pil_img, output_dict, threshold=0.7, id2label=None, display_mask=False, mask=None
13
+ ):
14
+ keep = output_dict["scores"] > threshold
15
+ boxes = output_dict["boxes"][keep].tolist()
16
+ scores = output_dict["scores"][keep].tolist()
17
+ labels = output_dict["labels"][keep].tolist()
18
+ if id2label is not None:
19
+ labels = [id2label[x] for x in labels]
20
+
21
+ fig, ax = plt.subplots(figsize=(12, 12))
22
+ ax.imshow(pil_img)
23
+ if display_mask and mask is not None:
24
+ # Convert the mask image to a numpy array
25
+ mask_arr = np.asarray(mask)
26
+
27
+ # Create a new mask with white objects and black background
28
+ new_mask = np.zeros_like(mask_arr)
29
+ new_mask[mask_arr > 0] = 255
30
+
31
+ # Convert the numpy array back to a PIL Image
32
+ new_mask = Image.fromarray(new_mask)
33
+
34
+ # Display the new mask as a semi-transparent overlay
35
+ ax.imshow(new_mask, alpha=0.5, cmap='viridis')
36
+
37
+ colors = COLORS * 100
38
+ for score, (xmin, ymin, xmax, ymax), label, color in zip(
39
+ scores, boxes, labels, colors
40
+ ):
41
+ ax.add_patch(
42
+ plt.Rectangle(
43
+ (xmin, ymin),
44
+ xmax - xmin,
45
+ ymax - ymin,
46
+ fill=False,
47
+ color=color,
48
+ linewidth=2,
49
+ )
50
+ )
51
+ ax.text(
52
+ xmin,
53
+ ymin,
54
+ f"{score:0.2f}",
55
+ fontsize=8,
56
+ bbox=dict(facecolor="yellow", alpha=0.5),
57
+ )
58
+ ax.axis("off")
59
+ return fig2img(fig)
60
+
61
+
62
+ def visualize_attention_map(pil_img, attention_map):
63
+ # Get the attention map for the last layer
64
+ attention_map = attention_map[-1].detach().cpu()
65
+
66
+ # Get the number of heads
67
+ n_heads = attention_map.shape[1]
68
+
69
+ # Calculate the average attention weight for each head
70
+ avg_attention_weight = torch.mean(attention_map, dim=1).squeeze()
71
+
72
+ # Resize the attention map
73
+ resized_attention_weight = F.interpolate(
74
+ avg_attention_weight.unsqueeze(0).unsqueeze(0),
75
+ size=pil_img.size[::-1],
76
+ mode="bicubic",
77
+ ).squeeze().numpy()
78
+
79
+ # Create a grid of subplots
80
+ fig, axes = plt.subplots(nrows=1, ncols=n_heads, figsize=(n_heads*4, 4))
81
+
82
+ # Loop through the subplots and plot the attention for each head
83
+ for i, ax in enumerate(axes.flat):
84
+ ax.imshow(pil_img)
85
+ ax.imshow(attention_map[0,i,:,:].squeeze(), alpha=0.7, cmap="viridis")
86
+ ax.set_title(f"Head {i+1}")
87
+ ax.axis("off")
88
+
89
+ plt.tight_layout()
90
+
91
+ return fig2img(fig)
92
+ # attention_map = attention_map[-1].detach().cpu()
93
+ # avg_attention_weight = torch.mean(attention_map, dim=1).squeeze()
94
+ # avg_attention_weight_resized = (
95
+ # F.interpolate(
96
+ # avg_attention_weight.unsqueeze(0).unsqueeze(0),
97
+ # size=pil_img.size[::-1],
98
+ # mode="bicubic",
99
+ # )
100
+ # .squeeze()
101
+ # .numpy()
102
+ # )
103
+
104
+ # plt.imshow(pil_img)
105
+ # plt.imshow(avg_attention_weight_resized, alpha=0.7, cmap="viridis")
106
+ # plt.axis("off")
107
+ # fig = plt.gcf()
108
+ # return fig2img(fig)