Skip to content

Day 63: std.multi_array_list模块:多维数组

在许多计算领域,如科学计算、机器学习和图像处理中,我们经常需要处理多维数组(也称为张量或NDArray)。虽然Zig的数组类型[N][M]T可以创建编译时固定大小的多维数组,但我们常常需要一个在运行时维度和大小都可以动态调整的结构。

std.MultiArrayList就是为此而设计的。它提供了一个动态的、多维的、类型擦除的数组实现,非常适合存储和操作矩阵、张量或任何其他结构化的多维数据。

std.MultiArrayList的核心思想是将多维结构扁平化存储在一个单一的ArrayList(u8)中,同时维护一个描述其“形状”(shape)和“步长”(strides)的元数据。

  • Shape: 一个描述每个维度大小的切片。例如,一个3x4的矩阵,其形状是{3, 4}。
  • Strides: 一个切片,描述在每个维度上移动一个单位需要跳过多少个字节。

初始化一个MultiArrayList需要一个分配器和你想存储的元素类型:

const std = @import("std");
pub fn main() !void {
var gpa = std.heap.GeneralPurposeAllocator(.{}){};
defer _ = gpa.deinit();
const allocator = gpa.allocator();
// 创建一个用于存储f32的多维数组
var matrix = std.MultiArrayList(f32).init(allocator);
defer matrix.deinit();
// ...
}

MultiArrayList提供了get和set方法,通过一个索引切片来访问特定位置的元素。

  • get(indices: []const usize) T: 获取指定索引处的元素。
  • set(indices: []const usize, value: T): 设置指定索引处的元素。
// ...
// 将其重塑为一个2x3的矩阵,所有元素初始化为0
try matrix.resize(.{ 2, 3 });
// 设置 (1, 2) 位置的元素
matrix.set(&.{ 1, 2 }, 3.14);
// 获取 (1, 2) 位置的元素
const value = matrix.get(&.{ 1, 2 });
std.debug.print("Value at (1, 2): {d}\n", .{value});
// 索引必须与维度匹配,否则会panic
// matrix.get(&.{1}); // panics

MultiArrayList支持在第一个维度上进行append操作,这对于逐行或逐批次添加数据非常有用。

  • append(count: usize) !void: 在第0维上追加count个未初始化的新“行”或“切片”。
  • appendSlice(slice: []const T): 追加一个与MultiArrayList除第0维外形状相同的切片数据。
// 创建一个空的二维数组,但指定后续维度的形状
var image = try std.MultiArrayList(u8).initWithCapacity(allocator, .{0, 2}); // 0行, 2列
// 追加一行
try image.append(1);
image.set(&.{0, 0}, 255);
image.set(&.{0, 1}, 0);
// 追加另一行
try image.append(1);
image.set(&.{1, 0}, 0);
image.set(&.{1, 1}, 255);
// image现在是一个2x2的数组
std.debug.print("Shape: {any}\n", .{image.shape});

MultiArrayList非常适合用于表示机器学习中的张量。

const std = @import("std");
// 创建一个 2x3x4 的3D张量
var tensor = std.MultiArrayList(f32).init(std.testing.allocator);
defer tensor.deinit();
try tensor.resize(.{ 2, 3, 4 });
// 填充数据
var counter: f32 = 0.0;
for (0..2) |i| {
for (0..3) |j| {
for (0..4) |k| {
tensor.set(&.{i, j, k}, counter);
counter += 1.0;
}
}
}
std.debug.print("Value at (1, 1, 1): {d}\n", .{tensor.get(&.{1, 1, 1})});

创建一个MultiArrayList(u8)来表示一个10x10的灰度图像。编写一个函数,将图像中所有像素值大于128的设置为255,否则设置为0(二值化)。

  • 形状变化(Reshaping) MultiArrayList的resize方法可以改变数组的形状。如果新形状的元素总数与原来相同,resize是一个O(1)操作,因为它只修改元数据而不移动实际数据。如果元素总数改变,则可能需要重新分配内存。

  • 内存布局 默认情况下,MultiArrayList使用C语言风格的行主序(row-major)布局。这意味着最后一个维度的元素在内存中是连续的。这对于缓存性能很重要。

  • 与ArrayList(ArrayList(T))的区别 使用嵌套的ArrayList来模拟多维数组(ArrayList(ArrayList(T)))会导致多次、分散的内存分配,访问性能也较差。MultiArrayList将所有数据存储在一个连续的内存块中,从而获得更好的缓存局部性和性能。

std.MultiArrayList为Zig带来了处理多维数据的强大能力。它通过将数据与形状/步长元数据分离的抽象,提供了极大的灵活性,同时保持了连续内存布局带来的高性能。虽然它的API比简单的ArrayList更复杂,但它为进行高级数值计算、数据分析和机器学习等任务提供了一个坚实、高效的基础。