This is a package for PyTorch training, and this project provides a large number of high-level training APIs and low-level interfaces to meet all training needs
Project description
dltrain
这是一个 PyTorch 训练包,本项目提供了大量的高级训练 API 和低级接口,满足所有训练需求
安装仅在Shell中使用pip安装即可
pip install dltrain
利用该工具包我们可用通过如下结构构建训练代码
Example
example .1 在iris数据集上使用多层感知器进行分类
from dltrain import TaskBuilder, SimpleTrainer
builder = TaskBuilder('iris')
builder.base.use_epoch(100).use_batch_size(8).use_device('cuda')
builder.model.use_mlp(4, 3)
builder.delineator.use_random_split(builder.dataset.use_iris())
builder.criterion.use_cross_entropy()
SimpleTrainer().run(builder.build())
example .1.1 使用Accuracy对模型进行验证
from dltrain import TaskBuilder, SimpleTrainer
builder = TaskBuilder('iris')
builder.base.use_epoch(100).use_batch_size(8).use_device('cuda')
builder.model.use_mlp(4, 3)
builder.delineator.use_random_split(builder.dataset.use_iris())
builder.criterion.use_cross_entropy()
builder.evaluation_handler.add_accuracy()
SimpleTrainer().run(builder.build())
ps: 如需指定只计算训练/测试则调用.add_accuracy()中内嵌role='train/eval'即可。如不需要输出最后的绘图则内嵌参数drawable=False即可。
example .2 使用自建模型在mnist手写集数字上进行分类
from dltrain import TaskBuilder, SimpleTrainer
import torch.nn as nn
class Model(nn.Module):
def __init__(self):
super().__init__()
def make_layer(i, o, k, s, p):
return nn.Sequential(
nn.Conv2d(i, o, k, s, p),
nn.BatchNorm2d(o),
nn.ReLU()
)
self.features = nn.Sequential(
make_layer(1, 3, 7, 1, 0),
make_layer(3, 16, 7, 1, 0),
make_layer(16, 64, 7, 1, 0),
make_layer(64, 128, 7, 1, 0)
)
self.classifier = nn.Sequential(
nn.Linear(128 * 4 * 4, 1000),
nn.Linear(1000, 10)
)
def forward(self, data):
data = self.features(data)
data = data.reshape(data.shape[0], -1)
data = self.classifier(data)
return data
builder = TaskBuilder('mnist')
builder.model.use_model(Model())
builder.base.use_device('cuda')
builder.optimizer.use_adam()
builder.criterion.use_cross_entropy()
builder.evaluation_handler.add_accuracy()
builder.delineator.use_train_eval(builder.dataset.use_mnist('./dataset', True),
builder.dataset.use_mnist('./dataset', False))
SimpleTrainer().run(builder.build())
模型向导对象的use_model接口允许用户传入自己的模型
example .2.1 使用torchvision的自带模型
from dltrain import TaskBuilder, SimpleTrainer
builder = TaskBuilder('mnist')
builder.model.use_pytorch_model('resnet18', num_classes=10)
builder.base.use_device('cuda')
builder.optimizer.use_adam()
builder.criterion.use_cross_entropy()
builder.evaluation_handler.add_accuracy()
builder.delineator.use_train_eval(builder.dataset.use_mnist('./dataset', True),
builder.dataset.use_mnist('./dataset', False))
SimpleTrainer().run(builder.build())
正如所说的,模型向导提供了use_pytorch_model接口,允许用户直接调用torchvision.models下的所有模型并封装到一个名为PyTorchNativeCNN的类型下,该类型会让模型自适应所有的数据集输入输出
Suppose Model In TorchVision
__Model__ = [
googlenet, alexnet,
resnet18, resnet34, resnet50, resnet101, resnet152,
vgg11, vgg13, vgg16, vgg19, vgg11_bn, vgg13_bn, vgg16_bn, vgg19_bn,
vit_b_16, vit_h_14, vit_b_32, vit_l_16, vit_l_32,
mobilenet_v2, mobilenet_v3_small, mobilenet_v3_large,
efficientnet_v2_s, efficientnet_v2_l, efficientnet_v2_m, efficientnet_b0, efficientnet_b1, efficientnet_b2,
efficientnet_b3, efficientnet_b4, efficientnet_b5, efficientnet_b6, efficientnet_b7,
densenet121, densenet161, densenet169, densenet201,
regnet_x_8gf, regnet_x_1_6gf, regnet_y_8gf, regnet_y_400mf, regnet_y_128gf, regnet_y_1_6gf, regnet_x_3_2gf,
regnet_x_16gf, regnet_x_32gf, regnet_x_400mf, regnet_y_800mf, regnet_x_800mf, regnet_y_3_2gf, regnet_y_16gf,
regnet_y_32gf,
shufflenet_v2_x0_5, shufflenet_v2_x1_0, shufflenet_v2_x1_5, shufflenet_v2_x2_0,
swin_b, swin_t, swin_s, swin_v2_b, swin_v2_t, swin_v2_s,
mnasnet0_5, mnasnet1_0, mnasnet1_3, mnasnet0_75
]
default about the arguments
| 参数名称 | 默认值 |
|---|---|
| optimizer | Sgd(lr=0.01,momentum=0,dampening=0,weight_decay=0) |
| scheduler | User-set |
| criterion* | None,Must be specified by the user |
| model* | None,Must be specified by the user |
| epoch | 10 |
| batch_size | 16 |
| seed | 3407 |
| device | cpu |
| save_checkpoint | False |
| start_checkpoint | None,User-set |
| delineator* | None,Must be specified by the user |
| forward | SimpleForward |
| trainer | SimpleTrainer |
Version Log
- 0.0.1
1、构造整体模型框架 - 0.0.2
1、加入了TaskBuilder方便构造模型 - 0.1.2
1、加入了InjectForward(可以通过InjectWizard注入训练时策略,如实时检测模型参数梯度的分布等)
2、将models.py封装到models包,添加了许多拆箱可用模型 - 0.1.3
1、引入错误处理机制,防止由于EventHandler设置问题导致的结果保存错误
2、引入VectorSequenceDataset
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
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 dltrain-0.1.4.tar.gz.
File metadata
- Download URL: dltrain-0.1.4.tar.gz
- Upload date:
- Size: 21.6 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/5.0.0 CPython/3.10.9
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
b08f60c2b1236b8b213142a1bb70cd42915ba7ca8267b0f4e9e43b49fc726fd9
|
|
| MD5 |
288621a6757e27651888f5dab71699ad
|
|
| BLAKE2b-256 |
9a287bb92e7435a6727165833e4bfc05748eed71592f7e66f23732aef9748d61
|
File details
Details for the file dltrain-0.1.4-py3-none-any.whl.
File metadata
- Download URL: dltrain-0.1.4-py3-none-any.whl
- Upload date:
- Size: 26.4 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/5.0.0 CPython/3.10.9
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
c5fbedee0df11ea73baa1daef8413b07107bb53dafdcedb7cd61cf0a0774890a
|
|
| MD5 |
d6e56702badd1b5798fc32c7d28df66d
|
|
| BLAKE2b-256 |
be9d9fe952a58fb5f72833569c88a8b103e2a3c27e8aac415876c60b122bb109
|