Implementation:FlagOpen FlagEmbedding MiniCPM Reranker Config
| Knowledge Sources | |
|---|---|
| Domains | Model_Configuration, Layer_Wise_Training, MiniCPM |
| Last Updated | 2026-02-09 00:00 GMT |
Overview
Configuration class for LayerWiseMiniCPM models enabling layer-wise reranking with multiple prediction heads.
Description
LayerWiseMiniCPMConfig extends the standard MiniCPM configuration with parameters for layer-wise reranking:
Core MiniCPM parameters:
- Standard transformer configuration (vocab_size, hidden_size, num_layers, attention_heads)
- MiniCPM-specific features: scale_emb, dim_model_base, scale_depth for efficient scaling
- RoPE (Rotary Position Embeddings) with optional scaling strategies (linear, dynamic)
- Flash Attention 2 support detection
Layer-wise reranking additions:
- start_layer: First layer to attach a reranking head (default: 8)
- head_multi: Whether to use multiple heads across layers (default: True)
- head_type: Type of reranking head - "simple" for basic classification (default: "simple")
This enables training models that make predictions at multiple intermediate layers, allowing:
- Earlier exit for efficiency (don't need full forward pass)
- Layer-wise knowledge distillation from deeper to shallower layers
- Multi-granularity relevance judgments
The configuration validates RoPE scaling parameters and automatically enables Flash Attention 2 if available.
Usage
Use this configuration when training MiniCPM models for reranking with layer-wise prediction heads for improved efficiency and distillation.
Code Reference
Source Location
- Repository: FlagOpen_FlagEmbedding
- File: research/llm_reranker/finetune_for_layerwise/configuration_minicpm_reranker.py
- Lines: 1-208
Signature
class LayerWiseMiniCPMConfig(PretrainedConfig):
def __init__(self, vocab_size=32000, hidden_size=4096,
intermediate_size=11008, num_hidden_layers=32,
num_attention_heads=32, num_key_value_heads=None,
hidden_act="silu", max_position_embeddings=2048,
initializer_range=0.02, rms_norm_eps=1e-6,
use_cache=True, pad_token_id=None, bos_token_id=1,
eos_token_id=2, pretraining_tp=1, tie_word_embeddings=True,
rope_theta=10000.0, rope_scaling=None,
attention_bias=False, attention_dropout=0.0,
scale_emb=1, dim_model_base=1, scale_depth=1,
start_layer=8, head_multi=True, head_type="simple",
**kwargs)
Import
from research.llm_reranker.finetune_for_layerwise.configuration_minicpm_reranker import LayerWiseMiniCPMConfig
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
| start_layer | int | No | First layer for reranking head (default: 8) |
| head_multi | bool | No | Use multiple heads across layers (default: True) |
| head_type | str | No | Head architecture type (default: "simple") |
| num_hidden_layers | int | No | Total transformer layers (default: 32) |
| hidden_size | int | No | Hidden dimension (default: 4096) |
| rope_scaling | Dict | No | RoPE scaling configuration {"type": "linear"/"dynamic", "factor": float} |
Outputs
| Name | Type | Description |
|---|---|---|
| config | LayerWiseMiniCPMConfig | Configuration object with validated parameters |
Usage Examples
from research.llm_reranker.finetune_for_layerwise.configuration_minicpm_reranker import LayerWiseMiniCPMConfig
# Create configuration for layer-wise reranker
config = LayerWiseMiniCPMConfig(
vocab_size=32000,
hidden_size=2048,
num_hidden_layers=24,
num_attention_heads=32,
# Layer-wise reranking settings
start_layer=8, # Start making predictions from layer 8
head_multi=True, # Use heads at layers 8, 9, ..., 24
head_type="simple", # Simple classification head
# MiniCPM-specific
scale_emb=1,
dim_model_base=256,
scale_depth=1.4,
# RoPE scaling for longer contexts
rope_scaling={"type": "linear", "factor": 2.0},
max_position_embeddings=4096
)
# Save configuration
config.save_pretrained("./minicpm_reranker_config")
# Load configuration
config = LayerWiseMiniCPMConfig.from_pretrained("./minicpm_reranker_config")
# Access layer-wise settings
print(f"Reranking heads from layer {config.start_layer} to {config.num_hidden_layers}")
print(f"Multiple heads: {config.head_multi}")
print(f"Flash Attention 2: {config._attn_implementation}")