pytorch张量的new_zeros方法介绍

在 PyTorch 中,Tensor.new_zeros 是一种用于创建与现有张量形状或设备匹配的新张量的方法。该方法生成一个全为零的张量,且其数据类型、设备等属性与调用它的张量一致,除非另行指定。


new_zeros 方法的语法

复制代码
Tensor.new_zeros(size, *, dtype=None, device=None, requires_grad=False)

参数说明

  • size (tuple)

    指定新张量的形状。例如 (2, 3) 表示创建一个形状为 2x3 的张量。

  • dtype (torch.dtype, 可选)

    指定新张量的数据类型。如果未指定,将与原张量的数据类型一致。

  • device (torch.device, 可选)

    指定新张量所在的设备(如 CPU 或 GPU)。如果未指定,将与原张量所在的设备一致。

  • requires_grad (bool, 可选)

    指定新张量是否需要计算梯度(默认为 False)。


new_zeros 的特性

  • 新张量与原张量具有相同的设备默认数据类型(除非显式更改)。
  • 新张量的内容为全零。

使用示例

1. 创建与现有张量形状匹配的零张量

复制代码
import torch

x = torch.ones(2, 3, device='cuda')  # 创建一个形状为 (2, 3) 的张量
zeros = x.new_zeros((2, 3))          # 创建一个全零张量,与 x 具有相同形状和设备
print(zeros)
# 输出(在 GPU 上):
# tensor([[0., 0., 0.],
#         [0., 0., 0.]], device='cuda:0')

2. 创建具有不同形状的零张量

复制代码
x = torch.ones(4, 5)
zeros = x.new_zeros((2, 3))  # 创建一个形状为 (2, 3) 的零张量
print(zeros)
# 输出:
# tensor([[0., 0., 0.],
#         [0., 0., 0.]])

3. 指定数据类型

复制代码
x = torch.ones(3, 3, dtype=torch.float32)
zeros = x.new_zeros((2, 2), dtype=torch.int32)  # 显式指定数据类型
print(zeros)
# 输出:
# tensor([[0, 0],
#         [0, 0]], dtype=torch.int32)

4. 指定设备

复制代码
x = torch.ones(2, 2, device='cuda')
zeros = x.new_zeros((3, 3), device='cpu')  # 在 CPU 上创建新张量
print(zeros)
# 输出:
# tensor([[0., 0., 0.],
#         [0., 0., 0.],
#         [0., 0., 0.]])

与其他创建零张量的方法的对比

  1. torch.zeros

    zeros = torch.zeros((2, 3))

    • 独立于已有张量。
    • 需要显式指定数据类型和设备。
  • Tensor.new_zeros

    zeros = x.new_zeros((2, 3))

  • 与现有张量 x 共享设备和默认数据类型。


常见应用场景

  1. 快速创建与输入张量匹配的零张量 在深度学习中,可能需要创建与现有张量形状和设备匹配的零张量。例如,用于初始化中间结果或辅助计算。

  2. 动态操作 当输入张量的形状、设备不固定时,可以使用 new_zeros 动态生成匹配的零张量,无需手动指定设备或数据类型。


总结

Tensor.new_zeros 是一个高效、方便的方法,适合在动态模型或设备敏感的代码中使用。它避免了显式管理设备和数据类型的麻烦,有助于提高代码的简洁性和可维护性。

相关推荐
老胖闲聊3 小时前
Python Copilot【代码辅助工具】 简介
开发语言·python·copilot
Blossom.1183 小时前
使用Python和Scikit-Learn实现机器学习模型调优
开发语言·人工智能·python·深度学习·目标检测·机器学习·scikit-learn
曹勖之4 小时前
基于ROS2,撰写python脚本,根据给定的舵-桨动力学模型实现动力学更新
开发语言·python·机器人·ros2
scdifsn4 小时前
动手学深度学习12.7. 参数服务器-笔记&练习(PyTorch)
pytorch·笔记·深度学习·分布式计算·数据并行·参数服务器
DFminer4 小时前
【LLM】fast-api 流式生成测试
人工智能·机器人
lyaihao4 小时前
使用python实现奔跑的线条效果
python·绘图
郄堃Deep Traffic5 小时前
机器学习+城市规划第十四期:利用半参数地理加权回归来实现区域带宽不同的规划任务
人工智能·机器学习·回归·城市规划
ai大师5 小时前
(附代码及图示)Multi-Query 多查询策略详解
python·langchain·中转api·apikey·中转apikey·免费apikey·claude4
海盗儿5 小时前
Attention Is All You Need (Transformer) 以及Transformer pytorch实现
pytorch·深度学习·transformer
GIS小天5 小时前
AI+预测3D新模型百十个定位预测+胆码预测+去和尾2025年6月7日第101弹
人工智能·算法·机器学习·彩票