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

请求将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 15:13:03