You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TypeScript自定义Tensor类:如何验证多维数组的同构形状?

实现TypeScript版Tensor构造函数(同形状验证+扁平化)

核心思路

要实现类似numpy.array的功能,核心是递归遍历输入数组,同时完成三件事:

  • 收集扁平化后的一维数据
  • 提取每个维度的长度(形状)
  • 验证所有同级子元素的结构/长度完全一致(同构性)

具体实现步骤

1. 递归验证形状并扁平化数据

写一个辅助函数,递归遍历输入,记录当前层级的预期长度,一旦发现子元素不符合预期就抛出错误,同时收集所有叶子节点数据。

function validateAndFlatten(arr: any[]): { shape: number[], data: number[] } {
    const shape: number[] = [];
    const data: number[] = [];

    // 递归遍历函数,跟踪当前元素和层级深度
    function traverse(current: any, depth: number): void {
        // 叶子节点(非数组):加入数据,验证层级一致性
        if (!Array.isArray(current)) {
            if (shape.length !== depth) {
                throw new Error("输入数组结构不一致,存在混合层级元素");
            }
            data.push(current as number);
            return;
        }

        // 记录当前维度长度(首次进入该层级时)
        if (depth === shape.length) {
            shape.push(current.length);
        } else {
            // 验证当前维度所有子元素长度一致
            if (current.length !== shape[depth]) {
                throw new Error(`第${depth+1}维度存在长度不一致的子数组,预期长度${shape[depth]}, 实际${current.length}`);
            }
        }

        // 递归处理子元素
        for (const item of current) {
            traverse(item, depth + 1);
        }
    }

    traverse(arr, 0);
    return { shape, data };
}

2. 实现Tensor类和array函数

基于上面的辅助函数,封装Tensor类和入口函数:

class Tensor {
    public readonly shape: number[];
    public readonly data: number[];

    constructor(shape: number[], data: number[]) {
        this.shape = shape;
        this.data = data;
    }

    // 静态方法,模拟numpy.array的调用方式
    static array(input: any[]): Tensor {
        const { shape, data } = validateAndFlatten(input);
        return new Tensor(shape, data);
    }
}

3. 测试示例

// 示例1:有效输入
try {
    const tensor1 = Tensor.array([[[1,2],[3,4]],[[5,6],[7,8]]]);
    console.log(tensor1.shape); // [2,2,2]
    console.log(tensor1.data); // [1,2,3,4,5,6,7,8]
} catch (e) {
    console.error(e.message);
}

// 示例2:子数组长度不一致
try {
    Tensor.array([[[1],[3,4]],[[5,6],[7,8]]]);
} catch (e) {
    console.error(e.message); // 第3维度存在长度不一致的子数组,预期长度1, 实际2
}

// 示例3:结构不一致(混合数组和非数组元素)
try {
    Tensor.array([[[1],2],[3,4]]);
} catch (e) {
    console.error(e.message); // 输入数组结构不一致,存在混合层级元素
}

学习建议

  1. 吃透递归遍历逻辑:这是处理多维数组的核心,重点理解层级跟踪和一致性校验的时机
  2. 掌握数组扁平化的多种实现:除了递归,还可以尝试迭代法,对比不同方式的优劣
  3. 理解张量的核心概念:形状、维度、多维索引与一维存储的映射关系(这是Tensor类后续扩展索引功能的基础)
  4. 熟悉TypeScript类型判断:比如Array.isArray()、类型断言的正确使用,避免类型安全问题

内容的提问来源于stack exchange,提问作者Mhhany

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.28 02:20:17