Build hybrid Mamba-Transformer models combining efficient Mamba-2 layers with standard attention to achieve 6x higher inference throughput while maintaining reasoning accuracy on long-context tasks.
Standardmäßig ist der Prompt ausgewählt, der zuerst die Quelle prüft. Sie können zu einem direkten Befehl wechseln oder eine lokale Kopie herunterladen.
Quelldateien prüfen
Lesen Sie SKILL.md und alle von SkillsMP angezeigten Begleitdateien, bevor Sie sich für eine Installation entscheiden.
Mit Codex oder Claude installieren Kopieren Sie diesen Prompt, fügen Sie ihn in Codex, Claude oder einen anderen Assistant ein und lassen Sie die Skill-Seite prüfen und installieren.
Ein direkter Befehl überspringt den Prüf-Prompt. Prüfen Sie die Quelle, bevor Sie ihn ausführen.
Build hybrid Mamba-Transformer models combining efficient Mamba-2 layers with standard attention to achieve 6x higher inference throughput while maintaining reasoning accuracy on long-context tasks.
Nemotron Nano 2: Hybrid Mamba-Transformer for Efficient Reasoning
Core Concept
Nemotron-Nano-9B-v2 achieves state-of-the-art accuracy with significantly improved inference throughput by replacing most self-attention layers in Transformers with Mamba-2 layers. This hybrid approach maintains the reasoning capabilities of full Transformers while gaining the linear-time complexity benefits of Mamba. The architecture processes extended reasoning traces efficiently, achieving up to 6x higher throughput on reasoning workloads while supporting 128k token context on a single GPU.
Architecture Overview
Hybrid Layer Composition: Mix of Mamba-2 and Transformer self-attention layers
Strategic Attention Placement: Preserve full attention for critical reasoning steps
Model Compression: Pre-train on 20 trillion tokens then compress to target size
FP8 Training: Optimize memory efficiency during pre-training and inference
Long-Context Support: Enable 128k token processing within single GPU memory constraints
Implementation Steps
1. Design Hybrid Layer Mixing Strategy
Determine which layers use Mamba vs. Transformer attention:
defcreate_hybrid_architecture(
total_layers: int,
mamba_ratio: float = 0.7,
attention_positions: list[int] = None) -> list[str]:
"""
Design layer composition with Mamba and Transformer layers.
Default: ~70% Mamba for efficiency, ~30% attention for critical reasoning
Attention layers placed strategically at early, middle, and late stages
"""if attention_positions isNone:
# Default: preserve attention at key positions for information bottlenecks
attention_positions = [0, total_layers // 2, total_layers - 1]
layer_types = []
for i inrange(total_layers):
i attention_positions:
layer_types.append()
i < (total_layers * mamba_ratio):
layer_types.append()
:
layer_types.append()
layer_types
if
in
"attention"
elif
int
"mamba"
else
"attention"
return
2. Implement Mamba-2 Layers
Use selective state space models for efficient sequence processing:
defcreate_mamba2_layer(
hidden_size: int,
state_size: int = 16,
expand_factor: int = 2) -> "Mamba2Layer":
"""
Create a Mamba-2 layer combining selective state space model with gating.
"""classMamba2Layer:
def__init__(self):
self.input_projection = Linear(hidden_size, hidden_size * expand_factor)
self.state_matrix = StateSpaceMatrix(hidden_size * expand_factor, state_size)
self.output_projection = Linear(hidden_size * expand_factor, hidden_size)
self.gate = Linear(hidden_size * expand_factor, hidden_size * expand_factor)
defforward(self, x):
# Project input
x_expanded = self.input_projection(x)
# Apply selective SSM
y = self.state_matrix(x_expanded)
# Apply gating mechanism
gated = y * sigmoid(self.gate(x_expanded))
# Project back to hidden size
output = self.output_projection(gated)
return output
return Mamba2Layer()
3. Configure Attention Layers
Maintain standard Transformer attention at strategic positions:
defcreate_attention_layer(
hidden_size: int,
num_heads: int = 8,
max_seq_length: int = 128000) -> "AttentionLayer":
"""
Standard multi-head attention, optimized for long sequences.
"""classAttentionLayer:
def__init__(self):
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.q_proj = Linear(hidden_size, hidden_size)
self.k_proj = Linear(hidden_size, hidden_size)
self.v_proj = Linear(hidden_size, hidden_size)
self.out_proj = Linear(hidden_size, hidden_size)
defforward(self, x, attention_mask=None):
batch_size, seq_len, _ = x.shape
# Project to Q, K, V
q = self.q_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim)
k = self.k_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim)
v = self.v_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim)
# Compute attention with optional long-context optimization
attention_scores = torch.matmul(q, k.transpose(-2, -1)) / sqrt(self.head_dim)
if attention_mask isnotNone:
attention_scores = attention_scores + attention_mask
attention_weights = softmax(attention_scores, dim=-1)
output = torch.matmul(attention_weights, v)
# Reshape and project
output = output.view(batch_size, seq_len, -1)
returnself.out_proj(output)
return AttentionLayer()
4. Pre-training and Compression
Implement the Minitron compression strategy:
defpretrain_and_compress(
base_config: dict,
target_size: str = "9B",
pretrain_tokens: int = 20_000_000_000,
fp8_enabled: bool = True) -> "CompressedModel":
"""
Pre-train large model then compress to target size using Minitron strategy.
"""# Step 1: Pre-train larger base modelprint(f"Pre-training on {pretrain_tokens:,} tokens...")
base_model = create_hybrid_model(
hidden_size=base_config["hidden_size"],
num_layers=base_config["num_layers"],
num_heads=base_config["num_heads"]
)
if fp8_enabled:
base_model = convert_to_fp8(base_model)
# Train on diverse data corpus
base_model = train_model(base_model, dataset, pretrain_tokens)
# Step 2: Apply Minitron compressionprint(f"Compressing to {target_size}...")
target_hidden_size = map_size_to_hidden_dim(target_size)
target_num_layers = map_size_to_layers(target_size)
compressed = compress_model(
base_model,
target_hidden_size=target_hidden_size,
target_num_layers=target_num_layers,
method="layer_dropping_and_projection"
)
return compressed
5. Configure Long-Context Inference
Enable efficient processing of extended sequences:
defconfigure_long_context_inference(
model: "HybridModel",
max_tokens: int = 128000,
device: str = "cuda") -> dict:
"""
Set up inference for long-context processing on single GPU.
"""
config = {
"max_sequence_length": max_tokens,
"use_kv_cache": True,
"kv_cache_dtype": torch.float8_e4m3fn, # FP8 for memory efficiency"attention_implementation": "flash_attention_2", # Optimized kernels"mamba_scan_mode": "hardware_optimized",
"device": device,
"gradient_checkpointing": False# Only for training
}
# Apply configuration
model.config.update(config)
model = model.to(device)
if fp8_enabled:
model = quantize_model(model, "fp8")
return config