Key Configuration Parameters for Needle 2's Architecture: A Complete Guide
The TransformerConfig dataclass in needle/model/architecture.py defines 20+ parameters that control model size, attention mechanics, engram memory, and quantization behavior.
Needle 2's architecture is fully configurable through a centralized configuration system. The TransformerConfig dataclass located in needle/model/architecture.py (lines 38–82) exposes every hyperparameter that determines transformer depth, attention patterns, memory hashing schemes, and numerical precision. Understanding these parameters is essential for customizing Needle 2 for specific hardware constraints or task requirements.
Core Model Size and Embedding Parameters
The foundational dimensions of Needle 2 are controlled by parameters that define the embedding space and sequence capacity.
Vocabulary and Hidden Dimensions
vocab_size(line 60): Sets the total number of tokens the embedding layer can represent.d_model(line 61): Determines the base dimensionality of hidden states throughout the network. This parameter scales linearly with memory usage and compute requirements.num_layers(line 66): Defines the total transformer blocks stacked in the architecture. The default preset uses 27 layers.
Sequence Handling
max_seq_len(line 67): Hard limit on input sequence length before truncation or padding is required.pad_token_id(line 68): Token ID reserved for padding sequences tomax_seq_len, which influences the attention padding mask generation.
Attention Mechanism and Memory Configuration
Needle 2 supports grouped query attention (GQA) and novel engram-based memory systems, configured through specialized parameters.
Multi-Head and Grouped Query Attention
num_heads(line 63): Number of parallel attention heads for standard self-attention.num_kv_heads(line 64): Number of key/value heads when using grouped query attention. Settingnum_kv_heads < num_headsreduces KV-cache memory during inference while maintaining query parallelism.
Engram Memory System
Needle 2 introduces engram memory blocks that use n-gram hashing for extended context modeling:
engram_orders(line 72): Tuple defining n-gram orders (e.g.,(2, 3)) used in the hashing scheme.engram_slots(line 74): Capacity of the engram hash table—the number of distinct memory slots available.engram_layers(line 75): Specific transformer layers that host engram memory blocks (e.g.,(2, 15)places engrams after layers 2 and 15).engram_heads(line 73): Number of attention heads allocated to engram processing; defaults to a heuristic based ond_model.
Positional Embeddings and Optimization
rope_theta(line 69): Base frequency for Rotary Positional Embeddings (RoPE), controlling the rotational angle encoding.flash(line 71): Boolean toggle for FlashAttention implementation, enabling memory-efficient attention on compatible GPUs.mhc_lanes(line 76): Number of multi-head communication lanes in the MHC block, facilitating cross-head information exchange.contrastive_dim(line 68): Dimensionality of the contrastive projection head for representation learning tasks.
Quantization and Performance Tuning
Needle 2 provides granular control over numerical precision and memory optimization through quantization and checkpointing parameters.
Bit-Width Quantization Controls
Located sequentially at lines 77–80, these parameters manage compression:
kv_window(line 77): Sliding window size for KV-budgeting; set to0to let the budget planner auto-configure.kv_bits(line 78): Bit-width for key/value quantization (e.g.,8for 8-bit KV cache).act_bits(line 79): Precision for activation quantization.weight_bits(line 80): Precision for weight quantization (supports string values like"4"for 4-bit weights).
Training and Compilation Optimization
dtype(line 70): Global data type for model parameters:bfloat16,float32, orfloat16.remat(line 81): Enables gradient checkpointing (rematerialization) to trade compute for memory during training.scan_unroll(line 82): Unroll factor for JAXscanloops that iterate over transformer layers, affecting compilation and execution speed.
Configuration Presets and Model Instantiation
The needle/model/architecture.py module includes predefined parameter bundles accessed via the PRESETS dictionary (lines 38–44).
Two official presets are available:
PRESETS = {
"needle": dict(d_model=768, num_heads=12, num_kv_heads=6, num_layers=27,
engram_layers=(2, 15)),
"base": dict(d_model=512, num_heads=8, num_kv_heads=4, num_layers=27,
engram_layers=(2, 15))
}
Instantiating a Custom Configuration
Create and modify configurations programmatically:
from needle.model.architecture import TransformerConfig, PRESETS
# Method 1: Custom configuration from scratch
cfg = TransformerConfig(**{
"d_model": 1024,
"num_heads": 16,
"num_kv_heads": 8,
"num_layers": 30,
"max_seq_len": 4096,
"dtype": "bfloat16",
"flash": True,
"engram_orders": (2, 3, 4),
"kv_window": 0,
"kv_bits": 8,
"act_bits": 8,
"weight_bits": "4",
})
# Method 2: Load preset and modify specific fields
cfg = TransformerConfig(**PRESETS["needle"])
cfg.num_layers = 36 # Extend depth while keeping other defaults
Building the Model
Pass the configuration to SimpleAttentionNetwork defined in the same module:
from needle.model.architecture import SimpleAttentionNetwork
import jax.numpy as jnp
model = SimpleAttentionNetwork(config=cfg)
# Forward pass example
tokens = jnp.array([[1, 5, 23, 7, 0, 0]]) # shape: (batch, seq_len)
logits = model(tokens) # → (batch, seq_len, vocab_size)
The configuration propagates through needle/model/quantize.py for applying fake quantization operations and needle/model/export.py for respecting dtype settings during model serialization.
Summary
TransformerConfiginneedle/model/architecture.pyserves as the single source of truth for Needle 2's architecture, defining 20+ parameters spanning lines 38–82.- Model sizing is controlled by
d_model,num_layers,vocab_size, andmax_seq_len. - Attention architecture uses
num_headsandnum_kv_headsfor GQA, withflashenabling optimized GPU kernels. - Engram memory parameters (
engram_orders,engram_slots,engram_layers) configure n-gram hashing-based external memory. - Quantization is managed through
kv_bits,act_bits, andweight_bitsfor inference optimization. - Presets provide battle-tested defaults:
"needle"(768-dim) and"base"(512-dim), expandable via dictionary unpacking.
Frequently Asked Questions
Where is the TransformerConfig class defined in the Needle repository?
The TransformerConfig dataclass is defined in needle/model/architecture.py starting at line 38, alongside the PRESETS dictionary. This file also contains the SimpleAttentionNetwork class that consumes these configurations.
What is the difference between num_heads and num_kv_heads in Needle 2?
num_heads (line 63) defines the total parallel attention heads for queries, while num_kv_heads (line 64) specifies how many key and value heads are shared across those queries. When num_kv_heads < num_heads, Needle 2 implements grouped query attention (GQA), reducing KV-cache memory by sharing keys and values across multiple query heads.
How does the engram memory system work in Needle 2?
The engram system uses n-gram hashing to extend context beyond standard attention windows. Parameters engram_orders (line 72) define which n-gram sizes to hash, engram_slots (line 74) sets the hash table capacity, and engram_layers (line 75) specifies which transformer layers insert these memory blocks. This allows the model to retrieve information from distant context via hash-based lookup rather than full attention.
What quantization options does Needle 2 support?
Needle 2 supports post-training quantization through kv_bits (line 78) for cache compression, act_bits (line 79) for activation quantization, and weight_bits (line 80) for weight compression down to 4-bit precision. The dtype parameter (line 70) sets the base compute precision (e.g., bfloat16), while flash (line 71) enables optimized kernels that maintain speed under quantization.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →