仓库:MuZiCul/CCRS(原创)· 语言:Python · 定位:Chinese Character Recognition System
⚠️ 重要更正:README 中的代码示例写的是 PyTorch(
nn.Module、torch),但实际代码使用的是 TensorFlow/Keras。以下分析基于真实源码(models/cnn_model.py、train.py)。
基于 CNN 的中文手写汉字识别系统,覆盖数据预处理、模型训练、预测分类的完整流程,并提供 Flask Web 界面。
models/cnn_model.py 使用 Keras Sequential:
import tensorflow as tf
from tensorflow.keras import layers, models
def build_model(input_shape, num_classes):
model = models.Sequential([
# 第一个卷积块
layers.Conv2D(32, (3, 3), activation='relu', input_shape=input_shape),
layers.BatchNormalization(),
layers.MaxPooling2D((2, 2)),
# 第二个卷积块
layers.Conv2D(64, (3, 3), activation='relu'),
layers.BatchNormalization(),
layers.MaxPooling2D((2, 2)),
# 第三个卷积块
layers.Conv2D(128, (3, 3), activation='relu'),
layers.BatchNormalization(),
layers.MaxPooling2D((2, 2)),
# 全连接层
layers.Flatten(),
layers.Dropout(0.5),
layers.Dense(1024, activation='relu'),
layers.Dropout(0.5),
layers.Dense(num_classes, activation='softmax')
])
return model
设计要点:每个卷积块都配 BatchNormalization(加速收敛、稳定训练),全连接层用 双重 Dropout(0.5) 抑制过拟合,输出层 softmax 多分类。
train.py 的 train_model() 有几个实用的工程细节:
if not use_gpu:
tf.config.set_visible_devices([], 'GPU') # 强制 CPU 训练
else:
gpus = tf.config.list_physical_devices('GPU')
for gpu in gpus:
tf.config.experimental.set_memory_growth(gpu, True) # 按需分配显存
# 打印 GPU 详细信息
for gpu in gpus:
gpu_details = tf.config.experimental.get_device_details(gpu)
training_control = {'should_pause': False, 'should_stop': False}
这是个很实用的设计——Web 界面可以随时暂停/中止训练,而不是让训练循环跑到底。配合 progress_callback 回调实时回报进度。
logging.basicConfig(
handlers=[logging.FileHandler('training.log'), logging.StreamHandler()]
)
文件持久化 + 控制台实时输出。
CCRS/
├── app.py # Flask Web 入口
├── train.py # 训练脚本(TensorFlow + 进度回调 + 暂停/停止控制)
├── predict.py # 预测脚本
├── config.py # 配置
├── models/
│ ├── cnn_model.py # 模型定义(Keras Sequential)
│ └── default/ # 默认模型存放
├── utils/ # 工具函数(数据集加载等)
├── data/ # 数据
├── static/ templates/ # 前端
└── requirements.txt
README 详细描述了这些进阶主题(部分属于设计文档、未在核心训练代码中完整实现):
alpha=1, gamma=2,解决汉字类别不平衡lr = lr0 * 0.5 * (1 + cos(π*(epoch-5)/(num_epochs-5)))patience=7 监控验证损失prune.global_unstructured + L1Unstructured我做了几处工程化处理:Flask Web 界面 + 进度回调 + 暂停/停止控制 + 日志双输出 + GPU 显存优化。模型本身是标准的三层 CNN + BatchNorm + Dropout。
文档与代码存在不一致(README 写 PyTorch、代码是 TensorFlow)是需要注意的问题——提醒我们:写技术博客/文档时要以源码为准。