Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -204,7 +204,7 @@ cython_debug/
tempCodeRunnerFile.py

# Downloaded datasets and model cache
data/
/data/

# Ruff stuff:
.ruff_cache/
Expand Down
2 changes: 1 addition & 1 deletion Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ install:
lint:
python -m black belt
python -m isort belt
python -m ruff check belt tests
python -m ruff check belt --fix

train:
python -m belt.supervised.softmax --config configs/supervised_softmax.yaml
Expand Down
2 changes: 2 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@ ML Tasks are broken into pipelines allowing models to be used with multiple pipe
```bash
# Softmax classifier
python -m belt classifier
# Seq2Seq text translation
python -m belt translation
```

Configs live in `configs/`. Each pipeline has a YAML controlling data, model, training, and output settings.
Expand Down
3 changes: 3 additions & 0 deletions belt/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,16 +10,19 @@
import argparse

from belt.supervised.classifier import deploy as softmax_train
from belt.supervised.generate import deploy as generate_train
from belt.supervised.translation import deploy as translation_train

_DEFAULT_CONFIGS: dict[str, str] = {
"classifier": "configs/softmax_classifier.yaml",
"translation": "configs/translation_small.yaml",
"generate": "configs/translation_small.yaml",
}

PIPELINES = {
"classifier": softmax_train,
"translation": translation_train,
"generate": generate_train,
}


Expand Down
54 changes: 54 additions & 0 deletions belt/data/image.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
"""Image Datasets

- CIFAR-10
- Fashion-MNIST

Author(s)
---------
Daniel Nicolas Gisolfi <dgisolfi3@gatech.edu>
"""

from torch.utils.data import DataLoader
from torchvision.datasets import CIFAR10, FashionMNIST

from belt.data.iris import SupervisedData
from belt.utils import DATA_DIR

# https://docs.pytorch.org/tutorials/beginner/basics/data_tutorial.html
fashion_mnist_class_map = {
0: "T-Shirt",
1: "Trouser",
2: "Pullover",
3: "Dress",
4: "Coat",
5: "Sandal",
6: "Shirt",
7: "Sneaker",
8: "Bag",
9: "Ankle Boot",
}
CLASS_NAMES = fashion_mnist_class_map.values()


def load_fashion_mnist(batch_size=64, **kwargs):
train = FashionMNIST(root=DATA_DIR, train=True, download=True)
test = FashionMNIST(root=DATA_DIR, train=False, download=True)

return SupervisedData(
train_loader=DataLoader(train, batch_size=batch_size, shuffle=True),
val_loader=DataLoader(test, batch_size=batch_size, shuffle=False),
test_loader=DataLoader(test, batch_size=batch_size, shuffle=False),
class_names=CLASS_NAMES,
)


def load_cifar10(batch_size=64, **kwargs):
train = CIFAR10(root=DATA_DIR, train=True, download=True)
test = CIFAR10(root=DATA_DIR, train=False, download=True)

return SupervisedData(
train_loader=DataLoader(train, batch_size=batch_size, shuffle=True),
val_loader=DataLoader(test, batch_size=batch_size, shuffle=False),
test_loader=DataLoader(test, batch_size=batch_size, shuffle=False),
class_names=CLASS_NAMES,
)
67 changes: 42 additions & 25 deletions belt/data/translation.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
"""Translation Datasets

Use hugging face datasets and tokenizers to create
Use hugging face datasets and tokenizers to create
language word pairings for translation task

Author(s)
Expand All @@ -10,15 +10,7 @@

import os
from dataclasses import dataclass
from pathlib import Path

DATA_DIR = Path(__file__).parents[2] / "data"
os.environ.setdefault("HF_HOME", str(DATA_DIR))
if (DATA_DIR / "datasets").exists():
os.environ.setdefault("HF_DATASETS_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")

import torch
from datasets import load_dataset
from torch.utils.data import DataLoader, Dataset
from transformers import AutoTokenizer, PreTrainedTokenizerBase
Expand All @@ -39,6 +31,7 @@ class TranslationData:
source_lang: str
target_lang: str


class TranslationDataset(Dataset):
def __init__(self, pairs, source_tokenizer, target_tokenizer, max_length):
self.pairs = pairs
Expand Down Expand Up @@ -68,22 +61,24 @@ def __getitem__(self, index):
return src_ids, tgt_ids


def _translation_loader(pairs, source_tokenizer, target_tokenizer, max_length, batch_size, shuffle):
def _translation_loader(
pairs, source_tokenizer, target_tokenizer, max_length, batch_size, shuffle
):
dataset = TranslationDataset(pairs, source_tokenizer, target_tokenizer, max_length)
return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle)



def lang_pairs(row, source_lang, target_lang):
translation = row.get("translation")
if not isinstance(translation, dict):
translation = row
source = translation.get(source_lang)

source = translation.get(source_lang)
target = translation.get(target_lang)

return str(source), str(target)


def load_translation_data(
dataset_name="multi30k",
dataset_config=None,
Expand All @@ -108,18 +103,29 @@ def load_translation_data(
os.environ.pop("TRANSFORMERS_OFFLINE", None)

model_name = f"Helsinki-NLP/opus-mt-{source_lang}-{target_lang}"
source_tokenizer = AutoTokenizer.from_pretrained(model_name, force_download=no_cache)
source_tokenizer = AutoTokenizer.from_pretrained(
model_name, force_download=no_cache
)
# opus-mt is directional, one tokenizer covers both sides
target_tokenizer = source_tokenizer

if dataset_name not in SUPPORTED_DATASETS:
raise ValueError(f"Unsupported dataset '{dataset_name}'. Choose from: {SUPPORTED_DATASETS}")
raise ValueError(
f"Unsupported dataset '{dataset_name}'. Choose from: {SUPPORTED_DATASETS}"
)
download_mode = "force_redownload" if no_cache else "reuse_dataset_if_exists"
dataset_split = f"{split}[:{sample_size}]" if sample_size else split
raw = (
load_dataset(dataset_name, dataset_config, split=dataset_split, download_mode=download_mode)
load_dataset(
dataset_name,
dataset_config,
split=dataset_split,
download_mode=download_mode,
)
if dataset_config
else load_dataset(dataset_name, split=dataset_split, download_mode=download_mode)
else load_dataset(
dataset_name, split=dataset_split, download_mode=download_mode
)
)
if shuffle:
raw = raw.shuffle(seed=seed)
Expand All @@ -139,20 +145,31 @@ def load_translation_data(

return TranslationData(
train_loader=_translation_loader(
train_pairs, source_tokenizer, target_tokenizer, max_length, batch_size, shuffle=True
train_pairs,
source_tokenizer,
target_tokenizer,
max_length,
batch_size,
shuffle=True,
),
val_loader=_translation_loader(
val_pairs, source_tokenizer, target_tokenizer, max_length, batch_size, shuffle=False
val_pairs,
source_tokenizer,
target_tokenizer,
max_length,
batch_size,
shuffle=False,
),
test_loader=_translation_loader(
test_pairs, source_tokenizer, target_tokenizer, max_length, batch_size, shuffle=False
test_pairs,
source_tokenizer,
target_tokenizer,
max_length,
batch_size,
shuffle=False,
),
source_tokenizer=source_tokenizer,
target_tokenizer=target_tokenizer,
source_lang=source_lang,
target_lang=target_lang,
)




14 changes: 13 additions & 1 deletion belt/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
from dataclasses import dataclass
from typing import Any

from belt.data.image import load_cifar10, load_fashion_mnist
from belt.data.iris import SupervisedData, iris
from belt.data.translation import TranslationData, load_translation_data

Expand All @@ -26,6 +27,14 @@ class DatasetLoaders:
"iris": DatasetLoaders(
supervised=iris,
),
# Image Data
"fashion_mnist": DatasetLoaders(
supervised=load_fashion_mnist,
),
"cifar10": DatasetLoaders(
supervised=load_cifar10,
),
# Translation Data
"opus_books_en_fr": DatasetLoaders(
translation=load_translation_data,
),
Expand All @@ -37,7 +46,10 @@ class DatasetLoaders:

# Basic Registry for Models and other composable components
class Registry(dict[str, type]):
"""Maps string names to classes and supports decorator-based registration."""
"""decorator based registration

used by the pipelines to register available models to run
"""

def register(self, name: str):
def decorator(cls: type) -> type:
Expand Down
55 changes: 55 additions & 0 deletions belt/supervised/generate.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
"""Generate Pipeline

Author(s)
---------
Daniel Nicolas Gisolfi <dgisolfi3@gatech.edu>
"""

import torch
from torch import nn

from belt.data.image import load_fashion_mnist
from belt.registry import dataset_registry
from belt.supervised.models import supervised_model_registry
from belt.supervised.pipeline import TorchPipeline
from belt.utils import describe_device, get_device


class GeneratePipeline(TorchPipeline):
"""Supervised image pipeline"""

def setup(self, config):
self.device = get_device(config.get("device", "auto"))
self.device_label = describe_device(self.device)

dataset_name = config.get("dataset", "fashion_mnist")
dataset_loader = dataset_registry.get(dataset_name)
if dataset_loader and dataset_loader.supervised:
data = dataset_loader.supervised(**config.get("data", {}))
else:
data = load_fashion_mnist(**config.get("data", {}))

self.train_loader = data.train_loader
self.val_loader = data.val_loader
self.test_loader = data.test_loader
self.class_names = data.class_names

self.model = supervised_model_registry.build(
config.get("model_name", "ImageClassifier"),
**config["model"],
num_classes=len(self.class_names),
device=self.device,
).to(self.device)
self.criterion = nn.CrossEntropyLoss()
self.optimizer = torch.optim.AdamW(
self.model.parameters(),
lr=config["training"]["lr"],
weight_decay=config["training"]["weight_decay"],
)

def extra_metrics(self, _):
return {"class_names": self.class_names}


def deploy(config_path, overrides=None):
return GeneratePipeline().run(config_path, overrides)
Loading