Keras手写数字识别模型在React画布预测时始终输出相同结果排查
手写数字识别模型预测结果固定问题排查
问题描述
我是机器学习新手,使用MNIST数据集训练了一个简单的Keras神经网络用于识别手写数字,模型代码如下:
model = keras.Sequential([ keras.layers.Flatten(input_shape=(28,28)), keras.layers.Dense(100,activation="relu"), keras.layers.Dense(10,activation='sigmoid') ]) model.compile( optimizer='adam', loss = 'sparse_categorical_crossentropy', metrics=['accuracy'] ) model.fit(x_train,y_train,epochs=8)
随后开发了React应用,用户可在画布上绘图并点击提交按钮让模型进行预测。通过TensorFlow.js导入模型后发现,无论在画布上绘制什么内容,模型都会输出相同的预测结果。React代码如下:
import React, { useRef, useEffect } from 'react' import * as tf from '@tensorflow/tfjs'; async function generatePrediction(image) { let prediction = 0; const model = await tf.loadLayersModel('/model.json'); let step1 = tf.browser.fromPixels(image) .resizeNearestNeighbor([28, 28]).mean(2).toFloat().expandDims(0).div(255.0) prediction = model.predict(step1); let digits = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]; prediction.data() .then((data) => { console.log(data); let max_val = -1 let max_val_index = -1 for (let i = 0; i < data.length; i++) { if (data[i] > max_val) { max_val = data[i] max_val_index = i } } var result = digits[max_val_index]; console.log(result); }) } const Canvas = props => { const canvasRef = useRef(null) var image = new Image(); useEffect(() => { const canvas = canvasRef.current const ctx = canvas.getContext('2d') ctx.strokeRect(0, 0, canvas.width, canvas.height); //Our first draw let isDrawing = false; //Check if drawing //If drawing then make a line from the beginning to the mouse. const drawing = (e) => { if (!isDrawing) { return; } ctx.lineTo(e.offsetX, e.offsetY); ctx.stroke(); ctx.lineWidth = 10; } //Mouse down start draw const startDraw = () => { isDrawing = true; ctx.beginPath(); } //Set drawing back to false when done drawing const stopDraw = () => { isDrawing = false; } ///Mouse Event listeners canvas.addEventListener("mousemove", drawing); canvas.addEventListener("mouseup", stopDraw); canvas.addEventListener("mousedown", startDraw); setInterval(function () { image.src = canvas.toDataURL() }, 100); }, []) function makePrediction() { generatePrediction(image); } return ( <div> <canvas id="myCanvas" width="280" height="280" ref={canvasRef} {...props} /> <h1> The Guessed Number Is: </h1> <button onClick={makePrediction}>Submit</button> </div> ) } export default Canvas
即使重新绘制新数字并点击提交,预测值仍完全相同,怀疑是传入的图像格式存在问题,请问该问题的原因是什么?
问题原因分析
1. 图像明暗与训练集完全相反
MNIST数据集的手写数字是黑底白字(背景像素值为0,数字像素值为255),但你的画布是白底黑字(背景为白色,数字为黑色)。直接转灰度后,数字部分的像素值接近0,和模型训练时的输入特征分布完全相反,模型无法提取有效特征,因此输出固定结果。
2. 图片未加载完成就传入模型
你通过setInterval定时更新image.src,但调用generatePrediction时,图片可能还没完成加载。tf.browser.fromPixels处理未加载完成的图片会生成全0张量,模型每次接收的输入都是一样的,自然输出相同的预测结果。
3. 模型重复加载的潜在问题
每次调用generatePrediction都重新加载模型,不仅浪费资源,还可能导致模型状态不稳定,虽然这不是预测固定的直接原因,但会影响整体运行效率和一致性。
4. 画布绘图的颜色配置问题
画布默认的绘图颜色是黑色,背景为白色,进一步加剧了与MNIST训练数据的格式差异,让模型无法识别输入的数字特征。
内容的提问来源于stack exchange,提问作者aang88
相关产品推荐
相关产品推荐

