You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用torch_geometric时GPU利用率极低的问题排查与优化求助

问题分析与优化建议

核心原因分析

  1. 实时图构建引发CPU瓶颈
    你的Dataset在__getitem__中返回的是SMILES字符串,而非预先构建好的PyG Data对象,后续需在CPU实时完成SMILES到图的转换——这部分操作耗时且完全占用CPU,导致GPU长时间等待数据,出现"空转",利用率自然极低。

  2. 输入图数据未移至GPU
    虽然模型已通过model.to(device)部署到GPU,但构建的Batch对象默认留在CPU,若未显式将图数据(data1-data4)转移到GPU,模型会被迫在CPU执行计算,直接导致GPU资源闲置。

  3. 模型计算量不足
    你的GCN分支层数少、通道数小,加上分子图本身节点数有限,GPU的计算任务不饱和,很快就能完成一批数据的计算,然后等待下一批,无法充分发挥并行计算能力。

  4. 数据加载配置仍有优化空间
    仅设置num_workers=4和pin_memory=True不足以抵消实时图构建的开销,若CPU核心数充足,worker数量仍有提升空间,且缺少persistent_workers=True这类减少进程启动开销的配置。

具体优化建议

1. 预先生成并保存图数据

提前将所有SMILES转换为PyG Data对象并保存到磁盘,避免训练时实时转换:

# 预处理阶段(训练前执行)
from rdkit import Chem
from torch_geometric.data import Data

def smiles_to_graph(smiles):
    mol = Chem.MolFromSmiles(smiles)
    # 根据任务需求提取节点特征、边索引、边特征(示例需调整)
    node_features = ... 
    edge_index = ...    
    edge_features = ... 
    return Data(x=node_features, edge_index=edge_index, edge_attr=edge_features)

# 遍历数据集转换并保存
preprocessed_train_graphs = []
for idx, row in train_df.iterrows():
    graphs = [smiles_to_graph(row[f'buildingblock{i}_smiles']) for i in range(1,4)] + [smiles_to_graph(row['molecule_smiles'])]
    preprocessed_train_graphs.append(graphs)
torch.save(preprocessed_train_graphs, 'train_graphs.pt')

# 修改Dataset加载预存数据
class MoleculeGraphDataset(Dataset):
    def __init__(self, graphs_path, one_hot_encoded, device):
        super().__init__()
        self.graphs = torch.load(graphs_path)
        self.one_hot_encoded = one_hot_encoded
        self.device = device

    def __len__(self):
        return len(self.graphs)

    def __getitem__(self, idx):
        graphs = self.graphs[idx]
        one_hot = torch.tensor(self.one_hot_encoded[idx], dtype=torch.float, device=self.device)
        target = torch.tensor([train_df.iloc[idx]['binds']], dtype=torch.float, device=self.device)
        return (*graphs, one_hot, target)

2. 确保输入数据移至GPU

在训练循环或collate_fn中显式转移图数据:

# 方案1:在collate_fn中转移
def collate_fn(batch, device):
    transposed = list(zip(*batch))
    graphs = [Batch.from_data_list(graph_list).to(device) for graph_list in transposed[:4]]    
    one_hot_vectors = torch.stack(transposed[4], dim=0)
    targets = torch.stack(transposed[5], dim=0)
    return (*graphs, one_hot_vectors, targets)

# 初始化DataLoader时传入device
train_loader = DataLoader(
    train_dataset,
    batch_size=10000,
    shuffle=True,
    collate_fn=lambda x: collate_fn(x, device),
    num_workers=8,
    pin_memory=True,
    persistent_workers=True
)

# 方案2:在训练循环中转移
for data1, data2, data3, data4, one_hot, targets in train_loader:
    data1, data2, data3, data4 = data1.to(device), data2.to(device), data3.to(device), data4.to(device)
    optimizer.zero_grad()
    outputs = model(data1, data2, data3, data4, one_hot)
    # ...后续训练步骤

3. 增加模型计算量,饱和GPU

调整模型结构提升计算密度:

class MultiGraphGNN(torch.nn.Module):
    def __init__(self, num_node_features, num_edge_features, protein_features_dim):
        super(MultiGraphGNN, self).__init__()
        
        # 增加通道数提升计算量
        self.graph1_conv1 = GCNConv(num_node_features, 64)
        self.graph1_conv2 = GCNConv(64, 128)

        self.graph2_conv1 = GCNConv(num_node_features, 64)
        self.graph2_conv2 = GCNConv(64, 128)

        self.graph3_conv1 = GCNConv(num_node_features, 64)
        self.graph3_conv2 = GCNConv(64, 128)

        self.graph4_conv1 = GCNConv(num_node_features, 128)
        self.graph4_conv2 = GCNConv(128, 256)
        self.graph4_conv3 = GCNConv(256, 512)

        self.fc1 = nn.Linear(128 * 3 + 512 + protein_features_dim, 512)
        self.fc2 = nn.Linear(512, 256)
        self.fc3 = nn.Linear(256, 1)

    # forward方法保持不变...

也可替换为GATConv这类计算量更大的图卷积层,更好利用GPU并行优势。

4. 优化数据加载配置

train_loader = DataLoader(
    train_dataset,
    batch_size=10000,  # 若图节点数过少,可尝试按节点数动态调整batch size
    shuffle=True,
    collate_fn=collate_fn,
    num_workers=16,  # 设置为CPU核心数的1-2倍
    pin_memory=True,
    persistent_workers=True,  # 保持worker进程活跃,减少启动开销
    prefetch_factor=4  # 预取4批数据,缓解GPU等待
)

5. 按图大小分组,优化batch利用率

若数据集内图的节点数差异大,将大小相近的图分为一组,设置合适的batch size,确保每个batch的总节点数足够多,充分利用GPU显存和计算能力。

内容的提问来源于stack exchange,提问作者Giorgio

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 09:32:02