Argus-Colqwen3.5-9b-v0-bf16 / configuration_argus.py
abdoelsayed's picture
Initial bf16 release (companion to fp32 v0)
ebb3727 verified
Raw History Blame Contribute Delete
2.95 kB
"""Argus: Region-Aware Query-Conditioned Mixture of Experts for Visual Document Retrieval.
Config class. Subclasses the Qwen3.5-VL config and adds the Argus-specific
retrieval + MoE hyperparameters. Used by ``AutoConfig.from_pretrained`` via the
``auto_map`` field in ``config.json`` (requires ``trust_remote_code=True``).
"""
from __future__ import annotations
try:
from transformers.models.qwen3_5 import Qwen3_5Config as _BackboneConfig
except ImportError:
try:
from transformers.models.qwen3_5 import Qwen35Config as _BackboneConfig
except ImportError as exc:
raise ImportError(
"Argus requires a transformers build that exposes the Qwen3.5 VL "
"classes (transformers.models.qwen3_5). Upgrade to transformers "
">= 4.57.0.dev0."
) from exc
class ArgusConfig(_BackboneConfig):
"""Top-level config for Argus-Colqwen3.5-9B.
Holds the standard Qwen3.5-VL fields (text_config, vision_config, image
token ids, etc.) plus Argus-specific retrieval + MoE knobs:
- ``retrieval_dim``: output dimensionality of the multi-vector retrieval
head (``custom_text_proj``). Default: 768.
- ``num_specialists``: number of latent spatial experts in the MoE stack.
- ``top_k_experts``: sparsity of the router (top-k routing).
- ``region_size``: spatial pooling window (patches) for region tokens.
- ``router_layer_index``: hidden-state layer used as input to the router.
- ``router_temperature``: softmax temperature of the router.
- ``mask_non_image_embeddings``: zero out embedding positions that are
not image tokens at encode time (document side).
- ``shared_gate_init`` / ``specialist_gate_init``: logit-space init for
the gate scalars (sigmoid of these multiplies shared/specialist expert
contributions).
"""
model_type = "argus_colqwen35"
def __init__(
self,
retrieval_dim: int = 768,
num_specialists: int = 4,
top_k_experts: int = 2,
region_size: int = 4,
router_layer_index: int = -5,
router_temperature: float = 0.8,
router_noise_std: float = 0.0,
mask_non_image_embeddings: bool = True,
shared_gate_init: float = 0.0,
specialist_gate_init: float = 0.0,
**kwargs,
) -> None:
super().__init__(**kwargs)
self.retrieval_dim = int(retrieval_dim)
self.num_specialists = int(num_specialists)
self.top_k_experts = int(top_k_experts)
self.region_size = int(region_size)
self.router_layer_index = int(router_layer_index)
self.router_temperature = float(router_temperature)
self.router_noise_std = float(router_noise_std)
self.mask_non_image_embeddings = bool(mask_non_image_embeddings)
self.shared_gate_init = float(shared_gate_init)
self.specialist_gate_init = float(specialist_gate_init)
__all__ = ["ArgusConfig"]