File size: 3,134 Bytes
6ca1e94
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
import numpy as np
import random
import matplotlib.pyplot as plt
import matplotlib.patches as patches
from src.dataset import VOC_CLASSES

def plot_image_with_annotations(img_array, annotation_dict, img_width=None, img_height=None):
    fig, ax = plt.subplots(1)
    ax.imshow(img_array)

    if img_width is None:
        img_width = annotation_dict["size"]["width"]
    if img_height is None:
        img_height = annotation_dict["size"]["height"]

    for obj in annotation_dict["objects"]:
        x_min = obj["bndbox"]["x_min"]*img_width
        y_min = obj["bndbox"]["y_min"]*img_height
        x_max = obj["bndbox"]["x_max"]*img_width
        y_max = obj["bndbox"]["y_max"]*img_height

        width = x_max - x_min
        height = y_max - y_min

        rect = patches.Rectangle((x_min, y_min), width, height, linewidth=2, edgecolor='r', facecolor='none')
        ax.add_patch(rect)

        ax.text(x_min, y_min - 5, obj["name"], color='r', fontsize=12, weight='bold')

    plt.show()


def display_random_images_with_annotations(dataset, 
                                           num_images=5,
                                           display_shape : bool = True,
                                           seed : int = None):
    """
    Displays a random selection of images from the dataset along with their annotations.

    Args:
        dataset: A dataset object that provides access to images and their annotations.
        num_images: The number of random images to display.
        display_shape: If True, displays the shape of each image.
        seed: Random seed for reproducibility.
    """

    # Set the random seed for reproducibility
    if seed is not None:
        random.seed(seed)

    if num_images > 10:
        num_images = 10
        print("Warning: Displaying more than 10 images may clutter the output. Displaying only 10 images.")

    random_indices = random.sample(range(len(dataset)), num_images)

    for idx in random_indices:
        img_tensor, annotation_dict = dataset[idx]
        
        plot_image_with_annotations(img_tensor.permute(1, 2, 0).numpy(), annotation_dict, img_width=img_tensor.shape[2], img_height=img_tensor.shape[1])



def display_random_batch_images(batch_images, batch_boxes, batch_labels=None, num_images=4):
    fig, axes = plt.subplots(1, num_images, figsize=(15, 5))
    for i in range(num_images):
        idx = random.randint(0, len(batch_images) - 1)
        img = batch_images[idx].permute(1, 2, 0).cpu().numpy()  # Convert to HWC format and move to CPU
        boxes = batch_boxes[idx].cpu().numpy()

        axes[i].imshow(img)
        axes[i].set_title(f"Image {idx}")
        axes[i].axis('off')

        for j, box in enumerate(boxes):
            x1, y1, x2, y2 = box
            rect = plt.Rectangle((x1, y1), x2 - x1, y2 - y1, fill=False, color='red', linewidth=2)
            axes[i].add_patch(rect)

            if batch_labels is not None:
                labels = batch_labels[idx].cpu().numpy()

                axes[i].text(x1, y1 - 5, VOC_CLASSES[labels[j]], color='r', fontsize=12, weight='bold')

    plt.tight_layout()
    plt.show()