Repository navigation
Expand file tree
/
Copy pathcox_risk.py
More file actions
369 lines (277 loc) · 15 KB
/
Copy pathcox_risk.py
File metadata and controls
369 lines (277 loc) · 15 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
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
import pandas as pd
import numpy as np
from sklearn.decomposition import PCA
import os
from ukb_clomics.eval.cox_model import load_single_disease_data, CoxData, CoxPL
from ukb_clomics.eval.cox_model import Survivaldata
import pytorch_lightning as pl
from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint
from pytorch_lightning.loggers import CSVLogger
from torch.utils.data import DataLoader
import torch
import shutil
from ukb_clomics._cli._run_inference import infer_no_save
from pathlib import Path
import gc
import warnings
warnings.filterwarnings("ignore", category=pd.errors.DtypeWarning)
def load_indices(path_list):
"""
Load the indices from a list of paths
Args:
path_list (List[str]): list of paths to the index files
Returns:
set: set of indices
"""
indices = set()
for path in path_list:
df_indices = pd.read_csv(path, usecols=[0])
indices.update(df_indices.iloc[:,0].astype(int).tolist())
return indices
def load_features(cfg):
if isinstance(cfg.train_path, list):
X_train = pd.concat([pd.read_csv(path,index_col=0) for path in cfg.train_path], axis=0)
else:
raise ValueError('train_path must be a list of paths')
if isinstance(cfg.test_path, list):
X_test = pd.concat([pd.read_csv(path,index_col=0) for path in cfg.test_path], axis=0)
else:
raise ValueError('test_path must be a list of paths')
if cfg.use_indices is not None:
print('Loading the specified indices for evaluation ...')
indices_to_use = load_indices(cfg.use_indices)
# new_used_indices = used_indices.intersection(indices_to_use)
X_train_used_indices = X_train.index.intersection(indices_to_use)
X_test_used_indices = X_test.index.intersection(indices_to_use)
used_indices = X_train_used_indices.union(X_test_used_indices)
else:
if cfg.exclude_indices:
indices_to_exclude = load_indices(cfg.exclude_indices)
X_train_used_indices = X_train.index.difference(indices_to_exclude)
X_test_used_indices = X_test.index.difference(indices_to_exclude)
used_indices = X_train_used_indices.union(X_test_used_indices)
else:
used_indices = X_train.index.union(X_test.index)
X_train_used_indices = X_train.index
X_test_used_indices = X_test.index
X_train = X_train.loc[X_train_used_indices,:]
X_test = X_test.loc[X_test_used_indices,:]
return X_train, X_test, used_indices
def load_disease_related(cfg,used_indices):
# df_disease = pd.read_csv(cfg.disease_path,index_col=0,
# usecols=[ 'Unnamed: 0',
# 'participant.p131290',
# 'participant.p132082',
# 'participant.p130646',
# 'participant.p132034',
# 'participant.p131962',])
print('Loading disease data ...')
df_disease = pd.read_csv(cfg.disease_path,index_col=0,)
df_cov = pd.read_csv(cfg.cov_path,index_col=0)
# make it slightly faster by only keeping used indices
df_disease = df_disease.loc[used_indices,:]
df_cov = df_cov.loc[used_indices,:]
df_cov.columns = ['age','sex', 'bmi','assessment_date','death_date','center']
df_cov = df_cov[['age','sex', 'bmi', 'assessment_date', 'death_date']]
df_disease_code_map = pd.read_csv(cfg.disease_code_map_path)
return df_disease, df_cov, df_disease_code_map
def add_covariates(X_train, X_test, df_cov, used_indices):
df_cov_used_indices = df_cov.loc[used_indices,:][['age','sex', 'bmi',]]
# train intersection
X_train_used_indices = X_train.index.intersection(used_indices)
X_test_used_indices = X_test.index.intersection(used_indices)
X_train_cov = df_cov_used_indices.loc[X_train_used_indices,:]
X_test_cov = df_cov_used_indices.loc[X_test_used_indices,:]
X_train = pd.concat([X_train, X_train_cov], axis=1)
X_test = pd.concat([X_test, X_test_cov], axis=1)
return X_train, X_test
def cox_model_inference(cox_model,dataloader_test):
cox_model.eval()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
cox_model = cox_model.to(device)
all_log_hz = []
with torch.no_grad():
# a single batch
for (x,event,time) in dataloader_test:
x = x.to(device)
log_hz = cox_model(x) # log hazard of length n
all_log_hz.append(log_hz.cpu())
log_hz = torch.cat(all_log_hz).numpy()
return log_hz
def train_one_cox_model(
train_df_surv:Survivaldata,
test_df_surv:Survivaldata,
log_dir:str,
crt_disease_code:str,
batch_size:int=128,
max_epochs:int=20,
lr:float=1e-3,
save_cox_model=False,
):
os.makedirs(log_dir, exist_ok=True)
## train one cox model, can be used during development
cox_model = CoxPL(num_features=train_df_surv.features.shape[1],lr=lr)
callbacks = [
EarlyStopping(monitor='train_loss', mode="min", patience=5),
ModelCheckpoint(monitor='train_loss', mode="min", save_top_k=1,dirpath=log_dir,)]
trainer = pl.Trainer(max_epochs=max_epochs,
callbacks=callbacks,
logger=CSVLogger(save_dir=log_dir, name=f'cox_{crt_disease_code}'),
limit_val_batches=0,
num_sanity_val_steps=0,
enable_progress_bar=True,
enable_model_summary=False)
trainer.fit(cox_model, CoxData(DataLoader(train_df_surv, batch_size=batch_size, shuffle=True),
DataLoader(test_df_surv, batch_size=batch_size, shuffle=False)
))
# load the best model from the checkpoint
best_model = CoxPL.load_from_checkpoint(trainer.checkpoint_callback.best_model_path)
log_hz = cox_model_inference(best_model,DataLoader(test_df_surv, batch_size=batch_size, shuffle=False))
# remove the model log dir
if not save_cox_model:
shutil.rmtree(log_dir)
return log_hz
def run_single_disease(df_crt_disease, column, df_disease_code_map,
X_train, X_test,log_dir,save_cox_model=False,
max_epochs=20):
n_cases = df_crt_disease['binary'].sum()
if n_cases < 80:
return None, None
column = column.replace('participant.','')
crt_disease_code = df_disease_code_map.loc[df_disease_code_map['index']==column,'code'].values[0]
train_df = pd.concat((X_train,df_crt_disease.iloc[:,-2:]),axis=1).dropna()
test_df = pd.concat((X_test,df_crt_disease.iloc[:,-2:]),axis=1).dropna()
train_df_surv = Survivaldata(train_df.iloc[:,:-2].to_numpy().astype(np.float32),
train_df.iloc[:,-1].to_numpy().astype(np.float32),
train_df.iloc[:,-2].to_numpy().astype(int),
)
test_df_surv = Survivaldata(test_df.iloc[:,:-2].to_numpy().astype(np.float32),
test_df.iloc[:,-1].to_numpy().astype(np.float32),
test_df.iloc[:,-2].to_numpy().astype(int),
)
model_log_dir = f'{log_dir}/cox_model/{crt_disease_code}'
log_hz = train_one_cox_model(train_df_surv,test_df_surv,model_log_dir,crt_disease_code,
save_cox_model=save_cox_model,
max_epochs=max_epochs,)
return log_hz,test_df.index.values.astype(int)
def apply_pca(X_train, X_test,n_components):
train_indices = X_train.index
test_indices = X_test.index
pca = PCA(n_components=n_components)
pca.fit(X_train)
X_train = pca.transform(X_train)
X_test = pca.transform(X_test)
X_train = pd.DataFrame(X_train, index=train_indices)
X_test = pd.DataFrame(X_test, index=test_indices)
return X_train, X_test
def predict_risk(cfg):
os.makedirs(f'{cfg.log_dir}/risk', exist_ok=True)
X_train, X_test, used_indices = load_features(cfg)
if cfg.model_path is not None and cfg.model_name is not None:
# use the model to make inference
print('Making inference using the specified model ...')
all_embeddings_df = infer_no_save(cfg.model_path, cfg.model_name,
pd.concat([X_train,X_test],axis=0),
batch_size=cfg.batch_size)
if cfg.use_proteins is not None:
if cfg.all_proteins is None:
raise ValueError('Please provide the name of all proteins file when specifying use_proteins')
# use only the proteins in the list
use_proteins = open(cfg.use_proteins,'r').read().strip().split('\n')
all_proteins = open(cfg.all_proteins,'r').read().strip().split('\n')
all_embeddings_df.columns = all_proteins
all_embeddings_df = all_embeddings_df.loc[:,use_proteins]
if cfg.use_raw_data:
print('Adding raw features to the features ...')
train_embeddings = all_embeddings_df.loc[X_train.index,:]
test_embeddings = all_embeddings_df.loc[X_test.index,:]
X_train = pd.concat([X_train, train_embeddings], axis=1)
X_test = pd.concat([X_test, test_embeddings], axis=1)
else:
X_train = all_embeddings_df.loc[X_train.index,:]
X_test = all_embeddings_df.loc[X_test.index,:]
print(f"Training set size: {X_train.shape}, Test set size: {X_test.shape}")
if cfg.pca_dim is not None:
print('Applying PCA to the features ...')
X_train, X_test = apply_pca(X_train, X_test,n_components=cfg.pca_dim)
df_disease, df_cov, df_disease_code_map = load_disease_related(cfg,used_indices)
if cfg.use_covariates:
print('Adding covariates to the features ...')
X_train, X_test = add_covariates(X_train, X_test, df_cov, used_indices)
else:
print('Not using covariates ...')
if cfg.disease_code:
columns = df_disease_code_map[df_disease_code_map['code'].isin(cfg.disease_code)]['index'].values
columns = [f'participant.{column}' for column in columns]
else:
columns = df_disease.columns
if cfg.save_test_indices:
print('Indices of the test set will be saved for each disease ...')
else:
print('Indices of the test set will NOT be saved for each disease ...')
results = list()
for column in columns:
# for column in ['participant.p130792']:
df_crt_disease = load_single_disease_data(column,df_disease,df_disease_code_map,
df_cov,years=cfg.years,
)
disease_index = column.split('.')[-1]
code = df_disease_code_map[df_disease_code_map['index']==disease_index]['code'].values[0]
npy_path = Path(f'{cfg.log_dir}/{code}.npy')
# if f'{cfg.log_dir}/{code}.npy' in os.listdir(f'{cfg.log_dir}'):
if npy_path.exists():
print(f'Risk scores for disease code {code} already exist, skipping ...')
n_cases = df_crt_disease['binary'].sum()
n_samples = df_crt_disease.shape[0]
results.append((code,n_samples,n_cases))
continue
n_cases = df_crt_disease['binary'].sum()
n_samples = df_crt_disease.shape[0]
results.append((code,n_samples,n_cases))
print(f'Processing disease code {code} with {n_cases} cases ...')
log_hz,test_indices = run_single_disease(df_crt_disease, column, df_disease_code_map,
X_train, X_test,log_dir=cfg.log_dir,save_cox_model=cfg.save_cox_model,
max_epochs=cfg.max_epochs)
if log_hz is not None:
# save as npy
np.save(f'{cfg.log_dir}/{code}.npy',log_hz)
if cfg.save_test_indices:
np.save(f'{cfg.log_dir}/{code}.indices.npy',test_indices)
# --- FIX: CLEAR MEMORY ---
del df_crt_disease
del log_hz
if 'test_indices' in locals():
del test_indices
gc.collect() # Manually trigger garbage collection
results_df = pd.DataFrame(results, columns=['disease_code',f'{cfg.model_name}_n_samples',
f'{cfg.model_name}_n_cases'])
results_df.to_csv(f'{cfg.log_dir}/number_of_cases.csv', index=False)
return None
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--train_path", nargs='+', type=str, required=True)
parser.add_argument("--test_path", nargs='+', type=str, required=True)
parser.add_argument("--log_dir", type=str, required=True)
parser.add_argument("--batch_size", type=int, required=False, default=128)
## use model to make inference
parser.add_argument("--model_path", type=str, required=False, default=None)
parser.add_argument("--model_name", type=str, required=False, default=None)
parser.add_argument("--disease_path", type=str, required=False, default="data/raw_data/first_outcome.raw.csv")
parser.add_argument("--disease_code_map_path", type=str, required=False, default="data/ukb_disease_code_map.csv")
parser.add_argument("--use_indices", nargs="+", type=str, required=False, default=None)
parser.add_argument("--use_proteins", type=str, required=False, default=None) # a subset of predicted proteins to use
parser.add_argument("--all_proteins", type=str, required=False, default=None) # all predicted proteins
parser.add_argument('--exclude_indices', nargs="+", type=str, required=False, default=None)
parser.add_argument('--pca_dim', type=int, required=False, default=None)
parser.add_argument('--save_cox_model', type=bool, required=False, default=False)
parser.add_argument('--max_epochs', type=int, required=False, default=20)
parser.add_argument("--use_raw_data", nargs="+", type=str, required=False, default=None)
parser.add_argument("--cov_path", type=str, required=False, default='data/raw_data/covariants.raw.csv')
parser.add_argument("--disease_code", nargs='+',type=str, required=False, default=None,help="use certrain disease codes")
parser.add_argument("--use_covariates", type=bool, required=False, default=False)
parser.add_argument("--save_test_indices", type=bool, required=False, default=False)
parser.add_argument("--years", type=int, required=False, default=10)
cfg = parser.parse_args()
pl.seed_everything(42)
predict_risk(cfg)