一、引言
本文基于 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("========程序退出========")
输出:
