IEEE TMI IF=9.8 | MM-GTUNets:统一多模态图深度学习框架赋能脑疾病预测
引言
脑疾病预测正面临一个核心挑战:如何有效融合多模态数据(如影像与非影像信息)并捕捉其复杂的跨模态关联。现有图深度学习方法往往难以同时处理大规模图数据、充分利用非影像特征,并深入建模模态间的交互,导致预测性能受限。
一项发表于《IEEE Transactions on Medical Imaging》的研究提出了一个名为 MM-GTUNets 的统一多模态图深度学习框架。该框架通过创新的模态奖励表征学习动态构建群体图,并利用基于图Transformer的统一编码器进行自适应跨模态图学习,在自闭症谱系障碍(ASD) 和注意力缺陷多动障碍(ADHD) 的公开数据集上实现了优越的预测性能,为脑疾病诊断提供了新的强大工具。
基本信息
• 文章标题:MM-GTUNETS: UNIFIED MULTI-MODAL GRAPH DEEP LEARNING FOR BRAIN DISORDERS PREDICTION
• 期刊:IEEE Transactions on Medical Imaging
• 影响因子:9.8
• 发表时间:2025年3月28日
• 研究单位:上海海事大学、徐州医科大学附属连云港医院、上海科技大学、香港理工大学
• 论文地址:https://arxiv.org/abs/2406.14455v3
• 算力描述:运行于配备12块NVIDIA GeForce 4090 GPU的服务器
研究内容与方法
1. 数据集构建与预处理
- 成像数据(rs-fMRI)预处理
- 采用C-PAC/Athena管道完成rs-fMRI数据预处理,基于AAL脑图谱将大脑分割为116个感兴趣区域(ROI)
- 计算每个ROI的平均时间序列,通过Pearson相关系数得到功能连接(FC)矩阵
- 提取FC矩阵的上三角元素并展平为一维向量,再通过递归特征消除(RFE) 降维到指定维度:
def rfe_feature_selection(X, y, n_features=500): estimator = SVR(kernel="linear") selector = RFE(estimator, n_features_to_select=n_features, step=10) selector = selector.fit(X, y) return X[:, selector.support_]
- 非成像数据预处理
- 对分类型非成像数据(如性别、采集站点)进行序数编码,数值型数据直接转换为浮点型,拼接为一维特征向量
- 采用预训练变分自编码器(VAE) 重构非成像特征,对齐到成像特征的维度:
class VAE(nn.Module): def __init__(self, input_dim, latent_dim=500): super().__init__() self.encoder = nn.Sequential(nn.Linear(input_dim, 1024), nn.ReLU(), nn.Linear(1024, 512), nn.ReLU()) self.fc_mu = nn.Linear(512, latent_dim) self.fc_logvar = nn.Linear(512, latent_dim) self.decoder = nn.Sequential(nn.Linear(latent_dim, 512), nn.ReLU(), nn.Linear(512, 1024), nn.ReLU(), nn.Linear(1024, input_dim), nn.Sigmoid()) def reparameterize(self, mu, logvar): std = torch.exp(0.5*logvar) eps = torch.randn_like(std) return mu + eps*std def forward(self, x): h = self.encoder(x) mu = self.fc_mu(h) logvar = self.fc_logvar(h) z = self.reparameterize(mu, logvar) return self.decoder(z), mu, logvar
2. MM-GTUNets整体架构
MM-GTUNets是端到端多模态图深度学习框架,由三大核心模块组成:
- Modality-Rewarding Representation Learning (MRRL):动态构建群体图
- Adaptive Cross-Modal Graph Learning (ACMGL):捕捉跨模态与模态内的复杂关系
- 分类与正则化模块:实现疾病预测与模型约束
【MM-GTUNets整体框架图,注:端到端多模态图深度学习框架,包含MRRL、ACMGL与分类模块】
3. Modality-Rewarding Representation Learning (MRRL)
该模块用于对齐多模态特征并构建自适应奖励种群图,分为3个子步骤:
3.1 模态对齐
对成像与非成像特征进行维度对齐,消除模态间隙:
{Ximge=RFE(Ximg,d)Xnone=VAE(Xnon)
\begin{cases}
X^e_{img} = \text{RFE}(X_{img}, d) \\
X^e_{non} = \text{VAE}(X_{non})
\end{cases}
{Ximge=RFE(Ximg,d)Xnone=VAE(Xnon)
其中XimgX_{img}Ximg为原始成像特征,XnonX_{non}Xnon为原始非成像特征,ddd为对齐后的维度。
3.2 亲和度量奖励系统(AMRS)
基于Q-Learning动态学习非成像特征的贡献权重,构建非成像亲和图:
- 定义非成像特征的权重向量α=[α1,α2,...,αv]\alpha = [\alpha_1, \alpha_2, ..., \alpha_v]α=[α1,α2,...,αv],满足∑u=1vαu=1\sum_{u=1}^v \alpha_u=1∑u=1vαu=1且0<αu<10<\alpha_u<10<αu<1
- 维护奖励表RRR、惩罚表PPP、激励表MMM,通过以下规则更新表中元素:
{ru(ui,uj)={1,ui=uj∧yi=yj0,otherwisepu(ui,uj)={1,ui=uj∧yi≠yj0,otherwisemu(ui,uj)={1,ui=uj∧{yi,yj}∈testset0,otherwise \begin{cases} r_u(u_i, u_j) = \begin{cases}1, & u_i=u_j \land y_i=y_j \\0, & \text{otherwise}\end{cases} \\ p_u(u_i, u_j) = \begin{cases}1, & u_i=u_j \land y_i \neq y_j \\0, & \text{otherwise}\end{cases} \\ m_u(u_i, u_j) = \begin{cases}1, & u_i=u_j \land \{y_i,y_j\} \in \text{testset} \\0, & \text{otherwise}\end{cases} \end{cases} ⎩⎨⎧ru(ui,uj)={1,0,ui=uj∧yi=yjotherwisepu(ui,uj)={1,0,ui=uj∧yi=yjotherwisemu(ui,uj)={1,0,ui=uj∧{yi,yj}∈testsetotherwise - 计算非成像亲和图的邻接矩阵CCC:
Cij=Sigmoid(∑u=1vαu(βruRij+βpuPij+βmuMij)) C_{ij} = \text{Sigmoid}\left( \sum_{u=1}^v \alpha_u \left( \beta^u_r R_{ij} + \beta^u_p P_{ij} + \beta^u_m M_{ij} \right) \right) Cij=Sigmoid(u=1∑vαu(βruRij+βpuPij+βmuMij))
对应代码片段:def compute_affinity_matrix(R, P, M, alpha, beta_r, beta_p, beta_m): affinity = 0 for u in range(len(alpha)): affinity += alpha[u] * (beta_r[u]*R + beta_p[u]*P + beta_m[u]*M) return torch.sigmoid(affinity) - 采用Q-Learning优化权重α\alphaα,最大化状态-动作值函数:
Qπ(s,a)=Eπ[Gt∣St=s,At=a]=1N2∑u=1varg maxαu∑i=1N∑j=1NαuReLU(βruRij+βpuPij) Q^\pi(s,a) = \mathbb{E}_\pi \left[ G_t | S_t=s, A_t=a \right] = \frac{1}{N^2} \sum_{u=1}^v \argmax_{\alpha_u} \sum_{i=1}^N \sum_{j=1}^N \alpha_u \text{ReLU}\left( \beta^u_r R_{ij} + \beta^u_p P_{ij} \right) Qπ(s,a)=Eπ[Gt∣St=s,At=a]=N21u=1∑vαuargmaxi=1∑Nj=1∑NαuReLU(βruRij+βpuPij)
【亲和度量奖励系统图,注:基于Q-Learning的AMRS机制,动态调整非成像特征的贡献权重】
3.3 自适应奖励种群图(ARPG)构建
融合多模态特征与亲和图,构建最终的种群图:
- 融合多模态节点特征:Xb=Concat(Ximge,Xnone)X^b = \text{Concat}(X^e_{img}, X^e_{non})Xb=Concat(Ximge,Xnone)
- 计算节点相似度并结合非成像亲和图,得到ARPG的邻接矩阵AAA:
Aij=Sim(Xib,Xjb)⊙Cij,Sim(xi,xj)=exp(−[ρ(xi,xj)]22σ2) A_{ij} = \text{Sim}(X^b_i, X^b_j) \odot C_{ij}, \quad \text{Sim}(x_i,x_j) = \exp\left( -\frac{[\rho(x_i,x_j)]^2}{2\sigma^2} \right) Aij=Sim(Xib,Xjb)⊙Cij,Sim(xi,xj)=exp(−2σ2[ρ(xi,xj)]2)
其中ρ(⋅)\rho(\cdot)ρ(⋅)为相关距离函数,σ\sigmaσ为核宽度,⊙\odot⊙为元素级乘法 - 加入Monte Carlo边Dropout缓解过平滑:
def monte_carlo_edge_dropout(A, dropout_rate=0.3): mask = torch.bernoulli(torch.ones_like(A) * (1 - dropout_rate)) return A * mask
4. Adaptive Cross-Modal Graph Learning (ACMGL)
该模块用于捕捉跨模态与模态内的复杂关系,分为2个子步骤:
4.1 GTUNet编码器
结合Graph UNet的池化机制与Graph Transformer的全局注意力,提取模态特征:
- Graph Transformer(GT)层:更新节点特征,包含注意力计算与门控残差连接
- 注意力计算:
{qi(l)=Wq(l)hi(l)+bq(l),ki(l)=Wk(l)hi(l)+bk(l),vi(l)=Wv(l)hi(l)+bv(l)αij(l)=⟨qi(l),kj(l)⟩+eij∑u∈N(i)⟨qi(l),ku(l)⟩+eiuhˉi(l+1)=∑j∈N(i)αij(l)(vj(l)+eij) \begin{cases} q^{(l)}_i = W^{(l)}_q h^{(l)}_i + b^{(l)}_q, \quad k^{(l)}_i = W^{(l)}_k h^{(l)}_i + b^{(l)}_k, \quad v^{(l)}_i = W^{(l)}_v h^{(l)}_i + b^{(l)}_v \\ \alpha^{(l)}_{ij} = \frac{\langle q^{(l)}_i, k^{(l)}_j \rangle + e_{ij}}{\sum_{u \in \mathcal{N}(i)} \langle q^{(l)}_i, k^{(l)}_u \rangle + e_{iu}} \\ \bar{h}^{(l+1)}_i = \sum_{j \in \mathcal{N}(i)} \alpha^{(l)}_{ij} \left( v^{(l)}_j + e_{ij} \right) \end{cases} ⎩⎨⎧qi(l)=Wq(l)hi(l)+bq(l),ki(l)=Wk(l)hi(l)+bk(l),vi(l)=Wv(l)hi(l)+bv(l)αij(l)=∑u∈N(i)⟨qi(l),ku(l)⟩+eiu⟨qi(l),kj(l)⟩+eijhˉi(l+1)=∑j∈N(i)αij(l)(vj(l)+eij) - 门控残差连接避免过平滑:
{ri(l)=Wr(l)hi(l)+br(l)γi(l)=Sigmoid(Wg(l)[hˉi(l+1);ri(l);hˉi(l+1)−ri(l)])hi(l+1)=ReLU(LN((1−γi(l))hˉi(l+1)+γi(l)ri(l))) \begin{cases} r^{(l)}_i = W^{(l)}_r h^{(l)}_i + b^{(l)}_r \\ \gamma^{(l)}_i = \text{Sigmoid}\left( W^{(l)}_g \left[ \bar{h}^{(l+1)}_i; r^{(l)}_i; \bar{h}^{(l+1)}_i - r^{(l)}_i \right] \right) \\ h^{(l+1)}_i = \text{ReLU}\left( \text{LN}\left( (1-\gamma^{(l)}_i)\bar{h}^{(l+1)}_i + \gamma^{(l)}_i r^{(l)}_i \right) \right) \end{cases} ⎩⎨⎧ri(l)=Wr(l)hi(l)+br(l)γi(l)=Sigmoid(Wg(l)[hˉi(l+1);ri(l);hˉi(l+1)−ri(l)])hi(l+1)=ReLU(LN((1−γi(l))hˉi(l+1)+γi(l)ri(l)))
对应代码片段:
class GraphTransformerLayer(nn.Module): def forward(self, h, adj): q = self.q_proj(h) k = self.k_proj(h) v = self.v_proj(h) # 注意力分数计算 attn_score = torch.bmm(q, k.transpose(1,2)) / math.sqrt(q.size(-1)) attn_score += adj.unsqueeze(0) attn = torch.softmax(attn_score, dim=-1) # 特征更新 h_bar = torch.bmm(attn, v) # 门控残差连接 r = self.r_proj(h) concat = torch.cat([h_bar, r, h_bar - r], dim=-1) gamma = torch.sigmoid(self.g_proj(concat)) h_new = (1 - gamma) * h_bar + gamma * r h_new = self.norm(h_new) return torch.relu(h_new) - 注意力计算:
- gPool/gUnpool 池化/反池化:
- 下采样(gPool):选择top-k信息最丰富的节点,更新邻接矩阵与节点特征:
{idx=rank(δ,bk)Aˉ(l)=A(l)(idx,idx)Hˉ(l)=H(l)(idx,:)⊙(Sigmoid(δ(idx))1dT) \begin{cases} \text{idx} = \text{rank}(\delta, bk) \\ \bar{A}^{(l)} = A^{(l)}(\text{idx}, \text{idx}) \\ \bar{H}^{(l)} = H^{(l)}(\text{idx}, :) \odot \left( \text{Sigmoid}(\delta(\text{idx})) \mathbf{1}^T_d \right) \end{cases} ⎩⎨⎧idx=rank(δ,bk)Aˉ(l)=A(l)(idx,idx)Hˉ(l)=H(l)(idx,:)⊙(Sigmoid(δ(idx))1dT)
其中δ\deltaδ为节点特征在可学习向量上的投影,bkbkbk为保留的节点数 - 上采样(gUnpool):恢复节点特征到原尺寸:
H~(l+θ)=Distribute(H(l−θ),H(l+θ),idx(l−θ)) \tilde{H}^{(l+\theta)} = \text{Distribute}(H^{(l-\theta)}, H^{(l+\theta)}, \text{idx}^{(l-\theta)}) H~(l+θ)=Distribute(H(l−θ),H(l+θ),idx(l−θ))
对应代码片段:
def gpool(h, adj, k): delta = torch.matmul(h, self.pool_proj.weight.t()) + self.pool_proj.bias idx = torch.topk(delta, k, dim=1)[1] # 更新邻接矩阵 adj = torch.gather(torch.gather(adj, 1, idx.unsqueeze(2).repeat(1,1,adj.size(2))), 2, idx.unsqueeze(1).repeat(1,adj.size(1),1)) # 更新节点特征 h = torch.gather(h, 1, idx.unsqueeze(2).repeat(1,1,h.size(2))) delta = torch.gather(delta, 1, idx) h = h * torch.sigmoid(delta).unsqueeze(2) return h, adj, idx - 下采样(gPool):选择top-k信息最丰富的节点,更新邻接矩阵与节点特征:
【Graph Transformer架构图,注:GT层的注意力机制与残差连接,用于全局特征捕捉】
4.2 多模态注意力融合模块
融合模态特异性与共享特征,得到跨模态联合表示:
- 提取成像与非成像的模态特异性特征,以及共享特征:
{Zimgs=GTUNet(Ximge)Znons=GTUNet(Xnone)Zsh=12(Zimgs+Znons) \begin{cases} Z^s_{img} = \text{GTUNet}(X^e_{img}) \\ Z^s_{non} = \text{GTUNet}(X^e_{non}) \\ Z_{sh} = \frac{1}{2}(Z^s_{img} + Z^s_{non}) \end{cases} ⎩⎨⎧Zimgs=GTUNet(Ximge)Znons=GTUNet(Xnone)Zsh=21(Zimgs+Znons) - 计算各特征的注意力权重:
{τsh=tanh(WZsh+B)τimgs=tanh(WimgZimgs+Bimg)τnons=tanh(WnonZnons+Bnon) \begin{cases} \tau_{sh} = \tanh(W Z_{sh} + B) \\ \tau^s_{img} = \tanh(W_{img} Z^s_{img} + B_{img}) \\ \tau^s_{non} = \tanh(W_{non} Z^s_{non} + B_{non}) \end{cases} ⎩⎨⎧τsh=tanh(WZsh+B)τimgs=tanh(WimgZimgs+Bimg)τnons=tanh(WnonZnons+Bnon) - 融合得到最终联合表示ZZZ:
Z=τsh⊙Zsh+τimgs⊙Zimgs+τnons⊙Znons Z = \tau_{sh} \odot Z_{sh} + \tau^s_{img} \odot Z^s_{img} + \tau^s_{non} \odot Z^s_{non} Z=τsh⊙Zsh+τimgs⊙Zimgs+τnons⊙Znons
对应代码片段:class MultiModalFusion(nn.Module): def forward(self, z_img, z_non): z_sh = (z_img + z_non) / 2 # 计算注意力权重 tau_sh = torch.tanh(self.w_sh(z_sh)) tau_img = torch.tanh(self.w_img(z_img)) tau_non = torch.tanh(self.w_non(z_non)) # 特征融合 z = tau_sh * z_sh + tau_img * z_img + tau_non * z_non return z
5. 分类与正则化模块
5.1 疾病预测与贡献权重计算
- 采用MLP对跨模态联合表示ZZZ进行分类:y^=MLP(Z)\hat{y} = \text{MLP}(Z)y^=MLP(Z)
- 计算各模态的贡献权重,用于模型可解释性分析:
{ω=(ωimg,ωnon)=Softmax(tr(τimgsτimgs)tr(τshτsh),tr(τnonsτnons)tr(τshτsh))tr(A)=∑iAii \begin{cases} \omega = (\omega_{img}, \omega_{non}) = \text{Softmax}\left( \frac{\text{tr}(\tau^s_{img}\tau^s_{img})}{\text{tr}(\tau_{sh}\tau_{sh})}, \frac{\text{tr}(\tau^s_{non}\tau^s_{non})}{\text{tr}(\tau_{sh}\tau_{sh})} \right) \\ \text{tr}(A) = \sum_{i} A_{ii} \end{cases} {ω=(ωimg,ωnon)=Softmax(tr(τshτsh)tr(τimgsτimgs),tr(τshτsh)tr(τnonsτnons))tr(A)=∑iAii
5.2 目标函数
总损失包含交叉熵损失、图正则化损失与奖励正则化损失,约束模型训练:
Ltotal=Lce+ωimgLimgg+ωnon(Lnong+ηLr)
\mathcal{L}_{total} = \mathcal{L}_{ce} + \omega_{img}\mathcal{L}^g_{img} + \omega_{non}\left( \mathcal{L}^g_{non} + \eta \mathcal{L}_r \right)
Ltotal=Lce+ωimgLimgg+ωnon(Lnong+ηLr)
其中:
- 图正则化损失:约束图的平滑性与稀疏性
Lψg=λLsmhg+μLdeg,Lsmhg=12N2∑i,j=1NAij∥ziψ−zjψ∥22,Ldeg=−1N1Tlog(A⋅1) \mathcal{L}^g_\psi = \lambda \mathcal{L}^g_{smh} + \mu \mathcal{L}_{deg}, \quad \mathcal{L}^g_{smh} = \frac{1}{2N^2} \sum_{i,j=1}^N A_{ij} \| z^\psi_i - z^\psi_j \|^2_2, \quad \mathcal{L}_{deg} = -\frac{1}{N} \mathbf{1}^T \log(A \cdot \mathbf{1}) Lψg=λLsmhg+μLdeg,Lsmhg=2N21i,j=1∑NAij∥ziψ−zjψ∥22,Ldeg=−N11Tlog(A⋅1) - 奖励正则化损失:约束AMRS的Q-Learning优化过程
Lr=1Qπ(s,a) \mathcal{L}_r = \frac{1}{Q^\pi(s,a)} Lr=Qπ(s,a)1
对应代码片段:def total_loss(y_pred, y_true, z_img, z_non, A, omega_img, omega_non, eta, q_value): ce_loss = F.cross_entropy(y_pred, y_true) # 图正则化损失计算 def graph_regularization(z, adj): smh_loss = 0.5 * torch.sum(adj * torch.norm(z.unsqueeze(1) - z.unsqueeze(2), dim=-1)**2) / (z.size(0)**2) deg_loss = -torch.mean(torch.log(torch.sum(adj, dim=-1) + 1e-8)) return smh_loss * self.lambda_ + deg_loss * self.mu img_reg = graph_regularization(z_img, A) non_reg = graph_regularization(z_non, A) # 奖励正则化损失 r_loss = 1 / q_value # 总损失 total = ce_loss + omega_img * img_reg + omega_non * (non_reg + eta * r_loss) return total
实验结果分析
MM-GTUNets在脑疾病预测中的性能表现
以下图表展示了MM-GTUNets模型在ABIDE和ADHD-200两个公开数据集上的预测性能,并与多种基线方法进行了比较。评估指标包括准确率(ACC)、灵敏度(SEN)、特异性(SPE)和AUC。

- 总体性能优势:MM-GTUNets在两个数据集上的所有评估指标中均达到最优或次优水平。在ABIDE数据集上,其准确率(82.92%)和AUC(88.21%)均显著优于其他方法。在ADHD-200数据集上,其AUC(90.71%)表现最佳。
- 多模态数据的有效性:与仅使用成像数据的单模态方法相比,大多数多模态方法(包括MM-GTUNets)表现更优,证明了整合成像与非成像数据的价值。
- 图构建方法的稳定性:基于群体图的方法(如Pop-GCN、EV-GCN)通常比基于脑图的方法(如Brain-GNN、DGCN)具有更小的性能标准差,表明其在处理群体关联特征时更为稳定。
模态联合表示的可视化与消融研究
该部分通过可视化与消融实验,验证了模型关键组件的有效性。

-
特征区分度:通过t-SNE对模型学习到的模态联合表示Z进行降维可视化。结果显示,健康对照组与患者组的特征形成了两个界限清晰的簇,表明模型学到的多模态特征具有强大的判别能力。


-
非成像特征重建器的作用:消融研究表明,使用变分自编码器(VAE) 重建非成像特征显著提升了模型性能(ABIDE上ACC: 82.92%),优于使用MLP或普通自编码器(AE)的方案,证明了VAE在弥合模态差距方面的有效性。
-
编码器架构的影响:对比不同图编码器架构,采用图U-Net架构的GTUNet编码器取得了最佳性能,验证了其通过下采样过滤重要节点特征、结合局部与全局信息的优势。
模型组件贡献与可解释性分析
此部分分析了不同模态数据对预测的贡献,并探讨了模型的可扩展性。


-
模态贡献权重:可视化分析显示,静息态功能磁共振成像(rs-fMRI) 数据对预测的贡献最大。在非成像数据中,性别、年龄和采集站点的影响相对均衡,其中性别的影响略大,而年龄的影响则更为一致。
-
输入数据的影响:仅使用非成像数据时模型性能很差(ACC约52-58%),仅使用成像数据时性能中等(ACC约76-80%),而结合两者时达到最优(ACC约82-83%),表明非成像数据对成像数据起到了有效的补充作用。


-
可扩展性与硬件需求:随着图规模(采样比)增大,模型性能逐渐收敛并趋于稳定。同时,模型的浮点运算量、GPU内存占用和训练时间随图规模增长而增加,为实际部署提供了硬件需求参考。
优势与局限
优势
• 多模态融合能力强:模型通过模态奖励表示学习(MRRL) 与自适应跨模态图学习(ACMGL),有效整合成像与非成像数据,并动态学习各模态贡献权重,提升了脑疾病预测的准确性。
• 图结构学习优化:采用基于图Transformer的GTUNet编码器,结合图U-Net的下采样机制,能过滤重要节点特征并捕获全局与局部信息,适用于处理大规模复杂图数据。
• 可解释性支持:模型可可视化各模态(如rs-fMRI、性别、年龄、采集站点)在预测中的贡献权重,为临床决策提供了一定的可解释依据。
局限
• 计算资源要求较高:模型包含VAE预训练、图Transformer与多模态注意力融合等模块,随着图规模增大,训练时间与GPU内存消耗显著增加,部署成本较高。
• 对成像数据依赖性大:实验表明,仅使用非成像数据时模型分类能力很弱,性能高度依赖成像数据,在成像数据缺失或质量差的场景中可能受限。
•实时应用受限:框架基于转导学习,预测时需要处理整个图结构,难以支持对新样本的快速增量预测,不适用于需要实时决策的临床场景。
参考文献
- Disease prediction using graph convolutional networks: Application to Autism Spectrum Disorder and Alzheimer’s disease Parisot et al., 2018:该论文提出了基于图卷积网络的疾病预测方法,并构建了静态人口图,是本研究构建自适应奖励人口图(ARPG) 的重要基线。本研究提出的AMRS和MRRL模块旨在克服其固定相似性度量的局限性。
- Disease prediction with edge-variational graph convolutional networks Huang and Chung, 2022:本文提出了边变分图卷积网络(EV-GCN),能够动态调整人口图的边权重。本研究提出的自适应奖励人口图构造方法(MRRL)和AMRS系统,是对其自适应图学习思想的进一步发展和深化。
- Multi-Modal Graph Learning for Disease Prediction Zheng et al., 2022:该论文提出了一个多模态图学习(MMGL) 框架,用于脑疾病预测。本研究提出的MM-GTUNets框架,特别是其自适应跨模态图学习(ACMGL) 模块,借鉴并扩展了其多模态特征交互与融合的思路。
- Graph U-Nets Gao and Ji, 2019:本文提出了Graph U-Nets 架构,引入了图池化(gPool)与反池化(gUnpool)操作。本研究将其与图变换器(GT)结合,构建了GTUNet编码器,用于从大规模图数据中过滤关键节点特征并提取全局与局部信息。
- Do Transformers Really Perform Bad for Graph Representation? Ying et al., 2021:该论文探讨了Transformer在图表示学习中的应用,提出了图变换器(GT)层。本研究采用GT层作为GTUNet的核心组件,以利用其强大的全局上下文捕获和自注意力机制来学习节点间复杂关系。
更多推荐
所有评论(0)