【Scala PyTorch深度学习】PyTorch On Scala 系列课程 第二章 04 :张量高级操作【AI Infra 3.0】[PyTorch Scala硕士研一课程]

PyTorch Scala 高校计算机硕士研一课程
章节 2: 高级张量操作
在前面介绍的基础张量操作之上,本章将讲解更高级的张量操作方法。对张量结构、数据类型和设备放置进行精细控制,对于有效准备数据和实现复杂的深度学习模型来说非常重要。
在本章中,你将学会:
- 使用索引和切片方法,精确选择和修改张量元素。
- 使用
view()、reshape()和permute()重构张量,而不改变其数据。 - 使用
cat()、stack()、split()和chunk()组合多个张量或将单个张量分割成不同部分。 - 使用 PyTorch 的广播机制,对形状兼容但不同的张量之间执行操作。
- 处理不同的数值数据类型(例如
float、int),并相应地转换张量类型。 - 使用
.to(device)在 CPU 内存和 GPU 加速器之间传输张量。
掌握这些操作对于处理深度学习工作流中遇到的多样化数据格式和计算要求是必要的。
张量索引与切片
访问和修改张量的特定部分是处理深度学习数据时常有的需求。无论您是需要选择单个数据点、提取一批训练样本、裁剪图像补丁,还是挑选特定特征,PyTorch都提供了强大且灵活的索引和切片机制,类似于NumPy数组中的那些,但PyTorch的机制与GPU加速和自动微分功能相集成。
基本索引
访问张量元素最直接的方式是使用标准的Python整数索引。请记住,PyTorch张量与Python列表和NumPy数组一样,使用0作为起始索引。
对于一维张量,您可以使用其索引访问元素:
import torch.*
// 创建一个一维张量
val x_1d = torch.tensor(Seq(10, 11, 12, 13, 14))
println(f"原始一维张量:\n{x_1d}")
// 访问第一个元素
val first_element = x_1d(0)
println(f"\n第一个元素 (x_1d(0)): {first_element}, 类型: {type(first_element)}")
// 访问最后一个元素
val last_element = x_1d(-1)
println(f"最后一个元素 (x_1d(-1)): {last_element}")
// 修改一个元素
x_1d(1) = 110
println(f"\n修改后的张量:\n{x_1d}")
注意,访问单个元素会返回一个包含单个值的torch.Tensor(一个0维张量或标量),而不是一个标准的Python数字,除非您使用.item()明确提取它。元素的修改是原地进行的。
对于多维张量,您需要为每个维度提供索引,并用逗号分隔:
// 创建一个二维张量 (例如,一个小矩阵)
val x_2d = torch.tensor(Seq(Seq(1, 2, 3),
Seq(4, 5, 6),
Seq(7, 8, 9)))
println(f"原始二维张量:\n{x_2d}")
// 访问第0行第1列的元素
val element_0_1 = x_2d(0, 1)
println(f"\n在 [0, 1] 的元素: {element_0_1}")
// 访问整个第一行 (索引0)
val first_row = x_2d(0) // or x_2d(0, *)
println(f"\n第一行 (x_2d(0)): {first_row}")
// 访问整个第二列 (索引1)
val second_col = x_2d(*, 1) // or x_2d(*, 1)
println(f"第二列 (x_2d(*, 1)): {second_col}")
// 修改一个元素
x_2d(1, 1) = 55
println(f"\n修改后的二维张量:\n{x_2d}")
提供的索引数量少于维数时,会沿剩余维度选择一个完整的子张量。例如,x_2d[0] 会选择整个第一行。
张量切片
切片允许您沿张量维度选择一系列元素。语法是 start:stop:step,其中 start 是包含的,stop 是不包含的,而 step 定义了间隔。省略 start 默认为0,省略 stop 默认为维度的末尾,省略 step 默认为1。
// 创建一个一维张量
val y_1d = torch.arange(10) // Tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9])
println(f"原始一维张量: {y_1d}")
// 选择从索引2开始到(不包含)索引5的元素
val slice1 = y_1d(2 until 5)
println(f"\n切片 y_1d(2 until 5): {slice1}")
// 选择从开头到索引4的元素
val slice2 = y_1d(0 until 4)
println(f"切片 y_1d(0 until 4): {slice2}")
// 选择从索引6到末尾的元素
val slice3 = y_1d(6 until y_1d.size(0))
println(f"切片 y_1d(6 until y_1d.size(0)): {slice3}")
// 选择每隔一个的元素
val slice4 = y_1d(0 until y_1d.size(0) by 2)
println(f"切片 y_1d(0 until y_1d.size(0) by 2): {slice4}")
// 选择从索引1到7的元素,步长为2
val slice5 = y_1d(1 until 8 by 2)
println(f"切片 y_1d(1 until 8 by 2): {slice5}")
// 反转张量
val slice6 = y_1d(y_1d.size(0) - 1 until 0 by -1)
println(f"切片 y_1d(y_1d.size(0) - 1 until 0 by -1): {slice6}")
切片对多维张量的工作方式类似。您可以将整数索引和切片结合使用:
// 创建一个3x4张量
val x_2d = torch.tensor(Seq(Seq( 0, 1, 2, 3),
Seq( 4, 5, 6, 7),
Seq( 8, 9, 10, 11)))
println(f"原始二维张量:\n{x_2d}")
// 选择前两行以及第1和第2列
val sub_tensor1 = x_2d(0 until 2, 1 until 3)
println(f"\n切片 x_2d(0 until 2, 1 until 3):\n{sub_tensor1}")
// 选择所有行,但只选择最后两列
val sub_tensor2 = x_2d(*, -2 until x_2d.size(1))
println(f"\n切片 x_2d(*, -2 until x_2d.size(1)):\n{sub_tensor2}")
// 选择第一行,从第1列到末尾
val sub_tensor3 = x_2d(0, 1 until x_2d.size(1))
println(f"\n切片 x_2d(0, 1 until x_2d.size(1)):\n{sub_tensor3}")
// 选择第0行和第2行(使用步长),所有列
val sub_tensor4 = x_2d(0 until x_2d.size(0) by 2, *)
println(f"\n切片 x_2d(0 until x_2d.size(0) by 2, *):\n{sub_tensor4}")
原始张量 (x_2d)切片: x_2d[0:2, 1:3]012345678910111256
张量
x_2d使用x_2d[0:2, 1:3]进行切片的视觉表示。它选择第0和第1行,以及第1和第2列。
切片的一个重要特性(与某些其他形式的索引不同)是,返回的张量通常与原始张量共享底层存储。修改切片会修改原始张量。
println(f"修改切片前的原始 x_2d:\n{x_2d}")
// 获取一个切片
val sub_tensor = x_2d(0 until 2, 1 until 3)
# 修改切片
sub_tensor(0, 0) = 101
println(f"\n修改后的切片:\n{sub_tensor}")
println(f"\n修改切片后的原始 x_2d:\n{x_2d}") // 注意变化!
如果您需要一个不共享内存的副本,请在切片上使用 .clone(): sub_tensor_copy = x_2d[0:2, 1:3].clone()。
布尔索引 (遮罩)
您可以使用布尔张量来索引另一个张量。布尔张量的形状必须能够广播到被索引张量的形状(通常,它们的形状完全相同)。只有布尔张量中对应 True 值的元素(即“遮罩”)才会被选中。这对于根据条件筛选数据非常有用。
// 创建一个张量
val data = torch.tensor(Seq(Seq(1, 2), Seq(3, 4), Seq(5, 6)))
println(f"原始数据张量:\n{data}")
// 创建一个布尔遮罩 (例如,选择大于3的元素)
val mask = data > 3
println(f"\n布尔遮罩 (data > 3):\n{mask}")
// 应用遮罩
val selected_elements = data(mask)
println(f"\n通过遮罩选择的元素:\n{selected_elements}")
println(f"所选元素的形状: {selected_elements.shape}")
// 根据条件修改元素
data(data <= 3) = 0
println(f"\n将小于等于3的元素设置为零后的数据:\n{data}")
布尔索引通常返回一个包含所有选定元素的一维张量。与切片不同,它不保留原始形状。此外,布尔索引通常会创建一个副本,而不是一个视图。
您可以将布尔索引与其他形式结合使用。例如,根据应用于其中一列的条件选择行:
// 选择第一列大于2的行
val row_mask = data(:, 0) > 2
println(f"\n行遮罩 (data[:, 0] > 2): {row_mask}")
// 使用 ':' 选择所选行中的所有列
// Or simply: data[row_mask] - PyTorch 通常会推断出完整的行选择
val selected_rows = data(row_mask, *)
println(f"\n第一列大于2的行:\n{selected_rows}")
整数数组索引
除了单个整数和切片,您还可以使用列表或一维整数张量沿维度进行索引。这使得您可以按任意顺序选择元素,或多次选择相同的元素。
// 创建一个一维张量
val x = torch.arange(10, 20) // Tensor([10, 11, 12, 13, 14, 15, 16, 17, 18, 19])
println(f"原始一维张量: {x}")
// 注意索引2的重复
val indices = torch.tensor(Seq(0, 4, 2, 2))
println(f"\n使用索引 {indices} 选择的元素: {x(indices)}")
// 对于二维张量
val y = torch.arange(12).reshape(3, 4)
// [[ 0, 1, 2, 3],
// [ 4, 5, 6, 7],
// [ 8, 9, 10, 11]]
println(f"\n原始二维张量:\n{y}")
// 选择特定行
val row_indices = torch.tensor(Seq(0, 2))
val selected_rows = y(row_indices, *)
println(f"\n使用索引 {row_indices} 选择的行:\n{selected_rows}")
// 选择特定列
val col_indices = torch.tensor(Seq(1, 3))
val selected_cols = y(*, col_indices) // 从所有行中选择第1列和第3列
println(f"\n使用索引 {col_indices} 选择的列:\n{selected_cols}")
// 使用索引对选择特定元素
val row_idx = torch.tensor(Seq(0, 1, 2))
val col_idx = torch.tensor(Seq(1, 3, 0))
val selected_elements = y(row_idx, col_idx) // 选择 (0,1), (1,3), (2,0) -> [1, 7, 8]
println(f"\n使用 (row_idx, col_idx) 选择的特定元素:\n{selected_elements}")
与布尔索引类似,整数数组索引通常返回一个新的张量(一个副本),而不是原始张量存储的视图。输出的形状取决于索引方法。当选择完整的行或列时,其他维度会被保留。当为多个维度提供索引数组(例如 y[row_idx, col_idx])时,结果通常是一个对应于所选元素的一维张量。
掌握这些索引和切片技术能够精准地控制张量数据,为后续步骤中的数据准备、特征提取以及模型输入输出的操作奠定根基。
张量的重塑与维度调整
通常,你会发现现有张量的结构不太适合后续计算步骤,尤其是在将数据送入特定的神经网络层时。PyTorch 提供了灵活的工具,可以在不改变底层数据元素本身的情况下,改变张量的形状或调整其维度。用于这些操作的主要方法是:view()、reshape() 和 permute()。
使用 view() 和 reshape() 改变形状
view() 和 reshape() 都允许你改变张量的维度,前提是总元素数量保持不变。它们在将多维张量展平后传递给线性层,或增加/移除大小为1的维度等任务中非常有用。
使用 view()
view() 方法返回一个新的张量,该张量与原始张量共享相同的底层数据,但具有不同的形状。它非常高效,因为它避免了数据复制。然而,view() 要求张量在内存中是连续的。连续张量是指其元素在内存中按维度顺序连续存储,没有间隙的张量。大多数新创建的张量是连续的,但某些操作(如切片或使用 t() 进行转置)会产生非连续张量。
我们来看一个例子:
import torch.*
// 创建一个连续张量
val x = torch.arange(12) // tensor([ 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11])
println(f"原始张量: {x}")
println(f"原始形状: {x.shape}")
println(f"是否连续? {x.is_contiguous()}")
// 使用 view() 重塑
val y = x.view(3, 4)
println("\nview(3, 4) 后的张量:")
println(y)
println(f"新形状: {y.shape}")
println(f"与 x 共享存储吗? {y.storage().data_ptr() == x.storage().data_ptr()}") // 检查它们是否共享内存
println(f"y 是否连续? {y.is_contiguous()}")
// 尝试另一个视图
val z = y.view(2, 6)
println("\nview(2, 6) 后的张量:")
println(z)
println(f"新形状: {z.shape}")
println(f"与 x 共享存储吗? {z.storage().data_ptr() == x.storage().data_ptr()}")
println(f"z 是否连续? {z.is_contiguous()}")
你可以在 view() 调用中对一个维度使用 -1,PyTorch 将根据总元素数量和其它维度的尺寸自动推断出该维度的正确尺寸。
// 使用 -1 进行推断
val w = x.view(2, 2, -1) // 推断出最后一个维度为 3 (12 / (2*2) = 3)
println("\nview(2, 2, -1) 后的张量:")
println(w)
println(f"新形状: {w.shape}")
如果你尝试在非连续张量上调用 view(),你会得到一个 RuntimeError。
// view() 在非连续张量上失败的例子
val a = torch.arange(12).view(3, 4)
val b = a.t() // 转置操作会创建一个非连续张量
println(f"\nb 是否连续? {b.is_contiguous()}")
try:
val c = b.view(12)
catch (e: RuntimeException) =>
println(f"\n尝试 b.view(12) 时出错: {e}")
使用 reshape()
reshape() 方法的行为类似于 view(),但提供了更多灵活性。如果张量对于目标形状是连续的,它会尝试返回一个视图。如果无法返回视图(例如,因为原始张量在与新形状兼容的方式上不是连续的),reshape() 将把数据复制到一个新的、具有所需形状的连续张量中。这使得 reshape() 通常更安全、更通用,尽管如果发生复制,性能可能会降低。
我们再次查看使用 reshape() 的转置例子:
// 在非连续张量 'b' 上使用 reshape()
println(f"\n原始非连续张量 b:\n{b}")
println(f"b 的形状: {b.shape}")
println(f"b 是否连续? {b.is_contiguous()}")
// 即使 'b' 不连续,reshape 也能工作
val c = b.reshape(12)
println(f"\nb.reshape(12) 后的张量 c:\n{c}")
println(f"c 的形状: {c.shape}")
println(f"c 是否连续? {c.is_contiguous()}")
// 检查 'c' 是否与 'b' 共享存储。由于 reshape 可能进行了复制,所以它们很可能不共享。
println(f"与 b 共享存储吗? {c.storage().data_ptr() == b.storage().data_ptr()}")
// reshape 也可以用 -1 推断维度
val d = b.reshape(2, -1) // 推断出最后一个维度为 6
println(f"\nb.reshape(2, -1) 后的张量 d:\n{d}")
println(f"d 的形状: {d.shape}")
何时使用哪个方法?
- 如果你确定张量是连续的,并且希望确保不发生数据复制以获得最高性能,请使用
view()。如果连续性假设有误,请准备好处理可能的RuntimeError。 reshape()适用于连续和非连续张量。如果可能,它会返回一个视图,否则会创建一个副本。除非性能绝对关键且你能保证连续性,否则这通常是优选方法。
使用 permute() 调整维度顺序
view() 和 reshape() 通过重新安排元素在维度间的解释方式来改变形状,而 permute() 则明确地交换维度本身。它不改变总元素数量,也不改变每个轴上元素数量方面的形状,但它改变的是哪个轴对应哪个原始维度。
假设你有一个图像数据,存储为(通道,高,宽)格式,但为了特定的库或可视化需求,需要其格式为(高,宽,通道)。permute() 就是为此而设计的工具。你将所需的维度顺序作为参数提供。
// 创建一个三维张量(例如,表示通道、高、宽)
val image_tensor = torch.randn(3, 32, 32) // 通道,高,宽
println(f"原始形状: {image_tensor.shape}") // torch.Size([3, 32, 32])
// 调整为(高,宽,通道)
val permuted_tensor = image_tensor.permute(1, 2, 0) // 指定新顺序:维度 1,维度 2,维度 0
println(f"调整后的形状: {permuted_tensor.shape}") // torch.Size([32, 32, 3])
// permute 通常返回一个非连续的视图
println(f"permuted_tensor 是否连续? {permuted_tensor.is_contiguous()}")
// 调回原状
val original_again = permuted_tensor.permute(2, 0, 1) // 回到通道,高,宽
println(f"调回后的形状: {original_again.shape}") // torch.Size([3, 32, 32])
println(f"original_again 是否连续? {original_again.is_contiguous()}") // (可能仍然是非连续的)
// 检查存储共享
println(f"与原始张量共享存储吗? {original_again.storage().data_ptr() == image_tensor.storage().data_ptr()}")
和 view() 一样,permute() 返回一个与原始张量共享底层数据的张量。它不复制数据。然而,生成的张量通常不是连续的。如果你在调换维度后需要一个连续张量(例如,为了后续使用 view()),你可以链式调用 .contiguous() 方法:
// 使调整维度的张量连续
val contiguous_permuted = permuted_tensor.contiguous()
println(f"\ncontiguous_permuted 是否连续? {contiguous_permuted.is_contiguous()}")
// 现在可以安全地使用 view()
val flattened_permuted = contiguous_permuted.view(-1)
println(f"展平后的形状: {flattened_permuted.shape}")
掌握 view()、reshape() 和 permute() 让你能够精确控制张量的结构,这是将数据适配到不同 PyTorch 操作和模型层要求所需的一项必备技能。请记住这些权衡:view() 速度快但要求连续性,reshape() 灵活但可能会复制,而 permute() 交换维度而不复制,但通常会产生非连续张量。
张量的合并与分割
收藏
许多深度学习情境下,你需要将多个张量组合成一个,或将一个更大的张量拆分为更小的部分。这可能涉及汇总不同处理步骤的结果、准备数据批次或分离特征。PyTorch 提供了几个函数,用于有效合并和分割张量。
合并张量
组合张量是一个常见操作,尤其是在处理数据批次或合并特征表示时。PyTorch 提供了两种主要方式来合并张量:拼接 (torch.cat) 和堆叠 (torch.stack)。主要区别在于它们是沿着现有维度操作,还是引入一个新维度。
使用 torch.cat 进行拼接
torch.cat 函数沿着现有维度拼接一系列张量。序列中的所有张量必须形状相同(除了拼接维度),或者为空。
import torch.*
// 创建两个张量
val tensor_a = torch.randn(2, 3)
val tensor_b = torch.randn(2, 3)
println(f"Tensor A (Shape: {tensor_a.shape}):\n{tensor_a}")
println(f"Tensor B (Shape: {tensor_b.shape}):\n{tensor_b}\n")
// 沿着维度0(行)进行拼接
// 结果形状: (2+2, 3) = (4, 3)
val cat_dim0 = torch.cat((tensor_a, tensor_b), dim=0)
println(f"沿着维度0拼接 (形状: {cat_dim0.shape}):\n{cat_dim0}\n")
// 沿着维度1(列)进行拼接
// 张量必须在其他维度(维度0)上匹配
// 结果形状: (2, 3+3) = (2, 6)
val cat_dim1 = torch.cat((tensor_a, tensor_b), dim=1)
println(f"沿着维度1拼接 (形状: {cat_dim1.shape}):\n{cat_dim1}")
// 3D张量示例
val tensor_c = torch.randn(1, 2, 3)
val tensor_d = torch.randn(1, 2, 3)
// 沿着维度0(批次维度)进行拼接
// 结果形状: (1+1, 2, 3) = (2, 2, 3)
val cat_3d_dim0 = torch.cat((tensor_c, tensor_d), dim=0)
println(f"\n3D张量沿着维度0拼接 (形状: {cat_3d_dim0.shape})")
请注意,torch.cat 增加了指定维度的大小,同时保持其他维度不变。张量在所有维度上都必须大小匹配,除了你进行拼接的那个维度。
张量 A (2x3)张量 B (2x3)torch.cat((A, B), dim=0)(4x3)torch.cat((A, B), dim=1)(2x6)a11a12a13a21a22a23b11b12b13b21b22b23a11a12a13a21a22a23b11b12b13b21b22b23a11a12a13b11b12b13a21a22a23b21b22b23cluster_a+ 维度0+ 维度1cluster_bcluster_cat0cluster_cat1
torch.cat沿维度0和维度1对两个2x3张量进行拼接的视觉比较。
使用 torch.stack 进行堆叠
与 cat 不同,torch.stack 沿着一个新维度连接一系列张量。当你希望从单个示例创建批次或将相关张量分组时,这会很有用。为了 stack 能够工作,输入序列中的所有张量必须具有完全相同的形状。
import torch.*
// 创建两个形状相同的张量
val tensor_e = torch.arange(6).reshape(2, 3)
val tensor_f = torch.arange(6, 12).reshape(2, 3)
println(f"Tensor E (Shape: {tensor_e.shape}):\n{tensor_e}")
println(f"Tensor F (Shape: {tensor_f.shape}):\n{tensor_f}\n")
// 沿着新维度0进行堆叠
// 结果形状: (2, 2, 3)
val stack_dim0 = torch.stack((tensor_e, tensor_f), dim=0)
println(f"沿着新维度0堆叠 (形状: {stack_dim0.shape}):\n{stack_dim0}\n")
// 沿着新维度1进行堆叠
// 结果形状: (2, 2, 3)
val stack_dim1 = torch.stack((tensor_e, tensor_f), dim=1)
println(f"沿着新维度1堆叠 (形状: {stack_dim1.shape}):\n{stack_dim1}\n")
// 沿着新维度2(最后一个维度)进行堆叠
// 结果形状: (2, 3, 2)
val stack_dim2 = torch.stack((tensor_e, tensor_f), dim=2)
println(f"沿着新维度2堆叠 (形状: {stack_dim2.shape}):\n{stack_dim2}")
张量 E (2x3)张量 F (2x3)torch.stack((E, F), dim=0)(2x2x3)切片 0切片 1torch.stack((E, F), dim=1)(2x2x3)e11e12e13e21e22e23f11f12f13f21f22f23e11e12e13e21e22e23f11f12f13f21f22f23cluster_stack0_ecluster_stack0_fe11e12e13f11f12f13e21e22e23f21f22f23cluster_e堆叠 维度0堆叠 维度1cluster_fcluster_stack0cluster_stack1
torch.stack在dim=0和dim=1处插入新维度的视觉比较。请注意原始张量如何成为新张量中的切片。
选择 cat 还是 stack 取决于你是想沿着现有维度合并,还是创建一个新维度。cat 通常用于水平/垂直组合批次或特征,stack 则常用于从单个样本创建批次。
分割张量
正如你可以合并张量一样,你也经常需要将它们分开。这可能涉及将一个批次拆分回单个样本、将特征与标签分离或为并行处理划分数据。PyTorch 为这些任务提供了 torch.split 和 torch.chunk 函数。
使用 torch.split 按特定大小分割
torch.split 函数沿着指定维度将张量分割成块。你可以指定每个块的大小(如果你想要等份),或者提供一个包含每个所需块大小的列表。
import torch
// 创建一个要分割的张量
val tensor_g = torch.arange(12).reshape(6, 2)
println(f"原始张量 (形状: {tensor_g.shape}):\n{tensor_g}\n")
// 沿着维度0(行)按大小2分割成块
// 6行 / 2行/块 = 3块
val split_equal = torch.split(tensor_g, 2, dim=0)
println("分割成大小为2的等份(dim=0):")
for i, chunk <- split_equal:
println(f" 块 {i} (形状: {chunk.shape}):\n{chunk}")
println("-" * 20)
// 沿着维度0按大小 [1, 2, 3] 分割成块
// 总大小必须等于该维度的大小 (1 + 2 + 3 = 6)
val split_unequal = torch.split(tensor_g, List(1, 2, 3), dim=0)
println("\n分割成大小不等的块 [1, 2, 3](dim=0):")
for i, chunk <- split_unequal:
println(f" 块 {i} (形状: {chunk.shape}):\n{chunk}")
println("-" * 20)
// 沿着维度1(列)进行分割
// 形状: (6, 2)。沿着维度1按大小1分割成块
val split_dim1 = torch.split(tensor_g, 1, dim=1)
println("\n分割成大小为1的等份(dim=1):")
for i, chunk <- split_dim1:
// 使用 squeeze 移除大小为1的维度,以便更清晰地显示
println(f" 块 {i} (形状: {chunk.shape}):\n{chunk.squeeze()}")
torch.split 返回一个张量元组。如果你为 split_size_or_sections 参数提供一个整数,PyTorch 会沿着指定的 dim 将张量分割成该大小的块。如果维度大小不能被分割大小完全整除,最后一个块会更小。如果你提供一个大小列表,它们的总和必须等于被分割维度的大小。
使用 torch.chunk 按数量分割
另一种方法是,torch.chunk 沿着给定维度将张量分割成指定数量的块。PyTorch 会尝试使这些块的大小尽可能相等。与需要指定块大小的 torch.split 不同,chunk 只需指定所需的块数量。
import torch.*
// 创建一个张量
val tensor_h = torch.arange(10).reshape(5, 2) // 沿着维度0的大小为5
println(f"原始张量 (形状: {tensor_h.shape}):\n{tensor_h}\n")
// 沿着维度0分割成3个块
// 5行 / 3块 -> 大小将是 [2, 2, 1] (前几个块取 ceil(5/3)=2)
val chunked_tensor = torch.chunk(tensor_h, 3, dim=0)
println("分割成3个部分(dim=0):")
for i, chunk <- chunked_tensor:
println(f" 块 {i} (形状: {chunk.shape}):\n{chunk}")
println("-" * 20)
// 创建另一个张量
val tensor_i = torch.arange(12).reshape(3, 4) // 沿着维度1的大小为4
println(f"\n原始张量 (形状: {tensor_i.shape}):\n{tensor_i}\n")
// 沿着维度1分割成2个块
// 4列 / 2块 -> 大小将是 [2, 2] (ceil(4/2)=2)
val chunked_tensor_dim1 = torch.chunk(tensor_i, 2, dim=1)
println("分割成2个部分(dim=1):")
for i, chunk <- chunked_tensor_dim1:
println(f" 块 {i} (形状: {chunk.shape}):\n{chunk}")
当你知道想要多少个部分,而不关心维度大小是否能被均匀整除时,torch.chunk 很方便。当你需要大小精确且可能变化的块时,torch.split 提供了更多的控制。
掌握这些合并和分割操作很重要,可以帮助你有效处理数据,因为它会流经你的深度学习管线的不同阶段,从初始加载和预处理,到训练的批处理,以及模型输出的分析。
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)