High level interface for text applications using PyTroch RNN's.
Project description
GeNN
GeNN (generative neural networks) is a high-level interface for text applications using PyTorch RNN's.
Features
- Preprocessing:
- Parsing txt, json, and csv files.
- NLTK, regex and spacy tokenization support.
- GloVe and fastText pretrained embeddings, with the ability to fine-tune for your data.
- Architectures and customization:
- GPT-2 with small, medium, and large variants.
- LSTM and GRU, with variable size.
- Variable number of layers and batches.
- Dropout.
- Text generation:
- Random seed sampling from the n first tokens in all instances, or the most frequent token.
- Top-K sampling for next token prediction with variable K.
- Nucleus sampling for next token prediction with variable probability threshold.
Getting started
How to install
pip install genn
Prerequisites
- PyTorch 1.4.0
pip install torch==1.4.0
- Pytorch Transformers
pip install pytorch_transformers
- NumPy
pip install numpy
- fastText
pip install fasttext
Use the package manager pip to install genn.
Usage
from genn import Preprocessing, LSTMGenerator, GPT2
#LSTM example
ds = Preprocessing("data.txt")
gen = LSTMGenerator(ds, nLayers = 2,
batchSize = 16,
embSize = 64,
lstmSize = 16,
epochs = 20)
#Train the model
gen.run()
# Generate 5 new documents
print(gen.generate_document(5))
#GPT-2 example
gen = GPT2("data.txt",
taskToken = "Movie:",
epochs = 7,
variant = "medium")
#Train the model
gen.run()
#Generate 10 new documents
print(gen.generate_document(10))
For more examples on how to use Preprocessing, please refer to this file.
For more examples on how to use LSTMGenerator and GRUGenerator, please refer to this file.
For more examples on how to use GPT2, please refer to this file
Contributing
Pull requests are welcome. For major changes, please open an issue first to discuss what you would like to change.
License
Distributed under the MIT License. See LICENSE for more information.
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
abdoTheBest-0.9.3.tar.gz
(13.2 kB
view hashes)
Built Distribution
Close
Hashes for abdoTheBest-0.9.3-py3-none-any.whl
Algorithm | Hash digest | |
---|---|---|
SHA256 | ff3951614c90edfaf0831a51a05171cfbca2af18f0e642a44823006a21b8b16d |
|
MD5 | d18817bf5afdc807d7206d64a148a3f6 |
|
BLAKE2b-256 | c911225a1ae412d8bad561c219f1eb354bd2015a2715c83e77911c45cbb51b04 |