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

Android(Java)中PyTorch Tensor的squeeze与unsqueeze操作实现及TorchScript迁移问题

解决PyTorch Android中Tensor的squeeze/unsqueeze操作问题

首先明确:PyTorch Android(Java API)确实提供了squeeze和unsqueeze的对应方法,完全可以实现你在Python中做的维度调整操作。下面针对你的输入和输出场景给出具体实现:

一、输入Tensor的unsqueeze处理(对应Python的unsqueeze(0))

你的输入图像在Android中通过TensorImageUtils.bitmapToFloat32Tensor得到的是(3, 224, 416)维度的Tensor(CHW格式),和你Python中transpose后的维度一致。要添加batch维度变成(1, 3, 224, 416),可以直接调用Tensor的unsqueeze(int dim)方法:

// 原输入Tensor:(3, 224, 416)
final Tensor inputTensor = TensorImageUtils.bitmapToFloat32Tensor(mBitmap, 
    TensorImageUtils.TORCHVISION_NORM_MEAN_RGB, 
    TensorImageUtils.TORCHVISION_NORM_STD_RGB);

// 在第0维添加维度,得到(1, 3, 224, 416)的输入Tensor
final Tensor inputWithBatch = inputTensor.unsqueeze(0);

之后将inputWithBatch传入模型的forward方法即可:

Map<String, IValue> outTensors = mModule.forward(IValue.from(inputWithBatch)).toDictStringKey();

二、输出Tensor的squeeze处理(对应Python的squeeze())

你的模型输出是(1, 1, 224, 416)维度的Tensor,需要移除所有大小为1的维度得到(224, 416)的灰度图像数据。可以调用Tensor的squeeze()方法(无参数时会移除所有尺寸为1的维度):

// 从输出字典中获取对应Tensor(假设键为"output",根据你的模型输出键调整)
Tensor outputTensor = outTensors.get("output").toTensor();

// 移除所有尺寸为1的维度,得到(224, 416)的Tensor
Tensor squeezedOutput = outputTensor.squeeze();

// 转换为float数组后转成uint8格式
float[] outputFloatArr = squeezedOutput.getDataAsFloatArray();
byte[] prediction = new byte[outputFloatArr.length];
for (int i = 0; i < outputFloatArr.length; i++) {
    // 注意根据你的模型输出范围调整转换逻辑,这里假设输出是0-255的浮点值
    prediction[i] = (byte) Math.round(outputFloatArr[i]);
}

如果你只想移除特定维度(比如只移除第0维或第1维),也可以使用squeeze(int dim)方法,例如squeezedOutput = outputTensor.squeeze(0).squeeze(0);,效果和无参数的squeeze()一致。

三、额外注意事项

  • Tensor是资源密集型对象,使用完毕后记得调用close()方法释放内存,避免泄漏:
    inputTensor.close();
    inputWithBatch.close();
    outputTensor.close();
    squeezedOutput.close();
    
  • 确保模型的输入输出维度和你调整后的Tensor完全匹配,若维度不匹配可能会导致运行时错误或输出结果异常。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 10:42:33