Views
No views yet
1from braindecode.models import EEGNetv4
2from huggingface_hub import hf_hub_download
3from skorch import NeuralNet
4import torch.nn as nn
5import torch as th
6
7path_params = hf_hub_download(
8 repo_id='guido151/EEGNetv4',
9 filename='EEGNetv4_Lee2019_ERP/params.pt',
10)
11path_optimizer = hf_hub_download(
12 repo_id='guido151/EEGNetv4',
13 filename='EEGNetv4_Lee2019_ERP/optimizer.pt',
14)
15path_history = hf_hub_download(
16 repo_id='guido151/EEGNetv4',
17 filename='EEGNetv4_Lee2019_ERP/history.json',
18)
19path_criterion = hf_hub_download(
20 repo_id='guido151/EEGNetv4',
21 filename='EEGNetv4_Lee2019_ERP/criterion.pt',
22)
23
24model = EEGNetv4(
25 n_chans=19,
26 n_outputs=2,
27 n_times=128,
28)
29
30net = NeuralNet(
31 model,
32 criterion=nn.CrossEntropyLoss(weight=th.tensor([1, 1])),
33)
34net.initialize()
35net.load_params(
36 path_params,
37 path_optimizer,
38 path_criterion,
39 path_history,
40)
411def get_fid_model(model: EEGNetv4) -> nn.Module:
2 fid_model = deepcopy(model)
3 for i in range(len(fid_model)):
4 if i >= 14:
5 fid_model[i] = Identity()
6 fid_model.eval()
7 for param in fid_model.parameters():
8 param.requires_grad = False
9 return fid_model1def get_is_model(model: EEGNetv4) -> nn.Module:
2 is_model = deepcopy(model)
3 is_model.eval()
4 for param in is_model.parameters():
5 param.requires_grad = False
6 return is_model