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
相关产品推荐
相关产品推荐

