pytorch小记(六):pytorch中的clone和detach操作:克隆/复制数据 vs 共享相同数据但 与计算图断开联系

pytorch小记(六):pytorch中的clone和detach操作:克隆/复制数据 vs 共享相同数据但 与计算图断开联系


以下代码片段:

python 复制代码
self.x = x.clone().detach()  # 或 torch.tensor(x).float()

用于处理和复制张量 x,并根据需要使其与原始计算图断开联系或改变其数据类型。下面是逐部分详细解释。


1. x.clone()

  • 作用 :对张量 x 进行深拷贝,生成一个新的张量。
    • 新的张量和原始张量具有相同的数据,但存储在不同的内存空间。
    • 修改 clone() 的返回值不会影响原始张量。

示例:

python 复制代码
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = x.clone()

y[0] = 99.0
print(x)  # tensor([1., 2., 3.], grad_fn=<CloneBackward>)
print(y)  # tensor([99.,  2.,  3.])

2. x.detach()

  • 作用 :返回一个与 x 共享相同数据但 与计算图断开联系 的张量。
    • 通常用于阻止梯度计算。
    • 在神经网络中,如果你不希望某些操作影响反向传播时,会用到 detach()

示例:

python 复制代码
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = x.detach()

y[0] = 99.0  # y 的数据更改不会影响 x
print(x)  # tensor([1., 2., 3.], requires_grad=True)
print(y)  # tensor([99.,  2.,  3.])

使用场景:

detach() 在以下场景中非常有用:

  1. 阻止梯度传播:

    python 复制代码
    z = x.clone().detach()
    # z 不会参与反向传播,x 的梯度也不会受 z 的影响
  2. 保存模型状态或生成推断结果:

    python 复制代码
    with torch.no_grad():
        output = model(x)  # 临时禁用梯度计算

3. torch.tensor(x).float()

  • 作用 :将输入 x 转换为 PyTorch 张量,并将其数据类型强制为 torch.float32(默认浮点类型)。
  • 适用场景:
    • 输入可能是一个 Python 列表或 NumPy 数组时,用于将其转换为 PyTorch 张量。
    • 确保张量数据类型一致(某些模型或操作对数据类型有严格要求)。

示例:

python 复制代码
x = [[1, 2, 3], [4, 5, 6]]  # Python 列表
y = torch.tensor(x).float()  # 转为 torch.float32 类型的张量
print(y)
# tensor([[1., 2., 3.],
#         [4., 5., 6.]])

4. 两者的对比与结合

  • x.clone().detach()torch.tensor(x).float() 是不同的操作:

    1. x.clone().detach()
      • 复制一个现有张量,且与原始计算图断开。
      • 适用于 PyTorch 张量 x,不适用于列表或其他数据类型。
    2. torch.tensor(x).float()
      • 将输入转换为新的 PyTorch 张量,适用于从非张量对象(如列表、NumPy 数组)构造张量。
      • 转换过程中可以指定数据类型(如 .float())。
  • 结合使用

    如果需要复制一个张量、改变数据类型,并断开计算图,可以将两者结合:

    python 复制代码
    self.x = torch.tensor(x.clone().detach()).float()

使用场景

x.clone().detach()

  • x 是一个 PyTorch 张量,且需要:
    • 复制数据。
    • 与原始计算图断开。

torch.tensor(x).float()

  • x 是一个非 PyTorch 张量对象(如列表或 NumPy 数组),且需要:
    • 转换为 PyTorch 张量。
    • 确保数据类型为浮点型。

完整示例:

python 复制代码
import torch

# 输入张量
x = torch.tensor([[2.0, -1.0], [1.0, 1.0]], requires_grad=True)

# 使用 clone().detach()
y = x.clone().detach()
y[0, 0] = 99.0
print("x:", x)  # 原始张量不会改变
print("y:", y)  # 新张量修改了

# 使用 torch.tensor()
z = torch.tensor([[1, 2], [3, 4]]).float()
print("z:", z)  # 转换为浮点张量

总结

  • clone():深拷贝一个张量。
  • detach():断开张量与计算图的连接。
  • torch.tensor(x).float():将非张量数据转换为浮点型 PyTorch 张量。
  • 它们在不同场景下各有用途,可以单独使用或结合使用。
相关推荐
人工智能训练2 小时前
【极速部署】Ubuntu24.04+CUDA13.0 玩转 VLLM 0.15.0:预编译 Wheel 包 GPU 版安装全攻略
运维·前端·人工智能·python·ai编程·cuda·vllm
yaoming1682 小时前
python性能优化方案研究
python·性能优化
源于花海3 小时前
迁移学习相关的期刊和会议
人工智能·机器学习·迁移学习·期刊会议
码云数智-大飞3 小时前
使用 Python 高效提取 PDF 中的表格数据并导出为 TXT 或 Excel
python
DisonTangor4 小时前
DeepSeek-OCR 2: 视觉因果流
人工智能·开源·aigc·ocr·deepseek
薛定谔的猫19824 小时前
二十一、基于 Hugging Face Transformers 实现中文情感分析情感分析
人工智能·自然语言处理·大模型 训练 调优
发哥来了5 小时前
《AI视频生成技术原理剖析及金管道·图生视频的应用实践》
人工智能
biuyyyxxx5 小时前
Python自动化办公学习笔记(一) 工具安装&教程
笔记·python·学习·自动化
数智联AI团队5 小时前
AI搜索引领开源大模型新浪潮,技术创新重塑信息检索未来格局
人工智能·开源
极客数模5 小时前
【2026美赛赛题初步翻译F题】2026_ICM_Problem_F
大数据·c语言·python·数学建模·matlab