https://github.com/danielenricocahall/sparkformers

https://github.com/danielenricocahall/sparkformers

Science Score: 26.0%

This score indicates how likely this project is to be science-related based on various indicators:

  • CITATION.cff file
  • codemeta.json file
    Found codemeta.json file
  • .zenodo.json file
    Found .zenodo.json file
  • DOI references
  • Academic publication links
  • Committers with academic emails
  • Institutional organization owner
  • JOSS paper metadata
  • Scientific vocabulary similarity
    Low similarity (13.6%) to scientific vocabulary
Last synced: 11 months ago · JSON representation

Repository

Basic Info
  • Host: GitHub
  • Owner: danielenricocahall
  • License: mit
  • Language: Python
  • Default Branch: main
  • Size: 2.9 MB
Statistics
  • Stars: 1
  • Watchers: 0
  • Forks: 1
  • Open Issues: 1
  • Releases: 2
Created about 1 year ago · Last pushed about 1 year ago
Metadata Files
Readme Contributing License

README.md

Build Status license Supported Versions

Overview

img.png Welcome to Sparkformers, where we offer distributed training of Transformers models on Spark!

Motivation / Purpose

Derived from Elephas, however with HuggingFace removing support for Tensorflow, I decided to spin some of the logic off into its own separate project, and also rework the paradigm to support the Torch backend! The purpose of this project is to serve as an experimental backend for distributed training that may be more developergonomic compared to other solutions such as Ray. Additionally, Sparkformers offers the capability for distributed prediction, model calling, and generation (for causal/autoregressive models).

The project is currently in a beta/experimental state. While not yet production ready, I invite you to experiment, provide feedback, and/or even contribute!

Approach

Training: The current architecture utilizes federated averaging (FedAvg), meaning that each executor is trained on a subset of data, and the model weights are averaged across all executors after each epoch. The original model is then updated with the averaged weights, and then the process is repeated for the next epoch.

Inference: The input data is distributed across the executors, and each executor performs the inference on its subset of data. The results are then collected and returned to the driver.

Generation: Same as above, but with the generate method of the model.

Installation

To install, you can simply run: bash pip install sparkformers `

(or uv add, poetry add, etc. with whichever project dependency management tool you may use).

Examples

Note that all examples are also available in the examples directory.

Autoregressive (Causal) Language Model Training and Inference

```python from datasets import loaddataset from sklearn.modelselection import traintestsplit from sparkformers.sparkformer import Sparkformer from transformers import ( AutoTokenizer, AutoModelForCausalLM, ) import torch

batch_size = 16 epochs = 100

dataset = loaddataset("gfigueroa/wikitextprocessed") x = dataset["train"]["text"]

xtrain, xtest = traintestsplit(x, test_size=0.1)

model_name = "hf-internal-testing/tiny-random-gptj"

model = AutoModelForCausalLM.frompretrained(modelname) tokenizer = AutoTokenizer.frompretrained(modelname) tokenizer.padtoken = tokenizer.eostoken tokenizerkwargs = { "maxlength": 50, "padding": True, "truncation": True, "padding_side": "left", }

sparkformermodel = Sparkformer( model=model, tokenizer=tokenizer, loader=AutoModelForCausalLM, optimizerfn=lambda params: torch.optim.AdamW(params, lr=1e-3), tokenizerkwargs=tokenizerkwargs, num_workers=2, )

perform distributed training

sparkformermodel.train(xtrain, epochs=epochs, batchsize=batchsize)

perform distributed generation

generations = sparkformermodel.generate( xtest, maxnewtokens=10, numreturnsequences=1 )

decode the generated texts

generatedtexts = [ tokenizer.decode(output, skipspecial_tokens=True) for output in generations ]

for i, text in enumerate(generatedtexts): print(f"Original text {i}: {xtest[i]}") print(f"Generated text {i}: {text}") ```

Sequence Classification

```python from datasets import loaddataset from sklearn.modelselection import traintestsplit from torch import softmax

from sparkformers.sparkformer import Sparkformer from transformers import ( AutoTokenizer, AutoModelForSequenceClassification, ) import numpy as np import torch

batch_size = 16 epochs = 20

dataset = loaddataset("agnews") x = dataset["train"]["text"][:2000] y = dataset["train"]["label"][:2000]

xtrain, xtest, ytrain, ytest = traintestsplit(x, y, test_size=0.1)

model_name = "prajjwal1/bert-tiny"

model = AutoModelForSequenceClassification.frompretrained( modelname, numlabels=len(np.unique(y)), problemtype="singlelabelclassification", )

tokenizer = AutoTokenizer.frompretrained(modelname) tokenizerkwargs = {"padding": True, "truncation": True, "maxlength": 512}

sparkformermodel = Sparkformer( model=model, tokenizer=tokenizer, loader=AutoModelForSequenceClassification, optimizerfn=lambda params: torch.optim.AdamW(params, lr=2e-4), tokenizerkwargs=tokenizerkwargs, num_workers=2, )

perform distributed training

sparkformermodel.train(xtrain, ytrain, epochs=epochs, batchsize=batch_size)

perform distributed inference

predictions = sparkformermodel.predict(xtest) for i, pred in enumerate(predictions[:10]): probs = softmax(torch.tensor(pred), dim=-1) print(f"Example {i}: probs={probs.numpy()}, predicted={probs.argmax().item()}")

review the predicted labels

print([int(np.argmax(pred)) for pred in predictions]) ```

Token Classification (NER)

```python from sklearn.modelselection import traintestsplit from sparkformers.sparkformer import Sparkformer from transformers import ( AutoTokenizer, AutoModelForTokenClassification, ) from datasets import loaddataset import numpy as np import torch

batchsize = 5 epochs = 1 modelname = "hf-internal-testing/tiny-bert-for-token-classification"

model = AutoModelForTokenClassification.frompretrained(modelname) tokenizer = AutoTokenizer.frompretrained(modelname)

def tokenizeandalignlabels(examples): tokenizedinputs = tokenizer( examples["tokens"], truncation=True, issplitintowords=True ) labels = [] for i, label in enumerate(examples["nertags"]): wordids = tokenizedinputs.wordids(batchindex=i) previouswordidx = None labelids = [] for wordidx in wordids: if wordidx is None: labelids.append(-100) elif wordidx != previouswordidx: labelids.append(label[wordidx]) else: labelids.append(-100) previouswordidx = wordidx labels.append(labelids) tokenizedinputs["labels"] = labels return tokenized_inputs

dataset = loaddataset("conll2003", split="train[:5%]", trustremotecode=True) dataset = dataset.map(tokenizeandalignlabels, batched=True)

x = dataset["tokens"] y = dataset["labels"]

xtrain, xtest, ytrain, ytest = traintestsplit(x, y, test_size=0.1)

tokenizerkwargs = { "padding": True, "truncation": True, "issplitintowords": True, }

sparkformermodel = Sparkformer( model=model, tokenizer=tokenizer, loader=AutoModelForTokenClassification, optimizerfn=lambda params: torch.optim.AdamW(params, lr=5e-5), tokenizerkwargs=tokenizerkwargs, num_workers=2, )

sparkformermodel.train(xtrain, ytrain, epochs=epochs, batchsize=batch_size)

inputs = tokenizer(xtest, **tokenizerkwargs) distributedpreds = sparkformermodel(**inputs) print([int(np.argmax(x)) for x in np.squeeze(distributed_preds)])

```

TODO

  • [ ] Add support for distributed training of other model types (e.g., image classification, object detection, etc.)
  • [ ] Support training paradigms using Trainer, TrainingArguments, and DataCollater
  • [ ] Expose more configuration options
  • [ ] Consider simplifying the API further (e.g; builder pattern, providing the model string and push loader logic inside the Sparkformer class, etc.) > 💡 Interested in contributing? Check out the Local Development & Contributions Guide.

Owner

  • Name: Danny
  • Login: danielenricocahall
  • Kind: user
  • Location: Philadelphia, PA
  • Company: Disney Streaming Services

GitHub Events

Total
  • Create event: 6
  • Issues event: 1
  • Release event: 6
  • Issue comment event: 1
  • Public event: 1
  • Push event: 11
  • Fork event: 1
Last Year
  • Create event: 6
  • Issues event: 1
  • Release event: 6
  • Issue comment event: 1
  • Public event: 1
  • Push event: 11
  • Fork event: 1

Committers

Last synced: about 1 year ago

All Time
  • Total Commits: 71
  • Total Committers: 1
  • Avg Commits per committer: 71.0
  • Development Distribution Score (DDS): 0.0
Past Year
  • Commits: 71
  • Committers: 1
  • Avg Commits per committer: 71.0
  • Development Distribution Score (DDS): 0.0
Top Committers
Name Email Commits
daniel.cahall d****l@g****m 71

Issues and Pull Requests

Last synced: 11 months ago

All Time
  • Total issues: 1
  • Total pull requests: 0
  • Average time to close issues: N/A
  • Average time to close pull requests: N/A
  • Total issue authors: 1
  • Total pull request authors: 0
  • Average comments per issue: 4.0
  • Average comments per pull request: 0
  • Merged pull requests: 0
  • Bot issues: 0
  • Bot pull requests: 0
Past Year
  • Issues: 1
  • Pull requests: 0
  • Average time to close issues: N/A
  • Average time to close pull requests: N/A
  • Issue authors: 1
  • Pull request authors: 0
  • Average comments per issue: 4.0
  • Average comments per pull request: 0
  • Merged pull requests: 0
  • Bot issues: 0
  • Bot pull requests: 0
Top Authors
Issue Authors
  • danielenricocahall (1)
Pull Request Authors
Top Labels
Issue Labels
enhancement (1)
Pull Request Labels

Packages

  • Total packages: 1
  • Total downloads:
    • pypi 22 last-month
  • Total dependent packages: 0
  • Total dependent repositories: 0
  • Total versions: 10
  • Total maintainers: 1
pypi.org: sparkformers

Distributed deep learning for Hugging Face Transformers on Spark

  • Versions: 10
  • Dependent Packages: 0
  • Dependent Repositories: 0
  • Downloads: 22 Last month
Rankings
Dependent packages count: 9.0%
Average: 29.8%
Dependent repos count: 50.6%
Maintainers (1)
Last synced: 11 months ago

Dependencies

.github/workflows/ci.yaml actions
  • actions/checkout v3 composite
  • astral-sh/setup-uv v5 composite
pyproject.toml pypi
  • pyspark <=4.0.0
  • torch >=2.7.1
  • transformers <5.0.0
uv.lock pypi
  • aiohappyeyeballs 2.6.1
  • aiohttp 3.12.12
  • aiosignal 1.3.2
  • async-timeout 5.0.1
  • attrs 25.3.0
  • certifi 2025.4.26
  • cfgv 3.4.0
  • charset-normalizer 3.4.2
  • colorama 0.4.6
  • datasets 3.6.0
  • dill 0.3.8
  • distlib 0.3.9
  • exceptiongroup 1.3.0
  • execnet 2.1.1
  • filelock 3.18.0
  • findspark 2.0.1
  • frozenlist 1.7.0
  • fsspec 2025.3.0
  • hf-xet 1.1.3
  • huggingface-hub 0.33.0
  • identify 2.6.12
  • idna 3.10
  • iniconfig 2.1.0
  • jinja2 3.1.6
  • joblib 1.5.1
  • markupsafe 3.0.2
  • mock 5.2.0
  • mpmath 1.3.0
  • multidict 6.4.4
  • multiprocess 0.70.16
  • networkx 3.2.1
  • networkx 3.4.2
  • networkx 3.5
  • nodeenv 1.9.1
  • numpy 2.0.2
  • numpy 2.2.6
  • numpy 2.3.0
  • nvidia-cublas-cu12 12.6.4.1
  • nvidia-cuda-cupti-cu12 12.6.80
  • nvidia-cuda-nvrtc-cu12 12.6.77
  • nvidia-cuda-runtime-cu12 12.6.77
  • nvidia-cudnn-cu12 9.5.1.17
  • nvidia-cufft-cu12 11.3.0.4
  • nvidia-cufile-cu12 1.11.1.6
  • nvidia-curand-cu12 10.3.7.77
  • nvidia-cusolver-cu12 11.7.1.2
  • nvidia-cusparse-cu12 12.5.4.2
  • nvidia-cusparselt-cu12 0.6.3
  • nvidia-nccl-cu12 2.26.2
  • nvidia-nvjitlink-cu12 12.6.85
  • nvidia-nvtx-cu12 12.6.77
  • packaging 25.0
  • pandas 2.3.0
  • pep8 1.7.1
  • platformdirs 4.3.8
  • pluggy 1.6.0
  • pre-commit 4.2.0
  • propcache 0.3.2
  • py4j 0.10.9.9
  • pyarrow 20.0.0
  • pygments 2.19.1
  • pyspark 4.0.0
  • pytest 8.4.0
  • pytest-cache 1.0
  • pytest-pep8 1.0.6
  • pytest-spark 0.8.0
  • python-dateutil 2.9.0.post0
  • pytz 2025.2
  • pyyaml 6.0.2
  • regex 2024.11.6
  • requests 2.32.4
  • ruff 0.11.13
  • safetensors 0.5.3
  • scikit-learn 1.6.1
  • scikit-learn 1.7.0
  • scipy 1.13.1
  • scipy 1.15.3
  • setuptools 80.9.0
  • six 1.17.0
  • sparkformers 0.0.0
  • sympy 1.14.0
  • threadpoolctl 3.6.0
  • tokenizers 0.21.1
  • tomli 2.2.1
  • torch 2.7.1
  • tqdm 4.67.1
  • transformers 4.52.4
  • triton 3.3.1
  • ty 0.0.1a10
  • typing-extensions 4.14.0
  • tzdata 2025.2
  • urllib3 2.4.0
  • virtualenv 20.31.2
  • xxhash 3.5.0
  • yarl 1.20.1