pyinglie commited on
Commit
3bb0551
·
verified ·
1 Parent(s): c84cb76

Upload folder using huggingface_hub

Browse files
EEG-To-Text/datasets/.gitattributes ADDED
@@ -0,0 +1,114 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.lz4 filter=lfs diff=lfs merge=lfs -text
12
+ *.mds filter=lfs diff=lfs merge=lfs -text
13
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
14
+ *.model filter=lfs diff=lfs merge=lfs -text
15
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
16
+ *.npy filter=lfs diff=lfs merge=lfs -text
17
+ *.npz filter=lfs diff=lfs merge=lfs -text
18
+ *.onnx filter=lfs diff=lfs merge=lfs -text
19
+ *.ot filter=lfs diff=lfs merge=lfs -text
20
+ *.parquet filter=lfs diff=lfs merge=lfs -text
21
+ *.pb filter=lfs diff=lfs merge=lfs -text
22
+ *.pickle filter=lfs diff=lfs merge=lfs -text
23
+ *.pkl filter=lfs diff=lfs merge=lfs -text
24
+ *.pt filter=lfs diff=lfs merge=lfs -text
25
+ *.pth filter=lfs diff=lfs merge=lfs -text
26
+ *.rar filter=lfs diff=lfs merge=lfs -text
27
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
28
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
29
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
30
+ *.tar filter=lfs diff=lfs merge=lfs -text
31
+ *.tflite filter=lfs diff=lfs merge=lfs -text
32
+ *.tgz filter=lfs diff=lfs merge=lfs -text
33
+ *.wasm filter=lfs diff=lfs merge=lfs -text
34
+ *.xz filter=lfs diff=lfs merge=lfs -text
35
+ *.zip filter=lfs diff=lfs merge=lfs -text
36
+ *.zst filter=lfs diff=lfs merge=lfs -text
37
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
38
+ # Audio files - uncompressed
39
+ *.pcm filter=lfs diff=lfs merge=lfs -text
40
+ *.sam filter=lfs diff=lfs merge=lfs -text
41
+ *.raw filter=lfs diff=lfs merge=lfs -text
42
+ # Audio files - compressed
43
+ *.aac filter=lfs diff=lfs merge=lfs -text
44
+ *.flac filter=lfs diff=lfs merge=lfs -text
45
+ *.mp3 filter=lfs diff=lfs merge=lfs -text
46
+ *.ogg filter=lfs diff=lfs merge=lfs -text
47
+ *.wav filter=lfs diff=lfs merge=lfs -text
48
+ # Image files - uncompressed
49
+ *.bmp filter=lfs diff=lfs merge=lfs -text
50
+ *.gif filter=lfs diff=lfs merge=lfs -text
51
+ *.png filter=lfs diff=lfs merge=lfs -text
52
+ *.tiff filter=lfs diff=lfs merge=lfs -text
53
+ # Image files - compressed
54
+ *.jpg filter=lfs diff=lfs merge=lfs -text
55
+ *.jpeg filter=lfs diff=lfs merge=lfs -text
56
+ *.webp filter=lfs diff=lfs merge=lfs -text
57
+ # Video files - compressed
58
+ *.mp4 filter=lfs diff=lfs merge=lfs -text
59
+ *.webm filter=lfs diff=lfs merge=lfs -text
60
+ ZuCo/task1-SR/Matlab_files/resultsZAB_SR.mat filter=lfs diff=lfs merge=lfs -text
61
+ ZuCo/task1-SR/Matlab_files/resultsZDM_SR.mat filter=lfs diff=lfs merge=lfs -text
62
+ ZuCo/task1-SR/Matlab_files/resultsZDN_SR.mat filter=lfs diff=lfs merge=lfs -text
63
+ ZuCo/task1-SR/Matlab_files/resultsZGW_SR.mat filter=lfs diff=lfs merge=lfs -text
64
+ ZuCo/task1-SR/Matlab_files/resultsZJM_SR.mat filter=lfs diff=lfs merge=lfs -text
65
+ ZuCo/task1-SR/Matlab_files/resultsZJN_SR.mat filter=lfs diff=lfs merge=lfs -text
66
+ ZuCo/task1-SR/Matlab_files/resultsZJS_SR.mat filter=lfs diff=lfs merge=lfs -text
67
+ ZuCo/task1-SR/Matlab_files/resultsZKB_SR.mat filter=lfs diff=lfs merge=lfs -text
68
+ ZuCo/task1-SR/Matlab_files/resultsZKH_SR.mat filter=lfs diff=lfs merge=lfs -text
69
+ ZuCo/task1-SR/Matlab_files/resultsZKW_SR.mat filter=lfs diff=lfs merge=lfs -text
70
+ ZuCo/task1-SR/Matlab_files/resultsZMG_SR.mat filter=lfs diff=lfs merge=lfs -text
71
+ ZuCo/task1-SR/Matlab_files/resultsZPH_SR.mat filter=lfs diff=lfs merge=lfs -text
72
+ ZuCo/task2-NR/Matlab_files/resultsZAB_NR.mat filter=lfs diff=lfs merge=lfs -text
73
+ ZuCo/task2-NR/Matlab_files/resultsZDM_NR.mat filter=lfs diff=lfs merge=lfs -text
74
+ ZuCo/task2-NR/Matlab_files/resultsZDN_NR.mat filter=lfs diff=lfs merge=lfs -text
75
+ ZuCo/task2-NR/Matlab_files/resultsZGW_NR.mat filter=lfs diff=lfs merge=lfs -text
76
+ ZuCo/task2-NR/Matlab_files/resultsZJM_NR.mat filter=lfs diff=lfs merge=lfs -text
77
+ ZuCo/task2-NR/Matlab_files/resultsZJN_NR.mat filter=lfs diff=lfs merge=lfs -text
78
+ ZuCo/task2-NR/Matlab_files/resultsZJS_NR.mat filter=lfs diff=lfs merge=lfs -text
79
+ ZuCo/task2-NR/Matlab_files/resultsZKB_NR.mat filter=lfs diff=lfs merge=lfs -text
80
+ ZuCo/task2-NR/Matlab_files/resultsZKH_NR.mat filter=lfs diff=lfs merge=lfs -text
81
+ ZuCo/task2-NR/Matlab_files/resultsZKW_NR.mat filter=lfs diff=lfs merge=lfs -text
82
+ ZuCo/task2-NR/Matlab_files/resultsZMG_NR.mat filter=lfs diff=lfs merge=lfs -text
83
+ ZuCo/task2-NR/Matlab_files/resultsZPH_NR.mat filter=lfs diff=lfs merge=lfs -text
84
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYAC_NR.mat filter=lfs diff=lfs merge=lfs -text
85
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYAG_NR.mat filter=lfs diff=lfs merge=lfs -text
86
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYAK_NR.mat filter=lfs diff=lfs merge=lfs -text
87
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYDG_NR.mat filter=lfs diff=lfs merge=lfs -text
88
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYDR_NR.mat filter=lfs diff=lfs merge=lfs -text
89
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYFR_NR.mat filter=lfs diff=lfs merge=lfs -text
90
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYFS_NR.mat filter=lfs diff=lfs merge=lfs -text
91
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYHS_NR.mat filter=lfs diff=lfs merge=lfs -text
92
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYIS_NR.mat filter=lfs diff=lfs merge=lfs -text
93
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYLS_NR.mat filter=lfs diff=lfs merge=lfs -text
94
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYMD_NR.mat filter=lfs diff=lfs merge=lfs -text
95
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYMS_NR.mat filter=lfs diff=lfs merge=lfs -text
96
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYRH_NR.mat filter=lfs diff=lfs merge=lfs -text
97
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYRK_NR.mat filter=lfs diff=lfs merge=lfs -text
98
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYRP_NR.mat filter=lfs diff=lfs merge=lfs -text
99
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYSD_NR.mat filter=lfs diff=lfs merge=lfs -text
100
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYSL_NR.mat filter=lfs diff=lfs merge=lfs -text
101
+ ZuCo/task2-NR-2.0/Matlab_files/resultsYTL_NR.mat filter=lfs diff=lfs merge=lfs -text
102
+ ZuCo/task3-TSR/Matlab_files/resultsZAB_TSR.mat filter=lfs diff=lfs merge=lfs -text
103
+ ZuCo/task3-TSR/Matlab_files/resultsZDM_TSR.mat filter=lfs diff=lfs merge=lfs -text
104
+ ZuCo/task3-TSR/Matlab_files/resultsZDN_TSR.mat filter=lfs diff=lfs merge=lfs -text
105
+ ZuCo/task3-TSR/Matlab_files/resultsZGW_TSR.mat filter=lfs diff=lfs merge=lfs -text
106
+ ZuCo/task3-TSR/Matlab_files/resultsZJM_TSR.mat filter=lfs diff=lfs merge=lfs -text
107
+ ZuCo/task3-TSR/Matlab_files/resultsZJN_TSR.mat filter=lfs diff=lfs merge=lfs -text
108
+ ZuCo/task3-TSR/Matlab_files/resultsZJS_TSR.mat filter=lfs diff=lfs merge=lfs -text
109
+ ZuCo/task3-TSR/Matlab_files/resultsZKB_TSR.mat filter=lfs diff=lfs merge=lfs -text
110
+ ZuCo/task3-TSR/Matlab_files/resultsZKH_TSR.mat filter=lfs diff=lfs merge=lfs -text
111
+ ZuCo/task3-TSR/Matlab_files/resultsZKW_TSR.mat filter=lfs diff=lfs merge=lfs -text
112
+ ZuCo/task3-TSR/Matlab_files/resultsZMG_TSR.mat filter=lfs diff=lfs merge=lfs -text
113
+ ZuCo/task3-TSR/Matlab_files/resultsZPH_TSR.mat filter=lfs diff=lfs merge=lfs -text
114
+ stanfordsentiment/stanfordSentimentTreebank/dictionary.txt filter=lfs diff=lfs merge=lfs -text
EEG-To-Text/datasets/pickle/task1-SR-datasets.pickle ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f3bd2ea703a1da05b7314f17c64929a8c78307c4d2531768213bb48936702c2f
3
+ size 1208188754
EEG-To-Text/datasets/pickle/task2-NR-2.0-datasets.pickle ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b95de456907d7851cf1807fc7bc2b65ae5242a81988ba337362b069b0b668109
3
+ size 1758497764
EEG-To-Text/datasets/pickle/task2-NR-datasets.pickle ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fb37b0f3c71dcaab4e2d49f6a1f899d2b909ef1db0adfb382f6a248e5eb8a8a9
3
+ size 1093181866
EEG-To-Text/datasets/pickle/task3-TSR-datasets.pickle ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6dd61aa8c19bc9e98b9400edbf36cd5db4cf6434beeba0983af4bb85a632c06e
3
+ size 1051403647
EEG-To-Text/environment.yml ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ name: EEGToText
2
+ channels:
3
+ - pytorch
4
+ - anaconda
5
+ - conda-forge
6
+ - huggingface
7
+ dependencies:
8
+ - pytorch=1.9.0
9
+ - torchaudio=0.9.0
10
+ - cudatoolkit=11.1
11
+ - scipy=1.6.2
12
+ - h5py=2.10.0
13
+ - tqdm=4.62.0
14
+ - matplotlib=3.3.2
15
+ - transformers=4.6.1
16
+ - nltk=3.5
17
+ - pip=21.0.1
18
+ - pip:
19
+ - fuzzy-match==0.0.1
20
+ - rouge==1.0.0
EEG-To-Text/scripts/train_decoding.sh ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ python3 train_decoding.py --model_name BrainTranslator \
2
+ --task_name task1_task2_task3 \
3
+ --one_step \
4
+ --pretrained \
5
+ --not_load_step1_checkpoint \
6
+ --num_epoch_step1 20 \
7
+ --num_epoch_step2 30 \
8
+ --train_input EEG \
9
+ -lr1 0.00002 \
10
+ -lr2 0.00002 \
11
+ -b 1 \
12
+ -s ./checkpoints/decoding \
EEG-To-Text/train_decoding.py ADDED
@@ -0,0 +1,379 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import numpy as np
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.optim as optim
6
+ from torch.optim import lr_scheduler
7
+ from torch.utils.data import Dataset, DataLoader, RandomSampler, SequentialSampler
8
+ import pickle
9
+ import json
10
+ import matplotlib.pyplot as plt
11
+ from glob import glob
12
+ import time
13
+ import copy
14
+ from tqdm import tqdm
15
+ from transformers import BertLMHeadModel, BartTokenizer, BartForConditionalGeneration, BartConfig, BartForSequenceClassification, BertTokenizer, BertConfig, BertForSequenceClassification, RobertaTokenizer, RobertaForSequenceClassification, PegasusForConditionalGeneration, PegasusTokenizer, T5Tokenizer, T5ForConditionalGeneration, BertGenerationEncoder, BertGenerationDecoder, EncoderDecoderConfig, EncoderDecoderModel
16
+ from data import ZuCo_dataset
17
+ from model_decoding import BrainTranslator, BrainTranslatorNaive, T5Translator
18
+ from config import get_config
19
+
20
+ def train_model(dataloaders, device, model, criterion, optimizer, scheduler, num_epochs=25, checkpoint_path_best = './checkpoints/decoding/best/temp_decoding.pt', checkpoint_path_last = './checkpoints/decoding/last/temp_decoding.pt'):
21
+ # modified from: https://pytorch.org/tutorials/beginner/transfer_learning_tutorial.html
22
+ since = time.time()
23
+
24
+ best_model_wts = copy.deepcopy(model.state_dict())
25
+ best_loss = 100000000000
26
+
27
+ for epoch in range(num_epochs):
28
+ print('Epoch {}/{}'.format(epoch, num_epochs - 1))
29
+ print('-' * 10)
30
+
31
+ # Each epoch has a training and validation phase
32
+ for phase in ['train', 'dev']:
33
+ if phase == 'train':
34
+ model.train() # Set model to training mode
35
+ else:
36
+ model.eval() # Set model to evaluate mode
37
+
38
+ running_loss = 0.0
39
+
40
+ # Iterate over data.
41
+ for input_embeddings, seq_len, input_masks, input_mask_invert, target_ids, target_mask, sentiment_labels in tqdm(dataloaders[phase]):
42
+
43
+ # load in batch
44
+ input_embeddings_batch = input_embeddings.to(device).float()
45
+ input_masks_batch = input_masks.to(device)
46
+ input_mask_invert_batch = input_mask_invert.to(device)
47
+ target_ids_batch = target_ids.to(device)
48
+ """replace padding ids in target_ids with -100"""
49
+ target_ids_batch[target_ids_batch == tokenizer.pad_token_id] = -100
50
+
51
+ # zero the parameter gradients
52
+ optimizer.zero_grad()
53
+
54
+ # forward
55
+ # track history if only in train
56
+ with torch.set_grad_enabled(phase == 'train'):
57
+ seq2seqLMoutput = model(input_embeddings_batch, input_masks_batch, input_mask_invert_batch, target_ids_batch)
58
+
59
+ """calculate loss"""
60
+ # logits = seq2seqLMoutput.logits # 8*48*50265
61
+ # logits = logits.permute(0,2,1) # 8*50265*48
62
+
63
+ # loss = criterion(logits, target_ids_batch_label) # calculate cross entropy loss only on encoded target parts
64
+ # NOTE: my criterion not used
65
+ loss = seq2seqLMoutput.loss # use the BART language modeling loss
66
+
67
+ # """check prediction, instance 0 of each batch"""
68
+ # print('target size:', target_ids_batch.size(), ',original logits size:', logits.size(), ',target_mask size', target_mask_batch.size())
69
+ # logits = logits.permute(0,2,1)
70
+ # for idx in [0]:
71
+ # print(f'-- instance {idx} --')
72
+ # # print('permuted logits size:', logits.size())
73
+ # probs = logits[idx].softmax(dim = 1)
74
+ # # print('probs size:', probs.size())
75
+ # values, predictions = probs.topk(1)
76
+ # # print('predictions before squeeze:',predictions.size())
77
+ # predictions = torch.squeeze(predictions)
78
+ # # print('predictions:',predictions)
79
+ # # print('target mask:', target_mask_batch[idx])
80
+ # # print('[DEBUG]target tokens:',tokenizer.decode(target_ids_batch_copy[idx]))
81
+ # print('[DEBUG]predicted tokens:',tokenizer.decode(predictions))
82
+
83
+ # backward + optimize only if in training phase
84
+ if phase == 'train':
85
+ # with torch.autograd.detect_anomaly():
86
+ loss.sum().backward()
87
+ optimizer.step()
88
+
89
+ # statistics
90
+ running_loss += loss.sum().item() * input_embeddings_batch.size()[0] # batch loss
91
+ # print('[DEBUG]loss:',loss.item())
92
+ # print('#################################')
93
+
94
+
95
+ if phase == 'train':
96
+ scheduler.step()
97
+
98
+ epoch_loss = running_loss / dataset_sizes[phase]
99
+
100
+ print('{} Loss: {:.4f}'.format(phase, epoch_loss))
101
+
102
+ # deep copy the model
103
+ if phase == 'dev' and epoch_loss < best_loss:
104
+ best_loss = epoch_loss
105
+ best_model_wts = copy.deepcopy(model.state_dict())
106
+ '''save checkpoint'''
107
+ torch.save(model.state_dict(), checkpoint_path_best)
108
+ print(f'update best on dev checkpoint: {checkpoint_path_best}')
109
+ # with torch.set_grad_enabled(False):
110
+ # traced_model_1 = torch.jit.trace(model, (torch.rand(1, 56, 840).to(device), torch.randint(1, 56).to(device), torch.rand(1, 56).to(device), torch.rand(1, 56).to(device)))
111
+ # traced_model_32 = torch.jit.trace(model, (torch.rand(32, 56, 840).to(device), torch.randint(32, 56).to(device), torch.rand(32, 56).to(device), torch.rand(32, 56).to(device)))
112
+ # torch.jit.save(traced_model_1, checkpoint_path_best[:-3]+'_1_jit.pt')
113
+ # torch.jit.save(traced_model_32, checkpoint_path_best[:-3]+'_32_jit.pt')
114
+ print()
115
+
116
+ time_elapsed = time.time() - since
117
+ print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))
118
+ print('Best val loss: {:4f}'.format(best_loss))
119
+ torch.save(model.state_dict(), checkpoint_path_last)
120
+ print(f'update last checkpoint: {checkpoint_path_last}')
121
+
122
+ # load best model weights
123
+ model.load_state_dict(best_model_wts)
124
+ return model
125
+
126
+ def show_require_grad_layers(model):
127
+ print()
128
+ print(' require_grad layers:')
129
+ # sanity check
130
+ for name, param in model.named_parameters():
131
+ if param.requires_grad:
132
+ print(' ', name)
133
+
134
+ if __name__ == '__main__':
135
+ args = get_config('train_decoding')
136
+
137
+ ''' config param'''
138
+ dataset_setting = 'unique_sent'
139
+
140
+ num_epochs_step1 = args['num_epoch_step1']
141
+ num_epochs_step2 = args['num_epoch_step2']
142
+ step1_lr = args['learning_rate_step1']
143
+ step2_lr = args['learning_rate_step2']
144
+
145
+ batch_size = args['batch_size']
146
+
147
+ model_name = args['model_name']
148
+ # model_name = 'BrainTranslatorNaive' # with no additional transformers
149
+ # model_name = 'BrainTranslator'
150
+
151
+ # task_name = 'task1'
152
+ # task_name = 'task1_task2'
153
+ # task_name = 'task1_task2_task3'
154
+ # task_name = 'task1_task2_taskNRv2'
155
+ task_name = args['task_name']
156
+ train_input = args['train_input']
157
+ print("train_input is:", train_input)
158
+ save_path = args['save_path']
159
+ if not os.path.exists(save_path):
160
+ os.makedirs(save_path)
161
+
162
+ skip_step_one = args['skip_step_one']
163
+ load_step1_checkpoint = args['load_step1_checkpoint']
164
+ use_random_init = args['use_random_init']
165
+ device_ids = [0] # device setting
166
+
167
+ if use_random_init and skip_step_one:
168
+ step2_lr = 5*1e-4
169
+
170
+ print(f'[INFO]using model: {model_name}')
171
+
172
+ if skip_step_one:
173
+ save_name = f'{task_name}_finetune_{model_name}_skipstep1_b{batch_size}_{num_epochs_step1}_{num_epochs_step2}_{step1_lr}_{step2_lr}_{dataset_setting}_{train_input}'
174
+ else:
175
+ save_name = f'{task_name}_finetune_{model_name}_2steptraining_b{batch_size}_{num_epochs_step1}_{num_epochs_step2}_{step1_lr}_{step2_lr}_{dataset_setting}_{train_input}'
176
+
177
+ if use_random_init:
178
+ save_name = 'randinit_' + save_name
179
+
180
+ save_path_best = os.path.join(save_path, 'best')
181
+ if not os.path.exists(save_path_best):
182
+ os.makedirs(save_path_best)
183
+
184
+ output_checkpoint_name_best = os.path.join(save_path_best, f'{save_name}.pt')
185
+
186
+ save_path_last = os.path.join(save_path, 'last')
187
+ if not os.path.exists(save_path_last):
188
+ os.makedirs(save_path_last)
189
+
190
+ output_checkpoint_name_last = os.path.join(save_path_last, f'{save_name}.pt')
191
+
192
+ # subject_choice = 'ALL
193
+ subject_choice = args['subjects']
194
+ print(f'![Debug]using {subject_choice}')
195
+ # eeg_type_choice = 'GD
196
+ eeg_type_choice = args['eeg_type']
197
+ print(f'[INFO]eeg type {eeg_type_choice}')
198
+ # bands_choice = ['_t1']
199
+ # bands_choice = ['_t1','_t2','_a1','_a2','_b1','_b2','_g1','_g2']
200
+ bands_choice = args['eeg_bands']
201
+ print(f'[INFO]using bands {bands_choice}')
202
+
203
+
204
+
205
+ ''' set random seeds '''
206
+ seed_val = 312
207
+ np.random.seed(seed_val)
208
+ torch.manual_seed(seed_val)
209
+ torch.cuda.manual_seed_all(seed_val)
210
+
211
+
212
+ ''' set up device '''
213
+ # use cuda
214
+ if torch.cuda.is_available():
215
+ # dev = "cuda:3"
216
+ dev = args['cuda']
217
+ else:
218
+ dev = "cpu"
219
+ # CUDA_VISIBLE_DEVICES=0,1,2,3
220
+ device = torch.device(dev)
221
+ print(f'[INFO]using device {dev}')
222
+ print()
223
+
224
+ ''' set up dataloader '''
225
+ whole_dataset_dicts = []
226
+ if 'task1' in task_name:
227
+ dataset_path_task1 = '/datasets/pickle/task1-SR-datasets.pickle'
228
+ with open(dataset_path_task1, 'rb') as handle:
229
+ whole_dataset_dicts.append(pickle.load(handle))
230
+ if 'task2' in task_name:
231
+ dataset_path_task2 = '/datasets/pickle/task2-NR-datasets.pickle'
232
+ with open(dataset_path_task2, 'rb') as handle:
233
+ whole_dataset_dicts.append(pickle.load(handle))
234
+ if 'task3' in task_name:
235
+ dataset_path_task3 = '/datasets/pickle/task3-TSR-datasets.pickle'
236
+ with open(dataset_path_task3, 'rb') as handle:
237
+ whole_dataset_dicts.append(pickle.load(handle))
238
+ if 'taskNRv2' in task_name:
239
+ dataset_path_taskNRv2 = '/datasets/pickle/task2-NR-2.0-datasets.pickle'
240
+ with open(dataset_path_taskNRv2, 'rb') as handle:
241
+ whole_dataset_dicts.append(pickle.load(handle))
242
+
243
+ print()
244
+
245
+ """save config"""
246
+ cfg_dir = './config/decoding/'
247
+
248
+ if not os.path.exists(cfg_dir):
249
+ os.makedirs(cfg_dir)
250
+
251
+ with open(os.path.join(cfg_dir,f'{save_name}.json'), 'w') as out_config:
252
+ json.dump(args, out_config, indent = 4)
253
+
254
+ if model_name in ['BrainTranslator','BrainTranslatorNaive']:
255
+ tokenizer = BartTokenizer.from_pretrained('facebook/bart-large')
256
+
257
+ elif model_name == 'PegasusTranslator':
258
+ tokenizer = PegasusTokenizer.from_pretrained('google/pegasus-xsum')
259
+
260
+ elif model_name == 'T5Translator':
261
+ tokenizer = T5Tokenizer.from_pretrained("t5-large")
262
+ #tokenizer.set_prefix_tokens(language='english')
263
+
264
+ # train dataset
265
+ train_set = ZuCo_dataset(whole_dataset_dicts, 'train', tokenizer, subject = subject_choice, eeg_type = eeg_type_choice, bands = bands_choice, setting = dataset_setting, test_input=train_input)
266
+ # dev dataset
267
+ dev_set = ZuCo_dataset(whole_dataset_dicts, 'dev', tokenizer, subject = subject_choice, eeg_type = eeg_type_choice, bands = bands_choice, setting = dataset_setting, test_input=train_input)
268
+ # test dataset
269
+ # test_set = ZuCo_dataset(whole_dataset_dicts, 'test', tokenizer, subject = subject_choice, eeg_type = eeg_type_choice, bands = bands_choice, setting = dataset_setting)
270
+
271
+
272
+ dataset_sizes = {'train': len(train_set), 'dev': len(dev_set)}
273
+ print('[INFO]train_set size: ', len(train_set))
274
+ print('[INFO]dev_set size: ', len(dev_set))
275
+ # print('[INFO]test_set size: ', len(test_set))
276
+
277
+ # train dataloader
278
+ train_dataloader = DataLoader(train_set, batch_size = batch_size, shuffle=True, num_workers=4)
279
+ # dev dataloader
280
+ val_dataloader = DataLoader(dev_set, batch_size = 1, shuffle=False, num_workers=4)
281
+ # dataloaders
282
+ dataloaders = {'train':train_dataloader, 'dev':val_dataloader}
283
+
284
+ ''' set up model '''
285
+ if model_name == 'BrainTranslator':
286
+ if use_random_init:
287
+ config = BartConfig.from_pretrained('facebook/bart-large')
288
+ pretrained = BartForConditionalGeneration(config)
289
+ else:
290
+ pretrained = BartForConditionalGeneration.from_pretrained('facebook/bart-large')
291
+
292
+ model = BrainTranslator(pretrained, in_feature = 105*len(bands_choice), decoder_embedding_size = 1024, additional_encoder_nhead=8, additional_encoder_dim_feedforward = 2048)
293
+
294
+ elif model_name == 'BrainTranslatorNaive':
295
+ pretrained = BartForConditionalGeneration.from_pretrained('facebook/bart-large')
296
+ model = BrainTranslatorNaive(pretrained, in_feature = 105*len(bands_choice), decoder_embedding_size = 1024, additional_encoder_nhead=8, additional_encoder_dim_feedforward = 2048)
297
+
298
+ elif model_name == 'PegasusTranslator':
299
+ pretrained = PegasusForConditionalGeneration.from_pretrained('google/pegasus-xsum')
300
+ model = BrainTranslator(pretrained, in_feature = 105*len(bands_choice), decoder_embedding_size = 1024, additional_encoder_nhead=8, additional_encoder_dim_feedforward = 2048)
301
+
302
+ elif model_name == 'T5Translator':
303
+ pretrained = T5ForConditionalGeneration.from_pretrained("t5-large")
304
+ model = T5Translator(pretrained, in_feature = 105*len(bands_choice), decoder_embedding_size = 1024, additional_encoder_nhead=8, additional_encoder_dim_feedforward = 2048)
305
+
306
+ model.to(device)
307
+ model = torch.nn.DataParallel(model, device_ids=device_ids)
308
+
309
+ ''' training loop '''
310
+
311
+ ######################################################
312
+ '''step one trainig: freeze most of BART params'''
313
+ ######################################################
314
+
315
+ # closely follow BART paper
316
+ if model_name in ['BrainTranslator','BrainTranslatorNaive', 'PegasusTranslator', 'T5Translator']:
317
+ for name, param in model.named_parameters():
318
+ if param.requires_grad and 'pretrained' in name:
319
+ if ('shared' in name) or ('embed_positions' in name) or ('encoder.layers.0' in name):
320
+ continue
321
+ else:
322
+ param.requires_grad = False
323
+
324
+ elif model_name == 'BertGeneration':
325
+ for name, param in model.named_parameters():
326
+ if param.requires_grad and 'pretrained' in name:
327
+ if ('embeddings' in name) or ('encoder.layer.0' in name):
328
+ continue
329
+ else:
330
+ param.requires_grad = False
331
+
332
+
333
+ if skip_step_one:
334
+ if load_step1_checkpoint:
335
+ stepone_checkpoint = 'path_to_step_1_checkpoint.pt'
336
+ print(f'skip step one, load checkpoint: {stepone_checkpoint}')
337
+ model.load_state_dict(torch.load(stepone_checkpoint))
338
+ else:
339
+ print('skip step one, start from scratch at step two')
340
+ else:
341
+
342
+ ''' set up optimizer and scheduler'''
343
+ optimizer_step1 = optim.SGD(filter(lambda p: p.requires_grad, model.parameters()), lr=step1_lr, momentum=0.9)
344
+
345
+ exp_lr_scheduler_step1 = lr_scheduler.StepLR(optimizer_step1, step_size=20, gamma=0.1)
346
+
347
+ ''' set up loss function '''
348
+ criterion = nn.CrossEntropyLoss()
349
+
350
+ print('=== start Step1 training ... ===')
351
+ # print training layers
352
+ show_require_grad_layers(model)
353
+ # return best loss model from step1 training
354
+ model = train_model(dataloaders, device, model, criterion, optimizer_step1, exp_lr_scheduler_step1, num_epochs=num_epochs_step1, checkpoint_path_best = output_checkpoint_name_best, checkpoint_path_last = output_checkpoint_name_last)
355
+
356
+ ######################################################
357
+ '''step two trainig: update whole model for a few iterations'''
358
+ ######################################################
359
+ for name, param in model.named_parameters():
360
+ param.requires_grad = True
361
+
362
+ ''' set up optimizer and scheduler'''
363
+ optimizer_step2 = optim.SGD(model.parameters(), lr=step2_lr, momentum=0.9)
364
+
365
+ exp_lr_scheduler_step2 = lr_scheduler.StepLR(optimizer_step2, step_size=30, gamma=0.1)
366
+
367
+ ''' set up loss function '''
368
+ criterion = nn.CrossEntropyLoss()
369
+
370
+ print()
371
+ print('=== start Step2 training ... ===')
372
+ # print training layers
373
+ show_require_grad_layers(model)
374
+
375
+ '''main loop'''
376
+ trained_model = train_model(dataloaders, device, model, criterion, optimizer_step2, exp_lr_scheduler_step2, num_epochs=num_epochs_step2, checkpoint_path_best = output_checkpoint_name_best, checkpoint_path_last = output_checkpoint_name_last)
377
+
378
+ # '''save checkpoint'''
379
+ # torch.save(trained_model.state_dict(), os.path.join(save_path,output_checkpoint_name))