如何将JavaScript手写数字像素数组传入PyScript实现KNN识别?
解决JavaScript与PyScript的像素数组通信问题
核心思路
通过PyScript将Python的数字识别函数暴露给JavaScript,JavaScript直接调用该函数并传递处理后的像素数组,无需通过文件中转(原代码中的下载图片步骤可保留为可选功能)。
第一步:修改JavaScript代码
- 修正像素值格式:
getGridValues函数直接返回数值数组(无需转二进制),匹配Python端的数据格式。 - 实现数据传递逻辑:调用PyScript暴露的Python预测函数,传递展平后的像素数组并接收识别结果。
修改后的JavaScript代码:
document.addEventListener('DOMContentLoaded', function() { const canvas = document.getElementById('gridCanvas'); const ctx = canvas.getContext('2d'); const gridSize = 28; const cellSize = canvas.width / gridSize; let isDrawing = false; // 绘制网格 function drawGrid() { ctx.beginPath(); for (let i = 1; i < gridSize; i++) { const pos = i * cellSize; ctx.moveTo(pos, 0); ctx.lineTo(pos, canvas.height); ctx.moveTo(0, pos); ctx.lineTo(canvas.width, pos); } ctx.strokeStyle = '#ccc'; ctx.stroke(); } // 清空画布 function clearGrid() { ctx.clearRect(0, 0, canvas.width, canvas.height); drawGrid(); } // 获取28x28像素数组(255=未绘制,0=已绘制) function getGridValues() { const gridValues = []; for (let y = 0; y < gridSize; y++) { const row = []; for (let x = 0; x < gridSize; x++) { const imageData = ctx.getImageData(x * cellSize, y * cellSize, 1, 1).data; const isLightColor = imageData[0] + imageData[1] + imageData[2] > 255 * 3 / 2; row.push(isLightColor ? 255 : 0); // 直接返回数值,无需二进制转换 } gridValues.push(row); } return gridValues; } // 传递像素数组给PyScript并获取预测结果 function sendGridValues(gridValues) { const flattened = gridValues.flat(); // 展平成一维数组,匹配KNN输入格式 window.predict_digit(flattened).then(result => { alert(`识别的数字是:${result}`); console.log(`识别结果:${result}`); }).catch(error => { console.error('预测失败:', error); }); } // 处理绘制逻辑 function handleDraw(event) { if (!isDrawing) return; const x = Math.floor(event.offsetX / cellSize); const y = Math.floor(event.offsetY / cellSize); ctx.fillRect(x * cellSize, y * cellSize, cellSize, cellSize); } // 初始化画布 drawGrid(); // 绑定事件 canvas.addEventListener('mousedown', () => isDrawing = true); canvas.addEventListener('mousemove', handleDraw); canvas.addEventListener('mouseup', () => isDrawing = false); canvas.addEventListener('mouseleave', () => isDrawing = false); document.getElementById('clearButton').addEventListener('click', clearGrid); // 点击"Procesar"直接传递数据,无需下载图片 document.getElementById('processButton').addEventListener('click', () => sendGridValues(getGridValues())); });
第二步:修改PyScript代码
- 提前训练模型:页面加载时完成MNIST数据加载和模型训练,避免每次预测重复加载数据。
- 暴露预测函数:将数字识别函数挂载到JavaScript的
window对象,允许JS直接调用。
修改后的PyScript代码:
import time import numpy as np from PIL import Image import js # 用于和JavaScript交互 DATA_DIR = r"C:/Users/lucas/Downloads/OCR/" TEST_DATA_FILENAME = DATA_DIR + "t10k-images.idx3-ubyte" TEST_LABELS_FILENAME = DATA_DIR + "t10k-labels.idx1-ubyte" TRAIN_DATA_FILENAME = DATA_DIR + "train-images.idx3-ubyte" TRAIN_LABELS_FILENAME = DATA_DIR + "train-labels.idx1-ubyte" DEBUG = True def bytes_to_int(byte_data): return int.from_bytes(byte_data, "big") def read_labels(filename, n_max_labels=None): labels = [] with open(filename, "rb") as f: _ = f.read(4) n_labels = bytes_to_int(f.read(4)) if n_max_labels: n_labels = n_max_labels for _ in range(n_labels): label = f.read(1) labels.append(label) return labels count = 0 def read_images(filename, n_max_images=None): global count images = [] with open(filename, "rb") as f: _ = f.read(4) n_images = bytes_to_int(f.read(4)) if n_max_images: n_images = n_max_images n_rows = bytes_to_int(f.read(4)) n_columns = bytes_to_int(f.read(4)) for _ in range(n_images): image = [] for _ in range(n_rows): row = [] for _ in range(n_columns): count += 1 pixel = f.read(1) row.append(pixel) image.append(row) images.append(image) return images def aplanar_lista(l): return [pixel for sublist in l for pixel in sublist] def pasar_lista_unidimensional(X): return [aplanar_lista(sample) for sample in X] def dist(x, y): return sum((bytes_to_int(x_i) - bytes_to_int(y_i)) **2 for x_i,y_i in zip(x,y))**0.5 def distancia_entre_samples(X_train, test_sample): return [dist(train_sample, test_sample) for train_sample in X_train] def most_frequent_element(list): return max(list, key=list.count) def knn(X_train, y_train, X_test, k=3): y_pred = [] for test_sample in X_test: training_distances = distancia_entre_samples(X_train, test_sample) sorted_distance_indices = [pair[0] for pair in sorted(enumerate(training_distances), key=lambda x: x[1])] candidates = [bytes_to_int(y_train[idx]) for idx in sorted_distance_indices[:k]] top_candidate = most_frequent_element(candidates) y_pred.append(top_candidate) return y_pred # 页面加载时提前训练模型 start_time = time.time() X_train = read_images(TRAIN_DATA_FILENAME, 1000) y_train = read_labels(TRAIN_LABELS_FILENAME, 1000) X_train = pasar_lista_unidimensional(X_train) end_time = time.time() print(f"模型训练完成,耗时:{round(end_time-start_time,5)} 秒") print(f"迭代次数:{count}") # 暴露给JavaScript的预测函数 def predict_digit(pixel_array): # 将JS传来的数值数组转换为训练数据使用的bytes格式 test_sample = [int(p).to_bytes(1, 'big') for p in pixel_array] X_test = [test_sample] y_pred = knn(X_train, y_train, X_test, k=5) return y_pred[0] # 将函数挂载到window对象,让JavaScript可以调用 js.window.predict_digit = predict_digit
关键说明
- 数据格式匹配:Python训练数据中的像素是
bytes类型,因此JS传递的数值需要转换为bytes后再输入KNN算法。 - 性能优化:提前训练模型并保存
X_train和y_train,避免每次预测重复加载MNIST数据,提升响应速度。 - 冗余步骤处理:原代码中的下载图片逻辑可保留作为手绘数字导出功能,但预测时无需通过文件中转,直接传递数组效率更高。
内容的提问来源于stack exchange,提问作者Luca Siegel
相关产品推荐
相关产品推荐

