(self, out_dim: int, args=None, hidden: int = 128, n_layers: int = 2, dropout: float = 0.1)
| 10 | |
| 11 | |
| 12 | def __init__(self, out_dim: int, args=None, hidden: int = 128, n_layers: int = 2, dropout: float = 0.1): |
| 13 | super().__init__() |
| 14 | self.out_dim = out_dim |
| 15 | self.hidden = hidden |
| 16 | self.n_layers = n_layers |
| 17 | self.dropout = nn.Dropout(p=dropout) |
| 18 | |
| 19 | self.type_embeddings = nn.ParameterDict() |
| 20 | |
| 21 | self.diag_emb = nn.Embedding(len(DIAGNOSIS_VOCAB) + 1, hidden) |
| 22 | self.desc_emb = nn.Embedding(len(DESCRIPTOR_VOCAB) + 1, hidden) |
| 23 | self.node_norm = nn.LayerNorm(hidden) |
| 24 | |
| 25 | self.convs = nn.ModuleList([ |
| 26 | HeteroGraphConv({}, aggregate='sum') for _ in range(n_layers) |
| 27 | ]) |
| 28 | |
| 29 | self.proj = nn.Linear(hidden, out_dim) |
| 30 | |
| 31 | def _ensure_type_embeddings(self, g: dgl.DGLHeteroGraph, device): |
| 32 | for ntype in g.ntypes: |
nothing calls this directly
no outgoing calls
no test coverage detected