diff --git a/README.md b/README.md index 3c4611b..2cd9013 100644 --- a/README.md +++ b/README.md @@ -151,19 +151,10 @@ FlashHead models use a custom architecture name (e.g., `FlashHeadQwen3VLForCondi ### 🏗️ Supported Architectures -See most recent architectures in [_FLASHHEAD_ARCHITECTURES](https://github.com/embedl/flash-head/blob/master/src/flash_head/__init__.py): -```python -_FLASHHEAD_ARCHITECTURES = { - "FlashHeadLlamaForCausalLM": "vllm.model_executor.models.llama:LlamaForCausalLM", - "FlashHeadQwen3ForCausalLM": "vllm.model_executor.models.qwen3:Qwen3ForCausalLM", - "FlashHeadQwen3VLForConditionalGeneration": "vllm.model_executor.models.qwen3_vl:Qwen3VLForConditionalGeneration", - "FlashHeadQwen3_5ForCausalLM": "vllm.model_executor.models.qwen3_5:Qwen3_5ForCausalLM", - "FlashHeadQwen3_5MoeForCausalLM": "vllm.model_executor.models.qwen3_5:Qwen3_5MoeForCausalLM", - "FlashHeadQwen3_5ForConditionalGeneration": "vllm.model_executor.models.qwen3_5:Qwen3_5ForConditionalGeneration", - "FlashHeadQwen3_5MoeForConditionalGeneration": "vllm.model_executor.models.qwen3_5:Qwen3_5MoeForConditionalGeneration", - "FlashHeadGemma3ForCausalLM": "vllm.model_executor.models.gemma3:Gemma3ForCausalLM", -} -``` +The `FlashHead` prefix is only a safety trip — it makes stock vLLM refuse the model when the plugin is missing. When the plugin is installed it simply **strips the prefix** and lets vLLM load the standard base architecture (`FlashHead` → ``), so **any architecture vLLM already supports works**. + +FlashHead itself activates from `flash_head_cache_dir` in the model config, independent of the architecture name. As a result a model that keeps its **standard** architecture name (no `FlashHead` prefix) but ships a `flash_head_cache_dir` also gets FlashHead applied — the prefix is purely about the no-plugin safety error. + ## 📤 Publishing FlashHead Models diff --git a/src/flash_head/__init__.py b/src/flash_head/__init__.py index 5acaabd..ab1af50 100644 --- a/src/flash_head/__init__.py +++ b/src/flash_head/__init__.py @@ -41,36 +41,74 @@ def register(): logger.info("[FlashHead] Plugin registered") -def _register_architectures(): - """Register FlashHead model architectures so vLLM recognizes them. +FLASHHEAD_ARCH_PREFIX = "FlashHead" + + +def _strip_flashhead_prefix(architectures): + """Map any 'FlashHead' architecture to its base '' name. + + Returns ``(architectures, changed)`` where ``changed`` is True if at + least one name was rewritten. Names without the prefix - and the bare + prefix itself - pass through untouched. + """ + stripped = [] + changed = False + for arch in architectures: + if ( + isinstance(arch, str) + and arch.startswith(FLASHHEAD_ARCH_PREFIX) + and len(arch) > len(FLASHHEAD_ARCH_PREFIX) + ): + stripped.append(arch[len(FLASHHEAD_ARCH_PREFIX):]) + changed = True + else: + stripped.append(arch) + return stripped, changed + - Models published with architectures like 'FlashHeadQwen3VLForConditionalGeneration' - will fail to load without this plugin -- giving a clear error instead of - silently falling back to the slow standard lm_head path. +def _register_architectures(): + """Let vLLM load any 'FlashHead' model by stripping the prefix. + + The 'FlashHead' prefix on a published model's architecture is a safety + trip, not a distinct model class: stock vLLM (without this plugin) does + not recognize 'FlashHead' and refuses to load, so users can never + silently fall back to the slow dense lm_head. Once THIS plugin is + installed the prefix has served its purpose, so we transparently strip it + and let vLLM resolve the standard '' architecture with its own + (lazy) model class. + + FlashHead activation is driven entirely by 'flash_head_cache_dir' in the + model config -- independent of the architecture name -- so this works for + ANY base architecture vLLM supports, with no hand-maintained allow-list. """ try: from vllm import ModelRegistry + except Exception as e: # vLLM not importable -- nothing to wrap + logger.debug("[FlashHead] Architecture shim skipped: %s", e) + return - # Map FlashHead architecture names to their base vLLM model classes. - # Uses lazy string imports to avoid premature CUDA initialization. - # The FlashHead interception happens via the LogitsProcessor patch, - # not via a custom model class, so we just need vLLM to accept the - # architecture name and load the base model. - _FLASHHEAD_ARCHITECTURES = { - "FlashHeadLlamaForCausalLM": "vllm.model_executor.models.llama:LlamaForCausalLM", - "FlashHeadQwen3ForCausalLM": "vllm.model_executor.models.qwen3:Qwen3ForCausalLM", - "FlashHeadQwen3VLForConditionalGeneration": "vllm.model_executor.models.qwen3_vl:Qwen3VLForConditionalGeneration", - "FlashHeadQwen3_5ForCausalLM": "vllm.model_executor.models.qwen3_5:Qwen3_5ForCausalLM", - "FlashHeadQwen3_5MoeForCausalLM": "vllm.model_executor.models.qwen3_5:Qwen3_5MoeForCausalLM", - "FlashHeadQwen3_5ForConditionalGeneration": "vllm.model_executor.models.qwen3_5:Qwen3_5ForConditionalGeneration", - "FlashHeadQwen3_5MoeForConditionalGeneration": "vllm.model_executor.models.qwen3_5:Qwen3_5MoeForConditionalGeneration", - "FlashHeadGemma3ForCausalLM": "vllm.model_executor.models.gemma3:Gemma3ForCausalLM", - } - - supported = ModelRegistry.get_supported_archs() - for fh_arch, model_cls_path in _FLASHHEAD_ARCHITECTURES.items(): - if fh_arch not in supported: - ModelRegistry.register_model(fh_arch, model_cls_path) - logger.info("[FlashHead] Registered architecture %s", fh_arch) - except Exception as e: - logger.debug("[FlashHead] Architecture registration skipped: %s", e) + def _wrap(method_name): + original = getattr(ModelRegistry, method_name, None) + if original is None or getattr(original, "_flashhead_wrapped", False): + return + + def _patched(architectures, *args, **kwargs): + archs = ( + [architectures] + if isinstance(architectures, str) + else list(architectures) + ) + stripped, changed = _strip_flashhead_prefix(archs) + if changed: + logger.info( + "[FlashHead] Loading %s as base architecture %s", + archs, stripped, + ) + return original(stripped, *args, **kwargs) + return original(archs, *args, **kwargs) + + setattr(_patched, "_flashhead_wrapped", True) + setattr(ModelRegistry, method_name, _patched) + + for _name in ("inspect_model_cls", "resolve_model_cls"): + _wrap(_name)