robana commited on
Commit
c0d66fc
·
1 Parent(s): 074176d

adding efficientnet weights + config

Browse files
Files changed (3) hide show
  1. README.md +15 -0
  2. base_config.yaml +80 -0
  3. deepecg_tokenizer_efficientnet.pt +3 -0
README.md ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ## DeepECG-Tok_EfficientNetV2_77_Classes
2
+
3
+ - DeepECG-Tok_EfficientNetV2_77_Classes is a repository for the EfficientNetV2 weights that classify embeddings from DeepECG-Tok into 77 classes.
4
+ - base_config.yaml file contains the parameters for the training and inference.
5
+
6
+ - Paths to be updated:
7
+ - base_checkpoint_path: path to the checkpoints directory
8
+ - pretrained_tokenizer_path: path to the pretrained tokenizer checkpoint
9
+ - inference_dataset_path: path to the inference dataset
10
+ - inference_checkpoint_path: path to the inference checkpoint
11
+
12
+ - To run inference, run the following command:
13
+ ```bash
14
+ bash scripts/runner.sh --base_config config/linear_probing/base_config.yaml --selected_gpus 0,1 --use_wandb true --run_mode inference
15
+ ```
base_config.yaml ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # pipeline project
2
+ pipeline_project: !!str "ECG_Tokenizer_Linear_Probing"
3
+ model_name: !!str "ECG_Tokenizer_Wrapper"
4
+ runner_name: !!str "ECG_Tokenizer_Linear_Probing_Runner"
5
+ run_mode: !!str "inference"
6
+ num_epochs: !!int 10
7
+ seed: !!int 42
8
+ base_checkpoint_path: !!str "/volume/ECG_tokenizer/outputs/"
9
+ pretrained_tokenizer_path: !!str "/volume/ECG_tokenizer/checkpoints/ECG_tokenizer_latest/best_model_epoch_10.pt"
10
+ inference_dataset_path: !!str "/volume/ECG_tokenizer/outputs/preprocessed_bert_output.parquet"
11
+ inference_checkpoint_path: !!str "/volume/ECG_tokenizer/checkpoints/ECG_Tokenizer_Linear_Probing/ECG_Tokenizer_Linear_Probing/5zg01bx6_20250824-041452/checkpoint_epoch_10.pt"
12
+ # wandb parameters
13
+ wandb_project: !!str "ECG_Tokenizer_Linear_Probing"
14
+ wandb_entity: !!str "mhi_ai"
15
+ use_wandb: !!bool false
16
+ # training parameters
17
+ lr: !!float 3e-4
18
+ scheduler_type: !!str "cosine_warm_restart"
19
+ lr_step_period: !!int 1
20
+ factor: !!float 0.3
21
+ optimizer: !!str "AdamW"
22
+ weight_decay: !!float 1e-4
23
+ step_size: !!int 1
24
+ gamma: !!float 0.1
25
+ num_warmup_percent: !!float 0.1
26
+ num_hard_restarts_cycles: !!float 0.5
27
+ warm_restart_tmult: !!int 2
28
+ # VQVAE parameters
29
+ decoder_mode: !!str "classification"
30
+ decoder_name: !!str "EfficientNetV2_Classifier_Decoder"
31
+ # Classification parameters
32
+ num_classes: !!int 77
33
+ criterion: !!str "bce_loss"
34
+ # Dataset parameters
35
+ train_dataset_path: !!str "/volume/ECG_tokenizer/output/MHI/mimic_mhi_psa_train_updated.parquet"
36
+ validation_dataset_path: !!str "/volume/ECG_tokenizer/output/MHI/mhi_psa_test_updated.parquet"
37
+ num_workers: !!int 32
38
+ batch_size: !!int 12
39
+ signal_path_column: !!str "ecg_path" # waveform_path_psa or waveform_path_original
40
+ # waveform parameters
41
+ waveform_length: !!int 2500
42
+ num_leads: !!int 12
43
+ normalize_waveforms: !!bool False
44
+ lead_stats:
45
+ I:
46
+ mean: !!float -0.982577
47
+ std: !!float 33.551235
48
+ II:
49
+ mean: !!float -0.683182
50
+ std: !!float 35.616466
51
+ III:
52
+ mean: !!float 0.299394
53
+ std: !!float 40.365378
54
+ aVR:
55
+ mean: !!float 0.832880
56
+ std: !!float 28.102812
57
+ aVL:
58
+ mean: !!float -0.640986
59
+ std: !!float 32.563651
60
+ aVF:
61
+ mean: !!float -0.191894
62
+ std: !!float 34.169092
63
+ V1:
64
+ mean: !!float -1.351383
65
+ std: !!float 54.328896
66
+ V2:
67
+ mean: !!float -1.366718
68
+ std: !!float 72.522918
69
+ V3:
70
+ mean: !!float -1.325758
71
+ std: !!float 66.522440
72
+ V4:
73
+ mean: !!float -1.091183
74
+ std: !!float 61.623552
75
+ V5:
76
+ mean: !!float -0.937534
77
+ std: !!float 57.341500
78
+ V6:
79
+ mean: !!float -0.882349
80
+ std: !!float 52.247728
deepecg_tokenizer_efficientnet.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cd9fed76ffc94d988dfe8611e5ce9fa12cc95f8ba9f0d227b4186db8f1f06d8e
3
+ size 173070264