How to Generate Text Using a Trained Transformer Model: Complete Implementation Guide
You generate text by tokenizing a prompt with tiktoken, loading a trained decoder-only Transformer checkpoint, and calling the model.generate() method to autoregressively sample tokens until you reach the desired output length.
This guide walks through the exact implementation in the FareedKhan-dev/train-llm-from-scratch repository, which provides a GPT-style decoder-only Transformer with causal self-attention for next-token prediction. The following steps cover the complete pipeline from loading trained weights to producing human-readable output.
Understanding the Generation Architecture
The repository implements a decoder-only Transformer (GPT-style architecture) defined in src/models/transformer.py. This model stacks multiple TransformerBlock layers (defined in src/models/transformer_block.py), where each block contains:
- Token and positional embeddings (
self.token_embedandself.pos_embed) - Causal self-attention (
src/models/attention.py) that restricts attention to previous tokens only - Feed-forward MLP (
src/models/mlp.py) with residual connections and layer normalization
The causal mask in the attention layer ensures the model maintains the autoregressive property required for text generation—each token prediction depends only on preceding tokens, never future positions.
Step-by-Step Text Generation Pipeline
1. Load the Tokenizer
The generation pipeline begins by converting your input prompt into token IDs using the tiktoken library. As implemented in scripts/generate_text.py, the code uses the r50k_base encoding to ensure vocabulary consistency with the training data.
import tiktoken
enc = tiktoken.get_encoding("r50k_base")
input_ids = enc.encode("Your prompt here")
2. Initialize the Model Architecture
Create a Transformer instance from src/models/transformer.py using the same hyperparameters defined in config/config.py. The configuration must match the trained checkpoint to avoid size mismatches during state dict loading.
from src.models.transformer import Transformer
cfg = {
"vocab_size": 50304,
"block_size": 1024,
"n_embed": 768,
"n_head": 12,
"n_layer": 12,
"dropout": 0.1,
}
model = Transformer(**cfg)
3. Restore Trained Weights
Load the checkpoint produced by scripts/train_transformer.py (a .pth file) and copy the parameters into the model. Move the model to your target device and set evaluation mode to disable dropout.
import torch
device = "cuda" if torch.cuda.is_available() else "cpu"
model.load_state_dict(torch.load("models/your_model.pth", map_location=device))
model = model.to(device)
model.eval()
4. Execute the Generation Loop
The core generation logic resides in Transformer.generate (lines 96-107 of src/models/transformer.py). This method accepts tokenized input and max_new_tokens, then repeatedly:
- Feeds the current sequence through the model
- Extracts logits for the last time-step
- Applies greedy sampling (
argmax) or temperature scaling - Appends the selected token ID to the sequence
- Continues until the specified number of new tokens is generated
with torch.no_grad():
# input_ids shape: (batch_size, seq_len)
generated = model.generate(input_ids, max_new_tokens=150)
5. Decode the Output
Convert the generated token IDs back to readable text using the tokenizer's decode method.
output_text = enc.decode(generated[0].tolist())
print(output_text)
Complete Python Implementation
This runnable example combines all steps from scripts/generate_text.py into a single function:
import torch
import tiktoken
from src.models.transformer import Transformer
def generate_text(
model_path: str,
prompt: str,
max_new: int = 100,
device: str = "cuda"
):
# Tokenize input
enc = tiktoken.get_encoding("r50k_base")
input_ids = torch.tensor(
enc.encode(prompt),
dtype=torch.long
)[None, ...].to(device)
# Initialize model architecture
cfg = {
"vocab_size": 50304,
"block_size": 1024,
"n_embed": 768,
"n_head": 12,
"n_layer": 12,
"dropout": 0.1,
}
model = Transformer(**cfg).to(device)
# Load trained weights
model.load_state_dict(
torch.load(model_path, map_location=device)
)
model.eval()
# Generate autoregressively
with torch.no_grad():
generated_ids = model.generate(
input_ids,
max_new_tokens=max_new
)[0]
# Decode to string
return enc.decode(generated_ids.tolist())
# Generate text
result = generate_text(
model_path="models/your_model.pth",
prompt="The future of AI is",
max_new=80,
device="cpu"
)
print(result)
Command-Line Interface
For direct terminal usage without writing Python scripts, invoke the provided CLI wrapper:
python scripts/generate_text.py \
--model_path models/your_model.pth \
--input_text "Once upon a time" \
--max_new_tokens 150
Key Implementation Details
Causal Self-Attention Mechanism
The src/models/attention.py file implements scaled dot-product attention with a causal mask (upper-triangular masking set to negative infinity). This prevents the model from attending to future token positions during training and inference, preserving the autoregressive generation flow.
Sampling Strategy
By default, the generate method uses greedy decoding (selecting the token with the highest logit value). To implement temperature sampling, modify the logits processing in src/models/transformer.py before the argmax operation by dividing by a temperature parameter (typically 0.7–1.0) and sampling from the resulting probability distribution.
Configuration Consistency
The vocab_size (50304), block_size (1024), and layer dimensions must exactly match the configuration used during training in scripts/train_transformer.py. Mismatched parameters will raise size mismatch errors when loading the state dictionary.
Summary
- Tokenization: Use
tiktoken.get_encoding("r50k_base")to convert prompts into token IDs that match the model's vocabulary. - Model Setup: Initialize the
Transformerclass fromsrc/models/transformer.pyand load weights from the.pthcheckpoint usingtorch.load. - Generation Process: Call
model.generate()to perform autoregressive next-token prediction with causal masking. - Decoding: Apply
tiktoken.decode()to convert the generated token sequence back to human-readable text. - Key Files: The pipeline relies on
src/models/transformer.py(generation logic),src/models/attention.py(causal masking), andscripts/generate_text.py(CLI wrapper).
Frequently Asked Questions
Why does the model only predict one token at a time?
The decoder-only Transformer uses causal masking to restrict attention to previous positions only. This autoregressive design requires feeding the generated token back into the model as input for the next prediction. While seemingly inefficient, this approach allows the model to maintain full context of the generated sequence and produce coherent, contextually appropriate continuations.
Can I use a different tokenizer than tiktoken?
The repository is specifically configured for tiktoken with the r50k_base encoding. While you could theoretically swap in other tokenizers like Hugging Face's AutoTokenizer, you would need to retrain the model from scratch using scripts/train_transformer.py with the new vocabulary, as the embedding dimensions in src/models/transformer.py are tied to the specific token ID space.
How do I implement temperature sampling for more creative outputs?
Modify the Transformer.generate method in src/models/transformer.py (around lines 96-107) to apply temperature scaling before sampling. Divide the logits by your temperature value (lower values like 0.7 produce focused output, higher values like 1.2 increase randomness), then use torch.multinomial to sample from the softmax distribution instead of using argmax.
What hardware requirements are needed for text generation?
You can generate text on either CPU or CUDA-enabled GPUs. The model supports device-agnostic loading via map_location in torch.load. For the default configuration (768 embedding dimensions, 12 layers, 12 heads), generation works efficiently on consumer GPUs with 8GB+ VRAM or modern CPUs with sufficient RAM to hold the model parameters (approximately 1-2GB for the checkpoint).
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →