-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathunimodal_text_only_baseline.py
More file actions
166 lines (129 loc) · 6.84 KB
/
Copy pathunimodal_text_only_baseline.py
File metadata and controls
166 lines (129 loc) · 6.84 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
154
155
156
157
158
159
160
161
162
163
164
165
166
import os
import sys
import argparse
import json
from datetime import datetime
import tqdm
from configs.default_configs import UNIMODAL_TEXT_ONLY_BASELINE_CONFIG
from dataset.eeg_report_data_loader import get_harvard_data_loader
from transformers import pipeline
from accelerate import Accelerator
from utils.utils import extract_json, clean_generation_for_json_parsing, seed_everything
class Tee:
def __init__(self, *files):
self.files = files
def write(self, data):
for f in self.files:
f.write(data)
f.flush()
def flush(self):
for f in self.files:
f.flush()
if __name__ == '__main__':
parser = argparse.ArgumentParser()
# Dataset Parameters
parser.add_argument('--site', type=str, default='S0002')
parser.add_argument('--split', type=str, default='test')
parser.add_argument('--split_type', type=str, default='random_split_data_by_patient')
parser.add_argument('--normalize_eeg_method', type=str, default='div_by_100')
parser.add_argument('--task', type=str, default='unimodal_text_only_baseline')
parser.add_argument('--load_eeg', type=bool, default=True)
# parser.add_argument('--batch_size', type=int, default=2)
parser.add_argument('--num_workers', type=int, default=8)
# Model Parameters
parser.add_argument('--model_name', type=str, default='meta-llama/Llama-3.1-8B-Instruct')
parser.add_argument('--max_new_tokens', type=int, default=2048)
# experiment parameters
parser.add_argument('--experiment_name', type=str, default=None)
parser.add_argument('--device', type=int, default=0)
parser.add_argument('--use_accelerate', action='store_true')
args = parser.parse_args()
if args.use_accelerate:
args.device = Accelerator().device
# seed everything
seed_everything(5)
# Create Save Directory
model_save_name = args.model_name.split('/')[-1]
if args.experiment_name:
save_dir = os.path.join(UNIMODAL_TEXT_ONLY_BASELINE_CONFIG['save_dir'], args.experiment_name, model_save_name)
else:
save_dir = os.path.join(UNIMODAL_TEXT_ONLY_BASELINE_CONFIG['save_dir'], model_save_name)
os.makedirs(save_dir, exist_ok=True)
os.makedirs(os.path.join(save_dir, 'generated_reports_json'), exist_ok=True)
os.makedirs(os.path.join(save_dir, 'generated_reports_txt'), exist_ok=True)
# save args
with open(os.path.join(save_dir, 'args.txt'), 'w') as f:
f.write(str(args.__dict__))
# Create Log File
log_file = os.path.join(save_dir, f'log_{datetime.now().strftime("%Y%m%d_%H%M%S")}.txt')
log_f = open(log_file, 'w')
sys.stdout = Tee(sys.stdout, log_f)
sys.stderr = Tee(sys.stderr, log_f)
print(f"Log file: {log_file}")
error_generated_reports = {'error_generated_reports': []}
if os.path.exists(os.path.join(save_dir, 'error_generated_reports.json')):
with open(os.path.join(save_dir, 'error_generated_reports.json'), 'r') as f:
error_generated_reports = json.load(f)
# Load Data
test_data_loader = get_harvard_data_loader(site=args.site,
split=args.split,
split_type=args.split_type,
normalize_eeg_method=args.normalize_eeg_method,
task=args.task,
load_eeg=args.load_eeg,
batch_size=1,
num_workers=args.num_workers)
# test data loader
for batch_idx, batch_dict in enumerate(test_data_loader):
for key in batch_dict.keys():
print(f'{key}: {batch_dict[key][0]}')
break
# Load Model
model_name = args.model_name
if args.use_accelerate:
llm_pipeline = pipeline('text-generation',
model=model_name,
device_map="auto")
else:
llm_pipeline = pipeline('text-generation',
model=model_name,
device=args.device)
# Generate Report
for batch_idx, batch_dict in tqdm.tqdm(enumerate(test_data_loader)):
deidentified_file_name = batch_dict['meta_data'][0]['DeidentifiedName(Reports)'].replace('.txt', '')
# check if the report is already generated
if os.path.exists(os.path.join(save_dir, 'generated_reports_json', f'GENERATED_REPORT_{deidentified_file_name}.json')):
print(f'Report {deidentified_file_name} already generated')
continue
report_generation_task_prompt = batch_dict['generated_prompt'][0]
# Generate Report
generated_reports = llm_pipeline(report_generation_task_prompt,
max_new_tokens=args.max_new_tokens,
return_full_text=False)
generated_report_temp = generated_reports[0]['generated_text']
# Save Generated Report
with open(os.path.join(save_dir, 'generated_reports_txt', f'GENERATED_REPORT_{deidentified_file_name}.txt'), 'w') as f:
f.write(generated_report_temp)
generated_report_json = extract_json(generated_report_temp)
# second attempt to parse
if generated_report_json is None:
print(f'RETRY: Error extracting JSON from generated report {deidentified_file_name}')
clean_generated_report_text = clean_generation_for_json_parsing(generated_report_temp)
generated_report_json = extract_json(clean_generated_report_text)
if generated_report_json is None:
print(f'Error extracting JSON from generated report {deidentified_file_name}')
# print(generated_report_temp)
error_generated_reports['error_generated_reports'].append({'DeidentifiedName(Reports)': deidentified_file_name, 'generated_report': generated_report_temp})
else:
# Save Generated Report
with open(os.path.join(save_dir, 'generated_reports_json', f'GENERATED_REPORT_{deidentified_file_name}.json'), 'w') as f:
json.dump(generated_report_json, f)
if batch_idx%10 == 0:
# Save Error Generated Reports
with open(os.path.join(save_dir, 'error_generated_reports.json'), 'w') as f:
json.dump(error_generated_reports, f)
# break
print(f'Total number of error generated reports: {len(error_generated_reports["error_generated_reports"])}')
print(f'Total number of generated reports in txt: {len(os.listdir(os.path.join(save_dir, "generated_reports_txt")))}')
print(f'Total number of generated reports in json: {len(os.listdir(os.path.join(save_dir, "generated_reports_json")))}')
print('Completed!')