-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
154 lines (139 loc) · 6.81 KB
/
Copy pathmain.py
File metadata and controls
154 lines (139 loc) · 6.81 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
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
import logging
import time
import datetime
import gc
import argparse
import torch
import torch.cuda
from src.server import *
from src.client import *
import src.datasets as my_datasets
from src.splitter import *
from src.utils import *
from src.dataset_bundle import *
from wilds.common.data_loaders import get_eval_loader
from wilds import get_dataset
from src.models import ResNet
"""
The main file function:
1. Load the hyperparameter dict.
2. Initialize logger
3. Initialize data (preprocess, data splits, etc.)
4. Initialize clients.
5. Initialize Server.
6. Register clients at the server.
7. Start the server.
"""
def main(args):
config_file = args.config_file
with open(config_file) as fh:
hparam = json.load(fh)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
server_config = hparam["server"]
client_config = hparam["client"]
data_config = hparam["dataset"]
global_config = hparam["global"]
seed = global_config['seed']
exp_id = global_config['id']
data_path = global_config['data_path']
set_seed(seed)
if not os.path.exists(data_path + "opt_dict/"): os.makedirs(data_path + "opt_dict/")
if not os.path.exists(data_path + "models/"): os.makedirs(data_path + "models/")
# 1. Preprocess some hyperparameters
data_config["num_shards"] = global_config["num_clients"]
server_config["batch_size"] = global_config["batch_size"]
client_config["batch_size"] = global_config["batch_size"]
data_config["seed"] = global_config["seed"]
# 2. Initialize logger.
# 2.1 modify log_path to contain current time
log_path = os.path.join(global_config["log_path"], "{}_{}_{}_{}".format(global_config['dataset_name'], client_config['algorithm'], global_config['id'], str(datetime.datetime.now().strftime("%Y-%m-%d_%H:%M:%S_%f")[:-3])))
os.makedirs(log_path)
# 2.2 set the configuration of global logger
logger = logging.getLogger(__name__)
logging.basicConfig(
filename=os.path.join(log_path, "FL.log"),
level=logging.INFO,
format="[%(levelname)s](%(asctime)s) %(message)s",
datefmt="%Y/%m/%d/ %I:%M:%S %p")
# 2.3 display and log experiment configuration
message = "\n[WELCOME] Unfolding configurations...!"
logging.info(message)
logging.info(hparam)
# initialize data
num_shards = data_config['num_shards']
iid = data_config['iid']
root_dir = data_config["data_path"]
if global_config['dataset_name'].lower() == 'pacs':
dataset = my_datasets.PACS(version='1.0', root_dir=root_dir, download=True)
elif global_config['dataset_name'].lower() == 'femnist':
dataset = my_datasets.FEMNIST(version='1.0', root_dir=root_dir, download=True)
elif global_config["dataset_name"].lower() == "officehome":
dataset = my_datasets.OfficeHome(root_dir=root_dir, download=True)
else:
dataset = get_dataset(dataset=global_config["dataset_name"].lower(), root_dir=root_dir, download=True)
print(dataset)
# if server_config['algorithm'] == "FedDG":
# # make it easier to hash fourier transformation
# indices = torch.arange(len(dataset)).reshape(-1,1)
# new_metadata_array = torch.cat((dataset.metadata_array, indices), dim=1)
# dataset._metadata_array = new_metadata_array
ds_bundle = eval(global_config["dataset_name"])(dataset, global_config['feature_dimension'])
total_subset = dataset.get_subset('train', transform=ds_bundle.train_transform)
try:
in_test_dataset = dataset.get_subset('id_test', transform=ds_bundle.test_transform)
except ValueError:
in_test_dataset, total_subset = RandomSplitter(ratio=0.2, seed=seed).split(total_subset)
in_test_dataset.transform = ds_bundle.test_transform
total_subset.transform = ds_bundle.train_transform
lodo_validation_dataset = dataset.get_subset('val', transform = ds_bundle.test_transform)
try:
in_validation_dataset = dataset.get_subset('id_val',transform = ds_bundle.test_transform)
except ValueError:
in_validation_dataset, in_test_dataset = RandomSplitter(ratio=0.5, seed=seed).split(total_subset)
in_validation_dataset.transform = ds_bundle.test_transform
in_test_dataset.transform = ds_bundle.test_transform
out_test_dataset = dataset.get_subset('test', transform=ds_bundle.test_transform)
in_train_dataloader = get_eval_loader(loader='standard', dataset=total_subset, batch_size=global_config["batch_size"])
out_test_dataloader = get_eval_loader(loader='standard', dataset=out_test_dataset, batch_size=global_config["batch_size"])
in_test_dataloader = get_eval_loader(loader='standard', dataset=in_test_dataset, batch_size=global_config["batch_size"])
lodo_validation_dataloader = get_eval_loader(loader='standard', dataset=lodo_validation_dataset, batch_size=global_config["batch_size"])
in_validation_dataloader = get_eval_loader(loader='standard', dataset=in_validation_dataset, batch_size=global_config["batch_size"])
sampler = RandomSampler(total_subset, replacement=True)
global_dataloader = DataLoader(total_subset, batch_size=global_config["batch_size"], sampler=sampler)
if num_shards == 1:
training_datasets = [total_subset]
elif num_shards > 1:
training_datasets = NonIIDSplitter(num_shards=num_shards, iid=iid, seed=seed).split(dataset.get_subset('train'), ds_bundle.groupby_fields, transform=ds_bundle.train_transform)
else:
raise ValueError("num_shards should be greater or equal to 1, we got {}".format(num_shards))
# initialize client
clients = []
for k in tqdm(range(global_config["num_clients"]), leave=False):
client = eval(client_config["algorithm"])(seed, exp_id, k, device, training_datasets[k], ds_bundle, client_config)
clients.append(client)
message = f"successfully initialize all clients!"
logging.info(message)
del message; gc.collect()
# initialize server (model should be initialized in the server. )
central_server = eval(server_config["algorithm"])(seed, exp_id, device, clients, ds_bundle, server_config)
central_server.init_model()
central_server.register_testloader({
"in_train": in_train_dataloader,
"in_val": in_validation_dataloader,
"lodo_val": lodo_validation_dataloader,
"in_test": in_test_dataloader,
"out_test": out_test_dataloader})
# do federated learning
central_server.fit()
# bye!
message = "...done all learning process!\n...exit program!"
logging.info(message)
time.sleep(3)
exit()
if __name__ == "__main__":
parser = argparse.ArgumentParser(description='PyTorch MAT Training')
parser.add_argument('--config_file', help='config file')
parser.add_argument('--resume_file', default=None)
parser.add_argument('--start_epoch', default=0)
args = parser.parse_args()
main(args)