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]这种情况,需要手动地更新一下。

浙公网安备 33010602011771号