-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathutils_file.py
More file actions
45 lines (37 loc) · 1.48 KB
/
Copy pathutils_file.py
File metadata and controls
45 lines (37 loc) · 1.48 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
import torch
import numpy as np
def average_models(global_model, local_models):
global_dict = global_model.state_dict()
for k in global_dict.keys():
global_dict[k] = torch.stack([local_model.state_dict()[k].float() for local_model in local_models], 0).mean(0)
global_model.load_state_dict(global_dict)
return global_model
def return_score(model):
"""
Returns the scores for the given model.
This function assumes that the model has an attribute `popup_scores`
or similar, which contains the scores for each layer or parameter.
Args:
model (torch.nn.Module): The model from which to extract scores.
Returns:
list: A list of scores (numpy arrays) for each parameter in the model.
"""
scores = []
for name, param in model.named_parameters():
if 'popup_scores' in name:
scores.append(param.data.cpu().numpy().flatten())
elif hasattr(param, 'scores'):
scores.append(param.scores.data.cpu().numpy().flatten())
return np.concatenate(scores) if scores else np.array([])
def return_weight(model):
"""
Returns the weights for the given model.
Args:
model (torch.nn.Module): The model from which to extract weights.
Returns:
list: A list of weights (numpy arrays) for each parameter in the model.
"""
weights = []
for name, param in model.named_parameters():
weights.append(param.data.cpu().numpy().flatten())
return np.concatenate(weights)