TrinityVLM-Nano / configuration_trinity_vlm.py
NyxKrage's picture
Upload folder using huggingface_hub
19f7733 verified
Raw
History Blame Contribute Delete
7.97 kB
from __future__ import annotations
from transformers.configuration_utils import PretrainedConfig
from transformers.utils import logging
logger = logging.get_logger(__name__)
class AfmoeConfig(PretrainedConfig):
"""
n_group (`int`, *optional*, defaults to 1):
Number of groups for routed experts.
topk_group (`int`, *optional*, defaults to 1):
Number of selected groups for each token(for each token, ensuring the selected experts is only within `topk_group` groups).
"""
model_type = "afmoe"
base_model_pp_plan = {
"embed_tokens": (["input_ids"], ["inputs_embeds"]),
"layers": (["hidden_states", "attention_mask"], ["hidden_states"]),
"norm": (["hidden_states"], ["hidden_states"]),
}
def __init__(
self,
num_hidden_layers: int = 32,
vocab_size: int = 200192,
hidden_size: int = 2048,
intermediate_size: int = 6144,
moe_intermediate_size=1408,
num_dense_layers=1,
num_attention_heads=16,
num_key_value_heads=None,
head_dim=128,
hidden_act="silu",
max_position_embeddings=16384,
initializer_range=0.02,
rms_norm_eps=1e-5,
use_cache=True,
tie_word_embeddings=False,
rope_theta=10000.0,
rope_scaling=None,
num_experts=64,
num_experts_per_tok=6,
num_shared_experts=2,
num_expert_groups=1,
num_limited_groups=1,
score_func="sigmoid",
route_norm=True,
route_scale=1.0,
global_attn_every_n_layers=4,
sliding_window=1024,
mup_enabled=False,
layer_types=None,
attention_dropout: float = 0.0,
n_group: int = 1,
topk_group: int = 1,
**kwargs,
):
self.vocab_size = vocab_size
self.max_position_embeddings = max_position_embeddings
self.hidden_size = hidden_size
self.intermediate_size = intermediate_size
self.num_hidden_layers = num_hidden_layers
self.num_dense_layers = num_dense_layers
self.num_attention_heads = num_attention_heads
self.head_dim = head_dim
self.hidden_act = hidden_act
self.initializer_range = initializer_range
self.rms_norm_eps = rms_norm_eps
self.use_cache = use_cache
self.rope_theta = rope_theta
self.rope_scaling = rope_scaling
self.rope_parameters = dict(rope_scaling) if rope_scaling is not None else None
# MoE specific
self.moe_intermediate_size = moe_intermediate_size
self.num_experts_per_tok = num_experts_per_tok
self.n_group = n_group
self.topk_group = topk_group
self.num_experts = num_experts
self.num_shared_experts = num_shared_experts
self.num_expert_groups = num_expert_groups
self.num_limited_groups = num_limited_groups
self.score_func = score_func
self.route_norm = route_norm
self.route_scale = route_scale
# Attention specific
self.attention_dropout = attention_dropout
self.global_attn_every_n_layers = global_attn_every_n_layers
self.sliding_window = sliding_window
self.layer_types = layer_types
if self.layer_types is None:
self.layer_types = [
"sliding_attention" if bool((i + 1) % global_attn_every_n_layers) else "full_attention" for i in range(self.num_hidden_layers)
]
self.validate_layer_type()
# muP specific
self.mup_enabled = mup_enabled
if num_key_value_heads is None:
num_key_value_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads
# Validate rope configs
if self.rope_scaling is not None and "type" in self.rope_scaling:
self.rope_scaling["rope_type"] = self.rope_scaling["type"]
if self.rope_parameters is not None and "type" in self.rope_parameters:
self.rope_parameters["rope_type"] = self.rope_parameters["type"]
self.standardize_rope_params()
self.validate_rope()
super().__init__(
tie_word_embeddings=tie_word_embeddings,
**kwargs,
)
class TrinityVLMConfig(PretrainedConfig):
model_type = "trinity_vlm"
is_composition = True
def __init__(
self,
text_config: dict | None = None,
vision_config: dict | None = None,
projector_hidden_dim: int | None = None,
vision_feature_dim: int | None = None,
image_seq_len: int | None = None,
enable_grouped_moe: bool = True,
output_router_logits: bool = False,
image_start_token: str = "<|vision_start|>",
image_end_token: str = "<|vision_end|>",
image_token: str = "<|image_pad|>",
image_start_token_id: int | None = None,
image_end_token_id: int | None = None,
image_token_id: int | None = None,
hidden_size: int | None = None,
vocab_size: int | None = None,
**kwargs,
) -> None:
text_config = dict(text_config or {})
vision_config = dict(
vision_config
or {
"enc_dim": 1152,
"enc_patch_size": 14,
"enc_n_layers": 27,
"enc_ff_dim": 4304,
"enc_n_heads": 16,
"proj_out_dim": 2048,
"crop_size": 378,
"in_channels": 3,
"max_crops": 12,
"overlap_margin": 4,
"proj_inner_dim": 8192,
"projector_hidden_dim": 2048,
}
)
if projector_hidden_dim is not None:
vision_config["projector_hidden_dim"] = projector_hidden_dim
if hidden_size is None:
hidden_size = text_config.get("hidden_size")
if vocab_size is None:
vocab_size = text_config.get("vocab_size")
kwargs.setdefault("bos_token_id", text_config.get("bos_token_id"))
kwargs.setdefault("eos_token_id", text_config.get("eos_token_id"))
kwargs.setdefault("pad_token_id", text_config.get("pad_token_id"))
if vision_feature_dim is None:
vision_feature_dim = vision_config.get("proj_out_dim", 2048)
if image_seq_len is None:
crop_size = int(vision_config.get("crop_size", 378))
patch_size = int(vision_config.get("enc_patch_size", 14))
image_seq_len = (crop_size // patch_size) ** 2
self.text_config = text_config
self.vision_config = vision_config
self.vision_feature_dim = vision_feature_dim
self.image_seq_len = image_seq_len
self.enable_grouped_moe = enable_grouped_moe
self.output_router_logits = output_router_logits
self.image_start_token = image_start_token
self.image_end_token = image_end_token
self.image_token = image_token
self.image_start_token_id = image_start_token_id
self.image_end_token_id = image_end_token_id
self.image_token_id = image_token_id
self.hidden_size = hidden_size
self.vocab_size = vocab_size
super().__init__(**kwargs)
@property
def projector_hidden_dim(self) -> int:
return int(self.vision_config.get("projector_hidden_dim", self.vision_feature_dim))
def get_text_config(self, decoder: bool = False):
del decoder
text_config = AfmoeConfig(**self.text_config)
text_config.vocab_size = self.vocab_size
text_config.bos_token_id = self.bos_token_id
text_config.eos_token_id = self.eos_token_id
text_config.pad_token_id = self.pad_token_id
text_config.packed_experts = True
text_config.enable_grouped_moe = self.enable_grouped_moe
text_config.output_router_logits = self.output_router_logits
return text_config
__all__ = ["AfmoeConfig", "TrinityVLMConfig"]