Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

51 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Классификатор ЭКГ сигналов

Оглавление

  1. Цель проекта
  2. Установка и запуск
  3. Описание данных
  4. Описание библиотеки
    4.1. Обработка датасета
    4.2. Архитектура нейросети
    4.2.1. ResNet
    4.2.2. ResNetMeta
    4.2.3. EarlyStopping
    4.2.4. Hydra
  5. Метрики
  6. Условия эксперимента
  7. Результаты
    7.1. Обучение на ЭКГ-сигналах
    7.2. Обучение с метаданными
  8. Обсуждение результатов
  9. Источники информации

Цель проекта

Целью данного проекта является исследование качества предсказания нейронной сети с архитектурой ResNet1d18 на датасете PTB-XL, а также влияние метаданных (рост, вес, пол, возраст, индекс массы тела) на точность предсказаний. Для исследования были выбраны 5 нарушений ритма: синусовый ритм, аритмия, брадикардия, тахикардия, фибрилляция предсердий.

Также была написана библиотека ecg_classification, которая позволяет воспроизводить данный эксперимент. Она позволяет загружать и предобрабатывать датасет PTB-XL, обучать и тестировать модель ResNet1d с последующим анализом результатов по множеству метрик.

Установка и запуск

  1. Установка:
git clone https://github.com/EntryFrager/ecg_classifier.git
cd ./ecg_classifier
  1. Установка зависимостей, создание виртуальной среды и загрузка датасета PTB-XL:
source ./scripts/setup_venv.sh
bash ./scripts/download_datasets.sh
  1. Сборка и одиночный запуск пакета:
pip install -e .
ecg_classifier
  1. Для запуска экспериментов можно использовать готовый скрипт:
python3 run_experiments.py
  1. Или пропискать команду запуска multirun Hydra:
ecg_classifier -m optimizer.lr=0.001,0.0001

В папку ./outputs сохраняются все запуски программы. Для каждого запуска создается отдельная подпапка, в которой хранятся:

  1. актуальные конфиги Hydra (config.yaml)
  2. логи выполнения программы (main.log)
  3. лучшая модель с подобранным threshold для бинаризации (save_best_models)
  4. summary для TensorBoard (./logs)

Описание данных

Датасет PTB-XL содержит 21837 записей от 18885 пациентов продолжительностью 10 секунд. Данные ЭКГ-сигналов были аннотированы двумя кардиологами в виде набора данных с несколькими метками, в котором диагностические метки были дополнительно объединены в супер- и подклассы. Набор данных охватывает широкий спектр диагностических классов, включая, в частности, большую часть записей о состоянии здоровья.

В столбце scp_codes (пример: {'NORM': 1.0, 'STTC': 0.9})содержится информация о заболеваниях пациента. Эта информация хранится в виде словаря, где ключ - заболевание, а значение - вероятность заболевания.

IMPORTANT: В данном проекте было принято считать, что если заболевание присутствует как ключ, то для него бинарная метка выставляется как 1, иначе 0.

Сами ЭКГ-сигналы хранятся в отдельных файлах, пути к которым указаны в столбцах filename_lr и filename_hr. В первом столбце хранятся названия файлов, в которые были записаны экг с частотой 100 Гц, а во втором - с частотой 500 Гц.

NOTE: Подробнее о PTB-XL вы можете прочитать здесь и здесь

Описание библиотеки

Обработка датасета

Класс ECGDataset наследуется от torch.utils.Dataset и позволяет предобработать датасет PTB-XL. Рассмотрим основные функции класса:

  1. _set_target_labels выставляет бинарные метки для исследуемых болезней.
  2. _process_ecg_signals выделяет нужные нам названия файлов с определенной частотой записи.
  3. _process_metadata обрабатывает конкретные метаданные, а именно вес, возраст и рост.
  4. Функция get_dataset реализует подготовку обучающей, валидационной и тестовой выборки.
  5. __getitem__ возвращает элемент датасета по индексу, загружая файл экг сигнала и производя его нормализацию и перевод в тензор.

Также были реализованы классы Compose, Normalize, ToTensor под данный датасет. Compose и ToTensor никак не отличаются от библиотечных реализаций PyTorch. Normalize же имеет другую структуру, так как sample представляет собой словарь, вида:

sample = {
    "ecg_signals": ecg_signals,
    "metadata": metadata,
    "pqrst_features": pqrst_features,
    "labels": labels
}

Поэтому он производит нормализацию по каждому ключу отдельно (кроме labels).

Архитектура нейросети

ResNet

Класс ResNet содержит архитектуру нейросети ResNet1d, включая реализации BasicBlock и BottleNeck. Данная архитектура поддерживает как и бинарную классификацию одного класса, так и MultiLabel классификации для нескольких классов.

Для эксперимента была выбрана архитектура ResNet1d18.

ResNetMeta

Для оценки влияния метаданных на точность предсказаний написан отдельный класс ResNetMeta. Архитектура такой нейросети состоит из двух ветвей и head слоя:

  1. Первая ветвь: ResNet1d без финального Linear слоя, для обучения ЭКГ-сигналов.
  2. Вторая ветвь: MLP для обработки метаданных.
  3. head слой: fully connected слой из двух Linear. Перед ним предсказания обоих ветвей конкатенируются и прогоняются через него.

EarlyStopping

Для избежания переобучения был реализован класс EarlyStopping. Он смотрит на то, как изменяется val_loss. Если изменения не происходят на протяжении скольких-то эпох(patience), то он возвращает флаг True, что можно использовать для остановки обучения. Если же происходит улучшение val_loss, то он сохраняет данную модель и threshold подобранный для нее как лучшие.

Функция train возвращает лучшую модель за все время обучения и соответствующий ей threshold.

Hydra

Для удобства работы с проектом была внедрена Hydra. Она позволяет хранить и менять значения гиперпараметров в конфигурационных файлах, а не менять их в коде. В результате это упорядочивает выходные данные. Также Hydra дает возможность запускать эксперименты при различных значениях гиперпараметров.

NOTE: Подробнее про Hydra вы можете прочитать здесь

Метрики

Подсчет и вывод всех метрик происходит в функции get_metrics. Рассмотрим метрики, которые особенно важны при классификации заболеваний.

Для оценки верности предсказания модели нельзя смотреть только на accuracy, так как она не учитывает:

  1. дисбаланс классов (вы можете проверить его при помощи функции get_stat())
  2. насколько критично FN (чем больше FN, тем больше будет больных пациентов, которым поставили неверный диагноз)

Поэтому были взяты такие метрики, как:

  1. sensitivity (доля истинно больных)
  2. specificity (доля истинно здоровых)
  3. precision (показывает сколько пациентов действительно больны)
  4. f1_score (баланс между sensitivity и precision)
  5. ROC AUC (Receiver Operating Characteristic Area Under Curve, способность модели правильно отделять классы при различных threshold)

В случае MultiLabel классификации также выводится информация о micro, macro усреднениях. Также выводится classification report от sklearn.

На выходе из нейросети мы получаем сырые логиты, поэтому они пропускаются через sigmoid, для преобразования в вероятности. Для бинаризации вероятностей в процессе обучения подбирается наилучший threshold, от чего напрямую зависят все метрики. Формула для подбора threshold следующая:

$$ \text{threshold} = \text{alpha} \times \text{sensitivity} + \text{beta} \times \text{specificity} $$ $$ \text{alpha} + \text{beta} = 1 $$ $$ 0 \leq \text{alpha} \leq 1 $$ $$ 0 \leq \text{beta} \leq 1 $$

Видно, что коэффициент alpha отвечает за то, с каким весом мы будем учитывать sensitivity, а beta - specificity. К примеру, при $alpha=0.8$, наша модель будет более чувствительна к пропущенным больным пациентам (соответственно более низкий FN).

Вывод для одной эпохи обучения выглядит таким образом:

Epoch 1/30:

Validation metrics:
Confusion matrix:
  TP  FP  TN  FN
1739 203 198  53
 TP  FP   TN  FN
 35 440 1661  57
 TP  FP   TN  FN
 59  16 2091  27
 TP  FP   TN  FN
 31  18 2111  33
 TP  FP   TN  FN
121  73 1964  35

Micro averaging:
sensitivity  0.906393
specificity  0.914530
precision    0.725777
f1 score     0.806091

Macro averaging:
sensitivity  0.659384
specificity  0.846491
precision    0.602437
f1 score     0.629625

ROC AUC: 0.8822

Classification report from sklearn:
              precision    recall  f1-score   support

           0       0.90      0.97      0.93      1792
           1       0.07      0.38      0.12        92
           2       0.79      0.69      0.73        86
           3       0.63      0.48      0.55        64
           4       0.62      0.78      0.69       156

   micro avg       0.73      0.91      0.81      2190
   macro avg       0.60      0.66      0.61      2190
weighted avg       0.83      0.91      0.86      2190
 samples avg       0.77      0.87      0.80      2190


train Loss: 0.3307
val Loss: 0.5607

NOTE: Более подробно с метриками вы можете ознакомиться здесь и здесь

Условия эксперимента

Для обучения модели используются EarlyStopping и Scheduler, чтобы при длительном переобучении (val_loss только повышается) прекратить обучение и своевременно понизить скорость обучения, чтобы избежать переобучения.

В качестве функции для оценки loss была взята функция BCEWithLogitsLoss. Её главный плюс в том, что она объединяет в себе sigmoid и BCELoss. Это делает её численно устойчивой и дает возможность независимо оценивать предсказания для различных классов при MultiLabel классификации.

Для обучения с метаданными были выбраны такие параметры, как:

  1. возраст
  2. вес
  3. рост
  4. пол
  5. индекс массы тела

В данном эксперименте производилась бинарная классификация по каждому классу отдельно. Для каждого класса происходило обучение на разных learning rate. Лучшая выбиралась по валидационным метрикам. Ниже в результатах приводятся метрики каждого класса на тестовом наборе данных лучшей модели за все время обучения.

Общие гиперпараметры для всех классов:

HyperParameters value
sampling rate при записи ЭКГ-сигналов 100Гц
Архитектура нейросети ResNet1d
epochs 30
batch_size 64
Optimizer Adam
Optimizer weight decay 1e-6
Loss function BCEWithLogitsLoss
Scheduler ReduceLROnPlateau
Scheduler mode min
Scheduler factor 0.6
Scheduler patience 3
EarlyStopping patience 8
dropout‑rate в head 0.25
dropout‑rate в backbone 0.1
dropout-rate в metadata (при добавлении метаданных) 0.1

В связи с тем, что синусовый ритм является обычным состоянием человека, то его бинарная классификация 1 фактически означает, что он здоров, а 0 - не здоров. То есть необходимо смотреть при обучении на specificity, чтобы получить низкий FP. Для этого коэффициенты alpha и beta отвечающие соответственно за sensitivity и specificity равны 0.2, 0.8. А для остальных классов наоборот: $alpha = 0.8$ и $beta = 0.2$

Специфичные гиперпараметры обучения на ЭКГ-сигналах

HyperParameters sinus arrit tach brad afib
Learning rate(LR) 1e-2 1e-4 1e-3 1e-4 1e-3

Специфичные гиперпараметры обучения на ЭКГ-сигналах с метаданными

HyperParameters sinus arrit tach brad afib
Learning rate(LR) 1e-2 1e-3 1e-3 1e-4 1e-4

Результаты

Обучение на ЭКГ-сигналах

metrics sinus arrit tach brad afib
sensitivity 0.91 0.86 0.98 0.95 0.91
specificity 0.76 0.23 0.94 0.83 0.97
precision 0.95 0.04 0.41 0.15 0.71
f1_score 0.93 0.09 0.57 0.26 0.80

Обучение на ЭКГ-сигналах с метаданными

metrics sinus arrit tach brad afib
sensitivity 0.88 0.59 0.98 0.92 0.94
specificity 0.83 0.79 0.96 0.89 0.96
precision 0.96 0.11 0.50 0.20 0.65
f1_score 0.92 0.18 0.66 0.33 0.77

NOTE: Более подробно метрики расписаны в файле

Обсуждение результатов

В ходе эксперимента были обучены 10 бинарных классификаторов (5 с использованием метаданных и 5 без).

Легко заметить, что при обучении без метаданных sensitivity для всех 5 классов довольно высокая(86% для аритмии и больше 90% для всех остальных классов), но specificity местами имеет довольно низкие значения (76% для синусового ритма, 23% для аритмии, 83% для брадикардии).

Добавление метаданных позволяет нам повысить specificity, не сильно вредя sensitivity. То есть при добавлении метаданных мы уравновешиваем sensitivity и specificity (причем гиперпараметры alpha и beta для подбора threshold в обоих случаях были одинаковыми).

Добавление метаданных на текущий момент не дало сильного уменьшения FN (где-то повысило), но оно сильно уменьшило количество FP (что также важно, ведь зачем нам отправлять на лечение здоровых людей). На данный момент можно еще как-то экспериментировать с архитектурой сети при использовании метаданных и с подбором более оптимальных гиперпараметров, так как текущий результат очевидно можно улучшить.

Вывод

В результате данного проекта были реализованы:

  1. Класс ECGDataset, позволяющий обрабатывать датасет PTB-XL
  2. Архитектура нейросети ResNet1d с добавлением метаданных и без них
  3. Гибкий метод подбора threshold для рассчета метрик и настройки нейросети под определенную клиническую задачу
  4. Получены хорошие метрики при обучении на ЭКГ-сигналах с добавлением метаданных и без них

Источники информации

  1. https://physionet.org/content/ptb-xl/1.0.3/
  2. https://pmc.ncbi.nlm.nih.gov/articles/PMC7248071/
  3. https://docs.pytorch.org/vision/main/models/resnet.html
  4. https://hydra.cc/docs/intro/
  5. https://habr.com/ru/articles/821547/
  6. https://education.yandex.ru/handbook/ml/article/metriki-klassifikacii-i-regressii

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages