-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathconfig.py
More file actions
65 lines (56 loc) · 2.75 KB
/
Copy pathconfig.py
File metadata and controls
65 lines (56 loc) · 2.75 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
"""
config.py — 超参数与全局配置
------------------------------
将所有可调参数集中到一处,避免"魔法数字"散落在各处。
后续只需修改此文件即可切换实验配置。
"""
import os
# ------------------------------------------------------------------ #
# 路径配置
# ------------------------------------------------------------------ #
# 项目根目录(dl_tutorial 2 目录)
BASE_DIR = os.path.abspath(os.path.dirname(__file__))
# 数据文件路径
DATA_PATH = os.path.join(BASE_DIR, 'data', 'train.csv')
# ------------------------------------------------------------------ #
# 数据集配置
# ------------------------------------------------------------------ #
DATA_CONFIG = {
'test_size': 0.3, # 测试集比例
'random_state': 42, # 随机种子,保证可复现
'normalize': True, # 是否对特征做归一化(MinMaxScaler)
}
# ------------------------------------------------------------------ #
# 网络结构配置
# ------------------------------------------------------------------ #
MODEL_CONFIG = {
'input_size': 784, # 输入层节点数(28×28=784)
'hidden_size': 50, # 隐藏层节点数(50 是老师推荐值,10 太小容易欠拟合)
'output_size': 10, # 输出层节点数(0-9共10类)
'weight_init_std': 0.01, # 权重初始化标准差
}
# ------------------------------------------------------------------ #
# 训练超参数配置
# ------------------------------------------------------------------ #
TRAIN_CONFIG = {
'learning_rate': 0.001, # 学习率(SGD 推荐 0.1;Adam 推荐 0.001)
'batch_size': 100, # 每次训练使用的小批量样本数
'num_epochs': 10, # 训练轮次
# 梯度计算方式:
# 'backprop' — 解析梯度(反向传播,已实现,速度快,推荐使用, 速度相比 numerical 1000x 提升)
# 'numerical' — 数值梯度(中心差分,慢但可用于验证反向传播的正确性)
'grad_method': 'backprop',
# 优化器类型(三种均已实现,可自由切换):
# 'sgd' — 随机梯度下降,简单直接
# 'momentum' — 动量法,收敛更快,推荐配合 lr=0.01
# 'adam' — Adam,最稳定,推荐配合 lr=0.001
'optimizer': 'adam',
}
# ------------------------------------------------------------------ #
# 日志 / 输出配置
# ------------------------------------------------------------------ #
LOG_CONFIG = {
'print_loss_every': 1, # 每次迭代都打印 loss(对齐老师代码,数值梯度时可看到训练进度)
'save_plot': True, # 是否保存训练曲线图
'plot_output_dir': os.path.join(os.path.dirname(__file__), 'outputs'),
}