File size: 4,311 Bytes
41c8683
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
import torch
import torchvision.transforms as transforms
import wandb
import numpy as np
from PIL import Image, ImageDraw, ImageFont

DISEASES = ["Atelectasis", "Calcification", "Cardiomegaly", "Consolidation", "Diffuse Nodule", "Effusion",
            "Emphysema", "Fibrosis", "Fracture", "Mass", "Nodule","Pleural Thickening", "Pneumothorax"]

COLOR_MAP = {
    "Atelectasis": (255, 0, 0),      
    "Calcification": (0, 255, 0),    
    "Cardiomegaly": (0, 0, 255),    
    "Consolidation": (255, 255, 0),  
    "Diffuse Nodule": (255, 165, 0), 
    "Effusion": (0, 255, 255),      
    "Emphysema": (255, 0, 255),      
    "Fibrosis": (128, 0, 128),       
    "Fracture": (255, 192, 203),     
    "Mass": (173, 255, 47),          
    "Nodule": (0, 128, 255),         
    "Pleural Thickening": (75, 0, 130), 
    "Pneumothorax": (255, 105, 180)  
}

def visualize_sample(img, target, alpha=0.3):

    img_pil = transforms.ToPILImage()(img).convert("RGB")
    overlay = Image.new("RGBA", img_pil.size, (0, 0, 0, 0))

    if not target or target.get("labels") is None or target["labels"].shape[0] == 0:
        return img_pil

    masks = target.get("masks", None)  # (num_classes, H, W)
    
    boxes = target.get("boxes", None)
    labels = target["labels"].numpy()  # (num_classes,)
    scores = target.get("scores", None) 

    if masks is not None and masks.numel() > 0:
        for i, mask in enumerate(masks):
            disease_idx = labels[i] - 1 
            if 0 <= disease_idx < len(DISEASES):
                disease_name = DISEASES[disease_idx]
                mask_color = COLOR_MAP[disease_name]
                
                mask = mask.squeeze(0)
                mask_pil = Image.fromarray((mask.cpu().numpy() * 255).astype(np.uint8), mode="L")
                color_mask = Image.new("RGBA", img_pil.size, mask_color + (int(255 * alpha),))
                overlay.paste(color_mask, (0, 0), mask_pil)

    img_pil = Image.alpha_composite(img_pil.convert("RGBA"), overlay)
    draw = ImageDraw.Draw(img_pil)
    font = ImageFont.load_default()

    if boxes is not None and boxes.numel() > 0:
        boxes = boxes.numpy()  # (num_classes, 4)
        for i, box in enumerate(boxes):
            disease_idx = labels[i] - 1
            if 0 <= disease_idx < len(DISEASES):  
                disease_name = DISEASES[disease_idx]
                color = COLOR_MAP[disease_name]
                x_min, y_min, x_max, y_max = box
                # draw.rectangle([x_min, y_min, x_max, y_max], outline=color, width=2)
                text = disease_name
                if scores is not None:
                    text += f" {scores[i]:.2f}"
                text_position = (x_min, max(y_min - 12, 0))
                draw.text(text_position, text, fill=color, font=font)

    return img_pil.convert("RGB")


def image_log(images, targets, preds, epoch, mode="val"):
    batch_size = len(images)
    num_samples = min(8, batch_size) 

    vis_images = []  
    for i in range(num_samples):
        img_sample = images[i].detach().cpu().numpy()
        target_sample = {k: v.cpu() for k, v in targets[i].items()}
        pred_sample = {k: v.cpu() for k, v in preds[i].items()} if preds[i] else {}
        
        img_sample = (np.transpose(img_sample, (1, 2, 0)) * 255).astype(np.uint8)
        target_draw = visualize_sample(img_sample, target_sample)
        output_draw =  visualize_sample(img_sample, pred_sample)
        
        row_image = np.concatenate([img_sample, np.array(target_draw), np.array(output_draw)], axis=0)
        vis_images.append(row_image)

    final_image = np.concatenate(vis_images, axis=1).astype(np.uint8)
    pil_image = Image.fromarray(final_image)
    wandb.log({f"{mode}_batch_visualization": wandb.Image(pil_image, caption=f"Epoch {epoch}")})


def image_save(images, preds):
    batch_size = len(images)
    num_samples = min(8, batch_size)

    vis_images = []  
    for i in range(num_samples):
        img_sample = images[i].detach().cpu().numpy()
        pred_sample = {k: v.cpu() for k, v in preds[i].items()} if preds[i] else {}
        
        img_sample = (np.transpose(img_sample, (1, 2, 0)) * 255).astype(np.uint8)
        output_draw =  visualize_sample(img_sample, pred_sample)
        vis_images.append(output_draw)

    return vis_images