Skip to main content

A lightweight PyTorch wrapper for fast and easy neural network training and evaluation.

Project description

Assess fast and easy your neural network

Importance of easy and quick assessment of quality

A lightweight PyTorch wrapper that can be used to fasten process of training and setting up arbitrary Neural Network to quickly test an idea/setup. The wrapper provides an interface for both standard neural networks and CNNs, but can be extended to any architecture, with built-in visualization and performance tracking capabilities. The wrapper is customizable and aims to be used on any dataset.

Usage

pip install simple_pytorch_wrapper

This package is aimed to be as simple as possible. The following example trains on the example dataset in this package, using a simple feedforward network.

from simple_pytorch_wrapper import PytorchWrapper, FNNGenerator, NetworkType
from simple_pytorch_wrapper import set_seed, load_language_digit_example_dataset
import torch.nn as nn

def main():
    # For reproducibility
    set_seed(0)

    # Example data
    X, Y = load_language_digit_example_dataset()

    # Transform data accordingly (vectorizing) + squeezing for batching
    X, Y = PytorchWrapper.vectorize_data(X, Y, NetworkType.FNN) 

    # Pytorch wrapper
    wrapper = PytorchWrapper(X, Y)  
    
    network = FNNGenerator(
    input_size=64*64,  # flattened value of image of example dataset
    output_size=10,    # number of outputs
    hidden_layers=[100, 500],  # sizes of hidden layers
    hidden_activations=[nn.ReLU(), nn.ReLU()],  # activation functions for each hidden layer
    )
    
    # Create custom pytorch network
    wrapper.upload_pyTorch_network(network) 

    wrapper.setup_training(batch_size=32, learning_rate=0.001, epochs=10) 

    # Training
    wrapper.train_network(plot=False)  # Enable plotting periodically during training

    # Visualization
    wrapper.visualize()

main()

You see that training a network from start to finish, with a clean visualization in less than 15 effective lines!

Upload your own network!

The network variable above is there to be defined by you. This is aimed to take in any Neural network design. The only caveat is to correctly shape your input. For a feedforward and CNN, these are already implemented by the vectorize_data(*args) function.

To get you started on the networks, there has been provided already NN generator classes for a feed-forward NN and a convolutional one. These can be used respectively by calling FNNGenerator(*args) and CNNGenerator(*args).

The function signature of the FNNGenerator is:

network = FNNGenerator(input_size: int,
              output_size: int,
              hidden_layers: List[int],
              hidden_activations: Optional[List[nn.Module]]
              )

While for the CNNGenerator, this is:

network = CNNGenerator(input_channels: int,
                      conv_layers: List[Dict[str, int]],  
                      fc_layers: List[int], 
                      output_size: int,
                      batch_size: Optional[int] = None,
                      use_pooling: bool = True)

Example runs

It can be hard to know the power of a tool, without having a solid example. By calling FNN_example_run() with the training params as function arguments, you get a run on the example data set for a FNN architecture. The CNN_example_run() does it with a CNN architecture.

Dataset

The data used for the example is the Sign Language Digits Dataset from the Turkey Ankara Ayrancı Anadolu High School students. This dataset is available at ardamavi/Sign-Language-Digits-Dataset.

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

simple_pytorch_wrapper-0.0.2.tar.gz (8.8 MB view details)

Uploaded Source

Built Distribution

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

simple_pytorch_wrapper-0.0.2-py3-none-any.whl (8.9 MB view details)

Uploaded Python 3

File details

Details for the file simple_pytorch_wrapper-0.0.2.tar.gz.

File metadata

  • Download URL: simple_pytorch_wrapper-0.0.2.tar.gz
  • Upload date:
  • Size: 8.8 MB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.0.1 CPython/3.12.6

File hashes

Hashes for simple_pytorch_wrapper-0.0.2.tar.gz
Algorithm Hash digest
SHA256 092f29ba8fb7e4ab3d1d191bcae290f67c5426d7d721eae5c391d83470bdc659
MD5 309db64a0614697bd308b4f06e8f7240
BLAKE2b-256 a206031860b6b7b7e6d5b95c7f2078001f63a2b9c892a1167ed6ce388d2965b4

See more details on using hashes here.

File details

Details for the file simple_pytorch_wrapper-0.0.2-py3-none-any.whl.

File metadata

File hashes

Hashes for simple_pytorch_wrapper-0.0.2-py3-none-any.whl
Algorithm Hash digest
SHA256 8cfc0b40ea183a06654f656dc1e7b8b79e4e9b9e265b9b35fe288e5b7d8ae46e
MD5 fc981c89acbedd46f1707f88e5d9b31c
BLAKE2b-256 c3e26f1e62bd0ea163ae38321310c008241e4f464035cd835b0a118aee755a3c

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