【論文抄読】即時心拍変動データから睡眠段階ラベリングを自動化する深層学習モデル

【第1回】自主論文抄読

読んだ論文は Deep learning for automated sleep staging using instantaneous heart rate である。

https://www.nature.com/articles/s41746-020-0291-x

論文について

  • 掲載誌: Nature 系(npj Digital Medicine)
  • 出版日: 2020年8月20日
  • 著者: Niranjan Sridhar 氏を含む9名。いずれもアメリカ・カリフォルニアの Verily Life Sciences(Google 系)所属

内容の要点は次のとおりである。

PSG(ポリソムノグラフィ)の ECG 信号から取り出した系列を入力 $x$、睡眠段階ラベルを正解 $y$ として、ECG のみから睡眠段階(4値)を予測する Dilated CNN を学習した、という研究である。

データについて

学習・検証に用いたのは、Sleep Heart Health Study(SHHS) と Multi-Ethnic Study of Atherosclerosis(MESA) である(いずれも NSRR に申請すれば得られる公開データ)。合わせて 1万件超の PSG が使われている。

データセット夜数
SHHS8,299 夜
MESA2,033 夜

2つのデータセットは混合され、学習 : 検証 : 評価 = 80 : 10 : 10 の割合で、被験者単位に分割された。学習・検証セットでモデルを構築し、評価セットはモデル構築には一切使っていない。

元ラベルは wake / N1 / N2 / N3 / REM の5クラスであったが、次のように4クラスへ集約している。

  • N1・N2 → Light Sleep
  • N3 → Deep Sleep(SHHS にあった N4 も Deep Sleep に統合)
  • wake / REM はそのまま

入力特徴

  • R 波検出には Pan–Tompkins based algorithm を使用
  • その後、$IBI$(inter-beat interval)から $IHR$(instantaneous heart rate)を計算
$$ IBI_i = t_{i+1} - t_i $$

これは $RRI$(R–R interval)と本質的に同じものである。

$$ IHR_i = \frac{1}{IBI_i} $$

$IBI$ の逆数が $IHR$(瞬時心拍数)である。つまりこの研究では $RRI$ そのものではなく、その逆数である $IHR$ を入力系列として使っている。

そのほかの前処理は次のとおりである。

  • 異常処理として、$5\,\mathrm{SD}$ を超える $IBI$ 区間を除外
  • 各夜ごとの平均・標準偏差で $IHR$ を正規化
  • 2 Hz でリサンプリング

入力窓の長さについては、睡眠段階ラベルが 30秒 epoch ごとに与えられており、その epoch の前後それぞれ49秒を合わせた 128秒=256サンプル(2 Hz)が1本の入力系列になる。

モデルアーキテクチャ

図1: モデルアーキテクチャ

図1. 論文のモデル全体像(入力の窓切り出しから、局所特徴抽出・時間方向の dilated 畳み込み、4クラス出力まで)

Dilated 畳み込みブロックを組み込んだ CNN を用いている。発想としては TCN に近い、という印象である。

TCN とは、中心窓を用いているため未来情報を使用して予測しており causal convolution の過去情報だけから推測するという点が明確に異なる。それ以外の1次元畳み込み、dilation で受容野を広げる、residual connection を使用する、時系列長を保つ、という特徴は共通する。

1夜を10時間とすると、IHR 系列は

1
2 Hz × 60秒 × 60分 × 10時間 = 72,000 点

となる。30秒 epoch は最大 1,200個。各30秒 epoch を中心に、長さ256点(128秒)の重複 segment を1個ずつ切り出すため、segment 集合の形状はバッチサイズを $B$ として $(B,\ 1200,\ 256)$ である。

1×1 畳み込みでチャネル数を増やす

  ---
title: 1×1畳み込みでチャネル数を増やす
---
flowchart TB
A("(B, 1200, 256)") --> R["reshape"]
R --> B0("(B × 1200, 1, 256)")
B0 --> B["Conv1d(k=1, d=1, out_channels=8)"]
B --> C("(B × 1200, 8, 256)")

style A fill:none
style B0 fill:none
style C fill:none

局所畳み込みブロック

このブロックでは次のルールが効いている。

  • Conv1d は padding='same' でサイズ変更なし
  • MaxPool(kernel=2, stride=2)で長さが半分
  • ブロック1回あたりチャネル数は2倍(フィルター数を倍にする)

結果として、系列長は 256 → 128 → 64 → 32($\frac{1}{2^3}$ 倍)、チャネル数は 8 → 16 → 32 → 64($2^3$ 倍)になる。

ただし、中間チャネル数 16・32 は論文本文に明記されておらず、再現上の仮定である。

  ---
title: 局所畳み込みblock(1回分)
---
flowchart TB
A("(B × 1200, 8, 256)") --> B["Conv1d(8→16, k=3, d=1, padding='same')"]
B --> BR["LeakyReLU"]
BR --> C["Conv1d(16→16, k=3, d=1, padding='same')"]
C --> CR["LeakyReLU"]
CR --> D["MaxPool1d(kernel=2, stride=2)"]
D --> ADD(("+"))
A --> R1["Residual downsampling"]
R1 --> R2["例: MaxPool1d(2) + Conv1d(k=1, 8→16)"]
R2 --> ADD
ADD --> E("(B × 1200, 16, 128)")

subgraph block
B
BR
C
CR
D
ADD
R1
R2
end
style A fill:none
style E fill:none

Flatten + 全結合

Flatten と全結合層で要素数を128に落とすため、形状は $(B \times 1200,\ 128)$ になる。

  ---
title: flatten+全結合層
---
flowchart
A("(B * 1200, 32, 64)") --> B["Flatten"]
B --> C["(B * 1200, 32 * 64)"]
C --> D["Dense(2048 → 128)"]
D --> E("(B * 1200, 128)")

style A fill:none
style C fill:none
style E fill:none

時間方向の Dilated 畳み込みへ

続いて、時間方向の dilated 畳み込みのために reshape と転置を行う。epoch 数 1,200 に対して畳み込む必要があるため、axis=1 と 2 を入れ替える。

  flowchart TB
A("(B × 1200, 128)") --> B["reshape"]
B --> C("(B, 1200, 128)")
C --> D["transpose(1, 2)"]
D --> E("(B, 128, 1200)")

style A fill:none
style C fill:none
style E fill:none

5層の時間方向 dilated 畳み込みを 2回繰り返す。dilation 側は padding='same' とし、系列長は $(B,\ 128,\ 1200)$ のまま変わらない。

padding の具体的な実装は論文に明記されていない。

  ---
title: 時間方向のdilated畳込みを2回繰り返す(下図は1回分)
---
flowchart
A["(B, 128, 1200)"] --> B["Conv1d(128→128, k=7, d=2)"]
B -- "LeakyReLU(α=0.15)" --> C["Conv1d(128→128, k=7, d=4)"]
C -- "LeakyReLU(α=0.15)" --> D["Conv1d(128→128, k=7, d=8)"]
D -- "LeakyReLU(α=0.15)" --> E["Conv1d(128→128, k=7, d=16)"]
E -- "LeakyReLU(α=0.15)" --> F["Conv1d(128→128, k=7, d=32)"]
F -- "Dropout(rate=0.2)" --> G["(B, 128, 1200)"]
A -- "Residual加算" --> G

subgraph block
B
C
D
E
F
end
style A fill:none
style G fill:none

クラス数への射影

最後に、128 の特徴量を4クラスへ落とすため 1×1 畳み込みを実行する。

  ---
title: 1x1畳込みでチャンネル数を減らす
---
flowchart TB
A("(B, 128, 1200)") --> B["Conv1D(k=1, d=1, filters=4)"]
B --> C("(B, 4, 1200)")

style A fill:none
style C fill:none

再度 axis=1, 2 で転置し、$(B,\ 1200,\ 4)$ となる。

学習まわりのパラメータ

  • すべての畳み込み層とドロップアウト層に L1 正則化
  • バッチ正規化は未使用
  • 損失関数は 平均クロスエントロピー
  • CNN ブロック数・バッチサイズ・学習率・weight decay・ドロップアウト率は、ハイパーパラメータ探索(1,000設定)で決定

最終的な設定は次のとおりである。

項目値
学習 step100万
バッチサイズ2
学習率$1\times10^{-4}$
weight decay の割合0.25
ドロップアウト率0.2

結果

800夜分(561人の被験者)の SHHS データセットのホールドアウトテストセットでは、全体としての4クラス正確度は 77%、Cohen の kappa 係数は 0.66 であった。993夜分(993人の被験者)の Physionet CinC データセットでは、全体としての4クラス正確度は 72%、Cohen の kappa 係数は 0.55 であった。

図2: 混同行列と評価結果

図2. (a) SHHSデータセットの混同行列の正規化版 (b) SHHSデータセットの混同行列のフルカウント版 (c) CinCデータセットの混同行列の正規化版 (d) CinCデータセットの混同行列のフルカウント版 (e) データセットサイズと全体のaccuracy

ホールドアウトされた SHHS データセットでは、1夜あたりの平均正確度は 77.3%(±8.8%)、kappa 係数は 0.65(±0.14) であった。Physionet CinC データセットでは、1夜あたりの平均正確度は 72.2%(±11.2%)、kappa 係数は 0.53(±0.17) であった。

図3: 1夜あたりの評価分布

考察

私が気になった箇所のみ抜粋。

著者のLSTM批判

「心拍時系列は、各時点の情報量は薄いが、睡眠周期のような長時間の文脈が重要なので、LSTMより dilated CNN の方が構造的に向いている」と著者は主張している。

LSTM は自然言語でよく使われるが、自然言語の場合、「彼は昨日、本屋で買った本を読んだ」のように、各単語そのものが多くの意味を持ち、重要な依存関係も数語〜数十語程度に収まることが多い。

しかし、30秒ごとの心拍特徴は、単独では情報量がそれほど多くなく、個々の30秒からは睡眠段階を強く決められないものの、「約90分の睡眠周期」「REM が夜の後半ほど増える」「睡眠段階が一定の順序で遷移する」といった長時間構造には重要な情報がある。

それにより、心拍特徴を扱う上では、より長時間構造を扱える dilation CNN が優れているというわけである。

また、LSTM は各時点を順番に処理するため、30秒単位の心拍系列に、「一過性の心拍上昇」「アーチファクト」「無呼吸イベント」「体動」「R波検出誤差」があると、それらが隠れ状態へ逐次取り込まれるが、LSTM はこういった局所ノイズに過剰に反応し、長期構造をうまく保持できない可能性があるとも考えているようだ。

また筆者は LSTM のパラメータ数は $N \cdot L$($N$ はバッチサイズ、$L$ は系列長)と述べているが、実際には LSTM 1層のパラメータ数は、おおむね $4(HI+H^2+H)$($I$ は入力特徴数、$H$ は hidden size)であり、系列長 $L$ には依存しないため、系列長 $L$ が増えてもパラメータ数は増えない。

項目LSTM で系列長が増えた場合
学習パラメータ数基本的に変わらない
forward 計算量ほぼ $O(L)$ で増える
中間 activation のメモリ学習時はほぼ $O(L)$ で増える
逐次処理時間増える
長距離勾配伝播難しくなる可能性
並列化CNN より難しい

TCN のパラメータ数は $C_{\mathrm{out}}(C_{\mathrm{in}}k+1)$ で、LSTM と同じく時系列長 $L$ には依存しない。

筆者の主張は LSTM ではパラメータ数が系列長に比例し増えることで計算量が膨大になるということであろうが、この点は正しくない。

実際のところ言えるのは、心拍由来の睡眠時系列は、個々の30秒区間の情報密度が低い一方、数十分から約90分に及ぶ長距離構造を含む。LSTM でもこの構造を理論上学習できるが、長い逐次経路を通じて情報を保持する必要があり、最適化や並列化が難しくなる。Dilated CNN は、パラメータ共有を保ちながら少数の層で広い受容野を形成できるため、本課題に適している可能性がある、ということであろう。

これまでの研究の中で最大のデータセット

著者は 1987夜(1748人の被験者)のデータセットはこれまでの研究の中で最大と述べている。

研究の限界

いくつか限界を挙げていたが、その中で特に私の研究と異なる点という意味で「一夜の心拍特徴をまるまる入力に要する点」は重要である。

また、この論文のモデルは、Wake、Light Sleep、および REM において良好な性能を示すが、専門家による参照睡眠段階(正解ラベル)と比較して徐波睡眠(Deep Sleep)を過小評価する傾向があるという。一つの可能性ある理由として、専門家の評価者間信頼度が Deep Sleep において最も低く、年齢および性別によって大きく変動することが知られていることがあるのではないか、とのこと。

この「評価者間信頼度が Deep Sleep において最も低く、年齢および性別によって大きく変動する」という点についての参考文献は、次の2件である。

Hugo で構築されています。
テーマ Stack は Jimmy によって設計されています。