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

如何在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循环,代码也更符合函数式编程的风格,可读性和维护性都更好。

额外的小修正

你原来的代码里有两个小错误,已经在上面的示例中修正了:

  1. super(DeepStatefulMVO, self).__init__()应该改成super().__init__()(对应你的模型类名Traveller);
  2. 最后的断言positions.shape == (inputs.shape[0], inputs.shape[2])是错误的,因为你的输出是2维的位置坐标,所以应该断言positions.shape == (inputs.shape[0], 2)。

备注:内容来源于stack exchange,提问作者quant

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 20:08:00