- Цель проекта
- Установка и запуск
- Описание данных
- Описание библиотеки
4.1. Обработка датасета
4.2. Архитектура нейросети
4.2.1. ResNet
4.2.2. ResNetMeta
4.2.3. EarlyStopping
4.2.4. Hydra - Метрики
- Условия эксперимента
- Результаты
7.1. Обучение на ЭКГ-сигналах
7.2. Обучение с метаданными - Обсуждение результатов
- Источники информации
Целью данного проекта является исследование качества предсказания нейронной сети с архитектурой ResNet1d18 на датасете PTB-XL, а также влияние метаданных (рост, вес, пол, возраст, индекс массы тела) на точность предсказаний. Для исследования были выбраны 5 нарушений ритма: синусовый ритм, аритмия, брадикардия, тахикардия, фибрилляция предсердий.
Также была написана библиотека ecg_classification, которая позволяет воспроизводить данный эксперимент. Она позволяет загружать и предобрабатывать датасет PTB-XL, обучать и тестировать модель ResNet1d с последующим анализом результатов по множеству метрик.
- Установка:
git clone https://github.com/EntryFrager/ecg_classifier.git
cd ./ecg_classifier- Установка зависимостей, создание виртуальной среды и загрузка датасета PTB-XL:
source ./scripts/setup_venv.sh
bash ./scripts/download_datasets.sh- Сборка и одиночный запуск пакета:
pip install -e .
ecg_classifier- Для запуска экспериментов можно использовать готовый скрипт:
python3 run_experiments.py- Или пропискать команду запуска multirun Hydra:
ecg_classifier -m optimizer.lr=0.001,0.0001В папку ./outputs сохраняются все запуски программы. Для каждого запуска создается отдельная подпапка, в которой хранятся:
- актуальные конфиги Hydra (config.yaml)
- логи выполнения программы (main.log)
- лучшая модель с подобранным threshold для бинаризации (save_best_models)
- 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 Гц.
Класс ECGDataset наследуется от torch.utils.Dataset и позволяет предобработать датасет PTB-XL. Рассмотрим основные функции класса:
- _set_target_labels выставляет бинарные метки для исследуемых болезней.
- _process_ecg_signals выделяет нужные нам названия файлов с определенной частотой записи.
- _process_metadata обрабатывает конкретные метаданные, а именно вес, возраст и рост.
- Функция get_dataset реализует подготовку обучающей, валидационной и тестовой выборки.
- __getitem__ возвращает элемент датасета по индексу, загружая файл экг сигнала и производя его нормализацию и перевод в тензор.
Также были реализованы классы Compose, Normalize, ToTensor под данный датасет. Compose и ToTensor никак не отличаются от библиотечных реализаций PyTorch. Normalize же имеет другую структуру, так как sample представляет собой словарь, вида:
sample = {
"ecg_signals": ecg_signals,
"metadata": metadata,
"pqrst_features": pqrst_features,
"labels": labels
}Поэтому он производит нормализацию по каждому ключу отдельно (кроме labels).
Класс ResNet содержит архитектуру нейросети ResNet1d, включая реализации BasicBlock и BottleNeck. Данная архитектура поддерживает как и бинарную классификацию одного класса, так и MultiLabel классификации для нескольких классов.
Для эксперимента была выбрана архитектура ResNet1d18.
Для оценки влияния метаданных на точность предсказаний написан отдельный класс ResNetMeta. Архитектура такой нейросети состоит из двух ветвей и head слоя:
- Первая ветвь: ResNet1d без финального Linear слоя, для обучения ЭКГ-сигналов.
- Вторая ветвь: MLP для обработки метаданных.
- head слой: fully connected слой из двух Linear. Перед ним предсказания обоих ветвей конкатенируются и прогоняются через него.
Для избежания переобучения был реализован класс EarlyStopping. Он смотрит на то, как изменяется val_loss. Если изменения не происходят на протяжении скольких-то эпох(patience), то он возвращает флаг True, что можно использовать для остановки обучения. Если же происходит улучшение val_loss, то он сохраняет данную модель и threshold подобранный для нее как лучшие.
Функция train возвращает лучшую модель за все время обучения и соответствующий ей threshold.
Для удобства работы с проектом была внедрена Hydra. Она позволяет хранить и менять значения гиперпараметров в конфигурационных файлах, а не менять их в коде. В результате это упорядочивает выходные данные. Также Hydra дает возможность запускать эксперименты при различных значениях гиперпараметров.
NOTE: Подробнее про Hydra вы можете прочитать здесь
Подсчет и вывод всех метрик происходит в функции get_metrics. Рассмотрим метрики, которые особенно важны при классификации заболеваний.
Для оценки верности предсказания модели нельзя смотреть только на accuracy, так как она не учитывает:
- дисбаланс классов (вы можете проверить его при помощи функции get_stat())
- насколько критично FN (чем больше FN, тем больше будет больных пациентов, которым поставили неверный диагноз)
Поэтому были взяты такие метрики, как:
- sensitivity (доля истинно больных)
- specificity (доля истинно здоровых)
- precision (показывает сколько пациентов действительно больны)
- f1_score (баланс между sensitivity и precision)
- ROC AUC (Receiver Operating Characteristic Area Under Curve, способность модели правильно отделять классы при различных threshold)
В случае MultiLabel классификации также выводится информация о micro, macro усреднениях. Также выводится classification report от sklearn.
На выходе из нейросети мы получаем сырые логиты, поэтому они пропускаются через sigmoid, для преобразования в вероятности. Для бинаризации вероятностей в процессе обучения подбирается наилучший threshold, от чего напрямую зависят все метрики. Формула для подбора threshold следующая:
Видно, что коэффициент alpha отвечает за то, с каким весом мы будем учитывать sensitivity, а beta - specificity. К примеру, при
Вывод для одной эпохи обучения выглядит таким образом:
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 классификации.
Для обучения с метаданными были выбраны такие параметры, как:
- возраст
- вес
- рост
- пол
- индекс массы тела
В данном эксперименте производилась бинарная классификация по каждому классу отдельно. Для каждого класса происходило обучение на разных 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. А для остальных классов наоборот:
| 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 (что также важно, ведь зачем нам отправлять на лечение здоровых людей). На данный момент можно еще как-то экспериментировать с архитектурой сети при использовании метаданных и с подбором более оптимальных гиперпараметров, так как текущий результат очевидно можно улучшить.
В результате данного проекта были реализованы:
- Класс ECGDataset, позволяющий обрабатывать датасет PTB-XL
- Архитектура нейросети ResNet1d с добавлением метаданных и без них
- Гибкий метод подбора threshold для рассчета метрик и настройки нейросети под определенную клиническую задачу
- Получены хорошие метрики при обучении на ЭКГ-сигналах с добавлением метаданных и без них