Download detection/model/visual_log.py from diing/AURAD: direct link, hf CLI and curl.
- Browser
- Download file 4.31 kB
-
https://huggingface.co/diing/AURAD/resolve/main/detection/model/visual_log.py
- Command line
-
hf download hf://diing/AURAD/detection/model/visual_log.py
-
curl -L -o visual_log.py https://huggingface.co/diing/AURAD/resolve/main/detection/model/visual_log.py
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 | |