AURAD / detection /model /visual_log.py
diing's picture
Upload folder using huggingface_hub
41c8683 verified
Raw History Blame Contribute Delete
4.31 kB
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