Pytorch-tensor维度的扩展,挤压,扩张
数据本身不发生改变,数据的访问方式发生了改变
1.维度的扩展
函数:
unsqueeze()
# a是一个4维的
a = torch.randn(4, 3, 28, 28)
print('a.shape\n', a.shape)
print('\n维度扩展(变成5维的):')
print('第0维前加1维')
print(a.unsqueeze(0).shape)
print('第4维前加1维')
print(a.unsqueeze(4).shape)
print('在-1维前加1维')
print(a.unsqueeze(-1).shape)
print('在-4维前加1维')
print(a.unsqueeze(-4).shape)
print('在-5维前加1维')
print(a.unsqueeze(-5).shape)
输出结果
a.shape
torch.Size([4, 3, 28, 28])
维度扩展(变成5维的):
第0维前加1维
torch.Size([1, 4, 3, 28, 28])
第4维前加1维
torch.Size([4, 3, 28, 28, 1])
在-1维前加1维
torch.Size([4, 3, 28, 28, 1])
在-4维前加1维
torch.Size([4, 1, 3, 28, 28])
在-5维前加1维
torch.Size([1, 4, 3, 28, 28])
注意,第5维前加1维,就会出错
# print(a.unsqueeze(5).shape)
# Errot:Dimension out of range (expected to be in range of -5, 4], but got 5)
连续扩维:
unsqueeze()
# b是一个1维的
b = torch.tensor([1.2, 2.3])
print('b.shape\n', b.shape)
print()
# 0维之前插入1维,变成1,2]
print(b.unsqueeze(0))
print()
# 1维之前插入1维,变成2,1]
print(b.unsqueeze(1))
# 连续扩维,然后再对某个维度进行扩张
print(b.unsqueeze(1).unsqueeze(2).unsqueeze(0).shape)
输出结果
b.shape
torch.Size([2])
tensor([[1.2000, 2.3000]])
tensor([[1.2000],
[2.3000]])
torch.Size([1, 2, 1, 1])
2.挤压维度
函数:
squeeze()
# 挤压维度,只会挤压shape为1的维度,如果shape不是1的话,当前值就不会变
c = torch.randn(1, 32, 1, 2)
print(c.shape)
print(c.squeeze(0).shape)
print(c.squeeze(1).shape) # shape不是1,不会变
print(c.squeeze(2).shape)
print(c.squeeze(3).shape) # shape不是1,不会变
输出结果
torch.Size([1, 32, 1, 2])
torch.Size([32, 1, 2])
torch.Size([1, 32, 1, 2])
torch.Size([1, 32, 2])
torch.Size([1, 32, 1, 2])
3.维度扩张
函数1:
expand():扩张到多少,
# shape的扩张
# expand():对shape为1的进行扩展,对shape不为1的只能保持不变,因为不知道如何变换,会报错
d = torch.randn(1, 32, 1, 1)
print(d.shape)
print(d.expand(4, 32, 14, 14).shape)
输出结果
torch.Size([1, 32, 1, 1])
torch.Size([4, 32, 14, 14])
函数2:
repeat()方法,扩张多少倍
d=torch.randn([1,32,4,5])
print(d.shape)
print(d.repeat(4,32,2,3).shape)
输出结果
torch.Size([1, 32, 4, 5])
torch.Size([4, 1024, 8, 15])
Pytorch-tensor维度的扩展,挤压,扩张的更多相关文章
- Pytorch Tensor 维度的扩充和压缩
维度扩展 x.unsqueeze(n) 在 n 号位置添加一个维度 例子: import torch x = torch.rand(3,2) x1 = x.unsqueeze(0) # 在第一维的位置 ...
- pytorch tensor 维度理解.md
torch.randn torch.randn(*sizes, out=None) → Tensor(张量) 返回一个张量,包含了从标准正态分布(均值为0,方差为 1)中抽取一组随机数,形状由可变参数 ...
- pytorch 中改变tensor维度的几种操作
具体示例如下,注意观察维度的变化 #coding=utf-8 import torch """改变tensor的形状的四种不同变化形式""" ...
- PyTorch中的C++扩展
今天要聊聊用 PyTorch 进行 C++ 扩展. 在正式开始前,我们需要了解 PyTorch 如何自定义module.这其中,最常见的就是在 python 中继承torch.nn.Module,用 ...
- [TensorFlow]Tensor维度理解
http://wossoneri.github.io/2017/11/15/[Tensorflow]The-dimension-of-Tensor/ Tensor维度理解 Tensor在Tensorf ...
- tensorflow中的函数获取Tensor维度的两种方法:
获取Tensor维度的两种方法: Tensor.get_shape() 返回TensorShape对象, 如果需要确定的数值而把TensorShape当作list使用,肯定是不行的. 需要调用Tens ...
- Pytorch 张量维度
Tensor类的成员函数dim()可以返回张量的维度,shape属性与成员函数size()返回张量的具体维度分量,如下代码定义了一个两行三列的张量: f = torch.randn(2, 3) pri ...
- Pytorch Tensor 常用操作
https://pytorch.org/docs/stable/tensors.html dtype: tessor的数据类型,总共有8种数据类型,其中默认的类型是torch.FloatTensor, ...
- Pytorch Tensor, Variable, 自动求导
2018.4.25,Facebook 推出了 PyTorch 0.4.0 版本,在该版本及之后的版本中,torch.autograd.Variable 和 torch.Tensor 同属一类.更确切地 ...
- tensor维度变换
维度变换是tensorflow中的重要模块之一,前面mnist实战模块我们使用了图片数据的压平操作,它就是维度变换的应用之一. 在详解维度变换的方法之前,这里先介绍一下View(视图)的概念.所谓Vi ...
随机推荐
- hadoop 3.3.5伪分布式集群部署以及遇到的问题解决
hadoop包下载 https://archive.apache.org/dist/hadoop/common/ 安装好jdk并配置环境变量 下载hadoop压缩包并放至 /data/hadoop目录 ...
- 1、eureka的注册流程
客户端注册到服务端是通过http请求的 涉及到多级缓存 register注册表 源码精髓:多级缓存设计思想 在拉取注册表的时候: 首先从ReadOnlyCacheMap里查缓存的注册表. 若没有,就找 ...
- Codeforces Round 923 (Div. 3)(A~F)
目录 A B C D E F A #include <bits/stdc++.h> #define int long long #define rep(i,a,b) for(int i = ...
- RIPEMD算法:多功能哈希算法的瑰宝
一.RIPEMD算法的起源与历程 RIPEMD(RACE Integrity Primitives Evaluation Message Digest)算法是由欧洲研究项目RACE发起,由Hans D ...
- WPF之资源
目录 WPF对象资源的定义和查找 动态.静态使用资源 向程序添加二进制资源 字符串资源 非字符串资源 使用Pack URI路径访问二进制资源 WPF不但支持程序级的传统资源,同时还推出了独具特色的对象 ...
- GoFrame 优化接口的错误码和异常的思路
前言 你是否想在使用 GoFrame 的过程中,拥有一个能打印异常堆栈,能自定义响应状态码,能统一处理响应数据的接口.如果你回答是,那么,请耐心看完本文,或许会对你有所启发.若文中由表达不当之处,恳请 ...
- 单麦克风AI降噪模块及解决方案
前记 随着以AI为核心的智能设备的广泛发展,语音这个非常重要的入口一直是很多厂商争夺的市场.作为音频采集的前端设备,能采集到的距离远,清晰度高,无噪声的信号是一个非常重要的能力.这样就对音频前端降 ...
- day02-SpringMVC映射请求数据
SpringMVC映射请求数据 1.获取参数值 在开发中,如何获取到 http://xxx/url?参数名1=参数值1&参数名2=参数值2 中的参数? 之前的案例中我们知道:提交的url的参数 ...
- KTL 一个支持C++14编辑公式的K线技术工具平台
K,K线,Candle蜡烛图. T,技术分析,工具平台 L,公式Language语言使用c++14,Lite小巧简易. 项目仓库:https://github.com/bbqz007/KTL 国内仓库 ...
- 三维模型OBJ格式轻量化的跨平台兼容性问题分析
三维模型OBJ格式轻量化的跨平台兼容性问题分析 三维模型的OBJ格式轻量化在跨平台兼容性方面具有重要意义,可以确保模型在不同平台和设备上的正确加载和渲染.本文将分析OBJ格式轻量化的跨平台兼容性技术, ...