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

如何将JavaScript手写数字像素数组传入PyScript实现KNN识别?

解决JavaScript与PyScript的像素数组通信问题

核心思路

通过PyScript将Python的数字识别函数暴露给JavaScript,JavaScript直接调用该函数并传递处理后的像素数组,无需通过文件中转(原代码中的下载图片步骤可保留为可选功能)。


第一步:修改JavaScript代码

  1. 修正像素值格式:getGridValues函数直接返回数值数组(无需转二进制),匹配Python端的数据格式。
  2. 实现数据传递逻辑:调用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代码

  1. 提前训练模型:页面加载时完成MNIST数据加载和模型训练,避免每次预测重复加载数据。
  2. 暴露预测函数:将数字识别函数挂载到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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 16:44:57