Skip to main content

Lightning Neural Compressor

License

This repository contains the implementation of the Lightning Neural Compressor. The main goal of this project is to provide Pytorch Lightning callbacks to use Intel® Neural Compressor. The callbacks aim at compressing a neural network so that it can be used on edge devices (i.e., mobile phones, raspberry pi, etc.). This project is a work in progress and is not ready for production use.

Current Status

The project is currently under development, starting with Quantization Aware Training, as the default callback has been deleted from Pytorch Lightning.

The project also supports Weight Pruning and should work at least with pruners related to the PytorchBasicPruner.

Installation

To install the dependencies for this project, use the following command to use pypi:

pip install -U lightning-nc

or directly by cloning the main branch:

git clone https://github.com/clementpoiret/lightning-nc
cd lightning-nc
pip install -e .

Usage

To use the Lightning Neural Compressor, import the callbacks from the lightning_nc module.

WARNING: Currently, the callbacks need the PyTorch model to be a nn.Module contained inside your LightningModule. This is not a huge limitation as the refactoring is easy and straightforward, such as:

import os

import lightning as L
import timm
import torch
import torch.nn.functional as F
from neural_compressor import QuantizationAwareTrainingConfig
from neural_compressor.config import Torch2ONNXConfig
from neural_compressor.training import WeightPruningConfig
from lightning_nc import QATCallback, WeightPruningCallback
from torch import Tensor, nn, optim, utils
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor


# Define your main model here
class VeryComplexModel(nn.Module):
    
    def __init__(self):
        super().__init__()
        self.backbone = timm.create_model("best_pretrained_model",
                                          pretrained=True)

        self.mlp = nn.Sequential(nn.Linear(self.backbone.num_features, 128),
                                 nn.ReLU(), nn.Linear(128, 10))

    def forward(self, x):
        return self.mlp(self.backbone(x))


# Then, define your LightningModule as usual
class Classifier(L.LightningModule):

    def __init__(self):
        super().__init__()

        # This is mandatory for the callbacks
        self.model = VeryComplexModel()

    def forward(self, x):
        return self.model(x)

    def training_step(self, batch, batch_idx):
        x, y = batch

        # This is just to use MNIST images on a pretrained timm model, you can skip that
        x = x.repeat(1, 3, 1, 1)
        x = F.interpolate(x, size=(224, 224))

        y_hat = self.forward(x)

        loss = F.cross_entropy(y_hat, y)

        return loss

    def configure_optimizers(self):
        optimizer = optim.Adam(self.parameters(), lr=1e-3)

        return [optimizer]


clf = Classifier()

# setup data
dataset = MNIST(os.getcwd(), download=True, transform=ToTensor())
train_loader = utils.data.DataLoader(dataset)

Now that everything is setup, the callbacks can be integrated into a PyTorch Lightning training routine:

# Define the configs for Pruning and Quantization
q_config = QuantizationAwareTrainingConfig()
p_config = WeightPruningConfig([{
    "op_names": ["backbone.*"],
    "start_step": 1,
    "end_step": 100,
    "target_sparsity": 0.5,
    "pruning_frequency": 1,
    "pattern": "4x1",
    "min_sparsity_ratio_per_op": 0.,
    "pruning_scope": "global",
}])

callbacks = [
    QATCallback(config=q_config),
    WeightPruningCallback(config=p_config),
]

trainer = L.Trainer(accelerator="gpu",
                    strategy="auto",
                    limit_train_batches=100,
                    max_epochs=1,
                    callbacks=callbacks)
trainer.fit(model=clf, train_dataloaders=train_loader)

Models can now be saved eaily such as:

clf.model.export(
    "model.onnx",
    Torch2ONNXConfig(
        dtype="int8",
        opset_version=17,
        quant_format="QOperator",  # or QDQ
        example_inputs=torch.randn(1, 3, 224, 224),
        input_names=["input"],
        output_names=["output"],
        dynamic_axes={
            "input": {
                0: "batch_size"
            },
            "output": {
                0: "batch_size"
            },
        },
    ))

Contributing

If you would like to contribute to this project, please submit a pull request. All contributions are welcome!

License

This project is licensed under the MIT License. See the LICENSE.md file for details.

Metadata

Release files for lightning-nc 0.0.3

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for lightning-nc 0.0.3
File Size Uploaded
lightning_nc-0.0.3.tar.gz 6.3 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for lightning-nc 0.0.3
File Interpreter ABI Platform
lightning_nc-0.0.3-py3-none-any.whl Python 3 none any Details

Total release size: 14.7 kB

Release files / lightning_nc-0.0.3.tar.gz

Download URL lightning_nc-0.0.3.tar.gz
Size 6.3 kB
Tags Source
SHA-256 checksum
How to use checksums
b4736b13726fe0b2efce90bf7921d43f73dadec450d51d393ef8f4a752c23ab8
BLAKE2b-256 checksum
How to use checksums
569fc3b0e91a6abf99b92a8d9d05485ec8ee2adbefed74fa581994d61ead0b71
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via poetry/1.7.0 CPython/3.11.6 Linux/6.6.1-arch1-1

Release files / lightning_nc-0.0.3-py3-none-any.whl

Download URL lightning_nc-0.0.3-py3-none-any.whl
Size 8.4 kB
Tags Python 3
SHA-256 checksum
How to use checksums
3d711e195128cf13bcd76830c19bdec1f248093d7671b718dd6afd84e56a5daa
BLAKE2b-256 checksum
How to use checksums
14fbc3b24a5df79410c96db314f5c97aa700424ff10fec5023090de04dba1bf3
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via poetry/1.7.0 CPython/3.11.6 Linux/6.6.1-arch1-1

Release history Release notifications | RSS feed

This release

0.0.3 This release

2 release files

0.0.2

2 release files

0.0.1

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page