gpjt/jax-with-mha-bias-no-dropout
Model Card for gpjt/jax-with-mha-bias-no-dropout
This model is gpjt/jax-with-mha-bias-no-dropout, a trained-from-scratch base model using the GPT-2-style architecture from Sebastian Raschka's book "Build a Large Language Model (from Scratch)".
The model was trained in JAX using a black-box reimplementation of Raschka's original PyTorch code, but the safetensors have been converted back to a PyTorch-compatible format for convenience -- if you run this model, it will be in PyTorch.
Model Details
Model Description
- Developed by: Giles Thomas, based on code by Sebastian Raschka
- Model type: GPT-2 style transformers-based causal LLM.
- License: Apache 2
- Parameters: 163,000,320
- Context length: 1,024
- Embedding dimensions: 768
- MHA heads: 12
- Layers: 12
- QKV bias: False
- Weight tying: False
Don't have high expectations for the model! It has only 163M parameters (the GPT-2 "small" size) and was trained on roughly the Chinchilla-optimal number of tokens (~20x the number of parameters), which means that it doesn't know many facts and is not terribly smart. If you want to do serious work, use a serious model (I like Qwen's). But if you want to build on this and see what you can do with a 2020-vintage LLM, please do feel free to play with it!
Model Sources
- Repository: https://github.com/gpjt/jax-gpt2-from-scratch for the training code, gpjt/ddp-base-model-from-scratch for the code used to run it from here.
- Blog post: Writing an LLM from scratch, part 34b -- from bigrams to GPT-2, one component at a time (in JAX). This is the third full LLM trained in the post, in the "Adding bias to the MHA output projections" section.
How to Get Started with the Model
You can download and run the model for inference directly:
from transformers import pipeline
pipe = pipeline("text-generation", model="gpjt/jax-with-mha-bias-no-dropout", trust_remote_code=True)
out = pipe(
"Every effort moves you",
max_new_tokens=20,
do_sample=True,
temperature=1.4,
top_k=25,
)
print(out[0]["generated_text"])Note that because it uses custom code, you'll need to set trust_remote_code to True.
It supports AutoTokenizer, AutoModel and AutoModelForCausalLM:
>>> from transformers import AutoTokenizer, AutoModel, AutoModelForCausalLM
>>> tokenizer = AutoTokenizer.from_pretrained("gpjt/jax-with-mha-bias-no-dropout")
>>> model = AutoModel.from_pretrained("gpjt/jax-with-mha-bias-no-dropout", trust_remote_code=True)
>>> llm_model = AutoModelForCausalLM.from_pretrained("gpjt/jax-with-mha-bias-no-dropout", trust_remote_code=True)You can also fine-tune it; this notebook has an example.
Again, don't expect too much from this model! It's a 163M-parameter GPT-2 one, trained on a limited number of tokens. It's both dumb and ignorant ;-)
Training Details
- Machine type: Local machine with an RTX 3090
- Tokens: 3,260,190,720 (Chinchilla-optimal of 20x parameters) rounded up to the nearest batch.
- Dataset: gpjt/fineweb-gpt2-tokens
- Micro-batch size: 6
- Global batch size: 96
- Dropout: 0.0
- Gradient clipping: 3.5
- Learning rate: 0.0014
- Schedule learning rate: True
- Weight decay: 0.01
