Team Ai
Apppublic

BioMike/clipsegmulticlass

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
model.py155 linesDownload Raw Back to src
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