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

PyTorch-TorchVision模型在Python与Kotlin中输出结果不一致排查

问题:Python与Kotlin端MobileNet特征提取结果差异巨大

我使用移除分类头的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()

验证步骤

  1. 先修复缩放逻辑,这是最核心的错误
  2. 打印Python和Kotlin预处理后tensor的前10个值,确认是否一致
  3. 确认模型输入的通道格式(CHW)和归一化逻辑完全匹配
  4. 确保两边模型都处于eval模式

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 03:07:05