-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathfusion_main.py
More file actions
101 lines (83 loc) · 3.04 KB
/
Copy pathfusion_main.py
File metadata and controls
101 lines (83 loc) · 3.04 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
from __future__ import absolute_import
from __future__ import print_function
import numpy as np
import argparse
import os
import imp
import re
from trainers.fusion_trainer import FusionTrainer
from trainers.mmtm_trainer import MMTMTrainer
from trainers.daft_trainer import DAFTTrainer
from ehr_utils.preprocessing import Discretizer, Normalizer
from datasets.ehr_dataset import get_datasets
from datasets.cxr_dataset import get_cxr_datasets
from datasets.fusion import load_cxr_ehr
from pathlib import Path
import torch
from arguments import args_parser
parser = args_parser()
# add more arguments here ...
args = parser.parse_args()
print(args)
if args.missing_token is not None:
from trainers.fusion_tokens_trainer import FusionTokensTrainer as FusionTrainer
path = Path(args.save_dir)
path.mkdir(parents=True, exist_ok=True)
seed = 1002
torch.manual_seed(seed)
np.random.seed(seed)
def read_timeseries(args):
path = f'{args.ehr_data_dir}/{args.task}/train/14991576_episode3_timeseries.csv'
ret = []
with open(path, "r") as tsfile:
header = tsfile.readline().strip().split(',')
assert header[0] == "Hours"
for line in tsfile:
mas = line.strip().split(',')
ret.append(np.array(mas))
return np.stack(ret)
discretizer = Discretizer(timestep=float(args.timestep),
store_masks=True,
impute_strategy='previous',
start_time='zero')
discretizer_header = discretizer.transform(read_timeseries(args))[1].split(',')
cont_channels = [i for (i, x) in enumerate(discretizer_header) if x.find("->") == -1]
normalizer = Normalizer(fields=cont_channels) # choose here which columns to standardize
normalizer_state = args.normalizer_state
if normalizer_state is None:
normalizer_state = 'normalizers/ph_ts{}.input_str:previous.start_time:zero.normalizer'.format(args.timestep)
normalizer_state = os.path.join(os.path.dirname(__file__), normalizer_state)
normalizer.load_params(normalizer_state)
ehr_train_ds, ehr_val_ds, ehr_test_ds = get_datasets(discretizer, normalizer, args)
cxr_train_ds, cxr_val_ds, cxr_test_ds = get_cxr_datasets(args)
train_dl, val_dl, test_dl = load_cxr_ehr(args, ehr_train_ds, ehr_val_ds, cxr_train_ds, cxr_val_ds, ehr_test_ds, cxr_test_ds)
with open(f"{args.save_dir}/args.txt", 'w') as results_file:
for arg in vars(args):
print(f" {arg:<40}: {getattr(args, arg)}")
results_file.write(f" {arg:<40}: {getattr(args, arg)}\n")
if args.fusion_type == 'mmtm':
trainer = MMTMTrainer(
train_dl,
val_dl,
args,
test_dl=test_dl
)
elif args.fusion_type == 'daft':
trainer = DAFTTrainer(train_dl,
val_dl,
args,
test_dl=test_dl)
else:
trainer = FusionTrainer(
train_dl,
val_dl,
args,
test_dl=test_dl
)
if args.mode == 'train':
print("==> training")
trainer.train()
elif args.mode == 'eval':
trainer.eval()
else:
raise ValueError("not Implementation for args.mode")