Training 仓库聚焦于 Transformer 模型训练/推理过程中的关键统计:为每一层生成权重、激活值与参数梯度的分布图,并输出对应的统计量。目前支持 GPT-2、LLaMA 以及 DiT。
- 多模型支持:内置 GPT-2、LLaMA 与 DiT 分析器。
- 分布统计一站式产出:一次前向/反向即可同时获得权重、激活、梯度的直方图与统计指标。
- 灵活的层选择:可通过参数筛选感兴趣的层,避免对全量模型进行昂贵分析。
- 快速可视化:为每个张量生成直方图与热力图(
*_heatmap.png),图像保存在weights/、activations/、gradients/子目录。
# 克隆仓库
git clone https://github.com/z-zanez/Visible-Distribution.git
cd Training
# 安装依赖
pip install -r requirements.txt
pip install -e .运行示例脚本:
python examples/analyze_model.py \
--model_type gpt2 \
--model_path ./models/gpt2 \
--output_dir ./outputs/gpt2-layer-stats \
--layers 0,1,2 \
--text "gpt2 test"对于 LLaMA(请确保模型文件已本地可用):
python examples/analyze_model.py \
--model_type llama \
--model_path ./models/Llama-3.2-1b \
--output_dir ./outputs/llama-layer-stats \
--layers 0,1 \
--local_files_only \
--text "llama test" \
--no_plots # 若只想保存JSON统计,可禁用图片对于 DiT-XL/2(默认使用随机 latent / timestep / label 触发一次接近扩散训练形式的前反向,不依赖 VAE):
PYTHONPATH=. python examples/analyze_model.py \
--model_type dit \
--model_path DiT-XL-2-256x256.pt \
--dit_model DiT-XL/2 \
--dit_repo_path /home/Shenchao/thuhpgc/zz/ArmTraining_best/DiT \
--image_size 256 \
--batch_size 1 \
--output_dir ./outputs/dit-xl2-layer-stats \
--layers 0,1,2 \
--no_heatmaps说明:
--model_path对 DiT 来说是 checkpoint 路径,或者预训练别名DiT-XL-2-256x256.pt/DiT-XL-2-512x512.pt。- 若使用预训练别名且未加
--local_files_only,脚本会自动下载到[DiT 仓库]/pretrained_models/。 --layers对 DiT 的0~27对应 28 个 Transformer blocks;额外的28对应final_layer。--text对 DiT 会被忽略,因为 DiT 输入不是文本而是合成 latent。- 若只想快速拿到每层权重/激活/梯度分布,
--batch_size 1、--no_heatmaps通常就够了。
生成的目录结构示例:
提示:默认会生成直方图与热力图;若需要快速调试可加上
--no_heatmaps、--no_gradients、--no_plots,或调小--heatmap-max-dim减少输出体积。
outputs/
├── layer_statistics.json # 权重/激活/梯度的统计量
├── activations/ # 每层激活分布图
├── gradients/ # 每层梯度分布图
└── weights/ # 每层权重分布图
training/
├── core/ # 基础抽象与统计工具
│ ├── models/base.py # 通用分析逻辑
│ └── utils/ # Hook 与统计函数
├── adapters/ # 架构特定实现(GPT-2 / LLaMA / DiT)
└── viz/ # 分布绘图工具
- Python 3.9+
- PyTorch 2.0+
- transformers 4.35+
- timm 0.9+
- matplotlib、seaborn、numpy
- Q:可以在线下载模型吗? GPT-2 / LLaMA 可通过 Hugging Face 加载;DiT 若传入预训练别名,会在未开启
--local_files_only时自动下载 checkpoint。 - Q:梯度统计为空? 请确认未使用
--no_gradients,并且模型允许反向传播(use_cache=False已自动处理)。 - Q:DiT 为什么不需要图片/VAE? 因为这里的目标是做层级统计,不是复现完整采样流程。DiT 分析器会直接构造随机 latent、随机 timestep 和随机类别标签,并对噪声预测头施加 MSE loss,从而在预训练权重上得到有意义的激活与梯度分布。
- Q:想自定义绘图? 可直接使用
training.viz.distribution.plot_tensor_distribution函数对任意张量绘图。