python:HandWriting Recognition using paddlepaddle

一、引言

本文基于 PaddlePaddle 与 DDD 分层架构,实现一个可运行的手写数字识别桌面应用。项目覆盖领域建模、应用服务编排、基础设施封装与 Tkinter 图形界面,并支持识别结果 CSV 导出与统计分析。

二、DDD 分层架构设计

项目采用 DDD 分层思想,将代码划分为 domain、application、infrastructure、common 四层,各层职责与文件清单如下:

三、领域层实现

领域层承载核心模型与业务规则,包括卷积网络实体与预测结果值对象,不依赖任何外部框架。

四、应用层实现

应用层负责业务流程编排,通过训练应用服务与预测应用服务串联领域模型、基础设施与接口层。

五、基础设施层实现

基础设施层封装数据集加载、模型持久化、CSV 导出与统计分析等能力,向上层提供仓储与数据访问实现。

六、通用支撑层实现

通用支撑层提供跨层复用的日志组件与统一异常处理装饰器,为各层提供横向支撑。

七、接口层与 GUI 实现

接口层包含主启动器与手写数字画板 GUI,负责用户交互,业务逻辑交由应用服务处理。

八、运行与调用

本节给出项目统一入口 main.py 的调用方式,以及训练、识别、CSV 导出与统计分析的完整运行流程。

九、总结

本文通过一个完整的手写数字识别项目,展示了 DDD 分层架构在 AI 桌面应用中的落地实践,并总结了各层协作方式与工程化收益。

项目结构:

项目采用 DDD 分层架构,各层职责与核心文件清单如下:

分层 职责说明 核心文件
domain(领域层) 承载领域核心模型与业务规则,不依赖任何外部框架,是系统最稳定的内核。 model_entity.py、value_object.py
application(应用层) 负责业务流程编排与用例驱动,协调领域模型、基础设施与接口层完成具体业务场景。 train_app_service.py、predict_app_service.py
infrastructure(基础设施层) 封装数据持久化、外部依赖与工具能力,向上层提供仓储与数据访问实现。 dataset_provider.py、model_repository.py、csv_exporter.py、csv_stat_analyzer.py
common(通用支撑层) 提供跨层复用的通用组件,如日志记录与统一异常处理,为各层提供横向支撑。 app_logger.py、exception_handler.py
python 复制代码
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:24
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : model_entity.py
# domain/model_entity.py
"""
领域实体:MNIST手写数字卷积模型实体
DDD:领域实体,封装神经网络结构,代表手写数字识别领域模型
"""
import paddle
import paddle.nn as nn
class MNISTModelEntity(nn.Layer):
"""
MNIST手写数字卷积网络领域实体
输入:1通道28*28灰度张量 [B,1,28,28]
输出:10维logits,对应0‑9分类
"""
def init(self):
super(MNISTModelEntity, self).init()
# 卷积层1:输入通道1,输出通道16,卷积核3,padding=1保持尺寸
self.conv1 = nn.Conv2D(in_channels=1, out_channels=16, kernel_size=3, padding=1)
# 激活函数ReLU
self.relu1 = nn.ReLU()
# 最大池化,核2步长2,尺寸减半
self.pool1 = nn.MaxPool2D(kernel_size=2, stride=2)
# self.pool = nn.MaxPool2D(2, 2)
# 卷积层2
self.conv2 = nn.Conv2D(in_channels=16, out_channels=32, kernel_size=3, padding=1)
self.relu2 = nn.ReLU()
self.pool2 = nn.MaxPool2D(kernel_size=2, stride=2)
    # 展平,把特征图转为一维向量
    self.flatten = nn.Flatten()
    # 全连接层1
    self.fc1 = nn.Linear(32 * 7 * 7, 128)
    # 输出层:10分类0‑9
    self.fc2 = nn.Linear(128, 10)
def forward(self, x: paddle.Tensor) -> paddle.Tensor:
"""
前向传播
:param x: 输入张量 shape [B,1,28,28]
:return: logits张量 shape [B,10],未做softmax
"""
x = self.conv1(x)
x = self.relu1(x)
x = self.pool1(x)
x = self.conv2(x)
x = self.relu2(x)
x = self.pool2(x)

x = self.flatten(x)
x = self.fc1(x)
x = self.fc2(x)
'''
x = self.pool(paddle.nn.functional.relu(self.conv1(x)))
x = self.pool(paddle.nn.functional.relu(self.conv2(x)))
x = paddle.flatten(x, start_axis=1)
x = paddle.nn.functional.relu(self.fc1(x))
x = self.fc2(x)
'''
return x
def set_state_dict(self, state_dict):
super().set_state_dict(state_dict)
def state_dict(self):
return super().state_dict()
encoding: utf-8
版权所有  2026 ©涂聚文有限公司™ ®
许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
描述:
Author    : geovindu,Geovin Du 涂聚文.
IDE       : PyCharm 2024.3.6 python 3.11
os        : windows 10
database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
Datetime  : 2026/9/3 16:24
User      :  geovindu
Product   : PyCharm
Project   : PyHandWritingRecognition
File      : value_object.py
domain/value_object.py
"""
DDD 值对象:预测结果VO,不可变对象,封装预测标签、各分类置信概率
"""
from dataclasses import dataclass
from typing import List
import time
@dataclass(frozen=True)
class PredictResultVO:
"""
预测返回值对象
:param pred_label: int, 预测数字 0‑9
:param confidence: float, 最大置信度 0~1
:param prob_list: List[float], 索引0‑9对应每个数字的softmax概率
:param timestamp: float 识别时间戳
"""
pred_label: int
confidence: float
prob_list: List[float]
timestamp: float
def get_prob_text(self) -> str:
"""
格式化输出全部0‑9置信度文本,用于GUI显示
:return:
"""
lines = []
for digit, p in enumerate(self.prob_list):
lines.append(f"{digit}:{p:.4f}")
return " ".join(lines)
def to_csv_row(self) -> list:
"""
转换为csv行数据
:return:
"""
row = [self.timestamp, self.pred_label, self.confidence]
row.extend(self.prob_list)
return row
encoding: utf-8
版权所有  2026 ©涂聚文有限公司™ ®
许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
描述:
Author    : geovindu,Geovin Du 涂聚文.
IDE       : PyCharm 2024.3.6 python 3.11
os        : windows 10
database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
Datetime  : 2026/9/3 16:19
User      :  geovindu
Product   : PyCharm
Project   : PyHandWritingRecognition
File      : app_logger.py
common/app_logger.py
common/app_logger.py
"""
日志组件:按级别输出,同时输出控制台+日志文件
日志文件按日期存放 logs/yyyy-mm-dd.log
日志级别 DEBUG < INFO < WARNING < ERROR < CRITICAL
"""
import logging
import os
from datetime import datetime
from typing import Optional
日志根目录
LOG_DIR = "logs"
if not os.path.exists(LOG_DIR):
os.makedirs(LOG_DIR, exist_ok=True)
def get_logger(name: str, level: int = logging.INFO) -> logging.Logger:
"""
获取logger实例
:param name: 日志器名称
:param level: 日志级别 logging.INFO / logging.DEBUG
:return: Logger对象
"""
logger = logging.getLogger(name)
logger.setLevel(level)
防止重复添加handler
if logger.handlers:
return logger
关键修复:使用普通减号 '-',不要使用 ‑ 软连字符
log_date_str = datetime.now().strftime("%Y-%m-%d")
log_file_path = os.path.join(LOG_DIR, f"{log_date_str}.log")
文件输出格式
file_formatter = logging.Formatter(
"%(asctime)s | %(levelname)s | %(name)s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S"
)
控制台输出格式
console_formatter = logging.Formatter(
"%(asctime)s | %(levelname)s | %(message)s",
datefmt="%H:%M:%S"
)
文件Handler
file_handler = logging.FileHandler(log_file_path, encoding="utf-8")
file_handler.setFormatter(file_formatter)
控制台Handler
console_handler = logging.StreamHandler()
console_handler.setFormatter(console_formatter)
logger.addHandler(file_handler)
logger.addHandler(console_handler)
return logger
encoding: utf-8
版权所有  2026 ©涂聚文有限公司™ ®
许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
描述:
Author    : geovindu,Geovin Du 涂聚文.
IDE       : PyCharm 2024.3.6 python 3.11
os        : windows 10
database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
Datetime  : 2026/9/3 16:23
User      :  geovindu
Product   : PyCharm
Project   : PyHandWritingRecognition
File      : exception_handler.py
common/exception_handler.py
"""
全局统一异常装饰器
捕获异常,记录日志,向上抛出业务友好异常信息
"""
import functools
from typing import Callable
from common.app_logger import get_logger
logger = get_logger("global_exception")
def global_exception_handler(raise_ui_msg: bool = True):
"""
统一异常处理装饰器
:param raise_ui_msg: True返回异常消息,用于UI弹窗提示
"""
def decorator(func: Callable):
@functools.wraps(func)
def wrapper(*args, **kwargs):
try:
return func(*args, **kwargs)
except Exception as e:
logger.error(f"执行函数[{func.name}]发生异常", exc_info=True)
调试阶段:直接raise原始异常,不要包装RuntimeError
raise
'''
def wrapper(*args, **kwargs):
try:
return func(*args, **kwargs)
except Exception as e:
logger.error(f"执行函数[{func.name}]发生异常", exc_info=True)
raise RuntimeError(f"业务异常:{str(e)}") from e
'''
return wrapper
return decorator
python 复制代码
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:26
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : csv_exporter.py
# infrastructure/csv_exporter.py
"""
基础设施层:识别结果CSV导出仓储
每次识别结果追加写入csv,自动创建output目录、表头
"""
import os
import csv
from domain.value_object import PredictResultVO
from common.app_logger import get_logger
 
logger = get_logger("csv_exporter")
 
OUTPUT_DIR = "output"
if not os.path.exists(OUTPUT_DIR):
    os.makedirs(OUTPUT_DIR, exist_ok=True)
 
 
class PredictCsvExporter:
    """
    预测结果CSV导出器
    """
    def __init__(self, csv_filename: str = "predict_records.csv"):
        """
        :param csv_filename:输出csv文件名
        :return:
        """
        self._file_path = os.path.join(OUTPUT_DIR, csv_filename)
        # csv表头
        self._headers = [
            "timestamp",
            "pred_label",
            "confidence",
            "prob_0", "prob_1", "prob_2", "prob_3", "prob_4",
            "prob_5", "prob_6", "prob_7", "prob_8", "prob_9"
        ]
        self._init_file()
 
    def _init_file(self):
        """
        文件不存在则创建并写入表头
        :return:
        """
        if os.path.exists(self._file_path):
            return
        with open(self._file_path, mode="w", encoding="utf‑8‑sig", newline="") as f:
            writer = csv.writer(f)
            writer.writerow(self._headers)
        logger.info(f"CSV记录文件初始化完成:{self._file_path}")
 
    def append_record(self, vo: PredictResultVO):
        """
        追加一条识别记录
        :param vo:预测结果值对象
        :return:
        """
        row = vo.to_csv_row()
        with open(self._file_path, mode="a", encoding="utf‑8‑sig", newline="") as f:
            writer = csv.writer(f)
            writer.writerow(row)
        logger.debug(f"识别记录写入CSV,预测数字:{vo.pred_label},置信度:{vo.confidence:.4f}")
 
 
 
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:34
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : csv_stat_analyzer.py
# infrastructure/csv_stat_analyzer.py
"""
基础设施层:CSV识别记录统计分析器
读取output下predict_records.csv,做简单统计:总识别条数、各数字识别次数、平均置信度
"""
import os
import csv
from typing import Dict, Optional, Tuple
from common.app_logger import get_logger
from infrastructure.csv_exporter import OUTPUT_DIR
 
logger = get_logger("csv_stat_analyzer")
 
 
class CsvStatAnalyzer:
    """识别记录统计分析"""
    def __init__(self, csv_filename: str = "predict_records.csv"):
        """
        :param csv_filename: csv记录文件名
        """
        self._file_path = os.path.join(OUTPUT_DIR, csv_filename)
 
    def is_file_exists(self) -> bool:
        """判断记录文件是否存在"""
        return os.path.exists(self._file_path)
 
    def get_stat(self) -> Optional[Tuple[int, Dict[int, int], float]]:
        """
        获取统计结果
        :return: (总条数, {数字:出现次数}, 平均置信度);文件不存在返回None
        """
        if not self.is_file_exists():
            logger.warning("统计分析:CSV记录文件不存在")
            return None
 
        total_count = 0
        digit_counter: Dict[int, int] = {i: 0 for i in range(10)}
        sum_confidence = 0.0
 
        try:
            with open(self._file_path, mode="r", encoding="utf-8-sig", newline="") as f:
                reader = csv.DictReader(f)
                for row in reader:
                    total_count += 1
                    pred_digit = int(row["pred_label"])
                    conf = float(row["confidence"])
                    digit_counter[pred_digit] += 1
                    sum_confidence += conf
        except Exception as e:
            logger.error(f"读取CSV统计异常:{e}", exc_info=True)
            return None
 
        avg_conf = sum_confidence / total_count if total_count > 0 else 0.0
        return total_count, digit_counter, avg_conf
 
    def build_stat_text(self) -> str:
        """
        构建用于GUI弹窗展示的统计文本
        """
        stat_result = self.get_stat()
        if stat_result is None:
            return "暂无识别记录,请先完成识别生成CSV文件。"
 
        total_count, digit_counter, avg_conf = stat_result
        lines = []
        lines.append(f"📊识别记录统计")
        lines.append(f"总识别记录数:{total_count}")
        lines.append(f"平均置信度:{avg_conf:.4f}")
        lines.append("-"*30)
        for d in range(10):
            lines.append(f"数字{d} : {digit_counter[d]} 次")
        return "\n".join(lines)
 
    def get_csv_filepath(self) -> str:
        """获取csv完整路径,用于打开文件"""
        return self._file_path
 
 
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:25
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : dataset_provider.py
# infrastructure/dataset_provider.py
"""
基础设施层:数据集提供器,封装MNIST数据集加载逻辑
隔离外部数据集依赖,上层应用不直接感知paddle.vision.datasets
"""
import paddle
import paddle.vision.transforms as T
from paddle.vision.datasets import MNIST
from common.app_logger import get_logger
 
logger = get_logger("dataset_provider")
 
 
class MnistDatasetProvider:
    """
    MNIST数据集加载提供器
    """
    @staticmethod
    def get_train_dataloader(batch_size: int = 128, shuffle: bool = True) -> paddle.io.DataLoader:
        """
        获取训练集DataLoader
        :param batch_size:批次大小
        :param shuffle:是否打乱顺序
        :return: DataLoader
        """
        transform = T.Compose([T.ToTensor()])
        train_ds = MNIST(mode="train", transform=transform)
        loader = paddle.io.DataLoader(train_ds, batch_size=batch_size, shuffle=shuffle)
        logger.info(f"训练数据集加载完成 batch_size={batch_size}")
        return loader
 
    @staticmethod
    def get_test_dataloader(batch_size: int = 128, shuffle: bool = False) -> paddle.io.DataLoader:
        """
        获取测试集DataLoader
        :param batch_size:批次大小
        :param shuffle:是否打乱
        :return: DataLoader
        """
        transform = T.Compose([T.ToTensor()])
        test_ds = MNIST(mode="test", transform=transform)
        loader = paddle.io.DataLoader(test_ds, batch_size=batch_size, shuffle=shuffle)
        logger.info(f"测试数据集加载完成 batch_size={batch_size}")
        return loader
 
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:26
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : model_repository.py
# infrastructure/model_repository.py
"""
基础设施层:模型仓储,负责模型权重持久化保存与加载
DDD仓储模式:隔离文件IO操作,领域层不感知存储路径细节
"""
import os
import paddle
from domain.model_entity import MNISTModelEntity
from common.app_logger import get_logger
 
logger = get_logger("model_repository")
 
 
class ModelRepository:
    """
    模型仓储,负责保存、加载模型参数
    """
    def __init__(self, model_path: str = "mnist_handwrite.pdparams"):
        """
        :param model_path:模型权重文件路径
        :return:
        """
        self._model_path = model_path
 
    def save(self, model_entity: MNISTModelEntity) -> None:
        """
        保存模型实体权重到磁盘
        :param model_entity: MNIST模型领域实体
        :return:
        """
        paddle.save(model_entity.state_dict(), self._model_path)
        logger.info(f"模型权重已保存至:{self._model_path}")
 
    def load(self, model_entity: MNISTModelEntity) -> bool:
        """
        加载权重到模型实体
        :param model_entity: MNIST模型领域实体
        :return: True加载成功;False文件不存在
        """
        '''
        if not os.path.exists(self._model_path):
            logger.warning(f"模型权重文件不存在:{self._model_path}")
            return False
        state_dict = paddle.load(self._model_path)
        model_entity.set_state_dict(state_dict)
        logger.info(f"模型权重加载完成:{self._model_path}")
        '''
        if not os.path.exists(self._model_path):
            logger.warning(f"模型权重文件不存在:{self._model_path}")
            return False
        state_dict = paddle.load(self._model_path)
        model_entity.set_state_dict(state_dict)
        logger.info(f"模型权重加载完成:{self._model_path}")
        return True
 
 
    def exists(self) -> bool:
        """
        判断权重文件是否存在
        :return:
        """
        return os.path.exists(self._model_path)
 
    def save_moodel(self, model):
        """
        保存模型权重
        :param model:
        :return:
        """
        save_path = "mnist_handwrite.pdparams"
        model.save_dict(save_path)
        logger.info(f"模型权重保存成功:{save_path}")
python 复制代码
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:30
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : predict_app_service.py
# application/predict_app_service.py
"""
应用层:预测应用服务
输入图像张量,调用领域模型,输出PredictResultVO值对象,返回0‑9全部置信概率
集成CSV导出
"""
import paddle
import paddle.nn.functional as F
import numpy as np
import time
from domain.model_entity import MNISTModelEntity
from domain.value_object import PredictResultVO
from infrastructure.model_repository import ModelRepository
from infrastructure.csv_exporter import PredictCsvExporter
from common.app_logger import get_logger
from common.exception_handler import global_exception_handler
 
logger = get_logger("predict_app_service")
 
 
class PredictAppService:
    """预测应用服务:加载模型,执行预测,输出带完整置信度的值对象"""
    def __init__(self):
        self._model = MNISTModelEntity()
        self._repo = ModelRepository()
        self._csv_exporter = PredictCsvExporter()
        # 加载权重
        load_ok = self._repo.load(self._model)
        if load_ok:
            self._model.eval()
        else:
            raise FileNotFoundError("预测服务启动失败,模型权重文件不存在,请先执行训练")
 
    @global_exception_handler()
    def predict(self, img_tensor: paddle.Tensor) -> PredictResultVO:
        """
        单张图像预测
        :param img_tensor: 输入张量 shape [1,1,28,28]
        :return: PredictResultVO 值对象,包含预测标签、最大置信度、0‑9全部概率列表
        """
        with paddle.no_grad():
            logits = self._model(img_tensor)
            # softmax转为概率 0‑1
            prob_tensor = F.softmax(logits, axis=1)
            pred_idx = paddle.argmax(prob_tensor, axis=1).numpy()[0]
            prob_np = prob_tensor.numpy()[0]
            max_confidence = float(prob_np[pred_idx])
            prob_list = [float(x) for x in prob_np]
 
        vo = PredictResultVO(
            pred_label=int(pred_idx),
            confidence=max_confidence,
            prob_list=prob_list,
            timestamp=time.time()
        )
        logger.info(f"完成预测,数字:{vo.pred_label},置信度:{vo.confidence:.4f}")
        # 导出CSV
        self._csv_exporter.append_record(vo)
        return vo
 
    @staticmethod
    def image_to_tensor(img_np: np.ndarray) -> paddle.Tensor:
        """
        numpy灰度图像转为模型输入张量
        :param img_np: shape(28,28) 归一化到0‑1的numpy数组
        :return: tensor shape [1,1,28,28]
        """
        tensor = paddle.to_tensor(img_np).unsqueeze(0).unsqueeze(0)
        return tensor
 
 
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:29
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : train_app_service.py
# application/train_app_service.py
"""
应用层:训练应用服务,编排训练流程
应用服务不包含领域逻辑,只做流程编排:数据集、模型实体、仓储协同
"""
import paddle
import paddle.nn as nn
from domain.model_entity import MNISTModelEntity
from infrastructure.dataset_provider import MnistDatasetProvider
from infrastructure.model_repository import ModelRepository
from common.app_logger import get_logger
from common.exception_handler import global_exception_handler
 
logger = get_logger("train_app_service")
 
 
class TrainAppService:
    """训练应用服务,封装完整训练业务流程"""
    def __init__(self):
        # 实例化领域模型实体
        self._model = MNISTModelEntity()
        # 模型仓储
        self._repo = ModelRepository()
        # 损失函数
        self._loss_func = nn.CrossEntropyLoss()
 
    @global_exception_handler()
    def run(self, epoch_num: int = 5, lr: float = 0.001) -> None:
        """
        执行完整训练流程
        :param epoch_num:训练轮数
        :param lr:学习率
        """
        logger.info("===== 开始MNIST模型训练 =====")
        train_loader = MnistDatasetProvider.get_train_dataloader(batch_size=128, shuffle=True)
        test_loader = MnistDatasetProvider.get_test_dataloader(batch_size=128, shuffle=False)
 
        optimizer = paddle.optimizer.Adam(learning_rate=lr, parameters=self._model.parameters())
 
        for epoch in range(epoch_num):
            self._model.train()
            total_loss = 0.0
            for batch_id, (img, label) in enumerate(train_loader):
                logits = self._model(img)
                loss = self._loss_func(logits, label)
                loss.backward()
                optimizer.step()
                optimizer.clear_grad()
                total_loss += loss.numpy()  #[0]
                if batch_id % 100 == 0:
                    logger.info(f"Epoch:{epoch}, Batch:{batch_id}, Loss:{loss.numpy():.4f}") #[0]
 
            avg_train_loss = total_loss / len(train_loader)
            logger.info(f"Epoch {epoch} 训练平均Loss:{avg_train_loss:.4f}")
 
            # 测试集评估
            self._model.eval()
            correct_count = 0
            total_count = 0
            test_loss = 0.0
            total_num = 0
            with paddle.no_grad():
                for img, label in test_loader:
                    out = self._model(img)
                    logits = self._model(img)
                    loss_t = self._loss_func(out, label)
                    test_loss += loss_t.numpy()
                    pred = paddle.argmax(logits, axis=1)
                    correct_count += paddle.sum(pred == label).numpy() #[0]
                    total_count += img.shape[0]
 
            avg_test_loss = test_loss / len(test_loader)
            acc = correct_count / total_count
            logger.info(f"Epoch {epoch} 测试集准确率:{acc:.4f}")
 
 
        # 全部epoch循环结束之后
        #self._model.save_dict("mnist_handwrite.pdparams")
        #logger.info("训练完成,模型权重已保存 mnist_handwrite.pdparams")
        # 全部epoch结束之后
        try:
            state = self._model.state_dict()
            paddle.save(state, self._repo._model_path)
            logger.info("训练完成,模型权重已保存")
        except Exception as e:
            logger.error("保存模型权重失败", exc_info=True)
            raise
        # 训练完成持久化保存
        #self._repo.save(self._model)
        #logger.info("===== 训练全部完成 =====")
python 复制代码
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:31
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : main_launcher.py
# interface/main_launcher.py
"""
接口层:主启动器窗体,菜单选择1训练、2打开GUI
"""
import tkinter as tk
from tkinter import messagebox
import subprocess
import sys
from application.train_app_service import TrainAppService
from common.app_logger import get_logger
 
logger = get_logger("main_launcher")
 
 
class MainLauncher:
    def __init__(self, root: tk.Tk):
        self.root = root
        self.root.title("DDD‑Paddle手写数字识别|启动器")
        self.root.geometry("460x260")
 
        # 创建菜单栏
        menubar = tk.Menu(root)
        func_menu = tk.Menu(menubar, tearoff=0)
        func_menu.add_command(label="1.执行模型训练", command=self.menu_train)
        func_menu.add_command(label="2.打开手写识别GUI", command=self.menu_open_gui)
        func_menu.add_separator()
        func_menu.add_command(label="退出程序", command=root.quit)
        menubar.add_cascade(label="功能菜单", menu=func_menu)
        root.config(menu=menubar)
 
        tip_text = tk.Label(root, text="DDD分层手写数字识别增强版\n\n菜单1:训练生成模型权重\n菜单2:打开画板识别\n✅输出全部置信概率\n✅自动记录识别结果到CSV\n✅日志输出logs目录\n⚠️必须先训练,再打开GUI",
                            font=("SimHei",10), justify="left")
        tip_text.pack(pady=50, padx=20)
 
    def menu_train(self):
        """菜单1:执行训练,新开子进程运行训练避免阻塞UI"""
        try:
            proc = subprocess.Popen([sys.executable, __file__, "--run‑train"])
            messagebox.showinfo("提示","训练任务已启动,请查看控制台输出日志!")
            logger.info("用户触发启动训练子进程")
        except Exception as e:
            logger.error(f"启动训练失败 {e}", exc_info=True)
            messagebox.showerror("启动训练失败", str(e))
 
    def menu_open_gui(self):
        """菜单2:启动手写识别GUI"""
        try:
            from interface.tk_gui_view import start_gui
            start_gui()
        except Exception as e:
            logger.error(f"打开GUI失败 {e}", exc_info=True)
            messagebox.showerror("打开GUI失败", str(e))
 
 
if __name__ == "__main__":
    import sys
    # 命令行参数:子进程直接执行训练
    if len(sys.argv) > 1 and sys.argv[1] == "--run‑train":
        svc = TrainAppService()
        svc.run(epoch_num=5, lr=0.001)
    else:
        main_win = tk.Tk()
        launcher = MainLauncher(main_win)
        main_win.mainloop()
 
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:30
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : tk_gui_view.py
# interface/tk_gui_view.py
"""
接口层:Tkinter GUI视图,只负责UI交互,业务交给PredictAppService
新增:打开CSV按钮、统计分析弹窗
"""
import tkinter as tk
from tkinter import messagebox
import numpy as np
import os
import platform
import subprocess
from PIL import ImageGrab, Image
from application.predict_app_service import PredictAppService
from infrastructure.csv_stat_analyzer import CsvStatAnalyzer
from common.app_logger import get_logger
from common.exception_handler import global_exception_handler
 
logger = get_logger("tk_gui_view")
 
 
class DrawDigitGUI:
    """手写数字画板GUI视图"""
    def __init__(self, root: tk.Tk):
        """
        :param root: tk主窗体
        """
        self.root = root
        self.root.title("手写数字识别|DDD增强版(置信度+CSV导出)")
        self.root.geometry("760x460")
 
        self.canvas_w = 400
        self.canvas_h = 400
        self.last_x = None
        self.last_y = None
 
        # 左侧画布
        self.canvas = tk.Canvas(root, bg="black", width=self.canvas_w, height=self.canvas_h)
        self.canvas.pack(side="left")
        # ✅正确写法
        self.canvas.bind("<B1-Motion>", self.on_draw)
        self.canvas.bind("<ButtonRelease-1>", self.on_up)
 
        # 右侧面板
        frame_right = tk.Frame(root)
        frame_right.pack(side="right", padx=12)
 
        self.label_result = tk.Label(frame_right, text="预测结果:\n?", font=("SimHei",24))
        self.label_result.pack(pady=10)
 
        self.label_conf = tk.Label(frame_right, text="置信度:0.0000", font=("SimHei",12))
        self.label_conf.pack(pady=5)
 
        self.label_all_prob = tk.Label(frame_right, text="0‑9概率:\n", font=("SimHei",10), wraplength=260)
        self.label_all_prob.pack(pady=10)
 
        btn_recognize = tk.Button(frame_right, text="识别数字", command=self.do_recognize, font=("SimHei",12), width=14)
        btn_recognize.pack(pady=4)
 
        btn_clear = tk.Button(frame_right, text="清空画布", command=self.clear_canvas, font=("SimHei",12), width=14)
        btn_clear.pack(pady=4)
 
        #====【新增按钮】====
        btn_open_csv = tk.Button(frame_right, text="打开历史CSV", command=self.open_csv_file, font=("SimHei",11), width=14)
        btn_open_csv.pack(pady=4)
 
        btn_show_stat = tk.Button(frame_right, text="查看统计分析", command=self.show_stat_dialog, font=("SimHei",11), width=14)
        btn_show_stat.pack(pady=4)
 
        # 初始化预测应用服务
        try:
            self.predict_service = PredictAppService()
        except FileNotFoundError as e:
            messagebox.showerror("初始化错误", str(e))
            logger.error(f"GUI初始化失败 {e}")
            self.predict_service = None
 
        # 统计分析实例
        self.stat_analyzer = CsvStatAnalyzer()
 
    def on_draw(self, event):
        """鼠标拖动绘制线条事件"""
        if self.last_x is None or self.last_y is None:
            self.last_x = event.x
            self.last_y = event.y
        self.canvas.create_line(self.last_x, self.last_y, event.x, event.y, fill="white", width=16, capstyle=tk.ROUND)
        self.last_x = event.x
        self.last_y = event.y
 
    def on_up(self, event):
        """鼠标松开事件"""
        self.last_x = None
        self.last_y = None
 
    def clear_canvas(self):
        """清空画布与显示文本"""
        self.canvas.delete("all")
        self.label_result.config(text="预测结果:\n?")
        self.label_conf.config(text="置信度:0.0000")
        self.label_all_prob.config(text="0‑9概率:\n")
 
    @global_exception_handler()
    def do_recognize(self):
        """执行识别:截图画布‑预处理‑调用应用服务‑渲染结果(含全部置信概率),自动写入CSV"""
        if self.predict_service is None:
            messagebox.showwarning("警告","预测服务未初始化,请先训练模型!")
            return
        # 获取画布截图
        x0 = self.canvas.winfo_rootx()
        y0 = self.canvas.winfo_rooty()
        x1 = x0 + self.canvas_w
        y1 = y0 + self.canvas_h
        img = ImageGrab.grab(bbox=(x0, y0, x1, y1))
        img = img.convert("L")
        img = img.resize((28, 28))
        img_np = np.array(img).astype(np.float32) / 255.0
 
        tensor = self.predict_service.image_to_tensor(img_np)
        vo = self.predict_service.predict(tensor)
 
        # 更新UI显示
        self.label_result.config(text=f"预测结果:\n{vo.pred_label}")
        self.label_conf.config(text=f"置信度:{vo.confidence:.4f}")
        self.label_all_prob.config(text=f"0‑9概率:\n{vo.get_prob_text()}")
        messagebox.showinfo("提示", f"识别完成,结果已写入output/predict_records.csv")
 
    def open_csv_file(self):
        """【新增】使用系统默认程序打开CSV文件"""
        filepath = self.stat_analyzer.get_csv_filepath()
        if not os.path.exists(filepath):
            messagebox.showwarning("提示", "CSV记录文件不存在,请先执行识别。")
            return
        try:
            if platform.system() == "Windows":
                os.startfile(filepath)
            elif platform.system() == "Darwin":
                subprocess.run(["open", filepath])
            else:
                subprocess.run(["xdg‑open", filepath])
            logger.info(f"打开CSV文件:{filepath}")
        except Exception as e:
            logger.error(f"打开CSV失败:{e}", exc_info=True)
            messagebox.showerror("错误", f"无法打开文件:{str(e)}")
 
    def show_stat_dialog(self):
        """【新增】弹出统计分析对话框"""
        stat_text = self.stat_analyzer.build_stat_text()
        messagebox.showinfo("识别记录统计分析", stat_text)
 
 
def start_gui():
    """启动GUI入口函数"""
    win = tk.Tk()
    app = DrawDigitGUI(win)
    win.mainloop()

调用:

python 复制代码
# encoding: utf-8
# 版权所有  2026 ©涂聚文有限公司™ ®
# 许可信息查看:言語成了邀功盡責的功臣,還需要行爲每日來值班嗎
# 描述:pip install paddlepaddle pillow
# Author    : geovindu,Geovin Du 涂聚文.
# IDE       : PyCharm 2024.3.6 python 3.11
# os        : windows 10
# database  : mysql 9.0 sql server 2019, postgreSQL 17.0  Oracle 21c Neo4j
# Datetime  : 2026/9/3 16:19
# User      :  geovindu
# Product   : PyCharm
# Project   : PyHandWritingRecognition
# File      : main.py
main.py 项目统一入口
from interface.main_launcher import MainLauncher
import tkinter as tk
from common.app_logger import get_logger
logger = get_logger("main")
if name == "main":
logger.info("========程序启动========")
window = tk.Tk()
app = MainLauncher(window)
window.mainloop()
logger.info("========程序退出========")

输出:

相关推荐
IT_Octopus35 分钟前
JSON 日志里的 `{“$ref“:“$.xxx“}`:从 fastjson 兼容包到原生 fastjson2 的迁移实录
开发语言·python·json
步行cgn44 分钟前
为何以继承方式引入SpringBoot
java·spring boot·后端
Mr_hou1 小时前
左手.NET右手Node:一个后端开发的双修之路
后端
wuyk5551 小时前
Python 零基础入门第九章:用户输入与 While 循环
开发语言·python
wno7041 小时前
Spring Boot配合Hibernate Validator参数校验
spring boot·后端·hibernate
兮动人1 小时前
Python变量与常量
开发语言·python·机器学习·python变量与常量
风筝在晴天搁浅1 小时前
解题代码清爽版
java·开发语言
2601_962300811 小时前
机器学习的理想基石
人工智能·python·机器学习·编程语言·数据处理
Java后端的Ai之路1 小时前
21、Python - 命令模式
开发语言·人工智能·python·命令模式·外观模式