--- license: apache-2.0 language: - en tags: - text-generation - small-models - chat - chatml - mla - jepa - experimental pipeline_tag: text-generation library_name: transformers datasets: - HuggingFaceFW/fineweb-edu --- # Escarda-86M ~86M decoder I trained from scratch. Chat checkpoint. Laptop / free CPU is the point. Architecture is SpikeWhale: MLA, n-gram engram memory, hyper-connections, HRM refine, JEPA + multi-token-prediction auxiliaries. Trained on **Modal credits** for the **Small Models, Big Adventures Hackathon**. I picked this one after a seed-controlled bake-off over 28 checkpoints plus a head-to-head battle test - coherence, instruction-following, and not collapsing into repetition. > **Live demo:** [Quazim0t0/Escarda-86M-Chat](https://huggingface.co/spaces/Quazim0t0/Escarda-86M-Chat) > > **Related:** base → [Quazim0t0/Escarda-86M-Base](https://huggingface.co/Quazim0t0/Escarda-86M-Base) > (use that if you want to SFT it yourself). > > Numbers: [Evaluation](#evaluation). --- ## Model summary | | | |---|---| | **Parameters** | ~85.7M (`tie_word_embeddings=True`) | | **Type** | Decoder-only autoregressive LM (`SpikeWhaleLM`, `model_type: spike_whale`) | | **Hidden size** | 640 | | **Layers** | 16 | | **Attention heads** | 10 (`head_dim=64`), 1 KV head (multi-query) | | **Context length** | 4096 tokens | | **Vocab size** | 16,512 (custom ChatML-aware tokenizer) | | **Positional encoding** | Decoupled RoPE (`theta=10000`) + NoPE split | | **Precision** | trained in float32 | | **License** | Apache-2.0 | --- ## Architecture Dense decoder. Flags match the released `config.json`. These models sit on **SpikeWhaleLM** (~86M, 16 layers, hidden 640, 4096 context, 16,512 vocab, tied embeddings). Same stack as below. ### Attention - Multi-head Latent Attention (MLA) + XSA `use_xsa=True`, `use_qk_norm=True` - **MLA-style low-rank projections**: queries and the output projection are LoRA-compressed (`q_lora_rank=128`, `o_lora_rank=128`). Attention params / KV stay small. - **Decoupled position encoding**: each head splits into a RoPE part (`qk_rope_head_dim=16`) and a NoPE part (`nope_head_dim=48`). Some of the head has rotary position; the rest does not. - **Multi-query attention**: `num_key_value_heads=1` - all query heads share one KV head. Cheaper KV cache. - **QK-norm** on the attention logits. ### Engram n-gram memory `use_engram=True` Hashes local n-grams (up to `engram_max_ngram=3`) into a learned table (`engram_table_size=4096`, `engram_num_heads=2`, `engram_compress_dim=32`) and gates that back into the residual (`engram_gate_init_bias=-1.0`, mostly off at init). Surface patterns without burning depth on them. ### Hash-lookup layers `num_hash_layers=2` - `MultiHeadHashLookup` next to the token embeddings. ### Hyper-Connections (not a plain residual) `use_hyper_connections=True` (`hc_mult=2`, `hc_sinkhorn_iters=20`, `hc_eps=1e-6`) Learned width-expanded connections, Sinkhorn-normalized routing, instead of a single identity add. ### HRM refinement `use_hrm_refine=True` (`hrm_refine_dim=128`, `hrm_refine_steps=1`) Small Hierarchical Reasoning Model block. Extra latent pass over hidden states before the output head. ### Feed-forward (MoE-capable, dense in this release) The block can do DeepSeek-style sparse MoE (`n_routed_experts=6`, `n_shared_experts=1`, `num_experts_per_tok=2`, `scoring_func=sqrtsoftplus`). This checkpoint is **dense** (`use_moe=False`, `moe_layers=[]`). Latency stays boring and predictable. ### Training-time auxiliary objectives No extra cost at inference: - **JEPA** (`use_jepa=True`, `jepa_pred_dim=256`, `jepa_horizon=1`, `jepa_loss_weight=0.1`) - Joint-Embedding Predictive loss on future latents. - **Multi-Token Prediction (MTP)** (`num_nextn_predict_layers=1`, `mtp_loss_weight=0.3`) - DeepSeek-V3-style extra head, more than one next token. - **z-loss** (`zloss_coef=1e-4`) for logit stability. > **JEPA vs HRM.** Escarda trains **both** HRM refine and JEPA > (`use_hrm_refine=True`, `use_jepa=True`). Byrne drops JEPA and keeps HRM only. --- ## Tokenizer & chat format **`SpikeTokenizer`**. Byte-level length-max (greedy longest-match), 16,512 vocab. Not BPE. Text → UTF-8 → latin-1 bytes → longest vocab key that fits. ChatML-aware. Atomic specials: `<|im_start|>`, `<|im_end|>`, ``/``, ``/``, tool-call markers, plus ``/``/``/``. Ships as a `PreTrainedTokenizer` (`spike_tokenizer.py`). Load with `AutoTokenizer.from_pretrained(..., trust_remote_code=True)`. `` (id 2) on every sequence. `<|im_end|>` and `` (id 3) end a turn. One turn: ``` <|im_start|>{role}\n{content}<|im_end|>\n ``` Generation starts after a trailing `<|im_start|>assistant\n`. --- ## Inference ChatML prompt, nucleus top-p 0.9, stop on `<|im_end|>`. This is what I used when it looked best: ```python import torch, torch.nn.functional as F from model_v2 import SpikeWhaleLM # custom architecture (ship with the repo) from spike_tokenizer import SpikeTokenizer from chat_format import format_chat, IM_END tok = SpikeTokenizer("tokenizer.json") model = SpikeWhaleLM.from_pretrained("Quazim0t0/Escarda-86M").eval() end_id = tok.convert_tokens_to_ids(IM_END) prompt = format_chat([{"role": "user", "content": "Explain photosynthesis in one sentence."}], add_generation_prompt=True) ids = torch.tensor(tok.encode(prompt)).unsqueeze(0) out = model(ids, use_cache=True); past = out.past_key_values; last = out.logits[0, -1] gen = [] for _ in range(120): p = F.softmax(last.float() / 0.3, -1) sp, si = p.sort(descending=True); cut = sp.cumsum(0) > 0.9 cut[1:] = cut[:-1].clone(); cut[0] = False; sp[cut] = 0 nxt = si[torch.multinomial(sp / sp.sum(), 1)].item() if nxt == end_id: break gen.append(nxt) out = model(torch.tensor([[nxt]]), past_key_values=past, use_cache=True) past = out.past_key_values; last = out.logits[0, -1] print(tok.decode(gen, skip_special_tokens=True)) ``` > Escarda is **not** a stock `transformers` model. You need `model_v2.py`, > `config.py`, `spike_tokenizer.py`, `chat_format.py`. Easiest path is the > [demo Space](https://huggingface.co/spaces/Quazim0t0/Escarda-86M-Chat). --- ## Evaluation Zero-shot multiple-choice, continuation log-likelihood on each task's val/test split. Standard error is binomial (`sqrt(p(1-p)/n)`). ### Language modeling byte_ppl is `exp(sum_NLL_nats / total_UTF8_bytes)` on WikiText-2 test (tokenizer-independent). BLiMP is `logprob(good) > logprob(bad)` on 12 paradigms × 150. | Metric | Value | |---|---| | WikiText-2 byte_ppl ↓ | 2.4898 | | BLiMP acc ↑ | **0.7483** | Chat checkpoint has the **best BLiMP** in the Escarda family even though [Base](https://huggingface.co/Quazim0t0/Escarda-86M-Base) has lower perplexity. PPL is not tracking capability here. ### Standard small-model suite | Task | acc | ± | acc_norm | ± | |---|---|---|---|---| | arc_easy | 0.3683 | 0.0099 | 0.3628 | 0.0099 | | arc_challenge | 0.1988 | 0.0117 | 0.2312 | 0.0123 | | hellaswag | 0.2845 | 0.0045 | 0.2928 | 0.0045 | | winogrande | 0.5067 | 0.0140 | - | - | | piqa | 0.5881 | 0.0115 | 0.5800 | 0.0115 | | openbookqa | 0.1600 | 0.0164 | 0.2720 | 0.0199 | | boolq | 0.4624 | 0.0087 | - | - | Random: arc/hellaswag/openbookqa ≈ 0.25; winogrande/boolq ≈ 0.50. At this size a lot of it sits near chance. Signal is mostly piqa (0.58) plus winogrande/boolq. ### ArithMark-2.0 ([AxiomicLabs](https://huggingface.co/datasets/AxiomicLabs/ArithMark-2.0)) Multiple-choice integer arithmetic (n = 2,500, chance = 0.25). | Metric | Value | |---|---| | acc | 0.2932 ± 0.0091 | | acc_norm | 0.2816 ± 0.0090 | Aggregate is flat. Underneath it is not. **~2× chance on multiplication and division**, at/below chance on add/sub: | Topic | acc_norm | n | | Difficulty | acc_norm | n | |---|---|---|---|---|---|---| | division | **0.5385** | 130 | | easy | 0.2872 | 1250 | | multiplication | **0.5278** | 144 | | medium | 0.2973 | 750 | | parentheses_two_ops | 0.3352 | 355 | | hard | 0.2440 | 500 | | mixed_two_ops | 0.2633 | 395 | | | | | | parentheses_three_ops | 0.2558 | 258 | | | | | | addition | 0.2323 | 538 | | | | | | mixed_three_ops | 0.2314 | 242 | | | | | | subtraction | 0.2009 | 438 | | | | | It actually learned multiplicative patterns. Not uniform guessing. --- ## What this is for / what it isn't Short chat, simple how-tos, definitions, drafting. Fine-tune it or run it on-device. I wanted something that stays coherent and follows instructions without costing anything. It is **86M**. Factual recall and multi-step arithmetic are weak and it will sound sure when it is wrong - check anything that matters. It repeats and drifts; short, bounded replies work better. English-centric. No safety / RLHF. Don't put it in a sensitive setting without your own guardrails. --- ## Training - **Compute:** Modal credits (Small Models, Big Adventures Hackathon). - **Pipeline:** from-scratch SpikeWhale pretrain, ChatML SFT, then an RL-prep stage. Released `rl_prep/final` came out of the 28-candidate bake-off + battle test. - **Objectives:** next-token CE + JEPA + MTP + z-loss. ### Token budget & scaling - **Tokens:** ~20B from-scratch (~28k steps), then ChatML SFT. - **Token/param:** ~233 (20B / 85.7M). About **11-12× Chinchilla's ~20-tokens/param**. Deliberately over-trained small model. Inference is the trade. Fitting Chinchilla's data term to this run's pretrain loss: `L(D) ≈ 2.611 + 77,715 · D^(-0.537)` (nats/token, R² = 0.92) From that: - **Compute-optimal for this 86M ≈ 4.3B** → 20B is **~4.6× past compute-optimal**. - **Diminishing-returns knee ≈ 22.5B** (where +1B buys < 0.005 nats). 20B lands right there. - **Parameter-bound, not data-bound** at 20B: capacity term (~0.82 nats) beats the data term (~0.54). Extra tokens do little. Doubling to 40B is projected ~0.07 nats lower loss (~7% PPL) with basically no downstream gain. Next lever is **more params, not more tokens**. *Caveats: single-size fit (irreducible loss + capacity floor folded into one constant). Cosine-LR decay inflates the fitted exponent, so treat β as an upper bound. Token counts are anchored to ~20B and scale linearly if that figure is off.* > ⚠️ **SFT was rushed.** Small SFT, thrown together for the hackathon > deadline. No real data mix. Weakest part of this release, not the base. > Re-SFT from > [Escarda-86M-Base](https://huggingface.co/Quazim0t0/Escarda-86M-Base) with > a cleaner set would almost certainly look better. Treat this checkpoint as > a rushed proof-of-concept. Use the base if you want to take it further. ## Acknowledgements Modal credits, Small Models, Big Adventures Hackathon. Apache-2.0. If you want to keep going on it, the base is the better starting point. ## Citation If you use this model, please cite: ```bibtex @misc{escarda86m, title = {Escarda-86M: A ~86M-parameter SpikeWhaleLM}, author = {Dean Byrne (Quazim0t0)}, year = {2026}, howpublished = {HuggingFace, \url{https://huggingface.co/Quazim0t0/Escarda-86M}}, note = {Quazim0t0/Escarda-86M} } ``` ## Update: format-blended SFT on the engram-repaired base This revision applies the (behavior-preserving) engram repair, then a short instruction/format SFT on a 60/25/15 blend of HuggingFaceTB/smoltalk, GSM8K-train (with '#### N' reasoning), and MMLU-style ('Answer: ') examples -- so chat fluency improves while the benchmark output-formats are preserved rather than overwritten. Held-out (test-split) before->after: MMLU acc 0.056->0.278, format 0.284->0.950; GSM8K '####' 0.005->0.750 Note: these are fluency + output-format gains. Benchmark *accuracy* remains near the floor for a model this size -- the SFT does not add reasoning ability.