-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
302 lines (257 loc) · 13.4 KB
/
Copy pathtrain.py
File metadata and controls
302 lines (257 loc) · 13.4 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
import logging
import os
import sys
import argparse
import numpy as np
import torch
from torch.utils.tensorboard import SummaryWriter
import monai
from monai.transforms import EnsureChannelFirst, Compose
from monai.data import DataLoader
import pandas as pd
from sklearn.metrics import roc_auc_score, f1_score, accuracy_score
import open_clip
from datetime import datetime
import wandb
import random
from dataset import Dataset
from model import Model, load_clip_to_cpu
from utils import set_seed, worker_init_fn, drop_id, session, rad_feat_ind
def main():
# Parse command line arguments
parser = argparse.ArgumentParser()
parser.add_argument('--text_prompt_len', type=int, default=6)
parser.add_argument('--vision_prompt_dep', type=int, default=10)
parser.add_argument('--vision_prompt_len', type=int, default=10)
parser.add_argument('--use_wandb', action='store_true', default=False)
parser.add_argument('--contrastive_loss_weight', type=float, default=0)
parser.add_argument('--orthogonal_loss_weight', type=float, default=0.1)
parser.add_argument('--batch_size', type=int, default=32)
parser.add_argument('--epochs', type=int, default=300)
parser.add_argument('--learning_rate', type=float, default=1e-3)
parser.add_argument('--seed', type=int, default=42)
args = parser.parse_args()
# Set random seed for reproducibility
SEED = args.seed
set_seed(SEED)
# Create run name based on parameters
run_name = f'conattn_sepnorm_vdep{args.vision_prompt_dep}_vlen{args.vision_prompt_len}_tlen{args.text_prompt_len}_coloss{args.contrastive_loss_weight}_orloss{args.orthogonal_loss_weight}_dualattn_labelsmooth04_lr{args.learning_rate}_weight12'
# Set up tensorboard writer
writer = SummaryWriter(f'runs_3d/{run_name}')
# Initialize wandb if specified
if args.use_wandb:
wandb.login(key='945c7addff384a579a9bafb12828e3bf1b040d64')
wandb.init(
project='GZ-Liver',
name=run_name,
config={
"seed": SEED,
"epochs": args.epochs,
"batch_size": args.batch_size,
"learning_rate": args.learning_rate,
"contrastive_loss_weight": args.contrastive_loss_weight,
"orthogonal_loss_weight": args.orthogonal_loss_weight,
"text_prompt_len": args.text_prompt_len,
"vision_prompt_dep": args.vision_prompt_dep,
"vision_prompt_len": args.vision_prompt_len,
}
)
wandb.save('train.py')
wandb.save('model.py')
wandb.save('dataset.py')
wandb.save('learnable.py')
wandb.save('utils.py')
# Configure logging
monai.config.print_config()
logging.basicConfig(stream=sys.stdout, level=logging.INFO)
# Define data paths
nii_path = './Dataset/Internal_nii2'
roi_path = './Dataset/Internal_roi'
label_df = pd.read_csv('./Dataset/label_v2.csv')
rad_feat_df = pd.read_csv('./Dataset/rad_feat.csv')
label_df = label_df.merge(rad_feat_df, on=['pathological number', 'label'], how='left')
# Prepare training data
train_images_id = label_df[label_df['train'] == 1]['hospital number'].values
train_images_id = [f for f in train_images_id if f not in drop_id]
train_images = [os.path.join(nii_path, str(f), f'{session}.nii.gz') for f in train_images_id]
train_segs = [os.path.join(roi_path, str(f), f'{session}.nrrd') for f in train_images_id]
train_labels = []
train_rad_feat = []
for f in train_images_id:
train_labels.append(label_df[label_df['hospital number'] == int(f)]['label'].values[0])
train_rad_feat.append(label_df[label_df['hospital number'] == int(f)][rad_feat_ind].values[0])
train_labels = np.array(train_labels, dtype=np.int64)
train_rad_feat = np.array(train_rad_feat, dtype=np.float32)
# Prepare validation data
valid_images_id = label_df[label_df['train'] == 0]['hospital number'].values
valid_images_id = [f for f in valid_images_id if f not in drop_id]
valid_images = [os.path.join(nii_path, str(f), f'{session}.nii.gz') for f in valid_images_id]
valid_segs = [os.path.join(roi_path, str(f), f'{session}.nrrd') for f in valid_images_id]
valid_labels = []
valid_rad_feat = []
for f in valid_images_id:
valid_labels.append(label_df[label_df['hospital number'] == int(f)]['label'].values[0])
valid_rad_feat.append(label_df[label_df['hospital number'] == int(f)][rad_feat_ind].values[0])
valid_labels = np.array(valid_labels, dtype=np.int64)
valid_rad_feat = np.array(valid_rad_feat, dtype=np.float32)
# Define transforms
train_transforms = Compose([EnsureChannelFirst()])
val_transforms = Compose([EnsureChannelFirst()])
# Create datasets
train_ds = Dataset(train_images, train_segs, train_labels, train_rad_feat, transform=train_transforms, seg_transform=train_transforms, train=True)
val_ds = Dataset(valid_images, valid_segs, valid_labels, valid_rad_feat, transform=val_transforms, seg_transform=val_transforms, train=False)
# Create data loaders
train_loader = DataLoader(train_ds, batch_size=args.batch_size, shuffle=True, num_workers=4, pin_memory=torch.cuda.is_available(), worker_init_fn=worker_init_fn)
val_loader = DataLoader(val_ds, batch_size=args.batch_size, num_workers=4, pin_memory=torch.cuda.is_available(), worker_init_fn=worker_init_fn)
# Setup device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Initialize models
ori_clip, _, _ = open_clip.create_model_and_transforms('hf-hub:wisdomik/GenMedClip')
clip = load_clip_to_cpu('./weights/open_clip_pytorch_model.bin', args.vision_prompt_dep, args.vision_prompt_len)
model = Model(clip.to(device), ori_clip.to(device), args.text_prompt_len).to(device)
# Define loss function and optimizer
loss_function = torch.nn.CrossEntropyLoss(weight=torch.tensor([1, 2]).float().to(device), label_smoothing=0.4)
optimizer = torch.optim.AdamW(model.parameters(), args.learning_rate)
# Training loop
val_interval = 1
best_metric = -1
best_metric_epoch = -1
epoch_loss_values = []
for epoch in range(args.epochs):
print("-" * 10)
print(f"epoch {epoch + 1}/{args.epochs}")
model.train()
epoch_loss = 0
epoch_contrastive_loss = 0
epoch_orthogonal_loss = 0
step = 0
train_prob_all, train_label_all = [], []
for batch_data in train_loader:
step += 1
inputs, segs, labels, rad_feat, valid_mask = batch_data[0].to(device), batch_data[1].to(device), batch_data[2].to(device), batch_data[3].to(device), batch_data[4].to(device)
inputs = inputs * segs
optimizer.zero_grad()
outputs = model(inputs, rad_feat, valid_mask)
classification_loss = loss_function(outputs, labels)
loss = classification_loss + args.contrastive_loss_weight*model.contrastive_loss + args.orthogonal_loss_weight*model.orthogonal_loss
loss.backward()
optimizer.step()
train_prob = torch.nn.functional.softmax(outputs, dim=1)
train_prob_all.append(train_prob.detach().to("cpu").numpy())
train_label_all.append(labels.to("cpu").numpy())
epoch_loss += classification_loss.item()
epoch_contrastive_loss += model.contrastive_loss.item()
epoch_orthogonal_loss += model.orthogonal_loss.item()
# Calculate training metrics
epoch_loss /= step
epoch_contrastive_loss /= step
epoch_orthogonal_loss /= step
epoch_loss_values.append(epoch_loss)
train_prob_all = np.concatenate(train_prob_all)
train_label_all = np.concatenate(train_label_all)
train_auc = roc_auc_score(train_label_all, train_prob_all[:, 1])
train_acc = accuracy_score(train_label_all, train_prob_all[:, 1].round())
train_f1 = f1_score(train_label_all, train_prob_all[:, 1].round())
print(f"epoch {epoch + 1} average loss: {epoch_loss:.4f}, contrastive loss: {epoch_contrastive_loss:.4f}, orthogonal loss: {epoch_orthogonal_loss:.4f}, train_auc: {train_auc:.4f}, train_acc: {train_acc:.4f}, train_f1: {train_f1:.4f}")
# Log training metrics
writer.add_scalar("Loss/train", epoch_loss, epoch)
writer.add_scalar("Loss/train_contrastive", epoch_contrastive_loss, epoch)
writer.add_scalar("Metrics/train_auc", train_auc, epoch)
writer.add_scalar("Metrics/train_acc", train_acc, epoch)
writer.add_scalar("Metrics/train_f1", train_f1, epoch)
# Validation
if (epoch + 1) % val_interval == 0:
model.eval()
val_prob_all_list, val_label_list = [], []
with torch.no_grad():
val_epoch_loss = 0
step = 0
metric_count_all = 0
num_correct_all = 0
image_id_list = []
for val_data in val_loader:
step += 1
val_images, segs, val_labels, rad_feat, valid_mask, image_id = val_data[0].to(device), val_data[1].to(device), val_data[2].to(device), val_data[3].to(device), val_data[4].to(device), val_data[5]
image_id_list.extend(image_id)
val_images = val_images * segs
val_outputs = model(val_images, rad_feat, valid_mask)
pred_all = val_outputs
val_loss = loss_function(pred_all, val_labels)
val_epoch_loss += val_loss.item()
val_prob_all = torch.nn.functional.softmax(pred_all, dim=1)
val_prob_all_list.append(val_prob_all.to("cpu").numpy())
val_label_list.append(val_labels.to("cpu").numpy())
value_all = torch.eq(pred_all.argmax(dim=1), val_labels)
metric_count_all += len(value_all)
num_correct_all += value_all.sum().item()
# Calculate validation metrics
val_epoch_loss /= step
print(f"epoch {epoch + 1} validation loss: {val_epoch_loss:.4f}")
val_prob_all = np.concatenate(val_prob_all_list)
val_label = np.concatenate(val_label_list)
auc_all = roc_auc_score(val_label, val_prob_all[:, 1])
acc_all = num_correct_all / metric_count_all
f1_all = f1_score(val_label, val_prob_all[:, 1].round())
# Save best model
if auc_all > best_metric:
best_metric = auc_all
best_metric_epoch = epoch + 1
torch.save(model.state_dict(), "best_metric_model_classification3d_array.pth")
print("saved new best metric model")
# Create results directory if it doesn't exist
os.makedirs('./results', exist_ok=True)
# Create DataFrame with results
results_df = pd.DataFrame({
'image_id': image_id_list,
'prob': val_prob_all[:, 1], # Probability of positive class
'label': val_label
})
# Save to CSV using run name
results_df.to_csv(f'./results/{run_name}.csv', index=False)
print(
"current epoch: {} current auc: {:.4f} current acc: {:.4f} current f1: {:.4f}. best AUC: {:.4f} at epoch {}".format(
epoch + 1, auc_all, acc_all, f1_all, best_metric, best_metric_epoch
)
)
# Log validation metrics
writer.add_scalar("Loss/val", val_epoch_loss, epoch)
writer.add_scalar("Metrics/val_auc", auc_all, epoch)
writer.add_scalar("Metrics/val_acc", acc_all, epoch)
writer.add_scalar("Metrics/val_f1", f1_all, epoch)
writer.add_scalar("Metrics/best_val_auc", best_metric, epoch)
# Log metrics to wandb
if args.use_wandb:
wandb.log({
"epoch": epoch + 1,
"train_loss": epoch_loss,
"train_contrastive_loss": epoch_contrastive_loss,
"train_orthogonal_loss": epoch_orthogonal_loss,
"train_auc": train_auc,
"train_acc": train_acc,
"train_f1": train_f1,
"val_loss": val_epoch_loss,
"val_auc": auc_all,
"val_acc": acc_all,
"val_f1": f1_all,
"best_val_auc": best_metric
})
if auc_all > best_metric:
wandb.run.summary["best_val_auc_all"] = best_metric
wandb.run.summary["best_epoch_all"] = best_metric_epoch
# Final summary
print(f"Training completed, best_metric: {best_metric:.4f} at epoch: {best_metric_epoch}")
# Record final results
writer.add_hparams(
{"lr": args.learning_rate, "batch_size": args.batch_size},
{
"best_val_auc": best_metric,
"best_epoch": best_metric_epoch,
}
)
if args.use_wandb:
wandb.run.summary["final_best_auc_all"] = best_metric
wandb.run.summary["final_best_epoch_all"] = best_metric_epoch
wandb.finish()
writer.close()
if __name__ == "__main__":
main()