【笔记】unsqueeze

unsqueeze是 PyTorch 中的一个方法,用于在指定位置插入一个维度为 1 的新维度。这个操作对于调整张量的形状非常有用,尤其是在需要匹配特定维度要求(例如模型输入或 `torchvision.utils.make_grid` 函数的要求)时。

理解 unsqueeze

假设你有一个形状为 2, 3 的二维张量:

python 复制代码
tensor = torch.randn(2, 3)
print(tensor.shape)  # 输出: torch.Size([2, 3])

如果你想要把这个张量变成三维的,比如形状变为 1, 2, 3,就可以使用 unsqueeze方法。你可以指定在哪一个维度上增加新的维度(从0开始计数)。

  • 在第0维增加新维度:tensor.unsqueeze(0)

  • 在第1维增加新维度:tensor.unsqueeze(1)

例如:

python 复制代码
# 在第0维增加新维度
new_tensor_0 = tensor.unsqueeze(0)
print(new_tensor_0.shape)  # 输出: torch.Size([1, 2, 3])

# 在第1维增加新维度
new_tensor_1 = tensor.unsqueeze(1)
print(new_tensor_1.shape)  # 输出: torch.Size([2, 1, 3])

应用场景

在我的代码上下文中,unsqueeze主要用于确保传入 `make_grid` 的张量具有正确的维度make_grid 需要输入是一个四维张量 (B, C, H, W),其中:

  • B表示批量大小(Batch Size)

  • C表示通道数(Channels)

  • H表示高度(Height)

  • W表示宽度(Width)

例如,如果有一个形状为 7, 224, 224的 mask 张量(即它只有三个维度),而你需要将其转换为四个维度的形式以满足 make_grid 的要求,你可以使用 unsqueeze(1)来在第二个维度(通道维度)上增加一个新的维度:

python 复制代码
masks = masks.unsqueeze(1)  # 将 [7, 224, 224] 转换为 [7, 1, 224, 224]

这样,mask 的形状就变成了 7, 1, 224, 224,符合 make_grid的输入要求。

相关推荐
学计算机的计算基4 小时前
操作系统内存管理全解:虚拟内存、页表、COW、malloc、OOM一篇搞定
java·笔记·算法
AOwhisky4 小时前
Python 学习笔记(第五期)——组合数据类型:列表、元组、集合与字典精讲——核心知识点自测与详解
开发语言·笔记·python·学习·云计算
砚凝霜4 小时前
软考网络工程师|第 2 章 信道延迟、传输介质、数据编码、数字调制、PCM 完整备考笔记
网络·笔记·pcm
遇乐的果园17 小时前
前端学习笔记-vue加载渲染优化
前端·笔记·学习
遇乐的果园18 小时前
前端学习笔记-vue状态管理优化
前端·笔记·学习
摇滚侠19 小时前
Java 全栈开发实战教程 课程笔记 29-33
笔记
茯苓gao19 小时前
嵌入式开发笔记:EtherCAT协议从硬件到软件完整配置指南——从零搭建一套EtherCAT通信系统
笔记·嵌入式硬件·学习
whyTeaFo20 小时前
GAMES101: Lecture 9: Shading 3(Texture Mapping cont.) ppt笔记
笔记
chase。21 小时前
【学习笔记】PointWorld:迈向通用机器人操控的3D世界模型
笔记·学习·机器人
星恒随风1 天前
C++ STL 栈详解:stack 的使用、经典题目与简单模拟实现
开发语言·数据结构·c++·笔记·学习