Deep Learning experiment assistant toolkit: config parsing, device setup, logging, checkpointing, CSV tracking, and more.
Project description
dlkit
Deep Learning experiment assistant toolkit — 深度学习实验辅助工具包
dlkit 是一个用于深度学习实验的辅助工具包,提供从配置解析、设备分配、实验目录管理、日志记录到训练状态保存/加载的全流程辅助功能。
安装
pip install dlkit-zihao
或从源码安装:
cd dlkit
pip install -e .
功能模块
| 模块 | 主要类/函数 | 说明 |
|---|---|---|
dlkit.pipeline |
StandardTrainingPipeline |
标准化实验流水线:配置解析、设备分配、种子设定、目录管理、日志、代码备份、CSV 文件名生成 |
dlkit.argparse_utils |
GroupArgparse |
多层级分组命令行参数解析器,支持 YAML/字典深度合并 |
dlkit.meters |
ScalarMeter, TimeMeter |
标量累加器与计时器 |
dlkit.tracker |
CSVExpTracker |
CSV 实验记录管理:多服务器聚合、查重、统计 |
dlkit.summary |
Summary |
JSON 格式标量指标记录 |
dlkit.utils |
set_seed, count_parameters_in_MB, get_gpus_memory_info, cur_time_str, get_caller_filename |
通用工具函数 |
快速开始
1. 使用 StandardTrainingPipeline
from dlkit import StandardTrainingPipeline
config = {
'train': {
'lr': 0.01,
'batch_size': 32,
}
}
# 保留参数可直接作为关键字参数传入,获得 IDE 自动补全
pipeline = StandardTrainingPipeline(
base_config=config,
gpu=0,
seed=42,
platform='6001',
workspace='experiments', # 工作区根目录
project_name='my_project', # 项目目录名
csv_dir='./csv', # CSV 文件保存目录
csv_tag='search', # CSV 文件名标签
)
args, config = pipeline.init(run_name='my_experiment')
pipeline.print('Training started', lr=args.train.lr)
pipeline.print(Epoch=1, Loss=0.45, Acc='92%')
# 自动生成 CSV 文件路径: './csv/train_search_6001.csv'
csv_path = pipeline.gen_csv_path(tag='search')
# 实验结束后可重命名目录
# pipeline.rename_run_name('Acc_92.5')
2. 使用 GroupArgparse
from dlkit import GroupArgparse
config = {
'model': {'type': 'resnet', 'layers': 50},
'train': {'lr': 0.01, 'epochs': 100},
}
gparser = GroupArgparse(base_config=config)
gparser.set_cur_group('train')
gparser.add_argument('--lr', type=float, default=0.01)
gparser.add_argument('--epochs', type=int, default=100)
args = gparser.parse_args(export_dataclass=True)
print(args.train.lr) # 支持点操作符访问
3. 使用 CSVExpTracker
from dlkit import StandardTrainingPipeline, CSVExpTracker
# 方式一:从 args 自动推导 CSV 路径(推荐)
pipeline = StandardTrainingPipeline(
platform='6001',
csv_dir='./csv',
csv_tag='search',
)
args, config = pipeline.init(run_name='my_experiment')
# 不传 root,自动从 args 读取 csv_dir / csv_tag / platform
# 生成路径: './csv/train_search_6001.csv'
tracker = CSVExpTracker(
args=args,
x_columns=['lr', 'batch_size'],
y_columns=['acc', 'loss'],
avoid_duplication=True,
platforms=['6001', '6002'],
)
# 方式二:直接指定 root
# tracker = CSVExpTracker(root='./csv/results.csv', ...)
# 查重
if not tracker.is_sampled(lr=0.01, batch_size=32, run_id=0):
# ... 训练模型 ...
tracker.save_csv({'lr': 0.01, 'batch_size': 32, 'acc': 0.925, 'loss': 0.05, 'run_id': 0})
# 统计聚合
tracker.count()
4. 使用 Meters
from dlkit import ScalarMeter, TimeMeter
loss_meter = ScalarMeter(reduction='mean')
timer = TimeMeter()
for batch in dataloader:
timer.update()
loss = train_batch(batch)
loss_meter.update(loss, num=batch_size)
print(f'Avg Loss: {loss_meter.avg}, Time: {timer.cost} min')
5. 使用 get_caller_filename
from dlkit import get_caller_filename
# 在 train.py 中调用
print(get_caller_filename()) # 'train'
print(get_caller_filename(stem=False)) # 'train.py'
相对原版 assistant.py 的修复
TimeMeter— 修复了start_time未初始化导致total属性崩溃的问题set_seed— 使用manual_seed_all覆盖所有 GPU(原版仅manual_seed)CSVExpTracker.save_csv— 修复了文档声称更新self.df但实际未实现的问题Summary.save_record— 修复了record_name=None时的逻辑不清晰问题StandardTrainingPipeline— 修复了enable_directory=False时self.print未定义的问题- 清理 — 移除了废弃的
_export_dataclass_schema1、未使用的globimport - 移除 checkpoint — 移除了
save_checkpoint/load_checkpoint(原版存在死代码 bug,且该功能可由 PyTorch 原生torch.save/torch.load替代) - 新增
get_caller_filename— 自动获取主调脚本文件名 - 新增
gen_csv_path— 自动生成[主调文件名]_[tag]_[platform].csv格式的 CSV 文件路径 - 新增
platform保留参数 — 用于标识当前运行所在的平台 - 重命名目录参数 —
exp_root→workspace(工作区根目录),task_name→project_name(项目名),含义更直观 CSVExpTracker支持args推导路径 —root改为可选,未指定时自动从args中读取csv_dir、csv_tag、platform推导路径,默认保存到./csv/目录- 新增
csv_dir、csv_tag保留参数 — 用于CSVExpTracker自动推导 CSV 路径
项目结构
dlkit/
├── pyproject.toml
├── README.md
└── src/
└── dlkit/
├── __init__.py
├── py.typed
├── utils.py # 通用工具函数 (含 get_caller_filename)
├── meters.py # ScalarMeter, TimeMeter
├── tracker.py # CSVExpTracker
├── summary.py # Summary
├── argparse_utils.py # GroupArgparse
└── pipeline.py # StandardTrainingPipeline (含 gen_csv_path)
发布到 PyPI
cd dlkit
pip install build twine
python -m build
twine upload dist/*
License
MIT
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
dlkit_zihao-0.2.0.tar.gz
(33.1 kB
view details)
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file dlkit_zihao-0.2.0.tar.gz.
File metadata
- Download URL: dlkit_zihao-0.2.0.tar.gz
- Upload date:
- Size: 33.1 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.13.5
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
1c4a1c3f04e64dc3b74e81a4eb7e6f2719a822a720a22c323e6871aed8836d47
|
|
| MD5 |
32e6fe2494d2162407fc1181d870759c
|
|
| BLAKE2b-256 |
e7beb3d7315ec2710e3f2164559bee2e67f98008d80d1668bf673077a37ed78a
|
File details
Details for the file dlkit_zihao-0.2.0-py3-none-any.whl.
File metadata
- Download URL: dlkit_zihao-0.2.0-py3-none-any.whl
- Upload date:
- Size: 34.0 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.13.5
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
d9b7b7216e67c2afe8aa69763efa58fd074c0d71b0976f6a593d69fcf1cb2575
|
|
| MD5 |
8176997661ce5217f82fb513c20eac59
|
|
| BLAKE2b-256 |
5223670a86de74bd012330dea4998d696a2ecb5f9056acf0ede7861e8c76eec1
|