"""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"]