You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
feat(models): add GPT-2 (gpt2) text model support #833
GPT-2 is the classic OpenAI decoder-only transformer and a foundational architecture that many derived checkpoints and educational/regression checkpoints still ship in. mlx-lm supports it, but mlxcel does not yet load the gpt2 architecture. Porting it gives mlxcel a small, cheap-to-run reference model that is useful for smoke-testing the text path and unlocks the family of GPT-2 derivatives. The port is low effort because the block structure is close to the existing Llama-style skeleton once RMSNorm is swapped for LayerNorm and the Conv1D weight layout is handled.
mlx-community/gpt2 (and larger GPT-2 variants such as gpt2-medium / gpt2-large on the same org). Verify mlx-community checkpoint availability before validation.
Architecture notes
Classic GPT-2 decoder with these distinctive traits:
Learned absolute position embeddings (wpe) added to token embeddings (wte). No RoPE anywhere.
nn.LayerNorm (with bias), NOT RMSNorm.
Fused c_attn QKV projection and c_proj output projection stored in the HuggingFace Conv1D layout, which is transposed relative to nn.Linear. sanitize must transpose the c_attn, c_proj, and MLP c_fc weights so they load as standard linear layers.Correction (verified against the merged src/models/gpt2.rs, PR feat(models): add GPT-2 (gpt2) text model support #924, c7e9b6e0d): three details of that sentence are wrong. First, the transpose covers four projections per block, not three: the attention c_attn and c_proj plus the MLP c_fc and the MLP c_proj all load through Gpt2Layout::conv1d_linear. Second, it applies to the .weight tensors only, never to bias vectors: c_attn.bias ([3 * n_embd]), c_proj.bias and c_fc.bias are 1-D and are copied through untouched. Third, there is no sanitize in the shipped port. Gpt2Layout::detect probes the shape of h.0.attn.c_attn.weight once ([n_embd, 3 * n_embd] means the Conv1D layout and needs the transpose, [3 * n_embd, n_embd] means an already-transposed MLX conversion and does not), and the transpose is then applied per projection at load, with every other projection shape-checked against that one decision. Deciding by shape rather than by key prefix is what keeps h.N.attn.c_proj.weight correct, since it is square ([n_embd, n_embd]) and carries no layout signal of its own.
Tied output embedding (lm_head shares wte).
GELU MLP.
Implementation plan
Reuse the following mlxcel hooks:
Base the decoder block on the existing Llama-style block skeleton in src/models/llama3.rs, swapping RMSNorm for LayerNorm. LayerNorm already exists in mlxcel (used in the VLM encoders), so reuse that rather than adding a new norm.
Add a learned position-embedding table added to the token embeddings at the input boundary.
Add the Conv1D weight transpose (c_attn, c_proj, c_fc) in sanitize.Correction: applied at load rather than in a sanitize pass, and over four projections rather than three, per the architecture-note correction above. Detect the layout once from h.0.attn.c_attn.weight, then transpose the attention c_attn / c_proj and the MLP c_fc / c_proj weights, leaving the biases alone.
Effort: LOW.
Touchpoints & acceptance criteria
Delivered and closed. Implemented by PR #924, merged as c7e9b6e0d; mlxcel arch now reports the family. Test coverage: 25 unit tests in src/models/gpt2_tests.rs plus a detection test. Real-checkpoint validation: 40 greedy tokens from models/gpt2 reproduce the mlx-lm reference token id sequence exactly. Three real key layouts were additionally validated end to end: the bare raw layout, distilbert/distilgpt2 under a transformer. prefix, and mlx-community/gpt2-base-mlx under a model. prefix with pre-transposed weights. The unticked boxes below are the original pre-implementation plan and were not maintained during the work; treat the merged PR and its review thread as the record of what shipped.
Follow the checklist in docs/adding-models.md. Integration is complete only when the model loads and generates from a real checkpoint, not when modules compile in isolation.
Config struct + serde parse for the gpt2 config.
from_weights constructor.
Weight-key remap, including the Conv1D transpose of the attention c_attn / c_proj and the MLP c_fc / c_proj weights (weights only, decided once from a shape probe at load rather than in a sanitize pass; see the correction above).
Add the gpt2 arch-string arm to src/models/detection.rs.
Register the model in src/model_metadata.rs (for_each_model_registration!).
Add the TP/distributed inference arch-string if applicable.
_tests.rs unit tests beside the implementation.
Update docs/supported-models.md.
Validate on a real checkpoint via ./target/release/mlxcel generate and confirm mlxcel arch reports it.
Correction: the real-checkpoint validation step above originally read mlxcel list, which lists downloaded checkpoints in the local model store. The architecture registry is mlxcel arch, and that is what confirms the binary knows the family.
Summary
GPT-2 is the classic OpenAI decoder-only transformer and a foundational architecture that many derived checkpoints and educational/regression checkpoints still ship in. mlx-lm supports it, but mlxcel does not yet load the
gpt2architecture. Porting it gives mlxcel a small, cheap-to-run reference model that is useful for smoke-testing the text path and unlocks the family of GPT-2 derivatives. The port is low effort because the block structure is close to the existing Llama-style skeleton once RMSNorm is swapped for LayerNorm and the Conv1D weight layout is handled.Upstream reference
https://github.com/ml-explore/mlx-lm/blob/main/mlx_lm/models/gpt2.py
Public checkpoint
mlx-community/gpt2(and larger GPT-2 variants such as gpt2-medium / gpt2-large on the same org). Verify mlx-community checkpoint availability before validation.Architecture notes
Classic GPT-2 decoder with these distinctive traits:
wpe) added to token embeddings (wte). No RoPE anywhere.nn.LayerNorm(with bias), NOT RMSNorm.c_attnQKV projection andc_projoutput projection stored in the HuggingFace Conv1D layout, which is transposed relative tonn.Linear.Correction (verified against the mergedsanitizemust transpose thec_attn,c_proj, and MLPc_fcweights so they load as standard linear layers.src/models/gpt2.rs, PR feat(models): add GPT-2 (gpt2) text model support #924,c7e9b6e0d): three details of that sentence are wrong. First, the transpose covers four projections per block, not three: the attentionc_attnandc_projplus the MLPc_fcand the MLPc_projall load throughGpt2Layout::conv1d_linear. Second, it applies to the.weighttensors only, never to bias vectors:c_attn.bias([3 * n_embd]),c_proj.biasandc_fc.biasare 1-D and are copied through untouched. Third, there is nosanitizein the shipped port.Gpt2Layout::detectprobes the shape ofh.0.attn.c_attn.weightonce ([n_embd, 3 * n_embd]means the Conv1D layout and needs the transpose,[3 * n_embd, n_embd]means an already-transposed MLX conversion and does not), and the transpose is then applied per projection at load, with every other projection shape-checked against that one decision. Deciding by shape rather than by key prefix is what keepsh.N.attn.c_proj.weightcorrect, since it is square ([n_embd, n_embd]) and carries no layout signal of its own.wte).Implementation plan
Reuse the following mlxcel hooks:
src/models/llama3.rs, swapping RMSNorm for LayerNorm. LayerNorm already exists in mlxcel (used in the VLM encoders), so reuse that rather than adding a new norm.Add the Conv1D weight transpose (Correction: applied at load rather than in ac_attn,c_proj,c_fc) insanitize.sanitizepass, and over four projections rather than three, per the architecture-note correction above. Detect the layout once fromh.0.attn.c_attn.weight, then transpose the attentionc_attn/c_projand the MLPc_fc/c_projweights, leaving the biases alone.Effort: LOW.
Touchpoints & acceptance criteria
Delivered and closed. Implemented by PR #924, merged as
c7e9b6e0d;mlxcel archnow reports the family. Test coverage: 25 unit tests insrc/models/gpt2_tests.rsplus a detection test. Real-checkpoint validation: 40 greedy tokens frommodels/gpt2reproduce the mlx-lm reference token id sequence exactly. Three real key layouts were additionally validated end to end: the bare raw layout,distilbert/distilgpt2under atransformer.prefix, andmlx-community/gpt2-base-mlxunder amodel.prefix with pre-transposed weights. The unticked boxes below are the original pre-implementation plan and were not maintained during the work; treat the merged PR and its review thread as the record of what shipped.Follow the checklist in
docs/adding-models.md. Integration is complete only when the model loads and generates from a real checkpoint, not when modules compile in isolation.gpt2config.from_weightsconstructor.c_attn/c_projand the MLPc_fc/c_projweights (weights only, decided once from a shape probe at load rather than in asanitizepass; see the correction above).gpt2arch-string arm tosrc/models/detection.rs.src/model_metadata.rs(for_each_model_registration!)._tests.rsunit tests beside the implementation.docs/supported-models.md../target/release/mlxcel generateand confirmmlxcel archreports it.Correction: the real-checkpoint validation step above originally read
mlxcel list, which lists downloaded checkpoints in the local model store. The architecture registry ismlxcel arch, and that is what confirms the binary knows the family.Effort
LOW.