Repository navigation
Expand file tree
/
Copy pathevaluation.py
More file actions
65 lines (49 loc) · 1.87 KB
/
Copy pathevaluation.py
File metadata and controls
65 lines (49 loc) · 1.87 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
import torch
from tqdm import tqdm
import clip
from accelerate import Accelerator
@torch.no_grad()
def evaluate(accelerator: Accelerator, model, dataloader, texts):
"""
Evaluate zero-shot classification accuracy by comparing image features
against the precomputed text classifier weights.
For Multicard Training, return **GATHERED** metric for all processes.
Returns:
accuracy: float in [0, 1]
"""
model.eval()
all_preds = []
all_labels = []
for images, labels in tqdm(dataloader, desc="Evaluating", leave=False, disable=not accelerator.is_main_process):
# accelerate 自动处理
# images = images.to(device)
# labels = labels.to(device)
# [B, num_classes]
logits, _ = model(images, texts)
preds = logits.argmax(dim=-1)
all_preds.append(accelerator.gather_for_metrics(preds))
all_labels.append(accelerator.gather_for_metrics(labels))
all_preds = torch.cat(all_preds)
all_labels = torch.cat(all_labels)
accuracy = (all_preds == all_labels).sum().item() / all_labels.size(0) if all_labels.size(0) > 0 else 0.0
return accuracy
@torch.no_grad()
def evaluate_multiple(accelerator, model, datasets_dict):
"""
Evaluate the model on multiple datasets simultaneously.
Used to track zero-shot transfer performance during finetuning.
Args:
model: CLIP model
datasets_dict: {name: dataset_instance} dictionary
device: torch device
Returns:
results: {name: accuracy} dictionary
"""
results = {}
model.eval()
for name, dataset in datasets_dict.items():
texts = [dataset.template(x) for x in dataset.classnames]
texts = clip.tokenize(texts).to(accelerator.device)
acc = evaluate(accelerator, model, dataset.test_loader, texts)
results[name] = acc
return results