Skip to main content

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


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

dltrain-0.1.4.tar.gz (21.6 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

dltrain-0.1.4-py3-none-any.whl (26.4 kB view details)

Uploaded Python 3

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

Hashes for dltrain-0.1.4.tar.gz
Algorithm Hash digest
SHA256 b08f60c2b1236b8b213142a1bb70cd42915ba7ca8267b0f4e9e43b49fc726fd9
MD5 288621a6757e27651888f5dab71699ad
BLAKE2b-256 9a287bb92e7435a6727165833e4bfc05748eed71592f7e66f23732aef9748d61

See more details on using hashes here.

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

Hashes for dltrain-0.1.4-py3-none-any.whl
Algorithm Hash digest
SHA256 c5fbedee0df11ea73baa1daef8413b07107bb53dafdcedb7cd61cf0a0774890a
MD5 d6e56702badd1b5798fc32c7d28df66d
BLAKE2b-256 be9d9fe952a58fb5f72833569c88a8b103e2a3c27e8aac415876c60b122bb109

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page