From e73c49d624bf61953627bc61f7f685834dfd5d45 Mon Sep 17 00:00:00 2001 From: whjthu Date: Fri, 21 Aug 2026 23:23:14 +0800 Subject: [PATCH] feat: support Ktransformers, CPU-GPU MoE offload via FusedMoE layer --- csrc/layers/moe/experts/fused_moe_experts.cpp | 5 + csrc/layers/moe/fused_moe.cpp | 29 ++- csrc/layers/moe/fused_moe.hpp | 2 + csrc/layers/moe/kt_moe_callback.hpp | 59 +++++ .../deepseek/deepseek_decoder_layer.cpp | 2 +- .../deepseek_v2/deepseek_v2_decoder_layer.cpp | 2 +- csrc/models/deepseek_v2/deepseek_v2_moe.cpp | 31 ++- csrc/models/deepseek_v2/deepseek_v2_moe.hpp | 3 + .../qwen3_moe/qwen3_moe_sparse_moe_block.cpp | 3 + .../qwen3_next_sparse_moe_block.cpp | 3 + csrc/pybind11/bindings.cc | 39 ++++ python/infinilm/kt_integration.py | 216 ++++++++++++++++++ python/infinilm/modeling_utils.py | 40 +++- 13 files changed, 424 insertions(+), 10 deletions(-) create mode 100644 csrc/layers/moe/kt_moe_callback.hpp create mode 100644 python/infinilm/kt_integration.py diff --git a/csrc/layers/moe/experts/fused_moe_experts.cpp b/csrc/layers/moe/experts/fused_moe_experts.cpp index b456395d6..8febecfde 100644 --- a/csrc/layers/moe/experts/fused_moe_experts.cpp +++ b/csrc/layers/moe/experts/fused_moe_experts.cpp @@ -12,6 +12,11 @@ FusedMoeExperts::FusedMoeExperts(std::shared_ptr num_experts_ = model_config->get("num_experts"); hidden_size_ = model_config->get("hidden_size"); const size_t intermediate_size = model_config->get("moe_intermediate_size"); + if (model_config->get_or("use_kt_moe", false)) { + // KT offload: expert weights are filtered out at load time; skip + // allocating GPU-side packed weights entirely. + return; + } const auto dtype = model_config->get_dtype(); ASSERT(num_experts_ > 0); diff --git a/csrc/layers/moe/fused_moe.cpp b/csrc/layers/moe/fused_moe.cpp index fe1301a8a..9d4d5a426 100644 --- a/csrc/layers/moe/fused_moe.cpp +++ b/csrc/layers/moe/fused_moe.cpp @@ -1,5 +1,7 @@ #include "fused_moe.hpp" +#include "../moe/kt_moe_callback.hpp" + #include "dispatcher/dispatcher_factory.hpp" #include "ep/ep_config.hpp" #include "runner/cuda_fused_moe_runner.hpp" @@ -12,8 +14,15 @@ namespace infinilm::layers::moe { FusedMoE::FusedMoE(std::shared_ptr model_config, const infinicore::Device &device, - size_t layer_id) { - (void)layer_id; + size_t layer_id) + : layer_id_(layer_id), + skip_experts_(model_config->get_or("use_kt_moe", false)) { + if (skip_experts_) { + // KT offload: routed experts live on CPU (kt-kernel); skip building + // the GPU dispatcher/runner entirely. forward() consults the KT + // callback registry. + return; + } const EPConfig ep_config = make_ep_config(); const size_t num_experts = model_config->get("num_experts"); @@ -41,6 +50,22 @@ FusedMoE::FusedMoE(std::shared_ptr model_config, infinicore::Tensor FusedMoE::forward(const infinicore::Tensor &hidden_states, const TopKOutput &topk_output, const MoeWeights &weights) const { + // KT (KTransformers) branch: delegate routed-expert compute to CPU-GPU + // heterogeneous offload before touching GPU weights. + { + auto kt_cb = infinilm::layers::moe::KTMoECallbackRegistry::instance().get( + static_cast(layer_id_)); + if (kt_cb) { + return (*kt_cb)(hidden_states, topk_output.topk_weights, topk_output.topk_ids, + static_cast(layer_id_)); + } + if (skip_experts_) { + throw std::runtime_error( + "FusedMoE: use_kt_moe is enabled but no KT callback is registered for layer " + + std::to_string(layer_id_)); + } + } + auto dispatch_output = dispatcher_->dispatch(hidden_states, topk_output, workspace_); auto combine_input = runner_->run(dispatch_output, weights, workspace_); return dispatcher_->combine(combine_input, workspace_); diff --git a/csrc/layers/moe/fused_moe.hpp b/csrc/layers/moe/fused_moe.hpp index 81b3cad77..11078bf74 100644 --- a/csrc/layers/moe/fused_moe.hpp +++ b/csrc/layers/moe/fused_moe.hpp @@ -25,6 +25,8 @@ class FusedMoE final : public infinicore::nn::Module { std::shared_ptr dispatcher_; std::shared_ptr runner_; mutable MoeWorkspace workspace_; + size_t layer_id_{0}; + bool skip_experts_{false}; }; } // namespace infinilm::layers::moe diff --git a/csrc/layers/moe/kt_moe_callback.hpp b/csrc/layers/moe/kt_moe_callback.hpp new file mode 100644 index 000000000..a92c433cb --- /dev/null +++ b/csrc/layers/moe/kt_moe_callback.hpp @@ -0,0 +1,59 @@ +#pragma once +#include "infinicore/tensor.hpp" +#include +#include +#include +#include + +namespace infinilm::layers::moe { + +// Global registry of KTransformers MoE callbacks (one per layer_idx). +// +// Concurrency contract: +// - Callbacks are stored as immutable shared_ptr entries. get() copies the +// entry out under the lock and the user callback is invoked WITHOUT any +// lock held, so a callback that acquires the GIL can never deadlock +// against set()/clear() called from a GIL-holding thread. +// - get() is a single atomic lookup: no has()+call() TOCTOU window. +class KTMoECallbackRegistry { +public: + using CallbackFn = std::function; + + static KTMoECallbackRegistry &instance() { + static KTMoECallbackRegistry reg; + return reg; + } + + void set(int layer_idx, CallbackFn cb) { + std::lock_guard lk(mtx_); + cbs_[layer_idx] = std::make_shared(std::move(cb)); + } + + void clear() { + std::lock_guard lk(mtx_); + cbs_.clear(); + } + + // Returns nullptr when no callback is registered for this layer. + // The returned pointer stays valid even if set/clear run concurrently. + std::shared_ptr get(int layer_idx) { + std::lock_guard lk(mtx_); + auto it = cbs_.find(layer_idx); + return it == cbs_.end() ? nullptr : it->second; + } + + bool empty() { + std::lock_guard lk(mtx_); + return cbs_.empty(); + } + +private: + std::mutex mtx_; + std::unordered_map> cbs_; +}; + +} // namespace infinilm::layers::moe diff --git a/csrc/models/deepseek/deepseek_decoder_layer.cpp b/csrc/models/deepseek/deepseek_decoder_layer.cpp index 9ae418d0f..38ed19fe9 100644 --- a/csrc/models/deepseek/deepseek_decoder_layer.cpp +++ b/csrc/models/deepseek/deepseek_decoder_layer.cpp @@ -22,7 +22,7 @@ DeepseekDecoderLayer::DeepseekDecoderLayer(std::shared_ptr( - this->register_module("mlp", model_config, device)); + this->register_module("mlp", model_config, layer_idx, device)); } else { mlp_ = std::make_shared( this->register_module("mlp", model_config, device)); diff --git a/csrc/models/deepseek_v2/deepseek_v2_decoder_layer.cpp b/csrc/models/deepseek_v2/deepseek_v2_decoder_layer.cpp index 9a1e0c0b1..d5e26960e 100644 --- a/csrc/models/deepseek_v2/deepseek_v2_decoder_layer.cpp +++ b/csrc/models/deepseek_v2/deepseek_v2_decoder_layer.cpp @@ -24,7 +24,7 @@ DeepseekV2DecoderLayer::DeepseekV2DecoderLayer(std::shared_ptr= first_k_dense_replace && (moe_layer_freq == 0 || layer_idx % moe_layer_freq == 0); if (use_moe_) { - moe_mlp_ = this->register_module("mlp", model_config, device); + moe_mlp_ = this->register_module("mlp", model_config, layer_idx, device); } else { dense_mlp_ = this->register_module("mlp", model_config, device); } diff --git a/csrc/models/deepseek_v2/deepseek_v2_moe.cpp b/csrc/models/deepseek_v2/deepseek_v2_moe.cpp index a18351ac8..1f7e9bee9 100644 --- a/csrc/models/deepseek_v2/deepseek_v2_moe.cpp +++ b/csrc/models/deepseek_v2/deepseek_v2_moe.cpp @@ -1,4 +1,5 @@ #include "deepseek_v2_moe.hpp" +#include "../../layers/moe/kt_moe_callback.hpp" #include "../../global_state/global_state.hpp" #include "../../utils.hpp" @@ -131,9 +132,14 @@ infinicore::Tensor DeepseekV2Experts::forward(const infinicore::Tensor &hidden_s } DeepseekV2MoE::DeepseekV2MoE(std::shared_ptr model_config, - const infinicore::Device &device) { + size_t layer_idx, + const infinicore::Device &device) + : layer_idx_(layer_idx) { + skip_experts_ = model_config->get_or("use_kt_moe", false); INFINICORE_NN_MODULE_INIT(gate, model_config, device); - INFINICORE_NN_MODULE_INIT(experts, model_config, device); + if (!skip_experts_) { + INFINICORE_NN_MODULE_INIT(experts, model_config, device); + } const size_t n_shared_experts = model_config->get_or("n_shared_experts", 0); has_shared_experts_ = n_shared_experts > 0; @@ -151,7 +157,26 @@ infinicore::Tensor DeepseekV2MoE::forward(const infinicore::Tensor &hidden_state auto hidden_states_reshaped = hidden_states->view({shape[0] * shape[1], shape[2]}); auto [routing_weights, selected_experts] = gate_->forward(hidden_states_reshaped); - auto final_hidden_states = experts_->forward(hidden_states_reshaped, selected_experts, routing_weights)->view(shape); + + const auto &rank_info = infinilm::global_state::get_tensor_model_parallel_rank_info(); + infinicore::Tensor expert_output; + auto kt_cb = infinilm::layers::moe::KTMoECallbackRegistry::instance().get(static_cast(layer_idx_)); + if (kt_cb) { + if (rank_info.tp_size > 1) { + // Each rank would run KT on the full expert weights and the native path's + // partial-sum allreduce semantics do not apply -> explicit refusal. + throw std::runtime_error( + "DeepseekV2MoE: KT offload does not support tensor_parallel_size > 1"); + } + expert_output = (*kt_cb)(hidden_states_reshaped, routing_weights, selected_experts, static_cast(layer_idx_)); + } else if (skip_experts_) { + throw std::runtime_error( + "DeepseekV2MoE: use_kt_moe is enabled but no KT callback is registered for layer " + + std::to_string(layer_idx_)); + } else { + expert_output = experts_->forward(hidden_states_reshaped, selected_experts, routing_weights); + } + auto final_hidden_states = expert_output->view(shape); if (has_shared_experts_) { auto shared_out = shared_experts_->forward(hidden_states); final_hidden_states = infinicore::op::add(final_hidden_states, shared_out); diff --git a/csrc/models/deepseek_v2/deepseek_v2_moe.hpp b/csrc/models/deepseek_v2/deepseek_v2_moe.hpp index 11af3669c..f538a09d2 100644 --- a/csrc/models/deepseek_v2/deepseek_v2_moe.hpp +++ b/csrc/models/deepseek_v2/deepseek_v2_moe.hpp @@ -63,6 +63,7 @@ class DeepseekV2Experts : public infinicore::nn::Module { class DeepseekV2MoE : public infinicore::nn::Module { public: DeepseekV2MoE(std::shared_ptr model_config, + size_t layer_idx, const infinicore::Device &device); infinicore::Tensor forward(const infinicore::Tensor &hidden_states) const; @@ -72,6 +73,8 @@ class DeepseekV2MoE : public infinicore::nn::Module { INFINICORE_NN_MODULE(DeepseekV2Experts, experts); INFINICORE_NN_MODULE(DeepseekV2MLP, shared_experts); bool has_shared_experts_{false}; + size_t layer_idx_{0}; + bool skip_experts_{false}; }; } // namespace infinilm::models::deepseek_v2 diff --git a/csrc/models/qwen3_moe/qwen3_moe_sparse_moe_block.cpp b/csrc/models/qwen3_moe/qwen3_moe_sparse_moe_block.cpp index 771e9c971..974ae0212 100644 --- a/csrc/models/qwen3_moe/qwen3_moe_sparse_moe_block.cpp +++ b/csrc/models/qwen3_moe/qwen3_moe_sparse_moe_block.cpp @@ -43,6 +43,9 @@ infinicore::Tensor Qwen3MoeSparseMoeBlock::forward(const infinicore::Tensor &hid infinicore::Tensor(), }; + // KT (KTransformers) offload is handled inside FusedMoE::forward when + // a KT callback is registered for this layer. + auto final_hidden_states = fused_moe_->forward( hidden_states_reshaped, topk_output, diff --git a/csrc/models/qwen3_next/qwen3_next_sparse_moe_block.cpp b/csrc/models/qwen3_next/qwen3_next_sparse_moe_block.cpp index d3548c9fa..9303451bf 100644 --- a/csrc/models/qwen3_next/qwen3_next_sparse_moe_block.cpp +++ b/csrc/models/qwen3_next/qwen3_next_sparse_moe_block.cpp @@ -64,6 +64,7 @@ Qwen3NextSparseMoeBlock::Qwen3NextSparseMoeBlock(std::shared_ptrregister_module("gate", model_config, device); experts_ = this->register_module("experts", model_config, device); fused_moe_ = this->register_module("fused_moe", model_config, device, layer_idx); + (void)layer_idx; shared_expert_ = this->register_module("shared_expert", model_config, device); shared_expert_gate_ = this->register_module( "shared_expert_gate", @@ -86,6 +87,8 @@ infinicore::Tensor Qwen3NextSparseMoeBlock::forward(const infinicore::Tensor &hi selected_experts, infinicore::Tensor(), }; + // KT (KTransformers) offload is handled inside FusedMoE::forward when + // a KT callback is registered for this layer. auto routed_states = fused_moe_->forward( hidden_states_reshaped, topk_output, diff --git a/csrc/pybind11/bindings.cc b/csrc/pybind11/bindings.cc index 63846338b..be9f43d63 100644 --- a/csrc/pybind11/bindings.cc +++ b/csrc/pybind11/bindings.cc @@ -1,7 +1,10 @@ +#include "../layers/moe/kt_moe_callback.hpp" #include #include "cache/cache.hpp" #include "engine/engine.hpp" +#include +#include namespace py = pybind11; @@ -12,4 +15,40 @@ PYBIND11_MODULE(_infinilm, m) { infinilm::engine::bind_hook_registry(m); infinilm::engine::distributed::bind_dist_config(m); infinilm::engine::bind_infer_engine(m); + + // ---- KTransformers MoE offload integration ---- + // Register a Python callback (per layer) that receives + // (hidden, routing_weights, topk_ids, layer_idx) as infinicore tensors + // and must return the routed-expert output tensor (infinicore view). + m.def( + "set_kt_moe_callback", [](int layer_idx, py::function callback) { + infinilm::layers::moe::KTMoECallbackRegistry::instance().set(layer_idx, + [callback](const infinicore::Tensor &h, const infinicore::Tensor &w, + const infinicore::Tensor &i, int l) -> infinicore::Tensor { + // Translate Python exceptions at the boundary while the GIL is + // held: the C++ worker thread has no GIL and no Python frame to + // unwind into, and py::error_already_set must not escape it. + py::gil_scoped_acquire gil; + try { + return callback(h, w, i, l).cast(); + } catch (const py::error_already_set &e) { + throw std::runtime_error( + "KT MoE callback (layer " + std::to_string(l) + ") failed: " + e.what()); + } + }); + }, + py::arg("layer_idx"), py::arg("callback"), "Register a KTransformers MoE callback for a given layer."); + + m.def( + "clear_kt_moe_callbacks", []() { + infinilm::layers::moe::KTMoECallbackRegistry::instance().clear(); + }, + "Clear all KT MoE callbacks."); + + // Release registered py::functions at module teardown (while the + // interpreter is still alive) instead of process-exit static destruction. + m.add_object("_kt_moe_cleanup", + py::capsule(reinterpret_cast(1), "_kt_moe_cleanup", [](void *) { + infinilm::layers::moe::KTMoECallbackRegistry::instance().clear(); + })); } diff --git a/python/infinilm/kt_integration.py b/python/infinilm/kt_integration.py new file mode 100644 index 000000000..b3ebb26d7 --- /dev/null +++ b/python/infinilm/kt_integration.py @@ -0,0 +1,216 @@ +#!/usr/bin/env python3 +"""InfiniLM + KT (KTransformers) CPU-GPU heterogeneous MoE offload. + +Performance-critical callback design (verified on L20, INT4 Q4_K_M): + - Pre-allocated persistent GPU staging buffers (hidden/ids/weights/output) + - Raw cudaMemcpyAsync via ctypes on the default stream (no per-call alloc) + - Double-buffered output: layer N and N+1 use different buffers so the + async C++ consumer (residual add) can never race the next snapshot + - Output snapshot: KT reuses its internal output GPU buffer across calls, + so we must copy() into our own buffer before returning to InfiniLM + +Usage: + model = LLM(model_path=..., cache_type="paged", attn_backend="paged-attn", + enable_prefix_caching=False, ...) # config.json needs use_kt_moe=true + from infinilm.kt_integration import setup_kt_moe + setup_kt_moe(model, model_path=GGUF_DIR, method="LLAMAFILE", + num_experts=512, num_experts_per_tok=10, + hidden_size=2048, moe_intermediate_size=512, + num_hidden_layers=48, num_gpu_experts=0, + max_tokens=ENGINE_MAX_NUM_BATCHED_TOKENS) + outs = model.generate([...], ...) + from infinilm.lib import _infinilm; _infinilm.clear_kt_moe_callbacks() + +Requirements / sharp edges: + - max_tokens must be >= the largest forward batch (prefill!), i.e. the + engine's max_num_batched_tokens, NOT the decode batch size. + - enable_graph must be False (CUDA graph capture executes this Python + callback; capture would fail or replay stale outputs). + - Models with GDN/linear-attention (qwen3_next): ensure num_blocks//4 + (mamba cache pool) >= max_batch_size or concurrency gets serialized. + - Single GPU (device 0), tensor_parallel_size=1 only. +""" + +import ctypes +import logging + +import torch + +logger = logging.getLogger(__name__) + + +def _load_cudart(): + """Load the CUDA runtime library, preferring versioned fallbacks.""" + for name in ("libcudart.so", "libcudart.so.12", "libcudart.so.11.0"): + try: + return ctypes.CDLL(name) + except OSError: + continue + raise OSError("libcudart not found (looked for libcudart.so/.so.12/.so.11.0)") + + +def setup_kt_moe( + model, + model_path, + num_experts, + num_experts_per_tok, + hidden_size, + moe_intermediate_size, + num_hidden_layers, + num_gpu_experts=0, + method="LLAMAFILE", + cpuinfer_threads=8, + threadpool_count=1, + max_tokens=512, + chunked_prefill_size=512, + moe_layer_freq=1, + moe_layers=None, +): + """Create one KTMoEWrapper per MoE layer and register InfiniLM callbacks. + + Args: + model: the infinilm LLM object (used for config sanity checks). + model_path: GGUF directory for KT expert weights (INT4 Q4_K_M recommended). + num_gpu_experts: experts kept on GPU inside KT (0 = all experts on CPU). + method: KT backend. "LLAMAFILE" (GGUF) is the verified path. + cpuinfer_threads: CPU threads for KT compute pool. + max_tokens: staging buffer capacity; MUST be >= engine max_num_batched_tokens + (a larger prefill raises RuntimeError instead of corrupting memory). + chunked_prefill_size: KT wrapper's chunked prefill size (its own buffers). + moe_layers: explicit iterable of MoE layer indices (e.g. DSV2's + range(first_k_dense_replace, num_hidden_layers)). Defaults to + range(0, num_hidden_layers, moe_layer_freq). + + Returns: + Number of layers registered. + """ + import infinicore as ic + from kt_kernel import KTMoEWrapper + + from infinilm.lib import _infinilm + + # ---- sanity checks -------------------------------------------------- # + engine_cfg = getattr(model, "config", None) + if engine_cfg is not None and getattr(engine_cfg, "enable_graph", False): + raise RuntimeError( + "KT offload is incompatible with enable_graph=True " + "(CUDA graph capture would execute the Python KT callback)" + ) + + ne, topk, hs, mis = ( + num_experts, + num_experts_per_tok, + hidden_size, + moe_intermediate_size, + ) + if moe_layers is None: + moe_layers = list(range(0, num_hidden_layers, moe_layer_freq)) + else: + moe_layers = list(moe_layers) + + # ---- ctypes cudaMemcpyAsync on default stream (fast, no torch overhead) ---- + ca = _load_cudart() + ca.cudaMemcpyAsync.argtypes = [ + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_size_t, + ctypes.c_int, + ctypes.c_void_p, + ] + ca.cudaMemcpyAsync.restype = ctypes.c_int + D2D = 1 # cudaMemcpyDeviceToDevice + + # ---- Persistent staging buffers shared by all layers (sequential access) ---- + h_buf = torch.empty(max_tokens, hs, dtype=torch.bfloat16, device="cuda") + i_buf = torch.empty(max_tokens, topk, dtype=torch.int32, device="cuda") + w_buf = torch.empty(max_tokens, topk, dtype=torch.float32, device="cuda") + # Double-buffered output: consecutive layers alternate, so layer N+1's + # snapshot can never overwrite memory layer N's async consumer still reads. + out_bufs = [ + torch.empty(max_tokens, hs, dtype=torch.bfloat16, device="cuda") + for _ in range(2) + ] + stream_cache = [None] # resolved on first callback (worker thread) + + def make_cb(wr): + def cb(hidden, weights, ids, layer): + n = hidden.size(0) + if n > max_tokens: + raise RuntimeError( + f"KT staging overflow: forward batch {n} > max_tokens {max_tokens}. " + "Increase setup_kt_moe(max_tokens=...) to cover the engine's " + "max_num_batched_tokens (prefill batches), or lower " + "INFINILM_MAX_NUM_BATCHED_TOKENS." + ) + if hidden.size(1) != hs: + raise RuntimeError( + f"KT hidden size mismatch: got {hidden.size(1)}, expected {hs}" + ) + # infinicore -> torch staging (async, default stream) + rc = ca.cudaMemcpyAsync( + h_buf.data_ptr(), hidden.data_ptr(), n * hs * 2, D2D, 0 + ) + if rc != 0: + raise RuntimeError(f"KT cudaMemcpyAsync(hidden) failed: cudaError {rc}") + rc = ca.cudaMemcpyAsync( + i_buf.data_ptr(), ids.data_ptr(), n * topk * 4, D2D, 0 + ) + if rc != 0: + raise RuntimeError(f"KT cudaMemcpyAsync(ids) failed: cudaError {rc}") + rc = ca.cudaMemcpyAsync( + w_buf.data_ptr(), weights.data_ptr(), n * topk * 4, D2D, 0 + ) + if rc != 0: + raise RuntimeError( + f"KT cudaMemcpyAsync(weights) failed: cudaError {rc}" + ) + if stream_cache[0] is None: + stream_cache[0] = torch.cuda.current_stream().cuda_stream + out = wr.forward(h_buf[:n], i_buf[:n], w_buf[:n], stream_cache[0]) + # KT reuses its output buffer between calls: snapshot before returning. + ob = out_bufs[layer & 1][:n] + ob.copy_(out) + return ic.from_torch(ob)._underlying + + return cb + + # ---- Register per-layer wrappers; roll back on partial failure -------- # + loaded = 0 + try: + for layer_idx in moe_layers: + mask = torch.tensor( + [True] * num_gpu_experts + [False] * (ne - num_gpu_experts), + dtype=torch.bool, + ) + wr = KTMoEWrapper( + layer_idx, + ne, + topk, + hs, + mis, + mask, + cpuinfer_threads, + threadpool_count, + model_path, + chunked_prefill_size, + method=method, + ) + wr.load_weights(torch.arange(ne, dtype=torch.int32)) + _infinilm.set_kt_moe_callback(layer_idx, make_cb(wr)) + loaded += 1 + except Exception: + # Never leave a partially-registered model behind: the C++ side would + # fall through to a native path whose expert modules were skipped. + _infinilm.clear_kt_moe_callbacks() + raise + + logger.info( + "KT MoE ready: %d layers [%d..%d], %d GPU + %d CPU experts, method=%s", + loaded, + moe_layers[0], + moe_layers[-1], + num_gpu_experts, + ne - num_gpu_experts, + method, + ) + return loaded diff --git a/python/infinilm/modeling_utils.py b/python/infinilm/modeling_utils.py index b210b0c56..f4d9ad15d 100644 --- a/python/infinilm/modeling_utils.py +++ b/python/infinilm/modeling_utils.py @@ -189,6 +189,11 @@ def get_model_state_dict( return model_param_infini +def _kt_expert_key(key: str) -> bool: + """True for routed-expert weight keys skipped when use_kt_moe is enabled.""" + return ".mlp.experts." in key or ".mlp.fused_moe." in key + + def load_model_state_dict_by_file( model: infinicore.nn.Module, model_path: str, @@ -274,6 +279,12 @@ def load_model_state_dict_by_file( # --------------------------------------------------------- # # model_param_infini references torch.Tensor # --------------------------------------------------------- # + # Skip expert weights when KT offload is enabled + use_kt_moe = model.hf_config.get("use_kt_moe", False) + if use_kt_moe: + model_param = { + k: v for k, v in model_param.items() if not _kt_expert_key(k) + } model_param_infini = {} for key in model_param.keys(): model_param_infini[key] = infinicore.from_torch(model_param[key]) @@ -319,6 +330,13 @@ def load_model_state_dict_by_file( if key in model_key_set } + # Skip expert weights BEFORE conversion when KT offload is enabled + use_kt_moe = model.hf_config.get("use_kt_moe", False) + if use_kt_moe: + model_params = { + k: v for k, v in model_params.items() if not _kt_expert_key(k) + } + model_param_infini = {} for key in model_params.keys(): target_dtype = ( @@ -331,7 +349,8 @@ def load_model_state_dict_by_file( ) already_loaded_keys.append(key) - model.load_state_dict(model_param_infini, strict=True) + # strict=False under KT: expert keys are intentionally absent + model.load_state_dict(model_param_infini, strict=not use_kt_moe) infinicore.sync_device() del model_param_infini del model_params @@ -350,7 +369,15 @@ def load_model_state_dict_by_file( embed_tokens_torch_unscaled = None gc.collect() - check_parameters(model_keys, already_loaded_keys) + use_kt_moe = model.hf_config.get("use_kt_moe", False) + if use_kt_moe: + # Keep the safety net for non-expert weights; expert keys are expected missing. + check_parameters( + [k for k in model_keys if not _kt_expert_key(k)], + [k for k in already_loaded_keys if not _kt_expert_key(k)], + ) + else: + check_parameters(model_keys, already_loaded_keys) if not weights_processed: model.process_weights_after_loading() @@ -428,7 +455,14 @@ def load_model_state_dict_by_tensor( model.load_param("lm_head.weight", lm_head_tensor) already_loaded_keys.append("lm_head.weight") - check_parameters(model_keys, already_loaded_keys) + use_kt_moe = model.hf_config.get("use_kt_moe", False) + if use_kt_moe: + check_parameters( + [k for k in model_keys if not _kt_expert_key(k)], + [k for k in already_loaded_keys if not _kt_expert_key(k)], + ) + else: + check_parameters(model_keys, already_loaded_keys) t2 = time.time() print(f" load weights over! {(t2 - t1) * 1000} ms \n")