Tensorflow2.0笔记 - 修改形状和维度

本次笔记主要使用reshape,transpose,expand_dim,和squeeze对tensor的形状和维度进行操作。

复制代码
import tensorflow as tf
import numpy as np

tf.__version__

#tensor的shape和维数获取
#假设下面这个tensor表示4张28*28*3的图片
tensor = tf.random.uniform([4,28,28,3], minval=0, maxval=10, dtype=tf.int32)
print("tensor.shape:", tensor.shape)
print("tensor.ndim:", tensor.ndim)

#reshape成一个三维的tensor,将行和列的信息去掉,只保留pixel概念
print("=======reshape([4,28*28,3].shape=========\n", tf.reshape(tensor, [4,28*28,3]).shape)
#reshape里的参数中可以出现一个-1,表示自动计算省略掉的维度的大小
#还是上面的例子,将行和列的信息去掉,只保留pixel的概念
print("=======reshape([4,-1,3].shape=========\n", tf.reshape(tensor, [4,-1,3]).shape)
#将图片的行和列信息和RGB通道信息去掉,图片数据作为一个整体,等价于tf.reshape(tensor, [4, 28*28*3])
print("=======reshape([4,-1].shape=========\n", tf.reshape(tensor, [4,-1]).shape)

#transpose进行转置操作,会修改tensor的数据布局
tensor = tf.random.uniform([4,3,2,1], minval=0, maxval=9, dtype=tf.int32)
print(tensor.shape,tensor.ndim)
print(tensor)

#不带参数,表示整体转置,对所有维度进行转置
transpose = tf.transpose(tensor)
print("========Transpose without arg:", transpose.shape)
print(transpose)
#带参数,给出perm参数,表示原来的维度放到哪个位置
#第0个和第1个维度保留,交换最后两个维度
transpose = tf.transpose(tensor, perm=[0,1,3,2])
print("========Transpose by arg:", transpose.shape)
print(transpose)

#transpose的一个应用案例
#pytorch中,图片信息一般以[b,c,h,w]来表示,b表示batch数量,c表示像素通道数量,h,w表示图片的高度和宽度
#tensorflow中,图片信息一般以[b,h,w,c]来表示
#可以使用transpose进行pytorch和tensorflow格式的互转
#下面的tensor按照pytorch格式理解,两张5*5*3的图片
tensor = tf.random.uniform([2,3,5,5], minval=0, maxval=9, dtype=tf.int32)
print("=====PYTORCH data=====\n", tensor)
#通过transpose转换为tensorflow格式
transpose = tf.transpose(tensor, [0,2,3,1])
print("=====TENSORFLW data====\n", transpose)

#增加(expand)或减少(squeeze)维度
#假设下面的tensor表示4个班级,10个学生,5门科目的成绩
tensor = tf.random.normal([4,10,5])

#现在我们要增加一个学校的维度,使用expand_dims,会在指定axis的前面增加一个维度
#axis表示要在那个维度前面增加
expanded = tf.expand_dims(tensor, axis=0)
print("Expanded at dim0:", expanded.shape)

#在5门科目成绩维度前增加一个维度
expanded = tf.expand_dims(tensor, axis=2)
print("Expanded at dim2:", expanded.shape)

#在5门科目成绩维度后面增加一个维度
expanded = tf.expand_dims(tensor, axis=3)
print("Expanded at dim3:", expanded.shape)

#axis为负数的时候,和numpy索引给-1的情况是类似的,需要注意的是此时会在指定axis的后面增加一个维度
#在5门科目成绩维度前增加一个维度
expanded = tf.expand_dims(tensor, axis=-2)
print("Expanded at dim2:", expanded.shape)
#在最前面增加一个维度
expanded = tf.expand_dims(tensor, axis=-4)
print("Expanded at dim0:", expanded.shape)

#减少维度,仅用于去掉shape=1的维度,如果指定要去掉的维度shape大于1会报错
tensor = tf.zeros([1,2,1,1,3])
print("tensor.shape:", tensor.shape)
#上面的tensor,只有1个维度的位置可以去掉
squeezed = tf.squeeze(tensor)
print("Squeezed:", squeezed.shape)
#指定某个axis进行squeeze
squeezed = tf.squeeze(tensor, axis=0)
print("Squeezed:", squeezed.shape)
#axis为负数的情况
squeezed = tf.squeeze(tensor, axis=-2)
print("Squeezed:", squeezed.shape)

运行结果:

相关推荐
FakeOccupational2 小时前
【github 有趣项目】OpenPLC: 支持通用硬件的开源 PLC 软件平台‌
笔记
circuitsosk2 小时前
NL2SQL在工业级场景下的精度优化:Schema Linking + 动态Few-shot实战
人工智能·python·sql·大模型·nl2sql
9000AI2 小时前
9000AI如何工业化生产流量?高质量规模生产与矩阵化饱和覆盖
人工智能
字节数据平台3 小时前
iDA:从 ChatBI 到专业数据分析助手的演进之路
大数据·人工智能·机器学习·数据分析
warpdrivelabs3 小时前
Codex 开源 harness 全面了解
开发语言·人工智能
mit6.8243 小时前
微软如何交付企业级Agent
人工智能
Mininglamp_27183 小时前
明略科技携手海康机器人亮相世界机器人大会,以“Agent+具身“联合进入商业机器人场景
人工智能·科技·机器人·开源·agent·ai agent
2401_894915533 小时前
GEO 优化源码全解析:从搜索引擎到 AI 引擎的底层改写逻辑
java·服务器·前端·数据库·人工智能·分布式·搜索引擎
MobotStone4 小时前
从“听得懂”到“干得了”:工业大模型落地工厂的三层进化路线
人工智能
wangruofeng5 小时前
GLM-5.3-Flash 发布:追平 Opus 4.8 的智力,1/40 的价格,跑在国产芯片上
人工智能·aigc·chatglm (智谱)