Advertisement

PyTorch Basics torch.sum()

阅读量:

torch.sum()对输入的tensor数据的某一维度求和,一共两种用法

1.PyTorch中使用torch.sum函数来计算输入张量的总值。

2.通过调用torch.sum函数可对输入张量进行求和运算;该操作需指定dim参数以确定求和维度;同时可设置keepdim参数以决定输出是否保持与输入相同的维度结构;此外还需要指定数据类型的设置以确保运算结果符合预期需求;最终运算结果将返回一个Tensor对象。

  • input:接受一个张量作为输入
  • dim:指定求和的维度(可指定为一个整数或列表)
  • keepdim:在求和操作后会缩减该维度的空间大小;因此会去除该维上的元素数量(即长度);如果希望该维度在求和后仍然存在,则需要设置keepdim=True

举例说明:

复制代码
    import torch
    a = torch.ones((2, 3))
    print(a)
    
    a1 = torch.sum(a, dim=(0, 1))
    a2 = torch.sum(a, dim=0)
    a3 = torch.sum(a, dim=1)
    
    
    print(a1)
    print(a2)
    print(a3)
复制代码
    tensor(

全部评论 (0)

还没有任何评论哟~