forked from GitAhubI-Lover/classify_i_machine_learning
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsimple_snr.py
More file actions
74 lines (71 loc) · 3.68 KB
/
Copy pathsimple_snr.py
File metadata and controls
74 lines (71 loc) · 3.68 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
import copy
import torch.nn.functional as F
import numpy as np
import torch
import util.params as params
import random
from util.model_prepare import get_model_archi
from dataset.ads_b import get_dataloader
from util.traintest import train_model
import sys
import logging
import os
from util.traintest import validation_evaluation, validation_evaluation_all_snr
from dataset.ads_b import get_dataloader as get_adsb
from dataset.wifi_data import get_dataloader_snr as get_wifi
import torch.optim as optim
from scipy.io import savemat
log_format = '%(message)s'
logging.basicConfig(stream=sys.stdout, level=logging.INFO, format=log_format)
os.environ['NUMEXPR_MAX_THREADS'] = '16'
if __name__ == '__main__':
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f'Using device {device}')
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
torch.manual_seed(params.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(params.seed)
np.random.seed(params.seed)
random.seed(params.seed)
acc_snr_archi = []
for archi in params.model_archi:
loss_fn = torch.nn.CrossEntropyLoss(reduction="mean")
acc_snr = []
for i, snr in enumerate(params.snr_list):
print(f'Using net {archi}')
net = get_model_archi(device, archi)
optimizer = torch.optim.Adam(net.parameters(), lr=params.lr, weight_decay=1e-5)
# initialize ExponentialLR Scheduler, set the decay factor gamma as 0.95
scheduler = optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.95)
## add log
if not os.path.exists(
'../checkpoint/{}_{}_{}_{}'.format(params.signal_repre, archi, snr, params.metric)):
os.makedirs(
'../checkpoint/{}_{}_{}_{}/'.format(params.signal_repre, archi, snr, params.metric))
defense_enhanced_saver = f'../checkpoint/{params.signal_repre}_{archi}_{snr}_{params.metric}/'
fh = logging.FileHandler(os.path.join(defense_enhanced_saver, 'log.txt'))
fh.setFormatter(logging.Formatter(log_format))
logging.getLogger().addHandler(fh)
param_str = params.get_all_variables_as_string()
logging.info(param_str)
print(f'Starting evaluating snr {snr}dB')
if params.dataset == 'adsb':
data, val_x, val_y = get_adsb(params.signal_repre, np.arange(params.class_nums))
else:
data, val_x, val_y = get_wifi(params.signal_repre, np.arange(params.class_nums), snr=snr)
train_param = {'loss_fn': loss_fn, 'optimizer': optimizer, 'train_loader': data.train,
'validation_loader': data.test, 'device': device, 'num_epochs': params.nb_epochs,
'scheduler': scheduler, 'sparse_flag': params.sparse_flag,
'sparse_scale': params.sparse_scale,
'temperature': params.temperature, 'soft_target_loss_weight': params.soft_target_loss_weight,
'label_loss_weight': params.label_loss_weight}
if not params.initial_flag:
train_model(defense_enhanced_saver, net, train_param)
net.load_state_dict(torch.load(defense_enhanced_saver + 'model_best.pth.tar'))
val_acc, val_loss = validation_evaluation(net, 0, data.test, device, snr=snr, archi=archi,
return_features=True)
acc_snr.append(val_acc)
acc_snr_archi.append(acc_snr)
# savemat(f'./res_test/res_{params.snr_list[0]}_{params.snr_list[-1]}_{params.model_archi}.mat',
# {'data_list': acc_snr_archi})