Repository navigation
Expand file tree
/
Copy pathclient.py
More file actions
119 lines (98 loc) · 4.6 KB
/
Copy pathclient.py
File metadata and controls
119 lines (98 loc) · 4.6 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
117
118
119
from model import AutoFed
import torch
import numpy as np
from dataset import Traffic_Dataset
from torch.utils.data import DataLoader
# calculate MSE, RMSE, MAE, MAPE
def loss_func(y_hat, y, epsilon = 1e-5):
# mask = ~torch.isnan(y)
# mask = mask.float()
# mask = mask / torch.sum(mask)
#
# mse = torch.sum(mask * (y_hat-y)**2)
# rmse = torch.sqrt(mse)
# mae = torch.sum(mask * torch.abs(y_hat-y))
#
# mask = (y>1e-4).float()
# mask = mask / (epsilon+torch.sum(mask))
# mape = torch.sum(mask * torch.abs((y_hat-y)/(y+epsilon)))
mse = torch.mean((y_hat-y)**2)
rmse = torch.sqrt(mse)
mae = torch.mean(torch.abs(y_hat-y))
# mape = torch.mean(torch.abs((y_hat-y)/(y+epsilon)))
mask = (y>1e-4).float()
mask = mask / (epsilon+torch.sum(mask))
mape = torch.sum(mask * torch.abs((y_hat-y)/(y+epsilon)))
return mse, rmse, mae, mape
class Client():
def __init__(self, dataset, client_id, args):
super().__init__()
self.client_id = client_id
self.save_path = f"./{str(args.mode)}/{str(args.num_client)}/{client_id}.pth"
self.max_norm = args.max_grad_norm
self.num_nodes = dataset[0][0].shape[2]
device_name = "cuda:"+args.cuda
self.device = torch.device(device_name if torch.cuda.is_available() and not args.cpu else 'cpu')
model = AutoFed(node_num=self.num_nodes, history=args.history, horizon=args.horizon, dim_in=args.input_dim, dim_in_dec=args.input_dec_dim, dim_out=args.output_dim, dim_hidden=args.hidden_dim, cheb_k=args.cheb_k, embed_dim=args.hidden_dim, layer=args.layer)
self.model = model.to(self.device)
self.W = {key: value for key, value in self.model.named_parameters()}
train_data = Traffic_Dataset(dataset[0][0], dataset[0][1])
valid_data = Traffic_Dataset(dataset[1][0], dataset[1][1])
test_data = Traffic_Dataset(dataset[2][0], dataset[2][1])
self.train_loader = DataLoader(train_data, batch_size=args.batch_size, shuffle=True)
self.valid_loader = DataLoader(valid_data, batch_size=args.batch_size, shuffle=True)
self.test_loader = DataLoader(test_data, batch_size=args.batch_size, shuffle=True)
self.optim = torch.optim.Adam(self.model.parameters(), lr=args.lr, eps=args.epsilon)
self.scheduler = torch.optim.lr_scheduler.ExponentialLR(self.optim, gamma=args.gamma)
def iteration(self, mode="train", epochs=1):
mse_list, rmse_list, mae_list, mape_list, loss_list = [], [], [], [], []
if mode == "train":
self.model.train()
torch.set_grad_enabled(True)
data_iter = self.train_loader
elif mode == "val":
self.model.eval()
torch.set_grad_enabled(False)
data_iter = self.valid_loader
elif mode == "test":
self.model.eval()
torch.set_grad_enabled(False)
data_iter = self.test_loader
else:
print("Wrong Mode")
exit()
for _ in range(epochs):
for x, y in data_iter:
x, y = x.float().to(self.device), y.float().to(self.device)
x_dec = torch.zeros_like(x)
if mode == "train":
y_hat, loss_ae = self.model(x, x_dec, y)
mse, rmse, mae, mape = loss_func(y_hat, y)
alpha = loss_ae / mae
alpha = alpha.detach()
loss = mae + alpha * loss_ae
self.optim.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.max_norm)
self.optim.step()
else:
y_hat, loss_ae = self.model(x, x_dec)
mse, rmse, mae, mape = loss_func(y_hat, y)
alpha = loss_ae / mae
loss = mae + alpha * loss_ae
mse_list.append(mse.item())
rmse_list.append(rmse.item())
mae_list.append(mae.item())
mape_list.append(mape.item())
loss_list.append(loss.item())
if mode == "train":
self.scheduler.step()
return np.mean(mse_list), np.mean(rmse_list), np.mean(mae_list), np.mean(mape_list), np.mean(loss_list)
def save(self, model_path=None):
if model_path is None:
model_path = self.save_path
torch.save(self.model.state_dict(), model_path)
def update_weight(self, W):
for key in W:
if "pattern_mlp" in key and "batch_norm" not in key:
self.W[key].data = W[key].data.clone()