Hexo

凡事预则立,不预则废


  • Home

  • Tags

  • Archives

  • Navigation

  • Search

PyTorch——gather函数使用

  • 参考链接:
    • 一篇浅显易懂的博客:PyTorch中的高级索引方法——gather详解

gather函数形式

  • 包含 torch.gather 和 tensor.gather 两种形式,基本思路等价,他们的函数签名如下

    1
    2
    torch.gather(input, dim, index, *, sparse_grad=False, out=None)
    tensor.gather(dim, index, *, sparse_grad=False, out=None)
    • 注:在 PyTorch 的函数签名中,* 是 Python 语法中的一个特殊标记,用于表示强制关键字参数(keyword-only arguments)。这意味着在 * 之后的所有参数(如 sparse_grad 和 out)必须通过关键字(即使用参数名)来传递,而不能通过位置参数传递
  • 参数解释

    • input (Tensor): 输入张量
    • dim (int): 沿着哪个维度进行收集
    • index (LongTensor): 索引张量,包含要收集的元素的索引
    • sparse_grad (bool, 可选): 如果为True,梯度将是稀疏张量
    • out (Tensor, 可选): 输出张量(若该值不为 None,则会将返回值存储到 out 引用中,此时 out 和 返回值 是同一个对象)
  • 特别说明:gather 和普通的矩阵索引操作一样,操作支持反向传播


基本原理

  • 对于 3D 张量,gather操作可以表示为:

    1
    2
    3
    out[i][j][k] = input[index[i][j][k]][j][k]  # dim=0
    out[i][j][k] = input[i][index[i][j][k]][k] # dim=1
    out[i][j][k] = input[i][j][index[i][j][k]] # dim=2
  • 特别说明(记忆方法)

    • index[i][j][k] 用于索引 out[i][j][k]
    • dim=0 :index 指定行号(第0维) ,列号(第1维)和通道(第2维)保持不变
    • dim=1 :index 指定列号(第1维) ,行号(第0维)和通道(第2维)保持不变
    • dim=2 :index 指定通道(第2维) ,行号(第0维)和列号(第1维)保持不变
    • 特别地:输出形状 = index 形状

代码示例

  • 一维张量示例

    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    import torch

    # 基本用法
    input_tensor = torch.tensor([10, 20, 30, 40, 50])
    index_tensor = torch.tensor([0, 2, 4, 1])
    result = torch.gather(input_tensor, dim=0, index=index_tensor)
    print(result) # tensor([10, 30, 50, 20])

    # 使用tensor.gather方法
    result2 = input_tensor.gather(0, index_tensor)
    print(result2) # tensor([10, 30, 50, 20])
  • 二维张量示例

    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    12
    13
    14
    15
    16
    17
    18
    19
    20
    21
    22
    23
    24
    # 二维张量
    input_2d = torch.tensor([[1, 2, 3],
    [4, 5, 6],
    [7, 8, 9]])

    # 沿着dim=0收集(按行收集)
    index_2d = torch.tensor([[0, 1, 2],
    [2, 0, 1]])
    result_dim0 = torch.gather(input_2d, dim=0, index=index_2d)
    print("dim=0 结果:")
    print(result_dim0)
    # tensor([[1, 5, 9],
    # [7, 2, 6]])

    # 沿着dim=1收集(按列收集)
    index_2d_col = torch.tensor([[0, 2],
    [1, 0],
    [2, 1]])
    result_dim1 = torch.gather(input_2d, dim=1, index=index_2d_col)
    print("dim=1 结果:")
    print(result_dim1)
    # tensor([[1, 3],
    # [5, 4],
    # [9, 8]])

一些实际应用场景

  • 获取最大值,获取对应标签等实例:
    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    12
    13
    14
    15
    16
    17
    18
    19
    20
    21
    22
    # 场景1: 获取每行的最大值索引对应的值
    scores = torch.tensor([[0.1, 0.8, 0.3],
    [0.6, 0.2, 0.9],
    [0.4, 0.7, 0.1]])

    # 获取每行最大值的索引
    max_indices = torch.argmax(scores, dim=1, keepdim=True)
    print("最大值索引:", max_indices) # tensor([[1], [2], [1]])

    # 使用gather获取最大值
    max_values = torch.gather(scores, dim=1, index=max_indices)
    print("最大值:", max_values) # tensor([[0.8], [0.9], [0.7]])

    # 场景2: 根据标签获取对应的预测概率
    predictions = torch.tensor([[0.2, 0.3, 0.5],
    [0.1, 0.8, 0.1],
    [0.6, 0.2, 0.2]])
    labels = torch.tensor([2, 1, 0]) # 真实标签

    # 获取每个样本对应标签的预测概率
    label_probs = torch.gather(predictions, dim=1, index=labels.unsqueeze(1))
    print("标签概率:", label_probs) # tensor([[0.5], [0.8], [0.6]])

注意事项

  • 索引值必须在 [0, input.size(dim)) 范围内
  • 除了指定的dim维度外,input 和 index 的其他维度大小必须相同
  • 输出张量的形状与 index 张量相同

PyTorch——torch.no_grad的用法


整体说明

  • 在 PyTorch 中,torch.no_grad()可用作装饰器 @torch.no_grad() 或上下文管理器 with torch.no_grad()(两者形式不同,但作用相同),用于禁用梯度计算
  • 如果 PyTorch 版本 >= 1.9,可以考虑使用 torch.inference_mode() 来替代 torch.no_grad(),以获得更好的性能

torch.no_grad()的作用

  • torch.no_grad() 的主要作用是临时关闭自动求导机制(autograd)。在被装饰的函数或代码块中,所有涉及张量的操作都不会构建计算图(computation graph),从而节省内存和计算资源:
    • 自动求导机制 :PyTorch 默认会记录张量操作的历史信息(即计算图),以便支持反向传播(backward())来计算梯度
    • 关闭梯度计算 :在推理阶段或其他不需要梯度的场景下,关闭自动求导可以减少内存占用,提高运行效率

使用场景

模型推理(Inference)

  • 在推理阶段,我们只需要前向传播(forward pass),而不需要计算梯度。因此,可以使用 @torch.no_grad() 来优化性能
    1
    2
    3
    4
    5
    6
    7
    8
    9
    @torch.no_grad()
    def evaluate_model(model, test_loader):
    model.eval() # 设置模型为评估模式,改回训练模式可以调用 model.train()
    total_loss = 0
    for data, target in test_loader:
    output = model(data)
    loss = loss_function(output, target)
    total_loss += loss.item()
    return total_loss

更新模型参数时不计算梯度

  • 在某些情况下,我们需要手动更新模型参数(例如权重剪枝、量化等),但不希望这些操作影响梯度计算
    1
    2
    3
    4
    @torch.no_grad()
    def update_weights(model):
    for param in model.parameters():
    param.add_(1.0) # 在参数上加 1,不会记录到计算图中

计算评估指标时不计算梯度

  • 当计算评估指标(如准确率、F1 分数等)时,不需要梯度计算
    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    @torch.no_grad()
    def compute_accuracy(model, data_loader):
    correct = 0
    total = 0
    for inputs, labels in data_loader:
    outputs = model(inputs)
    _, predicted = torch.max(outputs, 1)
    total += labels.size(0)
    correct += (predicted == labels).sum().item()
    return correct / total

附录:装饰器和上下文管理器的示例

作为装饰器

  • 装饰整个函数,使其在执行期间禁用梯度计算
    1
    2
    3
    @torch.no_grad()
    def inference(model, input_data):
    return model(input_data)

作为上下文管理器

  • 仅在特定代码块中禁用梯度计算
    1
    2
    3
    def inference(model, input_data):
    with torch.no_grad():
    return model(input_data)

附录:推理场景 torch.inference_mode() 的使用

  • 从 PyTorch 1.9 开始,引入了 torch.inference_mode(),它是 torch.no_grad() 的更高效替代品,专门用于推理阶段。与 torch.no_grad() 相比:
    • 性能更高 :torch.inference_mode() 会跳过一些额外的检查,进一步提升性能
    • 不可嵌套 :torch.inference_mode() 不能像 torch.no_grad() 那样嵌套使用
    • 推荐使用 :如果只用于推理,建议优先使用 torch.inference_mode()
  • 示例:
    1
    2
    3
    4
    @torch.inference_mode()
    def evaluate_model(model, test_loader):
    model.eval()
    ...
1…139140141…352
San Ye

San Ye

Stay Hungry. Stay Foolish.

704 posts
53 tags
© 2026 San Ye
Powered by Hexo
|
Theme — NexT.Gemini v5.1.4