TypeScript中如何将tf.Tensor类型转换为number数值类型
TypeScript下TensorFlow.js计算1D张量余弦相似度类型报错解决
问题场景
使用TypeScript结合TensorFlow.js开发时,计算两个1D张量的余弦相似度过程中出现类型不匹配错误:
- 编写的余弦相似度计算函数预期返回
number类型,但实际调用张量同步取值方法时,返回值类型为number与多层嵌套number数组的联合类型 - 初始实现代码:
function calculateCosineSimilarity(abstractEmbedding: tf.Tensor1D | Array<number>, queryEmbedding: tf.Tensor1D | Array<number>): number{ const dotProd: tf.Tensor = tf.dot(abstractEmbedding, queryEmbedding); const lenAbstractEmbedding: tf.Tensor = tf.dot(abstractEmbedding, abstractEmbedding); const lenQueryEmbedding: tf.Tensor = tf.dot(queryEmbedding, queryEmbedding); const similarityScore: tf.Tensor = tf.div(dotProd, tf.mul(lenAbstractEmbedding,lenQueryEmbedding)); return similarityScore.arraySync(); }
- 编译阶段抛出的类型错误:
Type 'number | number[] | number[][] | number[][][] | number[][][][] | number[][][][][] | number[][][][][][]' is not assignable to type 'number'.
已知点积运算的返回维度随输入张量维度变化,但当前场景下两个1D张量的余弦相似度计算结果必然是0维标量,要求在不修改函数number类型的返回值声明的前提下,解决类型报错。
解决方法
报错的本质是tf.Tensor.arraySync()方法的类型定义为了兼容所有维度张量,默认返回多层级数组与数字的联合类型,无法自动推导当前结果为标量,有两种无侵入的修复方案:
方案1:使用标量取值的标准写法(推荐)
替换arraySync()为dataSync(),该方法会返回张量存储值的扁平化Float32Array类型结果,标量对应的数组长度固定为1,直接取第0位即可得到符合类型要求的number值:
// 替换原return语句即可 return similarityScore.dataSync()[0];
该方案不需要额外类型断言,类型校验完全安全,同步取值的性能和arraySync()一致,没有额外开销。
方案2:类型断言适配
如果需要保留arraySync()的调用方式,可以先将最终输出的张量显式断言为0维标量类型tf.Scalar,再对arraySync()的返回值做类型断言匹配number类型:
// 修改similarityScore声明和return语句 const similarityScore: tf.Scalar = tf.div( dotProd, tf.mul(lenAbstractEmbedding, lenQueryEmbedding) ) as tf.Scalar; return similarityScore.arraySync() as number;
注意:该方案依赖开发者手动保证运算输出为标量,如果后续运算逻辑改动输出高维张量,TypeScript无法捕获对应的类型错误。
额外逻辑修正
当前实现的余弦相似度公式存在错误:余弦相似度的分母为两个向量L2模长的乘积,现有代码直接对两个向量和自身的点积(即模长的平方)做乘法,没有开根号,最终计算出的结果范围不在余弦相似度标准的[-1,1]区间内,正确的计算逻辑应该为:
function calculateCosineSimilarity(abstractEmbedding: tf.Tensor1D | Array<number>, queryEmbedding: tf.Tensor1D | Array<number>): number{ const dotProd = tf.dot(abstractEmbedding, queryEmbedding); const lenAbstractEmbedding = tf.dot(abstractEmbedding, abstractEmbedding).sqrt(); const lenQueryEmbedding = tf.dot(queryEmbedding, queryEmbedding).sqrt(); const similarityScore = dotProd.div(lenAbstractEmbedding.mul(lenQueryEmbedding)); return similarityScore.dataSync()[0]; }
内容的提问来源于stack exchange,提问作者BJ Bae
相关产品推荐
相关产品推荐

