请求将Python深度学习Tensor处理代码转换为Kotlin代码
代码转换实现
1. Tensor转Multik MultiArray
PyTorch Android的Tensor需先转为Multik支持的MultiArray类型,才能使用其argmax等方法:
import org.jetbrains.kotlinx.multik.api.d4array import org.jetbrains.kotlinx.multik.api.mk import org.jetbrains.kotlinx.multik.ndarray.data.DType import org.pytorch.Tensor // 将PyTorch 4D Tensor转为Multik 4D Float数组 fun tensorToMultik4D(tensor: Tensor): mk.MultiArray<Float, DType> { val shape = tensor.shape() val data = tensor.dataAsFloatArray return mk.d4array(shape[0], shape[1], shape[2], shape[3]) { data[it] } }
2. 对应原Python逻辑的Kotlin实现
import org.jetbrains.kotlinx.multik.api.argmax import org.jetbrains.kotlinx.multik.api.astype import org.jetbrains.kotlinx.multik.api.permute import org.jetbrains.kotlinx.multik.api.squeeze import org.jetbrains.kotlinx.multik.ndarray.data.MultiArray import org.pytorch.Tensor const val IMAGE_SIZE = 350 const val half = 175 fun processOutput(outTensor: Tensor): MultiArray<Int, DType> { // 1. 执行argmax、维度转置、压缩、类型转换 val outMultik = tensorToMultik4D(outTensor) val argmaxResult = outMultik.argmax(dim = 1, keepDim = true) val permuted = argmaxResult.permute(0, 2, 3, 1) val squeezed = permuted[0].squeeze().astype<Int>() // 2. 获取输出的高宽 val rH = squeezed.shape[0] val rW = squeezed.shape[1] // 3. 创建扩展数组并填充中间区域 val extendedArray = mk.d2array(IMAGE_SIZE + rH, IMAGE_SIZE + rW) { 0 } // 替代Python的切片赋值,用循环实现区域填充 for (i in 0 until IMAGE_SIZE) { for (j in 0 until IMAGE_SIZE) { extendedArray[half + i, half + j] = squeezed[i, j] * 255 } } return extendedArray.copy() }
替代方案:纯PyTorch API实现
如果不想依赖Multik,可直接用PyTorch Android的API处理Tensor,再转原生数组操作:
import org.pytorch.Tensor fun processOutputWithTorchAPI(outTensor: Tensor): Array<IntArray> { // 用PyTorch API执行argmax、转置、维度选择 val argmaxTensor = outTensor.argmax(1, true) val permutedTensor = argmaxTensor.permute(0, 2, 3, 1).select(0, 0) val squeezedData = permutedTensor.dataAsIntArray val rH = permutedTensor.shape()[0] val rW = permutedTensor.shape()[1] // 转换为二维数组 val squeezedArray = Array(rH) { i -> IntArray(rW) { j -> squeezedData[i * rW + j] } } // 创建扩展数组并填充 val extendedArray = Array(IMAGE_SIZE + rH) { IntArray(IMAGE_SIZE + rW) { 0 } } for (i in 0 until IMAGE_SIZE) { for (j in 0 until IMAGE_SIZE) { extendedArray[half + i][half + j] = squeezedArray[i][j] * 255 } } return extendedArray }
Python转Kotlin迁移建议
- 优先用框架原生API:如果使用PyTorch Android,尽量直接用其
argmax、permute、select等方法,避免Tensor和第三方数组库的转换开销; - 数组操作替代:Kotlin无Python式切片赋值,需用循环或
copyInto实现区域填充;Multik库可简化多维数组操作,但要注意类型转换成本; - 类型显式声明:Kotlin需明确指定类型转换(如
astype<Int>()),不像numpy自动推导; - 性能优化:Android设备上减少频繁数组拷贝,复用内存;大数组操作可考虑NDK或TensorFlow Lite原生算子;
- 调试对齐:打印每一步的Tensor/数组形状、部分元素值,确保和Python端输出一致。
内容的提问来源于stack exchange,提问作者kiyoungNoona
相关产品推荐
相关产品推荐

