gpjt/jax-with-mha-bias-no-dropout-extended
Model Card for gpjt/jax-with-mha-bias-no-dropout-extended
This model is gpjt/jax-with-mha-bias-no-dropout-extended, 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.
This model was deliberately overtrained on approximately 40 tokens per parameter (over one epoch).
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,009,536
- 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 twice the Chinchilla-optimal number of tokens (~40x 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: Why do OpenAI's GPT-2 weights beat mine? Part three: testing overtraining This is the first model trained in the post, in the "The extended-train model" 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-extended", 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-extended")
>>> model = AutoModel.from_pretrained("gpjt/jax-with-mha-bias-no-dropout-extended", trust_remote_code=True)
>>> llm_model = AutoModelForCausalLM.from_pretrained("gpjt/jax-with-mha-bias-no-dropout-extended", 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: 6,520,381,440 (Twice the 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
