如何在PyTorch中避免CPU循环以处理输入依赖前序输出的场景?
如何在PyTorch中避免CPU循环以处理输入依赖前序输出的场景?
嘿,我完全懂你现在的困扰——这种每一步输入依赖上一步输出的自回归场景,用Python层面的for循环跑确实会把GPU利用率拉垮,毕竟每一步都要在Python解释器和CUDA内核之间来回切换,太浪费算力了。下面给你几个实用的解决方案,既能保留你需要的功能逻辑,又能把GPU利用率拉回正常水平:
方法1:用TorchScript编译循环(兼容PyTorch 1.0+)
TorchScript可以把你的Python循环转换成PyTorch的内部高效表示,直接在GPU端执行整个循环流程,彻底减少Python层面的开销。你只需要给模型的forward方法加上一个简单的装饰器就行:
import torch from typing import Optional, Tuple class Traveller(torch.nn.Module): def __init__( self, num_inputs: int, hidden_size: int, num_layers: int, dropout: float, ): super().__init__() # 修正了你原来的父类名称错误 self.lstm_layers = torch.nn.LSTM( input_size=num_inputs + 2, hidden_size=hidden_size, num_layers=num_layers, dropout=dropout, batch_first=True ) self.output_layer = torch.nn.Linear( in_features=hidden_size, out_features=2 ) self.activation = torch.nn.Tanh() @torch.jit.script_method # 用TorchScript编译整个forward函数 def forward( self, inputs: torch.Tensor, ) -> torch.Tensor: assert (inputs.shape[1] == 1), "window should always be 1 since we're not using sequences" hidden_state: Optional[Tuple[torch.Tensor, torch.Tensor]] = None prev_position = torch.zeros( (1, 1, 2), device=inputs.device, dtype=inputs.dtype ) positions_list = [] seq_len = inputs.size(0) # 这个循环现在会被编译成高效的GPU代码,不再需要Python逐步调度 for t in range(seq_len): current_inputs = inputs[t : t + 1] input_data = torch.cat( [prev_position, current_inputs], dim=2, ) lstm_out, hidden_state = self.lstm_layers(input_data, hidden_state) out = self.output_layer(lstm_out) position = self.activation(out) prev_position = position positions_list.append(position) positions = torch.cat(positions_list, dim=0).squeeze(dim=1) assert positions.shape == (inputs.shape[0], 2) # 修正了原来的断言错误 return positions
加上@torch.jit.script_method后,PyTorch会把整个forward函数(包括内部的循环)编译成独立于Python的机器码,循环会直接在GPU端连续执行,完全避免了Python和CUDA之间的频繁交互,GPU利用率会立刻回升。
方法2:用PyTorch 2.0+的torch.scan(推荐)
如果你用的是PyTorch 2.0及以上版本,torch.scan是解决这种递推场景的最优解。它是官方专门为“每一步依赖前一步输出”的序列计算设计的算子,内部做了极致的优化,代码也更简洁:
import torch from typing import Optional, Tuple class Traveller(torch.nn.Module): def __init__( self, num_inputs: int, hidden_size: int, num_layers: int, dropout: float, ): super().__init__() self.lstm_layers = torch.nn.LSTM( input_size=num_inputs + 2, hidden_size=hidden_size, num_layers=num_layers, dropout=dropout, batch_first=True ) self.output_layer = torch.nn.Linear( in_features=hidden_size, out_features=2 ) self.activation = torch.nn.Tanh() # 定义单步计算的函数:输入前一步状态+当前输入,输出新状态+当前输出 def step_fn(self, state: Tuple[Optional[Tuple[torch.Tensor, torch.Tensor]], torch.Tensor], current_input: torch.Tensor): hidden_state, prev_position = state # 拼接当前输入和上一步位置 input_data = torch.cat([prev_position, current_input.unsqueeze(0).unsqueeze(0)], dim=2) # 执行LSTM和输出计算 lstm_out, new_hidden = self.lstm_layers(input_data, hidden_state) out = self.output_layer(lstm_out) new_position = self.activation(out) # 返回新状态和当前位置 return (new_hidden, new_position), new_position def forward( self, inputs: torch.Tensor, ) -> torch.Tensor: assert (inputs.shape[1] == 1), "window should always be 1 since we're not using sequences" # 初始化初始状态:(LSTM隐藏状态, 初始位置) initial_state = ( None, torch.zeros((1, 1, 2), device=inputs.device, dtype=inputs.dtype) ) # 将输入序列拆分为单个时间步的张量 inputs_seq = inputs.unbind(0) # 用torch.scan自动执行整个递推过程 _, positions = torch.scan(self.step_fn, initial_state, inputs_seq) # 整理输出形状 positions = torch.cat(positions, dim=0).squeeze(dim=1) assert positions.shape == (inputs.shape[0], 2) return positions
torch.scan会自动把你的步进函数应用到整个序列上,内部用高效的CUDA实现,完全不需要Python循环,代码也更符合函数式编程的风格,可读性和维护性都更好。
额外的小修正
你原来的代码里有两个小错误,已经在上面的示例中修正了:
super(DeepStatefulMVO, self).__init__()应该改成super().__init__()(对应你的模型类名Traveller);- 最后的断言
positions.shape == (inputs.shape[0], inputs.shape[2])是错误的,因为你的输出是2维的位置坐标,所以应该断言positions.shape == (inputs.shape[0], 2)。
备注:内容来源于stack exchange,提问作者quant
相关产品推荐
相关产品推荐

