-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel.py
More file actions
84 lines (56 loc) · 3.14 KB
/
Copy pathmodel.py
File metadata and controls
84 lines (56 loc) · 3.14 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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import SAGEConv
class GCNModel(nn.Module):
def __init__(self, num_node_features, num_classes, num_clinical_features, graph_hidden_dim, attention_hidden_dim):
super(GCNModel, self).__init__()
self.num_clinical_features = num_clinical_features
self.ln1 = nn.LayerNorm(graph_hidden_dim)
self.gcn1 = SAGEConv(num_node_features * 2, graph_hidden_dim)
self.gcn2 = SAGEConv(graph_hidden_dim, 1)
self.x_attn_linear = nn.Linear(1, attention_hidden_dim)
self.clinical_linears = nn.ModuleList([
nn.Linear(1, attention_hidden_dim) for _ in range(num_clinical_features)
])
self.score_linears = nn.ModuleList([
nn.Linear(attention_hidden_dim * 2, 1) for _ in range(num_clinical_features)
])
self.fc1 = nn.Sequential(
nn.Linear(num_clinical_features + graph_hidden_dim, num_classes),
nn.Sigmoid()
)
def forward(self, data, clinical_features):
x, edge_index = torch.cat([data.phenotypes, torch.abs(data.phenotypes - data.prototypes)], dim=1), data.edge_index
num_nodes = x.shape[0]
x = self.gcn1(x, edge_index)
x = self.ln1(x)
x = F.relu(x)
# node attention features
x_attn = self.gcn2(x, edge_index) # [num_nodes, 1]
x_attn = self.x_attn_linear(x_attn) # [num_nodes, attention_hidden_dim]
x_attn = x_attn.unsqueeze(1) # [num_nodes, 1, attention_hidden_dim]
# clinical features
clinical_attn = clinical_features.view(1, -1, 1).expand(num_nodes, -1, -1) # [num_nodes, num_clinical_features, 1]
attention_weights_list = []
for i in range(self.num_clinical_features):
clinic_feat = clinical_attn[:, i:i + 1, :] # [num_nodes, 1, 1]
clinic_attn = self.clinical_linears[i](clinic_feat) # [num_nodes, 1, attention_hidden_dim]
combined = torch.cat([x_attn, clinic_attn], dim=-1) # [num_nodes, 1, attention_hidden_dim*2]
score = self.score_linears[i](combined) # [num_nodes, 1, 1]
attention_weight = torch.sigmoid(score)
attention_weights_list.append(attention_weight)
attention_weights_clinic = torch.cat(attention_weights_list, dim=1) # [num_nodes, num_clinical_features, 1]
attention_weights_sum = F.softmax(
attention_weights_clinic.sum(dim=1), # [num_nodes, 1]
dim=0
).transpose(0, 1) # [num_nodes, 1, 1]
x_pool = torch.matmul(attention_weights_sum, x) # [1, graph_hidden_dim]
combined_features = torch.cat((x_pool, clinical_features), dim=1)
output = self.fc1(combined_features)
# output - positive-class probability
# attention_weights_clinic - per-clinical-feature attention weights
# attention_weights_sum - global (node-level) attention weights
# combined_features - final feature vector used for prediction (graph + clinical)
return (output,
attention_weights_clinic, attention_weights_sum, combined_features)