PyTorch中同图深度多层并行运行方案咨询——以ROCKET论文实现场景为例
解决PyTorch中同深度多卷积层并行运行的问题
正好我之前复现ROCKET论文的时候也踩过这个串行循环的坑——10000个卷积挨个跑实在太磨人了!下面给你几个实用的方案,从最高效到最灵活的都有:
方案1:合并卷积核为单一大卷积层(最优,适用于同参数卷积)
如果你的所有卷积层都有相同的in_channels、kernel_size和stride,直接把它们合并成多输出通道的卷积层是最快的。PyTorch的Conv1d本身就支持一次性输出多个通道,本质上就是同时运行多个卷积核:
class MyModel(pl.LightningModule): def __init__(self, num_kernels=10000, in_channels=1, kernel_size=7): super().__init__() # 直接创建输出通道为10000的卷积层 self.big_conv = nn.Conv1d( in_channels=in_channels, out_channels=num_kernels, kernel_size=kernel_size, bias=True # 按ROCKET需求决定是否加bias ) # 手动随机初始化权重(ROCKET要求参数随机且不更新) nn.init.normal_(self.big_conv.weight, mean=0, std=1) nn.init.zeros_(self.big_conv.bias) # 冻结参数,避免训练时更新 self.big_conv.requires_grad_(False) def forward(self, x): # 一次卷积直接得到10000个特征图 features = self.big_conv(x) # shape: [batch_size, 10000, seq_len - kernel_size + 1] # 后续特征提取(比如ROCKET的PPV和Max) ppv = (features > 0).float().mean(dim=-1) max_val = features.max(dim=-1)[0] return torch.cat([ppv, max_val], dim=1)
这种方式完全利用了PyTorch的CUDA并行优化,没有Python循环的开销,速度是串行的几十倍甚至上百倍。
方案2:用torch.vmap实现自动并行(灵活,适用于不同参数的卷积)
如果你的卷积层有不同的kernel_size(ROCKET里确实会随机不同核大小),没法合并成单一卷积层,那可以用PyTorch 2.0+推出的torch.vmap——它能自动把循环中的张量操作并行化,而且不需要改太多原有代码:
import torch import torch.nn as nn import pytorch_lightning as pl class MyModel(pl.LightningModule): def __init__(self, num_kernels=10000, in_channels=1): super().__init__() self.convs = nn.ModuleList() for _ in range(num_kernels): # 随机生成不同的kernel_size(比如ROCKET里的2-9) kernel_size = torch.randint(2, 10, (1,)).item() conv = nn.Conv1d( in_channels=in_channels, out_channels=1, kernel_size=kernel_size, bias=True ) # 随机初始化并冻结参数 nn.init.normal_(conv.weight, mean=0, std=1) nn.init.zeros_(conv.bias) conv.requires_grad_(False) self.convs.append(conv) def forward(self, x): # 定义单个卷积的处理函数 def apply_conv(conv, input_x): return conv(input_x) # 使用vmap并行应用所有卷积层 # 注意:如果输出特征图长度不一致,需要先padding到相同长度或用自适应池化统一 features = torch.vmap(apply_conv, in_dims=(0, None))(self.convs, x) # 调整维度方便后续处理:[10000, batch_size, 1, seq_len] → [batch_size, 10000, seq_len] features = features.squeeze(2).transpose(0, 1) # 后续特征提取 ppv = (features > 0).float().mean(dim=-1) max_val = features.max(dim=-1)[0] return torch.cat([ppv, max_val], dim=1)
vmap会把ModuleList里的卷积层当成一个批次,自动在后台并行计算,避开了Python循环的GIL限制。如果核大小不同,记得先统一特征图尺寸。
方案3:分组批量处理(折中方案)
如果vmap的显存占用太高,你可以把10000个卷积分成若干组(比如每组100个),每组用方案1合并成一个卷积层,分组计算再合并结果。这样既兼顾效率,又能控制显存使用:
class MyModel(pl.LightningModule): def __init__(self, num_kernels=10000, in_channels=1, group_size=100): super().__init__() self.groups = nn.ModuleList() num_groups = (num_kernels + group_size - 1) // group_size for g in range(num_groups): # 每组内用相同参数(也可以组内用不同参数,再结合vmap) kernel_size = torch.randint(2, 10, (1,)).item() current_group_size = min
相关产品推荐
相关产品推荐

