-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy path__init__.py
More file actions
65 lines (55 loc) · 1.85 KB
/
Copy path__init__.py
File metadata and controls
65 lines (55 loc) · 1.85 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
""" """
import os
from typing import Optional, Tuple, Union
import torch
from torch import nn
from cfg import ModelCfg
from model import ECG_CRNN_CPSC2020, ECG_SEQ_LAB_NET_CPSC2020
_BASE_DIR = os.path.dirname(os.path.abspath(__file__))
__all__ = [
"load_model",
]
def load_model(which: str = "both") -> Union[nn.Module, Tuple[nn.Module, ...]]:
"""finished, checked,
Parameters:
-----------
which: str,
choice of the models
Returns:
--------
nn.Module, or sequence of nn.Module
"""
if torch.cuda.is_available():
device = torch.device("cuda")
else:
device = torch.device("cpu")
_which = which.lower()
if _which in ["both", "crnn"]:
crnn_cfg = ModelCfg.crnn
crnn_model = ECG_CRNN_CPSC2020(
classes=crnn_cfg.classes,
n_leads=crnn_cfg.n_leads,
input_len=4000,
config=crnn_cfg,
)
crnn_state_dict = torch.load(os.path.join(_BASE_DIR, "crnn_10s.pth"), map_location=device)
crnn_state_dict["clf.lin_0.weight"] = crnn_state_dict.pop("clf.weight")
crnn_state_dict["clf.lin_0.bias"] = crnn_state_dict.pop("clf.bias")
crnn_model.load_state_dict(crnn_state_dict)
crnn_model.eval()
if _which == "crnn":
return crnn_model
if _which in ["both", "seq_lab"]:
seq_lab_cfg = ModelCfg.seq_lab
seq_lab_model = ECG_SEQ_LAB_NET_CPSC2020(
classes=seq_lab_cfg.classes,
n_leads=seq_lab_cfg.n_leads,
input_len=4000,
config=seq_lab_cfg,
)
seq_lab_state_dict = torch.load(os.path.join(_BASE_DIR, "seq_lab_10s.pth"), map_location=device)
seq_lab_model.load_state_dict(seq_lab_state_dict)
seq_lab_model.eval()
if _which == "seq_lab":
return seq_lab_model
return crnn_model, seq_lab_model