pytorch geometric并行运算要注意的一些问题

如果有多gpu的情况,可以这样写:

from torch_geometric.nn import DataParallel

if torch.cuda.device_count() > 1:
  model = DataParallel(model, device_ids=[0,1])

这样模型就会变成pytorch geometric里的并行模型DataParallel Model,但随之而来的会有各种的限制和问题:

1. 数据集形式变化:DataLoader变为DataListLoader

需要自动将DataLoader变为DataListLoader,这样对于并行的模型,多张图的训练并不会在最开始就合并成一张大图。

from torch_geometric.loader import DataListLoader
for _, batch_data in enumerate(tqdm(loader, desc="Iteration", mininterval=30)):
    print(batch_data.type)
    pred = model(batch_data)

这里的batch_data是list类型的,进入model之后,才会合并成一张大图,注意如果batch_size设置为64,那就是64张图分别在多个gpu输入并进行训练,但并不一定会平均分为32,有可能会有差别。

2. 输入的限制

如果model的forward输入了多个参数是不被允许的,所以只能输入一个元素,用tuple打包的话,就没办法自动合并图了。

3. 模型的输出

from torch_geometric.loader import DataListLoader
for _, batch_data in enumerate(tqdm(loader, desc="Iteration", mininterval=30)):
    pred,batch = model(batch_data)

在这里,模型输出的batch就是每个节点对应的图id,源于batch_data.batch,注意这里尽管合并了,但是,分开的每一个gpu的batch都是从0开始,但返回后合并的batch并不会自动地重新更新batch,所以有可能会出现[0,0,0,1,1,0,0,1,1,1,2,2,2]这种情况,需要手动地更新一下。

posted @ 2023-01-11 13:26  阿莱慢慢来  阅读(668)  评论(0)    收藏  举报