SupraBrain-50M / README.md
LH-Tech-AI's picture
Update README.md
92872c3 verified
|
Raw
History Blame Contribute Delete
7.23 kB
metadata
license: apache-2.0
language:
  - en
pipeline_tag: text-generation
tags:
  - suprabrain
  - gated-deltanet
  - linear-attention
  - sliding-window-attention
  - custom-architecture
library_name: transformers

SupraBrain 50M v0.1

SupraBrain 50M v0.1 is an experimental 50-million parameter hybrid language model engineered by SupraLabs. It combines Gated DeltaNet linear recurrence with Sliding-Window Attention and Surprise-Gated update mechanisms, optimized using a custom Muon + AdamW hybrid optimizer schedule.

Model Summary

  • Developer: SupraLabs
  • Architecture: Hybrid Gated DeltaNet (3:1) + Sliding-Window Attention + Surprise Gating
  • Parameters: ~50M (Sub-50M budget constraint)
  • Vocabulary Size: 23,808 (GEMM-friendly: 186*128, Byte-Level BPE with Digit-Splitting)
  • Context Length: 1,024 tokens (Supports sliding-window attention)
  • Primary Training Data: FineWeb-Edu & Cosmopedia-v2 (5B tokens total)
  • License: Apache 2.0

Key Architectural Innovations

  1. Hybrid Layer Layout (3:1 Ratio):
    • Gated DeltaNet (GDN): 3 out of every 4 layers use Gated DeltaNet linear state-space recurrence for linear-time complexity and fast sequence processing.
    • Sliding-Window Attention (SWA): Every 4th layer incorporates localized attention (Window size = 256) with QK-Normalization to maintain strong long-range associative recall. Layer 19 features full global attention.
  2. Surprise-Gated Updates ($\beta_t$):
    • Implements a scale-invariant residual prediction mechanism (SurpriseBeta) that dynamically scales learning updates based on local sequence surprise/prediction error.
  3. Digit-Split Tokenizer:
    • Custom Byte-Level BPE tokenizer trained on FineWeb-Edu. Enforces single-digit splitting (individual_digits=True) to dramatically boost arithmetic and numerical reasoning performance in sub-100M parameter models.
  4. Half-Untied Head & Low-Rank Gates:
    • Utilizes an unembedding rank adapter (unembed_rank=32) and low-rank output gating (gdn_gate_rank=32) to conserve parameter count while maintaining model capacity in the core layers.
  5. Custom Muon + AdamW Hybrid Optimizer:
    • 2D weight matrices in the body are optimized using the Muon optimizer (Newton-Schulz orthogonalization updates), while embeddings, norms, and 1D vectors are updated via AdamW over a WSD (Warmup-Stable-Decay) schedule.

Benchmarks

Tasks Version Filter n-shot Metric Value Stderr
arc_challenge 1 none 0 acc 0.2065 ± 0.0118
none 0 acc_norm 0.2329 ± 0.0124
arc_easy 1 none 0 acc 0.4882 ± 0.0103
none 0 acc_norm 0.4255 ± 0.0101
boolq 2 none 0 acc 0.4223 ± 0.0086
hellaswag 1 none 0 acc 0.2914 ± 0.0045
none 0 acc_norm 0.3160 ± 0.0046
lambada_openai 1 none 0 acc 0.3072 ± 0.0064
none 0 perplexity 66.3967 ± 2.8867
openbookqa 1 none 0 acc 0.1920 ± 0.0176
none 0 acc_norm 0.3160 ± 0.0208
piqa 1 none 0 acc 0.6295 ± 0.0113
none 0 acc_norm 0.6175 ± 0.0113
sciq 1 none 0 acc 0.7020 ± 0.0145
none 0 acc_norm 0.5990 ± 0.0155
wikitext 2 none 0 bits_per_byte 1.0489 ± N/A
none 0 byte_perplexity 2.0689 ± N/A
none 0 word_perplexity 48.8008 ± N/A
winogrande 1 none 0 acc 0.4878 ± 0.0140

Model Configuration

Hyperparameter Value
hidden_size 384
num_hidden_layers 28
intermediate_size 1152 (2-Matrix MLP with RationalAct)
gdn_num_heads / dim 3 heads / 128 dim
attn_num_heads / kv_heads 6 query heads / 2 KV heads (GQA)
attn_window 256
max_position_embeddings 1024
mlp_act Per-channel learnable Rational Activation

Usage

Since SupraBrain uses a custom architecture without standard Hugging Face native integration, you must register the model class locally before loading it with AutoModelForCausalLM.

Quickstart (Inference Script)

First, download the modeling script:

wget https://huggingface.co/SupraLabs/SupraBrain-50M/resolve/main/modeling_suprabrain.py

Then, load the model:

import importlib
import torch
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer

# Import custom model and config classes from the model script (modeling_suprabrain.py)
sb_module = importlib.import_module("modeling_suprabrain")
SupraBrainConfig = sb_module.SupraBrainConfig
SupraBrainForCausalLM = sb_module.SupraBrainForCausalLM

# Register custom architecture with Hugging Face AutoClasses
AutoConfig.register("suprabrain", SupraBrainConfig)
AutoModelForCausalLM.register(SupraBrainConfig, SupraBrainForCausalLM)

model_id = "SupraLabs/SupraBrain-50M"

print("[*] Loading tokenizer and model...")
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
    model_id,
    trust_remote_code=True,
    torch_dtype=torch.bfloat16
).to("cuda")

# Prompt setup
prompt = "The mitochondrion produces"
inputs = tokenizer(prompt, return_tensors="pt").to("cuda")

# Generation with repetition penalty control
print("[*] Generating text...")
with torch.no_grad():
    outputs = model.generate(
        **inputs,
        max_new_tokens=100,
        temperature=0.7,
        top_p=0.9,
        do_sample=True,
        no_repeat_ngram_size=3,  # Prevents 3-gram repetitions
        pad_token_id=tokenizer.pad_token_id,
        eos_token_id=tokenizer.eos_token_id
    )

generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True)
print("\n--- Output ---")
print(generated_text)

Training Details

Training Pipeline & Schedule

  • Dataset: 5 Billion Tokens Total

  • Stable Phase (3.6B Tokens): FineWeb-Edu (sample-100BT)

  • Anneal Phase (1.4B Tokens): 65% FineWeb-Edu (Score $\ge 4.2$) + 35% Cosmopedia-v2

  • Schedule: Warmup-Stable-Decay (WSD) with square-root decay during the annealing phase.

  • Batch Size: Micro-batch size 16 with Gradient Accumulation 8 ($\approx 262,144$ tokens/step over sequence length 1024).

Hardware Requirements & Optimization

  • Dependencies: Optimized with flash-linear-attention (fla) for Gated DeltaNet kernels and PyTorch flex_attention for masked sliding-window operations.
  • FP32 Logit Chunking: Uses memory-checkpointed chunked Cross-Entropy loss to avoid VRAM allocation spikes.