From 00cb767457b7d415689e65105fcf5d893d80bfb6 Mon Sep 17 00:00:00 2001 From: Jeff Ma Date: Sun, 9 Aug 2026 03:28:39 +0000 Subject: [PATCH] allow tpu to import kimi_k3.common Signed-off-by: Jeff Ma --- vllm/models/kimi_k3/__init__.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/vllm/models/kimi_k3/__init__.py b/vllm/models/kimi_k3/__init__.py index 54f965eadaa7..52c8a6538e66 100644 --- a/vllm/models/kimi_k3/__init__.py +++ b/vllm/models/kimi_k3/__init__.py @@ -13,13 +13,20 @@ # The NVIDIA branch is the static default that type-checkers see; the ROCm # branch overrides it at runtime (kept type-compatible via type: ignore). -if TYPE_CHECKING or not current_platform.is_rocm(): +# TPU plugins import the shared ``common`` modules through this package, but +# register their own model classes. Do not eagerly import a GPU implementation. +if TYPE_CHECKING: from .nvidia.model import KimiK3ForConditionalGeneration, KimiLinearForCausalLM from .nvidia.mtp import KimiK3MTP -else: +elif current_platform.device_type == "tpu": + pass +elif current_platform.is_rocm(): from .amd.linear import KimiLinearForCausalLM # type: ignore[assignment] from .amd.model import KimiK3ForConditionalGeneration # type: ignore[assignment] from .amd.mtp import KimiK3MTP # type: ignore[assignment] +else: + from .nvidia.model import KimiK3ForConditionalGeneration, KimiLinearForCausalLM + from .nvidia.mtp import KimiK3MTP __all__ = [ "KimiK3ForConditionalGeneration",