molmo-clip-b16-olmo3 / configuration_molmo_olmo3.py
amitha's picture
Upload folder using huggingface_hub
d2b31cd verified
Raw History Blame Contribute Delete
4.86 kB
# coding=utf-8
"""Configuration for the Molmo-v1 (CLIP vision) VLM in HuggingFace format.
LLaVA-style composition:
- vision_tower : a HuggingFace vision encoder (CLIPVisionModel), referenced by
`vision_tower_name_or_path`; its weights are NOT stored in this
checkpoint (loaded from the referenced repo at load time).
- multi_modal_projector : SwiGLU image projector + a separate CLS Linear projector.
- language_model : a native transformers `Olmo3ForCausalLM` (text_config).
The vision encoder architecture is stored in `vision_config` (metadata only) so the
module can be constructed without network access; the trained vision *weights* live
in the referenced `vision_tower_name_or_path` repo.
"""
from transformers.configuration_utils import PretrainedConfig
from transformers.models.auto import CONFIG_MAPPING, AutoConfig
from transformers.models.olmo3.configuration_olmo3 import Olmo3Config
class MolmoOlmo3Config(PretrainedConfig):
model_type = "molmo_olmo3"
sub_configs = {"text_config": Olmo3Config, "vision_config": AutoConfig}
def __init__(
self,
text_config=None,
vision_config=None,
vision_tower_name_or_path="amitha/clip-vit-b16-datacomp-1b-medium-subset",
vision_trust_remote_code=False,
vision_feature_layer=-1, # Molmo vit_layers=[-1] -> last block output (pre-final-norm)
# DINOv3 applies a final LayerNorm to last_hidden_state that Molmo discards; strip it
# so the tower output is the pre-norm last-block features Molmo's connector consumes.
vision_strip_final_norm=False,
vision_final_norm_attr="norm",
projector_intermediate_size=11008,
projector_hidden_act="silu",
include_cls_token=True,
# lm_head covers the real vocab; ids >= lm_head_vocab_size are image
# placeholder tokens (never generation targets) and are masked to -inf.
lm_head_vocab_size=100352,
# Molmo's get_tokenizer pads to vocab_size=100278 (no padding tokens), so the
# 5 image special tokens land at 100278..100282 and index wte.embedding directly.
image_token_id=100280, # <im_patch>
image_start_token_id=100278, # <im_start>
image_end_token_id=100279, # <im_end>
image_col_token_id=100281, # <im_col>
image_prompt_token_id=100282, # <|image|>
bos_token_id=100257,
eos_token_id=100257,
pad_token_id=None,
tie_word_embeddings=False,
**kwargs,
):
# --- text config (native Olmo3) ---
if text_config is None:
text_config = {}
if isinstance(text_config, dict):
text_config = Olmo3Config(**text_config)
self.text_config = text_config
# --- vision config (architecture metadata for the referenced encoder) ---
if vision_config is None:
# Default to the CLIP ViT-B/16 (224) vision tower used by this VLM family.
vision_config = CONFIG_MAPPING["clip_vision_model"](
hidden_size=768,
intermediate_size=3072,
num_hidden_layers=12,
num_attention_heads=12,
num_channels=3,
image_size=224,
patch_size=16,
hidden_act="quick_gelu",
layer_norm_eps=1e-5,
)
elif isinstance(vision_config, dict):
vision_model_type = vision_config.get("model_type", "clip_vision_model")
vision_config = CONFIG_MAPPING[vision_model_type](**vision_config)
self.vision_config = vision_config
self.vision_tower_name_or_path = vision_tower_name_or_path
self.vision_trust_remote_code = vision_trust_remote_code
self.vision_feature_layer = vision_feature_layer
self.vision_strip_final_norm = vision_strip_final_norm
self.vision_final_norm_attr = vision_final_norm_attr
self.vision_hidden_size = getattr(vision_config, "hidden_size", 768)
self.text_hidden_size = self.text_config.hidden_size
self.projector_intermediate_size = projector_intermediate_size
self.projector_hidden_act = projector_hidden_act
self.include_cls_token = include_cls_token
self.lm_head_vocab_size = lm_head_vocab_size
self.image_token_id = image_token_id
self.image_start_token_id = image_start_token_id
self.image_end_token_id = image_end_token_id
self.image_col_token_id = image_col_token_id
self.image_prompt_token_id = image_prompt_token_id
super().__init__(
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
pad_token_id=pad_token_id,
tie_word_embeddings=tie_word_embeddings,
**kwargs,
)
__all__ = ["MolmoOlmo3Config"]