-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathselect_sample.py
More file actions
85 lines (62 loc) · 3.11 KB
/
Copy pathselect_sample.py
File metadata and controls
85 lines (62 loc) · 3.11 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
import numpy as np
import torch
import os
from models.construct import model_construct
def distribute_budget(budget, n):
avg_budget = budget / n
lower_budget = int(np.floor(avg_budget))
upper_budget = int(np.ceil(avg_budget))
num_upper = budget - lower_budget * n
num_lower = n - num_upper
budget_distribution = [lower_budget] * num_lower + [upper_budget] * num_upper
return budget_distribution
def entropy(probabilities):
normalized_prob = probabilities / torch.sum(probabilities, dim=1, keepdim=True)
entropy = -torch.sum(normalized_prob * torch.log2(normalized_prob), dim=1)
return entropy
def sort(data, args, idx_train, idx_val, device):
model = model_construct(args, args.model, data, device).to(device)
model.fit(data.x, data.edge_index, None, data.y, idx_train, idx_val, train_iters=args.epochs, verbose=False)
model.eval()
logits = model(data.x, data.edge_index)
unlabeled_idx = (torch.bitwise_not(data.test_mask) & torch.bitwise_not(data.train_mask)).nonzero().flatten()
full_label = torch.zeros_like(data.y)
full_label[idx_train] = data.y[idx_train]
full_label[unlabeled_idx] = torch.argmax(logits[unlabeled_idx], dim=1)
entropy_score = entropy(logits).detach()
sorted_unlabeled_idx = unlabeled_idx[entropy_score[unlabeled_idx].argsort(descending=True)]
num_classes = int(full_label.max()) + 1
sorted_samples = torch.empty(0, dtype=torch.long, device=device)
class_pointers = {i: 0 for i in range(num_classes)}
while len(sorted_samples) < len(unlabeled_idx):
all_classes_exhausted = True
for i in range(num_classes):
class_idx = (full_label[sorted_unlabeled_idx] == i).nonzero(as_tuple=True)[0]
if class_pointers[i] < len(class_idx):
sorted_samples = torch.cat([sorted_samples, sorted_unlabeled_idx[class_idx[class_pointers[i]]].unsqueeze(0)])
class_pointers[i] += 1
all_classes_exhausted = False
if len(sorted_samples) >= len(unlabeled_idx):
break
if all_classes_exhausted:
break
current_dir = os.getcwd()
save_dir = os.path.join(current_dir, "sorted_samples")
if not os.path.exists(save_dir):
os.makedirs(save_dir)
file_path = os.path.join(save_dir, f"{args.dataset}/sorted_samples.pt")
torch.save(sorted_samples, file_path)
return sorted_samples
def select(data, args, idx_train, idx_val, device):
current_dir = os.getcwd()
save_dir = os.path.join(current_dir, "sorted_samples")
file_path = os.path.join(save_dir, f"{args.dataset}/sorted_samples.pt")
if os.path.exists(file_path):
print(f"Loading selected samples from: {file_path}")
selected_samples = torch.load(file_path, map_location=device)
else:
print("sorted_samples.pt not found, calling sort() to generate it.")
selected_samples = sort(data, args, idx_train, idx_val, device)
budget = args.vs_number
selected_samples = selected_samples[:budget]
return selected_samples