PyTorch-TorchVision模型在Python与Kotlin中输出结果不一致排查
我使用移除分类头的MobileNet作为相似度搜索模型,以TorchScript格式保存。Python环境下相似度搜索结果正确,但Kotlin环境用相同图片提取的特征差异极大,怀疑是预处理环节问题但未解决。
代码与输出对比
Python代码
# Model and transform setup device = torch.device("cuda" if torch.cuda.is_available() else "cpu") image_transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) def load_trained_model(): model = torch.jit.load(MODEL_FILE, map_location=device) model.eval() return model def extract_features(pil_img, model): with torch.no_grad(): tensor = image_transform(pil_img).unsqueeze(0).to(device) features = model(tensor) if len(features.shape) > 2: features = features.view(features.size(0), -1) return features.cpu().numpy().astype(np.float32) @app.post("/extract_features") async def extract_image_features(image: UploadFile = File(...)): try: image_bytes = await image.read() with Image.open(BytesIO(image_bytes)) as img: processed_img = img.convert("RGB") raw_features = extract_features(processed_img, model)
Kotlin代码
fun preprocessImage(bitmap: Bitmap): Tensor { val rgbBitmap = if (bitmap.config != Bitmap.Config.ARGB_8888) { bitmap.copy(Bitmap.Config.ARGB_8888, true) } else { bitmap } val resizedBitmap = resizeWithAspectRatio(rgbBitmap, 256) val croppedBitmap = centerCrop(resizedBitmap, 224, 224) val mean = floatArrayOf(0.485f, 0.456f, 0.406f) val std = floatArrayOf(0.229f, 0.224f, 0.225f) return TensorImageUtils.bitmapToFloat32Tensor(croppedBitmap, mean, std) } private fun resizeWithAspectRatio(bitmap: Bitmap, targetSize: Int): Bitmap { val width = bitmap.width val height = bitmap.height val scale = if (width < height) { targetSize.toFloat() / width.toFloat() } else { targetSize.toFloat() / height.toFloat() } val newWidth = (width * scale).roundToInt() val newHeight = (height * scale).roundToInt() // Use createScaledBitmap with bilinear filtering (matches PyTorch default) return bitmap.scale(256, 256) } private fun centerCrop(bitmap: Bitmap, targetWidth: Int, targetHeight: Int): Bitmap { val width = bitmap.width val height = bitmap.height val startX = (width - targetWidth) / 2 val startY = (height - targetHeight) / 2 val validStartX = Math.max(0, startX) val validStartY = Math.max(0, startY) val validTargetWidth = Math.min(targetWidth, width - validStartX) val validTargetHeight = Math.min(targetHeight, height - validStartY) return Bitmap.createBitmap(bitmap, validStartX, validStartY, validTargetWidth, validTargetHeight) } fun extractFeatures(bitmap: Bitmap): FloatArray { if (model == null) { throw IllegalStateException("Model not loaded. Call loadModel() first.") } val inputTensor = preprocessImage(bitmap) val output = model!!.forward(IValue.from(inputTensor)) val outputTensor = output.toTensor() val features = outputTensor.dataAsFloatArray return features }
输出对比
- Kotlin输出前5个值:
[0.7993497, 0.30109355, 0.32214138, 0.47712356, 0.5185487] - Python输出前5个值:
[ 1.2595854 -0.07939269 -0.3717999 0.22528967 0.12919804]
核心问题排查与修复
1. 图片缩放逻辑完全错误
Kotlin的resizeWithAspectRatio方法中,前面计算了正确的保持宽高比的缩放比例和新尺寸,但最后直接返回bitmap.scale(256, 256),强制将图片缩放到256x256,完全破坏了原图片的宽高比。而Python的transforms.Resize(256)是保持宽高比,将短边缩放到256,长边按比例缩放,两者逻辑完全不一致,这是导致特征差异的最直接原因。
修复代码:
private fun resizeWithAspectRatio(bitmap: Bitmap, targetSize: Int): Bitmap { val width = bitmap.width val height = bitmap.height val scale = if (width < height) { targetSize.toFloat() / width.toFloat() } else { targetSize.toFloat() / height.toFloat() } val newWidth = (width * scale).roundToInt() val newHeight = (height * scale).roundToInt() // 保持宽高比缩放,匹配PyTorch的Resize行为 return Bitmap.createScaledBitmap(bitmap, newWidth, newHeight, true) }
2. 图像通道顺序可能不匹配
PyTorch的ToTensor()会将PIL的RGB图像转为CHW格式(通道在前,形状为[3,224,224]),而Android Bitmap默认是HWC格式(高度在前,通道在后)。部分版本的TensorImageUtils.bitmapToFloat32Tensor可能输出HWC格式,和模型输入要求的CHW不匹配,导致特征提取错误。
验证与修复:
打印预处理后tensor的shape,若为[1,224,224,3](HWC),则需要转换为CHW格式:
fun preprocessImage(bitmap: Bitmap): Tensor { // ... 其他代码保持不变 ... val hwcTensor = TensorImageUtils.bitmapToFloat32Tensor(croppedBitmap, mean, std) // 转换HWC到CHW,适配模型输入格式 val chwArray = FloatArray(3 * 224 * 224) val hwcArray = hwcTensor.dataAsFloatArray for (h in 0 until 224) { for (w in 0 until 224) { chwArray[0 * 224*224 + h*224 + w] = hwcArray[h*224*3 + w*3 + 0] chwArray[1 * 224*224 + h*224 + w] = hwcArray[h*224*3 + w*3 + 1] chwArray[2 * 224*224 + h*224 + w] = hwcArray[h*224*3 + w*3 + 2] } } return Tensor.fromBlob(chwArray, longArrayOf(1, 3, 224, 224)) }
3. 归一化逻辑需严格对齐
PyTorch的Normalize执行的是(像素值/255 - mean) / std,即先将[0,255]的像素值归一化到[0,1],再做均值减法和标准差除法。需要确认TensorImageUtils.bitmapToFloat32Tensor是否执行了相同的操作:
- 若该方法直接使用[0,255]的像素值做归一化,会导致计算结果偏差极大
- 可手动实现归一化确保逻辑一致:
fun preprocessImage(bitmap: Bitmap): Tensor { // ... 缩放、裁剪代码保持不变 ... val mean = floatArrayOf(0.485f, 0.456f, 0.406f) val std = floatArrayOf(0.229f, 0.224f, 0.225f) val pixels = IntArray(224 * 224) croppedBitmap.getPixels(pixels, 0, 224, 0, 0, 224, 224) val chwArray = FloatArray(3 * 224 * 224) for (i in pixels.indices) { val pixel = pixels[i] // 提取RGB并转为[0,1] val r = ((pixel shr 16) and 0xFF) / 255.0f val g = ((pixel shr 8) and 0xFF) / 255.0f val b = (pixel and 0xFF) / 255.0f // 执行归一化 chwArray[0 * 224*224 + i] = (r - mean[0]) / std[0] chwArray[1 * 224*224 + i] = (g - mean[1]) / std[1] chwArray[2 * 224*224 + i] = (b - mean[2]) / std[2] } return Tensor.fromBlob(chwArray, longArrayOf(1, 3, 224, 224)) }
4. 确保模型处于评估模式
Python中调用了model.eval()关闭了Dropout、BatchNorm等训练时的行为,Kotlin端加载模型时也需要设置评估模式,避免推理时的随机行为:
// 加载模型后添加 model?.eval()
验证步骤
- 先修复缩放逻辑,这是最核心的错误
- 打印Python和Kotlin预处理后tensor的前10个值,确认是否一致
- 确认模型输入的通道格式(CHW)和归一化逻辑完全匹配
- 确保两边模型都处于
eval模式
内容的提问来源于stack exchange,提问作者URFMODEG

