cute(2)Tensor基础

鱿鱼圈 Lv4

Ref

2 Tensor 基础

对应代码:02_tensor_basics.cu 纯 host 代码,不需要 GPU。

reed博士:Layout描述了数据的排列和底层存储位置关系,但Layout并没有指定存储。Tensor就是在Layout的基础上包含了存储,即Tensor = Layout + storage,数据存储的具体表现上可以是指针表达的数据或则是栈上数据(GPU上表现为寄存器)。cute中的Tensor并不同于深度学习框架中的Tensor(如Pytorch、Tensorflow),深度学习框架中的Tensor更强调数据的表达实体,通过Tensor实体与实体之间的计算产生新的Tensor实体,即多份数据实体,cute中的Tensor更多的是对Tensor进行分解和组合等操作,而这些操作多是对Layout的变换(只是逻辑层面的数据组织形式),底层的数据实体一般不变更。也就是说深度学习框架中的Tensor是通过Tensor产生新Tensor,cute中是对数据表达形式的变换,底层数据一般不变更,指变更表达的形式,这个表达形式的变更是通过之前文章介绍的Layout上的运算实现的。之所以有这些差别是因为深度学习框架中的Tensor是用来表达数据实体,cute中的Tensor是偏向描述的实体

核心概念

Tensor = 数据指针 + Layout

1
2
3
Tensor = (data_ptr, layout)
^ ^
数据在哪 怎么索引

数据和索引方式是分离的。同一段内存可以用不同的 layout 去解读,不搬运任何数据。

代码解析

第 1 节:从数组创建 Tensor

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
// ============================================================
// 1. 从数组创建 Tensor
// ============================================================
{
printf("=== 1. Create Tensor from array ===\n");

// 准备一个 host 数组
float data[12];
for (int i = 0; i < 12; ++i) data[i] = float(i);

// 创建一个 3x4 row-major Tensor
auto tensor = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(4, 1)));

printf(" tensor: ");
print(tensor); // 打印 tensor 的元信息(指针 + layout)
printf("\n");

// 访问元素
printf(" tensor(0, 0) = %.0f\n", tensor(0, 0)); // data[0*4 + 0] = 0
printf(" tensor(1, 2) = %.0f\n", tensor(1, 2)); // data[1*4 + 2] = 6
printf(" tensor(2, 3) = %.0f\n", tensor(2, 3)); // data[2*4 + 3] = 11

// 打印整个 tensor 的值
printf("\n print_tensor:\n");
print_tensor(tensor);
printf("\n");
}

img

1
2
3
4
float data[12];
for (int i = 0; i < 12; ++i) data[i] = float(i);

auto tensor = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(4, 1)));
  • 第一个参数:指针&data[0]),不能直接传数组名
  • 第二个参数:layout,决定怎么解读这段内存
  • 结果:3×4 row-major tensor

访问元素

1
2
3
tensor(0, 0)  // data[0*4 + 0*1] = data[0] = 0
tensor(1, 2) // data[1*4 + 2*1] = data[6] = 6
tensor(2, 3) // data[2*4 + 3*1] = data[11] = 11

常见错误make_tensor(data, layout) 传数组名会报错,必须传 &data[0] 指针。

打印函数

  • print(tensor) — 打印元信息(指针地址 + layout)
  • print_tensor(tensor) — 打印所有元素的值

第 2 节:同一数据,不同 Layout

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
// ============================================================
// 2. Tensor 的 layout 和 data 是独立的
// ============================================================
{
printf("=== 2. Same data, different layout ===\n");

float data[12];
for (int i = 0; i < 12; ++i) data[i] = float(i);

// 同一段内存,用不同 layout 解读
auto row_major = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(4, 1)));
auto col_major = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(1, 3)));

printf(" row_major(1, 2) = %.0f\n", row_major(1, 2)); // data[6] = 6
printf(" col_major(1, 2) = %.0f\n", col_major(1, 2)); // data[7] = 7

printf("\n row_major:\n");
print_tensor(row_major);
printf("\n col_major:\n");
print_tensor(col_major);
printf("\n");
}

img

1
2
3
float data[12];
auto row_major = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(4, 1)));
auto col_major = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(1, 3)));

同一段内存 data,两个 tensor 指向同一个地址,只是 layout 不同:

1
2
row_major(1, 2) = data[1*4 + 2*1] = data[6] = 6
col_major(1, 2) = data[1*1 + 2*3] = data[7] = 7

核心理解:Layout 改变的是"怎么看数据",不是"数据本身"。这就是为什么后面的 composeretile_D 等操作都是零开销——只改 layout 元数据,不搬数据。

第 3 节:Slice —— 用 _ 取整行/整列

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
// ============================================================
// 3. Slice:用 _ 取整行/整列
// ============================================================
{
printf("=== 3. Slice with Underscore ===\n");

float data[12];
for (int i = 0; i < 12; ++i) data[i] = float(i);

// 3x4 row-major
auto tensor = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(4, 1)));

// 取第 1 行(固定 M=1,保留 N 方向)
auto row1 = tensor(1, _);
printf(" tensor(1, _) = ");
print(row1);
printf("\n values: ");
for (int j = 0; j < size(row1); ++j) printf("%.0f ", row1(j));
printf("\n");

// 取第 2 列(保留 M 方向,固定 N=2)
auto col2 = tensor(_, 2);
printf(" tensor(_, 2) = ");
print(col2);
printf("\n values: ");
for (int i = 0; i < size(col2); ++i) printf("%.0f ", col2(i));
printf("\n\n");
}

img

1
2
3
4
auto tensor = make_tensor(&data[0], make_layout(make_shape(3, 4), make_stride(4, 1)));

auto row1 = tensor(1, _); // 固定 M=1,保留 N 方向 → 取第 1 行
auto col2 = tensor(_, 2); // 保留 M 方向,固定 N=2 → 取第 2 列
  • _(Underscore)表示"保留这个维度"
  • 结果是一个降维的 tensor(二维变一维)
1
2
tensor(1, _) = [4, 5, 6, 7]    ← 第 1 行的 4 个元素
tensor(_, 2) = [2, 6, 10] ← 第 2 列的 3 个元素

Slice 的本质

  • 固定某个维度的 index → 指针偏移到那个位置
  • 保留的维度 → 构成新的 layout
  • 不拷贝数据,只是换了一个起始指针 + 新 layout

第 4 节:local_tile —— 把大 Tensor 切成小 Tile

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
// ============================================================
// 4. local_tile:把大 Tensor 切成小 tile
// ============================================================
{
printf("=== 4. local_tile ===\n");

float data[48];
for (int i = 0; i < 48; ++i) data[i] = float(i);

// 6x8 row-major 矩阵
auto tensor = make_tensor(&data[0], make_layout(make_shape(6, 8), make_stride(8, 1)));

// 打印原始 tensor
printf(" 原始 tensor (6x8 row-major):\n");
print_tensor(tensor);
printf("\n");

// 按 (3, 4) 切 tile
auto tiled = local_tile(tensor, make_tile(Int<3>{}, Int<4>{}), make_coord(_, _));
printf(" tiled shape: ");
print(shape(tiled));
print_tensor(tiled);
printf("\n");
// shape = (3, 4, 2, 2)
// 前两维是 tile 内的坐标 (3, 4)
// 后两维是 tile 的编号 (6/3=2, 8/4=2)

printf(" tiled tensor (4D):\n");
print_tensor(tiled);
printf("\n");

// 打印所有 4 个 tile
for (int ti = 0; ti < 2; ++ti) {
for (int tj = 0; tj < 2; ++tj) {
auto tile = local_tile(tensor, make_tile(Int<3>{}, Int<4>{}), make_coord(ti, tj));
printf(" tile(%d,%d) shape: ", ti, tj);
print(shape(tile));
printf("\n");
print_tensor(tile);
printf("\n");
}
}
}

img

1
2
3
4
5
6
// 6×8 row-major 矩阵
auto tensor = make_tensor(&data[0], make_layout(make_shape(6, 8), make_stride(8, 1)));

// 按 (3, 4) 切 tile,保留所有 tile
auto tiled = local_tile(tensor, make_tile(Int<3>{}, Int<4>{}), make_coord(_, _));
// shape = (3, 4, 2, 2)

local_tile 的三个参数:

  1. 原始 tensor
  2. tile 的大小 (3, 4)
  3. 取哪个 tile 的坐标,_ 表示保留所有

结果 shape (3, 4, 2, 2) 的含义

1
2
前两维 (3, 4) = tile 内部的坐标(每个 tile 是 3×4)
后两维 (2, 2) = tile 的编号(6/3=2 个行方向 tile,8/4=2 个列方向 tile)

取具体某个 tile

1
2
auto tile = local_tile(tensor, make_tile(Int<3>{}, Int<4>{}), make_coord(ti, tj));
// shape = (3, 4),就是第 (ti, tj) 个 tile

可视化(6×8 矩阵切成 4 个 3×4 tile):

1
2
3
4
5
6
7
8
9
10
11
12
原始 6×8 矩阵:
+---+---+---+---+---+---+---+---+
| 0 1 2 3 | 4 5 6 7 |
| 8 9 10 11 | 12 13 14 15 |
| 16 17 18 19 | 20 21 22 23 |
+-----------------+-----------------+
| 24 25 26 27 | 28 29 30 31 |
| 32 33 34 35 | 36 37 38 39 |
| 40 41 42 43 | 44 45 46 47 |
+-----------------+-----------------+
tile(0,0) tile(0,1)
tile(1,0) tile(1,1)

local_tile 的内部原理:把每个维度拆成 (tile_内部, tile_编号),stride = tile_size × 原始 stride。不拷贝数据。

reef博士:该函数是Tensor中用户可以使用到的重要的函数,可以通过tile方法对tensor进行分块,通过local_tile可以实现从大的tensor中切取tile块,并且通过coord进行块的选取,以下代码展示了将维度为MNK的张量按照2x3x4的小块进行划分,取出其中的第(1, 2, 3)块。

img

如上图所示,A Tensor表达了4行6列的行优先的数据,分块的tile大小为2x2, local_tile会将A矩阵按照tile的单位进行分块,然后根据坐标选取出(1, 1)位置的分块,如图则取出了(1,1)位置的tensor得到右下角的结果。

1
2
Tensor tensor = make_tensor(ptr, make_shape(M, N, K));
Tensor tensor1 = local_tile(tensor, make_shape(2, 34), make_coord(1, 23));

第 5 节:make_tensor_like

1
2
3
auto tensor = make_tensor(&data[0], make_layout(make_shape(Int<3>{}, Int<4>{}),
make_stride(Int<4>{}, Int<1>{})));
auto new_tensor = make_tensor_like(tensor);

创建一个和 tensor 同 shape/layout 的新 tensor,数据分配在栈上。

要求:shape 必须是静态的(Int<N>{}),因为编译期需要知道大小才能在栈上分配数组。

用途:在 kernel 中创建临时的寄存器 tensor,用于类型转换等。

第 6 节:Tensor的局部数据提取local_partition

local partition和local tile类似,现将大的Tensor按照tile大小进行分块,分块后每一块取出coordinate指定的元素形成新的块,代码形式和具体实例参考下图:

img

1
2
Tensor tensor = make_tensor(ptr, ...);
Tensor tensor1 = local_partition(tensor, tile);

API 速查

函数 作用
make_tensor(ptr, layout) 从指针+layout 创建 tensor
make_tensor_like(tensor) 创建同 layout 的新 tensor(栈上分配)
tensor(i, j) 二维访问
tensor(i, _) Slice:取第 i 行
tensor(_, j) Slice:取第 j 列
local_tile(tensor, tile, coord) 按 tile 大小切分
print(tensor) 打印元信息
print_tensor(tensor) 打印所有值
size(tensor) 总元素数

易错点总结

问题 解决方法
make_tensor(data, layout) 报错 必须传指针 &data[0],不能传数组名
make_tensor_like 报错 shape 必须是静态 Int<N>{},不能是动态 int
local_tile 结果维度太多 正常,前半是 tile 内坐标,后半是 tile 编号
  • 标题: cute(2)Tensor基础
  • 作者: 鱿鱼圈
  • 创建于 : 2026-06-07 22:13:32
  • 更新于 : 2026-06-14 23:33:01
  • 链接: https://yuyanqi.com/2026/06/07/cute(2)Tensor基础/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论