AstralLM

AstralLM is a compact, efficient causal language model built for Hugging Face Transformers. It features a clean decoder-only transformer architecture with Grouped Query Attention (GQA), SwiGLU feed-forward networks, RoPE positional embeddings, and per-head QK-norm β€” delivering strong performance at a small parameter count (~138M).

Model Summary

Property Value
Architecture Decoder-only Transformer
Parameters ~138M
Vocabulary Size 32,768
Hidden Size 640
Layers 24
Attention Heads 8 (Q) / 4 (KV)
Head Dimension 80
FFN Size 1,920
Context Length 3,072 tokens
Positional Enc. RoPE (ΞΈ = 100,000)
Normalization RMSNorm (Ξ΅ = 1e-6)
Activation SwiGLU
Tied Embeddings Yes
dtype float32

Architecture Details

AstralLM is a decoder-only transformer with the following design choices:

Grouped Query Attention (GQA)

The model uses 8 query heads and 4 key/value heads (2:1 ratio), halving the KV cache memory footprint during inference without measurable quality degradation. Each head operates over a head dimension of 80.

QK-Norm

Per-head RMSNorm is applied to both query and key projections before the rotary embeddings. This stabilises attention logit magnitudes and helps training at scale.

Rotary Position Embeddings (RoPE)

Positions are encoded via complex-valued rotary embeddings with ΞΈ = 100,000, which improves length generalisation compared to the default ΞΈ = 10,000. Frequencies are precomputed and cached per device for efficient reuse across decoding steps.

SwiGLU Feed-Forward Network

Each block uses a gated MLP:

output = w_down( SiLU(w_gate(x)) βŠ™ w_up(x) )

with an expansion ratio of 3Γ— (640 β†’ 1,920).

Pre-Norm with RMSNorm

Both the attention and MLP sub-layers use pre-normalization (RMSNorm), which avoids the instability of post-norm and removes the need for a Ξ² bias term.

Embedding Scale

Input embeddings are multiplied by √hidden_size (β‰ˆ 25.3) to keep the residual stream magnitudes well-conditioned from the first layer.

KV Cache

The model uses Hugging Face's DynamicCache during inference. The cache is allocated automatically when use_cache=True (the default in inference mode).


Files

astral-lm/AstralLM/
β”œβ”€β”€ config.json                  # Model configuration
β”œβ”€β”€ configuration_astrallm.py    # AstralLMConfig class
β”œβ”€β”€ modelling_astrallm.py        # AstralLMForCausalLM implementation
β”œβ”€β”€ generation_config.json       # Default generation parameters
β”œβ”€β”€ tokenizer.json               # Tokenizer vocabulary & rules (fast tokenizer)
β”œβ”€β”€ tokenizer_config.json        # Tokenizer metadata
└── special_tokens_map.json      # Special token definitions

Quick Start

Installation

pip install transformers torch

Loading the Model

from transformers import AutoTokenizer, AutoModelForCausalLM

model_id = "astral-lm/AstralLM"

tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(model_id, trust_remote_code=True)
model.eval()

Note: trust_remote_code=True is required because the model registers a custom model_type (astrallm) via auto_map. The configuration and modelling code are shipped alongside the weights.


Proper Generation

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

model_id = "pihu21057w/astrallm"

tokenizer = AutoTokenizer.from_pretrained(
    model_id,
    trust_remote_code=True,
)

device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = (
    torch.bfloat16
    if torch.cuda.is_available() and torch.cuda.is_bf16_supported()
    else torch.float32
)

model = AutoModelForCausalLM.from_pretrained(
    model_id,
    trust_remote_code=True,
    dtype=dtype,
).to(device).eval()

prompt = "Narendra Modi is"
inputs = tokenizer(prompt, return_tensors="pt").to(device)

with torch.no_grad():
    output = model.generate(
        **inputs,
        max_new_tokens=80,
        do_sample=True,
        temperature=0.7,
        top_p=0.9,
        repetition_penalty=1.1,
        pad_token_id=tokenizer.eos_token_id,
        eos_token_id=tokenizer.eos_token_id,
        use_cache=True,
    )

print(tokenizer.decode(output[0], skip_special_tokens=True))

Text Generation (Greedy)

import torch

prompt = "The universe is vast and"
inputs = tokenizer(prompt, return_tensors="pt")

with torch.no_grad():
    output_ids = model.generate(
        **inputs,
        max_new_tokens=128,
        do_sample=False,
    )

print(tokenizer.decode(output_ids[0], skip_special_tokens=True))

Text Generation (Sampling)

with torch.no_grad():
    output_ids = model.generate(
        **inputs,
        max_new_tokens=256,
        do_sample=True,
        temperature=0.8,
        top_p=0.95,
        top_k=50,
        repetition_penalty=1.1,
    )

print(tokenizer.decode(output_ids[0], skip_special_tokens=True))

Streaming Generation

from transformers import TextStreamer

streamer = TextStreamer(tokenizer, skip_special_tokens=True)

with torch.no_grad():
    model.generate(
        **inputs,
        max_new_tokens=256,
        do_sample=True,
        temperature=0.8,
        streamer=streamer,
    )

Inference Tips

Reduced Precision

Running in bfloat16 or float16 cuts memory roughly in half with negligible quality loss:

model = AutoModelForCausalLM.from_pretrained(
    model_id,
    torch_dtype=torch.bfloat16,
    trust_remote_code=True,
).cuda()

Device Placement

# Single GPU
model = model.to("cuda")

# CPU-only
model = model.to("cpu")

Batch Inference

tokenizer.padding_side = "left"  # required for decoder-only batch generation

prompts = ["Tell me about stars.", "What is quantum computing?"]
inputs = tokenizer(prompts, return_tensors="pt", padding=True).to(model.device)

with torch.no_grad():
    outputs = model.generate(**inputs, max_new_tokens=128, do_sample=False)

for out in outputs:
    print(tokenizer.decode(out, skip_special_tokens=True))
    print("---")

Fine-Tuning

Full Fine-Tune

from transformers import AutoTokenizer, AutoModelForCausalLM, TrainingArguments, Trainer
from datasets import load_dataset

model_id = "astral-lm/AstralLM"
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(model_id, trust_remote_code=True)

dataset = load_dataset("your-dataset-here", split="train")

def tokenize(example):
    return tokenizer(
        example["text"],
        truncation=True,
        max_length=1024,
        padding="max_length",
    )

tokenized = dataset.map(tokenize, batched=True, remove_columns=dataset.column_names)

args = TrainingArguments(
    output_dir="./astrallm-finetuned",
    per_device_train_batch_size=4,
    gradient_accumulation_steps=8,
    num_train_epochs=3,
    learning_rate=2e-5,
    lr_scheduler_type="cosine",
    warmup_ratio=0.05,
    bf16=True,
    logging_steps=10,
    save_strategy="epoch",
)

trainer = Trainer(model=model, args=args, train_dataset=tokenized)
trainer.train()

Parameter-Efficient Fine-Tuning (LoRA)

from peft import get_peft_model, LoraConfig, TaskType

lora_config = LoraConfig(
    task_type=TaskType.CAUSAL_LM,
    r=16,
    lora_alpha=32,
    target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
    lora_dropout=0.05,
    bias="none",
)

model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# trainable params: ~1.6M || all params: ~138M || trainable%: ~1.16%

Configuration Reference

All configuration parameters are exposed through AstralLMConfig and can be overridden at load time:

from configuration_astrallm import AstralLMConfig

config = AstralLMConfig(
    vocab_size=32768,
    hidden_size=640,
    num_hidden_layers=24,
    num_attention_heads=8,
    num_key_value_heads=4,
    head_dim=80,
    intermediate_size=1920,
    max_position_embeddings=3072,
    rope_theta=100000.0,
    rms_norm_eps=1e-6,
    tie_word_embeddings=True,
    use_cache=True,
    z_loss_coeff=0.0,
)
Parameter Default Description
vocab_size 32768 Vocabulary size
hidden_size 640 Residual stream / embedding dimension
num_hidden_layers 24 Number of transformer blocks
num_attention_heads 8 Number of query heads
num_key_value_heads 4 Number of KV heads (GQA)
head_dim 80 Per-head dimension
intermediate_size 1920 FFN hidden dimension (SwiGLU gate + up projection width)
max_position_embeddings 3072 Maximum sequence length
rope_theta 100000.0 RoPE base frequency
rms_norm_eps 1e-6 Epsilon for RMSNorm numerical stability
tie_word_embeddings True Share input embedding and output projection weights
use_cache True Enable KV cache during generation
z_loss_coeff 0.0 Auxiliary z-loss coefficient (set > 0 to penalise large logits)
bos_token_id 1 Begin-of-sequence token ID
eos_token_id 2 End-of-sequence token ID
pad_token_id 0 Padding token ID
unk_token_id 3 Unknown token ID

Tokenizer

AstralLM uses a PreTrainedTokenizerFast with a vocabulary of 32,768 tokens and a maximum sequence length of 3,072 tokens.

Special Token Value ID
<|pad|> Padding 0
<|bos|> BOS 1
<|eos|> EOS 2
<|unk|> Unknown 3
# Inspect special tokens
print(tokenizer.all_special_tokens)
# ['<|pad|>', '<|bos|>', '<|eos|>', '<|unk|>']

Model Architecture Diagram

Input IDs
    β”‚
    β–Ό
[Token Embedding Γ— √d_model]
    β”‚
    β–Ό  ╔══════════════════════════════╗
    β”‚  β•‘  AstralLMBlock Γ— 24          β•‘
    β”‚  β•‘                              β•‘
    β”‚  β•‘  β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”   β•‘
    β”‚  β•‘  β”‚  Pre-RMSNorm         β”‚   β•‘
    β”‚  β•‘  β”‚  AstralLMAttention   β”‚   β•‘
    β”‚  β•‘  β”‚  β”œβ”€ QK-Norm (q, k)   β”‚   β•‘
    β”‚  β•‘  β”‚  β”œβ”€ RoPE             β”‚   β•‘
    β”‚  β•‘  β”‚  β”œβ”€ GQA (8Q / 4KV)   β”‚   β•‘
    β”‚  β•‘  β”‚  └─ SDPA             β”‚   β•‘
    β”‚  β•‘  β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”¬β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜   β•‘
    β”‚  β•‘      + Residual             β•‘
    β”‚  β•‘  β”Œβ”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”   β•‘
    β”‚  β•‘  β”‚  Pre-RMSNorm         β”‚   β•‘
    β”‚  β•‘  β”‚  AstralLMSwiGLUMLP   β”‚   β•‘
    β”‚  β•‘  β”‚  SiLU(gate) βŠ™ up     β”‚   β•‘
    β”‚  β•‘  β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”¬β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜   β•‘
    β”‚  β•‘      + Residual             β•‘
    β”‚  β•šβ•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•β•
    β”‚
    β–Ό
[Final RMSNorm]
    β”‚
    β–Ό
[LM Head (tied to embedding)]
    β”‚
    β–Ό
 Logits

Custom Code Registration

This model ships its own configuration_astrallm.py and modelling_astrallm.py. Hugging Face resolves these via auto_map in config.json:

"auto_map": {
  "AutoConfig": "configuration_astrallm.AstralLMConfig",
  "AutoModelForCausalLM": "modelling_astrallm.AstralLMForCausalLM"
}

Always pass trust_remote_code=True when loading this model. The code is self-contained and has no external dependencies beyond torch and transformers.


Requirements

Package Minimum Version
torch 2.0.0
transformers 4.40.0
peft 0.10.0 (optional, for LoRA)
datasets 2.0.0 (optional, for fine-tuning)

License

This model is released under the Apache 2.0 license. See LICENSE for details.

Downloads last month
322
Safetensors
Model size
0.2B params
Tensor type
F32
Β·
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support