BioMike/clipsegmulticlass
0
1from dataclasses import dataclass
2from typing import Optional, Tuple, Union, List
3from PIL import Image
4import PIL
5import torch
6import torch.nn as nn
7import torch.nn.functional as F
8from transformers import (
9 PreTrainedModel,
10 CLIPSegProcessor,
11 CLIPSegForImageSegmentation,
12)
13from transformers.modeling_outputs import ModelOutput
14
15from .config import ClipSegMultiClassConfig
16from sklearn.metrics import precision_score, recall_score, f1_score, accuracy_score
17import numpy as np
18from torch.utils.data import DataLoader
19from collections import defaultdict
20
21def flatten_outputs(preds, targets, num_classes):
22 """Flatten predictions and targets to 1D arrays, filter ignored labels."""
23 preds = preds.cpu().numpy().reshape(-1)
24 targets = targets.cpu().numpy().reshape(-1)
25
26 mask = (targets >= 0) & (targets < num_classes)
27 return preds[mask], targets[mask]
28
29def compute_metrics(all_preds, all_targets, num_classes, average="macro"):
30 y_pred = np.concatenate(all_preds)
31 y_true = np.concatenate(all_targets)
32
33 metrics = {
34 "accuracy": accuracy_score(y_true, y_pred),
35 "precision": precision_score(y_true, y_pred, average=average, zero_division=0),
36 "recall": recall_score(y_true, y_pred, average=average, zero_division=0),
37 "f1": f1_score(y_true, y_pred, average=average, zero_division=0),
38 }
39
40 return metrics
41
42
43@dataclass
44class ClipSegMultiClassOutput(ModelOutput):
45 loss: Optional[torch.FloatTensor] = None
46 logits: Optional[torch.FloatTensor] = None
47 predictions: Optional[torch.LongTensor] = None
48
49
50class ClipSegMultiClassModel(PreTrainedModel):
51 config_class = ClipSegMultiClassConfig
52 base_model_prefix = "clipseg_multiclass"
53
54 def __init__(self, config: ClipSegMultiClassConfig):
55 super().__init__(config)
56
57 self.config = config
58 self.class_labels = config.class_labels
59 self.num_classes = config.num_classes
60 self.processor = CLIPSegProcessor.from_pretrained(config.model)
61 self.clipseg = CLIPSegForImageSegmentation.from_pretrained(config.model)
62 self.loss_fct = nn.CrossEntropyLoss()
63
64 def forward(
65 self,
66 pixel_values: Optional[torch.Tensor] = None,
67 input_ids: Optional[torch.Tensor] = None,
68 labels: Optional[torch.Tensor] = None,
69 **kwargs
70 ) -> ClipSegMultiClassOutput:
71
72 if pixel_values is None or input_ids is None:
73 raise ValueError("Both `pixel_values` and `input_ids` must be provided.")
74
75 pixel_values = pixel_values.to(self.device)
76 input_ids = input_ids.to(self.device)
77
78 outputs = self.clipseg(pixel_values=pixel_values, input_ids=input_ids)
79 raw_logits = outputs.logits # shape: [B * C, H, W]
80
81 B = raw_logits.shape[0] // self.num_classes
82 C = self.num_classes
83 H, W = raw_logits.shape[-2:]
84
85 logits = raw_logits.view(B, C, H, W) # [B, C, H, W]
86 pred = torch.argmax(logits, dim=1) # [B, H, W]
87
88 loss = self.loss_fct(logits, labels.long()) if labels is not None else None
89
90 return ClipSegMultiClassOutput(
91 loss=loss,
92 logits=logits,
93 predictions=pred
94 )
95
96 @torch.no_grad()
97 def predict(self, images: Union[List, "PIL.Image.Image"]) -> torch.Tensor:
98 self.eval()
99 if isinstance(images, Image.Image):
100 images = [images]
101
102 inputs = self.processor(
103 images=[img for img in images for _ in self.class_labels],
104 text=self.class_labels * len(images),
105 return_tensors="pt",
106 padding=True,
107 truncation=True
108 ).to(self.device)
109
110 output = self.forward(
111 pixel_values=inputs["pixel_values"],
112 input_ids=inputs["input_ids"]
113 )
114 return output.predictions
115
116 def evaluate(self, dataloader: torch.utils.data.DataLoader) -> dict:
117 from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score
118 import numpy as np
119
120 self.eval()
121
122 all_preds = []
123 all_targets = []
124
125 with torch.no_grad():
126 for batch in dataloader:
127 pixel_values = batch["pixel_values"].to(self.device) # [B * C, 3, H, W]
128 input_ids = batch["input_ids"].to(self.device) # [B * C, T]
129 labels = batch["labels"].to(self.device) # [B, H, W]
130
131 outputs = self.forward(pixel_values=pixel_values, input_ids=input_ids)
132 preds = outputs.predictions # [B, H, W]
133
134 for pred, label in zip(preds, labels):
135 pred = pred.cpu().flatten()
136 label = label.cpu().flatten()
137
138 mask = label != 0
139 pred = pred[mask]
140 label = label[mask]
141
142 all_preds.append(pred)
143 all_targets.append(label)
144
145 y_pred = torch.cat(all_preds).numpy()
146 y_true = torch.cat(all_targets).numpy()
147
148 return {
149 "accuracy": accuracy_score(y_true, y_pred),
150 "precision": precision_score(y_true, y_pred, average="macro", zero_division=0),
151 "recall": recall_score(y_true, y_pred, average="macro", zero_division=0),
152 "f1": f1_score(y_true, y_pred, average="macro", zero_division=0),
153 }
154
155 