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