Java TF API中有没有更简便的方法将Tensor转换为Java数组?
TensorFlow Java 0.3.1+版本张量转原生数组简化方案
原有代码的繁琐来自多余的中间对象实例化步骤,TFM 0.3+版本中TFloat32等张量类本身就实现了FloatNdArray接口,同时官方提供了工具方法可以直接完成转换,不需要手动创建FloatNdArray、FloatDataBuffer中间载体,以下是两种常用场景的简化实现:
场景1:获取扁平化一维float数组
不管输出张量是多少维度,都可以直接读取为一维原生float数组:
try (TFloat32 out = (TFloat32) runner.run().get(0)) { float[] res = new float[(int) out.size()]; out.data().read(res); }
场景2:获取和张量维度匹配的多维原生数组
如果需要保留张量的维度结构,直接调用StdArrays提供的静态转换方法即可:
try (TFloat32 out = (TFloat32) runner.run().get(0)) { // 1维张量转float[] float[] arr1d = StdArrays.array1dOf(out); // 2维张量转float[][] float[][] arr2d = StdArrays.array2dOf(out); // 3维张量转float[][][] float[][][] arr3d = StdArrays.array3dOf(out); }
内容的提问来源于stack exchange,提问作者Sennri
相关产品推荐
相关产品推荐

