-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsegformer_eval.py
More file actions
116 lines (90 loc) · 3.89 KB
/
Copy pathsegformer_eval.py
File metadata and controls
116 lines (90 loc) · 3.89 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
import os
import numpy as np
from PIL import Image
from tqdm import tqdm
from sklearn.metrics import f1_score, accuracy_score, precision_score, recall_score
from transformers import SegformerImageProcessor, SegformerForSemanticSegmentation
import torch
import matplotlib.pyplot as plt
model_path = "segformer_output/checkpoint-940"
image_dir = "images/test"
mask_dir = "images/test_labels"
batch_size = 4
resize_dim = (640, 640)
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")
road_rgb = (128, 64, 128)
processor = SegformerImageProcessor.from_pretrained(model_path)
model = SegformerForSemanticSegmentation.from_pretrained(model_path).to(device)
model.eval()
def resize_image(image, size):
return image.resize(size, resample=Image.NEAREST)
def process_batch(batch_images, batch_masks, model, processor):
y_true, y_pred = [], []
batch_images = [resize_image(img, resize_dim) for img in batch_images]
batch_masks = [resize_image(mask, resize_dim) for mask in batch_masks]
inputs = processor(images=batch_images, return_tensors="pt", padding=True).to(device)
with torch.no_grad():
outputs = model(**inputs)
logits = outputs.logits
for pred_mask, mask in zip(logits, batch_masks):
pred_mask = torch.argmax(pred_mask, dim=0).cpu().numpy()
pred_mask_resized = resize_image(Image.fromarray(pred_mask.astype(np.uint8)), resize_dim)
pred_mask_resized = np.array(pred_mask_resized)
mask_array = np.array(mask.convert("RGB"))
road_mask = np.all(mask_array == road_rgb, axis=-1).astype(np.uint8)
y_true.extend(road_mask.flatten())
y_pred.extend(pred_mask_resized.flatten())
torch.cuda.empty_cache()
return y_true, y_pred
def evaluate_model(image_dir, mask_dir, model, processor, batch_size):
y_true, y_pred = [], []
images = sorted(os.listdir(image_dir))
num_samples = len(images)
for i in tqdm(range(0, num_samples, batch_size)):
batch_images = []
batch_masks = []
for j in range(i, min(i + batch_size, num_samples)):
image_path = os.path.join(image_dir, images[j])
mask_name = (images[j].split(".png")[0] + "_L.png")
mask_path = os.path.join(mask_dir, mask_name)
image = Image.open(image_path).convert("RGB")
mask = Image.open(mask_path).convert("RGB")
batch_images.append(image)
batch_masks.append(mask)
if batch_images:
batch_true, batch_pred = process_batch(batch_images, batch_masks, model, processor)
y_true.extend(batch_true)
y_pred.extend(batch_pred)
print(f"y_true length: {len(y_true)}, y_pred length: {len(y_pred)}")
precision = precision_score(y_true, y_pred, pos_label=1, average="binary")
recall = recall_score(y_true, y_pred, pos_label=1, average="binary")
f1 = f1_score(y_true, y_pred, pos_label=1, average="binary")
acc = accuracy_score(y_true, y_pred)
intersection = np.logical_and(np.array(y_true) == 1, np.array(y_pred) == 1).sum()
union = np.logical_or(np.array(y_true) == 1, np.array(y_pred) == 1).sum()
iou = intersection / union if union != 0 else 0
return {
"F1 Score (Road)": f1,
"Accuracy": acc,
"Precision (Road)": precision,
"Recall (Road)": recall,
"IoU (Road)": iou
}
def plot_metrics(results):
metrics = list(results.keys())
values = list(results.values())
plt.figure(figsize=(10, 6))
plt.bar(metrics, values, color="skyblue")
plt.ylim(0, 1)
plt.xlabel("Metrics")
plt.ylabel("Values")
plt.title("Model Evaluation Metrics")
plt.xticks(rotation=45)
plt.grid(axis="y", linestyle="--", alpha=0.7)
plt.tight_layout()
plt.show()
results = evaluate_model(image_dir, mask_dir, model, processor, batch_size)
for metric, value in results.items():
print(f"{metric}: {value:.4f}")
plot_metrics(results)