ADjayantan commited on
Commit
f822916
·
1 Parent(s): 3c3131b

Optimize CI benchmark test: Added automatic CI dry-run detection in train.py for fast execution

Browse files
Files changed (1) hide show
  1. train.py +13 -6
train.py CHANGED
@@ -161,7 +161,7 @@ def train_scaffold_fold(
161
  return model, fold_auroc, fold_auprc, y_tf, y_pf
162
 
163
 
164
- def run_5fold_ensemble_benchmark(use_tissue_conditioning: bool = True):
165
  # Instantiate single model to display exact parameter count
166
  temp_m = EpiADRNet(
167
  in_features=24, hidden_dim=1536, tissue_dim=1024, num_classes=10, num_gat_layers=12, num_heads=16,
@@ -176,7 +176,8 @@ def run_5fold_ensemble_benchmark(use_tissue_conditioning: bool = True):
176
  print(f" Parameters : {n_params:,} (~116.5M per fold) ", flush=True)
177
  print("=" * 70, flush=True)
178
 
179
- dataset = EpiADRDataset(repeat=4)
 
180
  smiles_ls = [s["smiles"] for s in dataset.samples]
181
 
182
  total_len = len(dataset)
@@ -192,16 +193,19 @@ def run_5fold_ensemble_benchmark(use_tissue_conditioning: bool = True):
192
  y_val_trues = []
193
  y_val_preds = []
194
 
195
- for fold in range(1, 6):
 
 
 
196
  val_start = (fold - 1) * fold_size
197
  val_end = fold * fold_size if fold < 5 else total_len
198
 
199
  val_idx = all_scaffold_idx[val_start:val_end]
200
  train_idx = all_scaffold_idx[:val_start] + all_scaffold_idx[val_end:]
201
 
202
- print(f" --> Running Fold {fold}/5 Scaffold Split (~116.5M Params | {mode_str})...", flush=True)
203
  model, f_auroc, f_auprc, y_t, y_p = train_scaffold_fold(
204
- fold, dataset, train_idx, val_idx, epochs=2, use_tissue_conditioning=use_tissue_conditioning
205
  )
206
  folds_models.append(model)
207
  fold_aurocs.append(f_auroc)
@@ -238,7 +242,10 @@ def run_5fold_ensemble_benchmark(use_tissue_conditioning: bool = True):
238
  if __name__ == "__main__":
239
  parser = argparse.ArgumentParser(description="EpiADR-Net v5 Training & Benchmark Pipeline")
240
  parser.add_argument("--baseline", action="store_true", help="Run molecule-only baseline without tissue conditioning")
 
241
  args, _ = parser.parse_known_args()
242
 
 
 
243
  use_tissue_conditioning = not args.baseline
244
- run_5fold_ensemble_benchmark(use_tissue_conditioning=use_tissue_conditioning)
 
161
  return model, fold_auroc, fold_auprc, y_tf, y_pf
162
 
163
 
164
+ def run_5fold_ensemble_benchmark(use_tissue_conditioning: bool = True, is_ci: bool = False):
165
  # Instantiate single model to display exact parameter count
166
  temp_m = EpiADRNet(
167
  in_features=24, hidden_dim=1536, tissue_dim=1024, num_classes=10, num_gat_layers=12, num_heads=16,
 
176
  print(f" Parameters : {n_params:,} (~116.5M per fold) ", flush=True)
177
  print("=" * 70, flush=True)
178
 
179
+ repeat_val = 1 if is_ci else 4
180
+ dataset = EpiADRDataset(repeat=repeat_val)
181
  smiles_ls = [s["smiles"] for s in dataset.samples]
182
 
183
  total_len = len(dataset)
 
193
  y_val_trues = []
194
  y_val_preds = []
195
 
196
+ max_folds = 1 if is_ci else 5
197
+ epochs = 1 if is_ci else 2
198
+
199
+ for fold in range(1, max_folds + 1):
200
  val_start = (fold - 1) * fold_size
201
  val_end = fold * fold_size if fold < 5 else total_len
202
 
203
  val_idx = all_scaffold_idx[val_start:val_end]
204
  train_idx = all_scaffold_idx[:val_start] + all_scaffold_idx[val_end:]
205
 
206
+ print(f" --> Running Fold {fold}/{max_folds} Scaffold Split (~116.5M Params | {mode_str})...", flush=True)
207
  model, f_auroc, f_auprc, y_t, y_p = train_scaffold_fold(
208
+ fold, dataset, train_idx, val_idx, epochs=epochs, use_tissue_conditioning=use_tissue_conditioning
209
  )
210
  folds_models.append(model)
211
  fold_aurocs.append(f_auroc)
 
242
  if __name__ == "__main__":
243
  parser = argparse.ArgumentParser(description="EpiADR-Net v5 Training & Benchmark Pipeline")
244
  parser.add_argument("--baseline", action="store_true", help="Run molecule-only baseline without tissue conditioning")
245
+ parser.add_argument("--quick", action="store_true", help="Run quick CI smoke test")
246
  args, _ = parser.parse_known_args()
247
 
248
+ import os
249
+ is_ci_env = os.getenv("CI") is not None or os.getenv("GITHUB_ACTIONS") is not None or args.quick
250
  use_tissue_conditioning = not args.baseline
251
+ run_5fold_ensemble_benchmark(use_tissue_conditioning=use_tissue_conditioning, is_ci=is_ci_env)