FlashGen
FlashGen 是基于昇腾(Ascend)硬件的扩散模型蒸馏加速框架,为昇腾NPU提供多种蒸馏和加速方法。
仓库结构
FlashGen/
├── flashgen/
│ ├── __init__.py
│ ├── configs/
│ ├── datasets/
│ ├── methods/
│ └── networks/
├── scripts/
├── tests/
└── train.py
环境要求
- 昇腾硬件:Atlas 800 训练服务器(Ascend 910B)或更高版本
- CANN:8.5.0 及以上版本
- Python:3.10 及以上
- torch / torch_npu:2.7.1 及以上
- FastVideo:0.2.0
依赖 FastVideo
FlashGen 只实现算法/网络层(flashgen/),训练循环、数据管线等基础设施全部复用 FastVideo。flashgen.entrypoint 在启动时固定 NPU 安全的注意力后端(Torch SDPA),再转发到 fastvideo.train.entrypoint.train。
当前对齐 FastVideo v0.2.0。拉取:
git clone -b v0.2.0 https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
注意:不能直接 pip install -e .。 FastVideo 的 pyproject.toml 默认会从 CUDA 源拉 torch==2.11.0 和仅限 CUDA 的 fastvideo-kernel,在昇腾环境会覆盖掉 torch_npu 或直接装不上。正确做法是先在 NPU 环境备好 torch / torch_npu,再跳过依赖解析安装:
pip install -e . --no-deps
FastVideo 运行所需的其余依赖(transformers、diffusers、einops 等纯 Python 包)按需手动 pip install,不要让它自动拉 torch / fastvideo-kernel。FlashGen 自身的轻量依赖见 requirements.txt。
安装
git clone <repo-url>
cd FlashGen
pip install -e .
贡献
欢迎贡献!详见 CONTRIBUTING.md。
安全
安全相关须知详见 SECURITY.md。
行为准则
参与本项目即表示您同意遵守 行为准则。
许可证
本项目使用 Mulan PSL v2 许可证。详见 LICENSE。