Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

5 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

激活函数对比实验

基于 PyTorch(MPS 加速)的深度学习对比实验,在 MNIST 和 CIFAR-10 数据集上系统比较 ReLU、Leaky ReLU、Swish、GELU 四种激活函数,分析其对训练效果、死亡神经元比例和激活值分布的影响。

实验内容

  • 激活函数:ReLU、Leaky ReLU、Swish(SiLU)、GELU
  • 数据集:MNIST、CIFAR-10(标准)、CIFAR-10(轻量配置)、CIFAR-10(深层无BN)
  • 网络结构:浅层 CNN(4 层卷积 + BatchNorm,约 130 万参数);深层 DeepCNN(8 层卷积,无 BN,约 1400 万参数),激活函数可插拔
  • 指标:训练/验证 Acc & Loss、死亡神经元比例(动态 + 最终统计)、激活值分布快照

环境要求

  • Python ≥ 3.11
  • uv(包管理)
  • Apple Silicon Mac(MPS)或 CUDA GPU,CPU 也可运行

安装

git clone <repo>
cd ML
uv sync --dev

首次运行会下载 PyTorch(~500MB)、MNIST(~11MB)、CIFAR-10(~170MB)。

使用

运行全部实验

uv run train

共 16 组实验(4 配置 × 4 激活函数),支持断点续跑——中断后重新运行会自动跳过已完成的组。

预计时长(M4 Pro):

  • MNIST × 4:约 20 分钟
  • CIFAR-10 浅层 × 4:约 60 分钟
  • CIFAR-10 轻量 × 4:约 20 分钟
  • CIFAR-10 深层 × 4:约 90 分钟

生成对比图

uv run plot

PNG(300 dpi)与 PDF(矢量)同时输出到 plots/

运行测试

uv run pytest tests/ -v

项目结构

ML/
├── pyproject.toml                     # 依赖与入口配置
├── src/activation_exp/
│   ├── cli.py                         # train / plot 入口
│   ├── runner.py                      # 实验调度
│   ├── model.py                       # CNN / DeepCNN 网络定义
│   ├── configs.py                     # 实验配置(16 组)
│   ├── data.py                        # 数据集加载
│   ├── metrics.py                     # 死亡神经元检测、激活分布采样
│   ├── trainer.py                     # 训练循环
│   ├── utils.py                       # 设备选择、随机种子
│   └── plotting/
│       ├── training_curves.py
│       ├── dead_neurons.py
│       └── activation_dist.py
├── plots/                             # 生成的图片(PNG + PDF)
├── results/                           # 原始实验数据(JSON + npz)
│   └── {dataset}_{activation}/
│       ├── metrics.json
│       └── activation_dist.npz
└── tests/
    ├── test_model.py
    ├── test_metrics.py
    └── test_trainer.py

输出文件说明

results/{dataset}_{activation}/metrics.json

{
  "config": {"dataset": "cifar10", "activation": "relu", "epochs": 30},
  "epoch":      [1, 2, ...],
  "train_loss": [...],
  "train_acc":  [...],
  "val_loss":   [...],
  "val_acc":    [...],
  "dead_ratio": [...]
}

results/{dataset}_{activation}/activation_dist.npz:每 5 epoch 采样一次第 2 卷积层后的激活值,key 为 epoch_1epoch_5epoch_10 等。

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages