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.