Replaces cherry-picked upstream_ref with complete source trees. xllm/ — Iluvatar official C++ inference engine (15MB, 1470 files) Complete: kernels → layers → models → runtime → scheduler → api Excluded: .git, binary images, third_party submodule checkouts ds_vllm/ — Iluvatar official vllm fork (8MB, 703 files) Included: csrc/ (ALL CUDA kernels), fused_moe/, qwen3_5 model, _custom_ops Excluded: tests, benchmarks, docs, examples (not needed for reference) Critical call chains now fully traceable: MoE: moe_topk_softmax_kernels.cuh → ixformer.h → fused_moe.cpp → layer GDN: qwen3_gated_delta_net_base.cpp → qwen3_5_gated_delta_net.cpp Attention: ixformer.h → xllm_paged_attention → attention.cpp
464 lines
16 KiB
C++
464 lines
16 KiB
C++
/* Copyright 2025 The xLLM Authors. All Rights Reserved.
|
|
Copyright 2024 The ScaleLLM Authors. All Rights Reserved.
|
|
|
|
Licensed under the Apache License, Version 2.0 (the "License");
|
|
you may not use this file except in compliance with the License.
|
|
You may obtain a copy of the License at
|
|
|
|
https://github.com/jd-opensource/xllm/blob/main/LICENSE
|
|
|
|
Unless required by applicable law or agreed to in writing, software
|
|
distributed under the License is distributed on an "AS IS" BASIS,
|
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
See the License for the specific language governing permissions and
|
|
limitations under the License.
|
|
==============================================================================*/
|
|
|
|
#include "model_registry.h"
|
|
|
|
#include <glog/logging.h>
|
|
|
|
#include <iostream>
|
|
#include <unordered_set>
|
|
|
|
#include "core/common/global_flags.h"
|
|
#include "models.h"
|
|
|
|
namespace {
|
|
|
|
// Safe logging macro to avoid crashes during static initialization
|
|
#define SAFE_LOG_WARNING(message) \
|
|
do { \
|
|
if (google::IsGoogleLoggingInitialized()) { \
|
|
LOG(WARNING) << message; \
|
|
} else { \
|
|
std::cerr << "WARNING: " << message << std::endl; \
|
|
} \
|
|
} while (0)
|
|
|
|
#define SAFE_LOG_ERROR(message) \
|
|
do { \
|
|
if (google::IsGoogleLoggingInitialized()) { \
|
|
LOG(ERROR) << message; \
|
|
} else { \
|
|
std::cerr << "ERROR: " << message << std::endl; \
|
|
} \
|
|
} while (0)
|
|
|
|
#define SAFE_LOG_INFO(message) \
|
|
do { \
|
|
if (google::IsGoogleLoggingInitialized()) { \
|
|
LOG(INFO) << message; \
|
|
} else { \
|
|
std::cerr << "INFO: " << message << std::endl; \
|
|
} \
|
|
} while (0)
|
|
|
|
} // anonymous namespace
|
|
|
|
namespace xllm {
|
|
|
|
namespace {
|
|
|
|
#if defined(USE_NPU)
|
|
constexpr char kAutoBackend[] = "AUTO";
|
|
constexpr char kAtbBackend[] = "ATB";
|
|
constexpr char kTorchBackend[] = "TORCH";
|
|
|
|
bool is_torch_only_model_type(const std::string& model_type) {
|
|
static const std::unordered_set<std::string> kTorchOnlyModelTypes = {
|
|
"qwen3_5",
|
|
"qwen3_5_text",
|
|
"qwen3_5_moe",
|
|
"qwen3_5_moe_text",
|
|
"qwen3_5_mtp",
|
|
"qwen3_5_moe_mtp",
|
|
"qwen3_next"};
|
|
return kTorchOnlyModelTypes.count(model_type) != 0;
|
|
}
|
|
#endif
|
|
|
|
} // namespace
|
|
|
|
bool resolve_model_registration(const std::string& model_type,
|
|
const std::string& requested_npu_kernel_backend,
|
|
std::string* effective_npu_kernel_backend,
|
|
std::string* resolved_name,
|
|
std::string* error_message) {
|
|
if (resolved_name == nullptr) {
|
|
if (error_message != nullptr) {
|
|
*error_message = "resolved_name must not be null";
|
|
}
|
|
return false;
|
|
}
|
|
|
|
#if defined(USE_NPU)
|
|
const std::string backend = requested_npu_kernel_backend.empty()
|
|
? kAutoBackend
|
|
: requested_npu_kernel_backend;
|
|
if (backend != kAutoBackend && backend != kAtbBackend &&
|
|
backend != kTorchBackend) {
|
|
if (error_message != nullptr) {
|
|
*error_message = "Unsupported --npu_kernel_backend=" + backend +
|
|
". Supported values: AUTO, ATB, TORCH.";
|
|
}
|
|
return false;
|
|
}
|
|
|
|
std::string effective_backend = backend;
|
|
if (backend == kAutoBackend) {
|
|
effective_backend =
|
|
is_torch_only_model_type(model_type) ? kTorchBackend : kAtbBackend;
|
|
} else if (model_type == "qwen3") {
|
|
// qwen3 supports both backends.
|
|
} else if (is_torch_only_model_type(model_type)) {
|
|
if (backend != kTorchBackend) {
|
|
if (error_message != nullptr) {
|
|
*error_message = "Model type " + model_type +
|
|
" only supports --npu_kernel_backend=TORCH.";
|
|
}
|
|
return false;
|
|
}
|
|
} else if (backend != kAtbBackend) {
|
|
if (error_message != nullptr) {
|
|
*error_message = "Model type " + model_type +
|
|
" only supports --npu_kernel_backend=ATB.";
|
|
}
|
|
return false;
|
|
}
|
|
|
|
if (effective_npu_kernel_backend != nullptr) {
|
|
*effective_npu_kernel_backend = effective_backend;
|
|
}
|
|
*resolved_name = (model_type == "qwen3" && effective_backend == kAtbBackend)
|
|
? "qwen3_atb"
|
|
: model_type;
|
|
return true;
|
|
#else
|
|
*resolved_name = model_type;
|
|
return true;
|
|
#endif
|
|
}
|
|
|
|
bool resolve_model_registration_name(const std::string& model_type,
|
|
std::string* resolved_name,
|
|
std::string* error_message) {
|
|
#if defined(USE_NPU)
|
|
return resolve_model_registration(model_type,
|
|
FLAGS_npu_kernel_backend,
|
|
nullptr,
|
|
resolved_name,
|
|
error_message);
|
|
#else
|
|
return resolve_model_registration(
|
|
model_type, "", nullptr, resolved_name, error_message);
|
|
#endif
|
|
}
|
|
|
|
ModelRegistry* ModelRegistry::get_instance() {
|
|
static ModelRegistry registry;
|
|
|
|
return ®istry;
|
|
}
|
|
|
|
void ModelRegistry::register_causallm_factory(const std::string& name,
|
|
CausalLMFactory factory) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].causal_lm_factory != nullptr) {
|
|
SAFE_LOG_WARNING("causal lm factory for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].causal_lm_factory = factory;
|
|
instance->model_backend_[name] = "llm";
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_rec_model_factory(const std::string& name,
|
|
RecModelFactory factory) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].rec_model_factory != nullptr) {
|
|
SAFE_LOG_WARNING("rec model factory for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].rec_model_factory = factory;
|
|
instance->model_backend_[name] = "rec";
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_causalvlm_factory(const std::string& name,
|
|
CausalVLMFactory factory) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].causal_vlm_factory != nullptr) {
|
|
SAFE_LOG_WARNING("causal vlm factory for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].causal_vlm_factory = factory;
|
|
instance->model_backend_[name] = "vlm";
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_mm_embedding_vlm_factory(
|
|
const std::string& name,
|
|
MMEmbeddingVLMFactory factory) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].mm_embedding_vlm_factory != nullptr) {
|
|
SAFE_LOG_WARNING("mm embedding vlm factory for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].mm_embedding_vlm_factory = factory;
|
|
instance->model_backend_[name] = "vlm";
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_dit_model_factory(const std::string& name,
|
|
DiTModelFactory factory) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].dit_model_factory != nullptr) {
|
|
SAFE_LOG_WARNING("DiT model factory for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].dit_model_factory = factory;
|
|
instance->model_backend_[name] = "dit";
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_input_processor_factory(
|
|
const std::string& name,
|
|
InputProcessorFactory factory) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].input_processor_factory != nullptr) {
|
|
SAFE_LOG_WARNING("input processor factory for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].input_processor_factory = factory;
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_image_processor_factory(
|
|
const std::string& name,
|
|
ImageProcessorFactory factory) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].image_processor_factory != nullptr) {
|
|
SAFE_LOG_WARNING("image processor factory for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].image_processor_factory = factory;
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_model_args_loader(const std::string& name,
|
|
ModelArgsLoader loader) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].model_args_loader != nullptr) {
|
|
SAFE_LOG_WARNING("model args loader for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].model_args_loader = loader;
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_quant_args_loader(const std::string& name,
|
|
QuantArgsLoader loader) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].quant_args_loader != nullptr) {
|
|
SAFE_LOG_WARNING("quant args loader for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].quant_args_loader = loader;
|
|
}
|
|
}
|
|
|
|
void ModelRegistry::register_tokenizer_args_loader(const std::string& name,
|
|
TokenizerArgsLoader loader) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
if (instance->model_registry_[name].tokenizer_args_loader != nullptr) {
|
|
SAFE_LOG_WARNING("tokenizer args loader for " << name
|
|
<< " already registered.");
|
|
} else {
|
|
instance->model_registry_[name].tokenizer_args_loader = loader;
|
|
}
|
|
}
|
|
|
|
CausalLMFactory ModelRegistry::get_causallm_factory(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].causal_lm_factory;
|
|
}
|
|
|
|
RecModelFactory ModelRegistry::get_rec_model_factory(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].rec_model_factory;
|
|
}
|
|
|
|
CausalVLMFactory ModelRegistry::get_causalvlm_factory(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].causal_vlm_factory;
|
|
}
|
|
|
|
MMEmbeddingVLMFactory ModelRegistry::get_mm_embedding_vlm_factory(
|
|
const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].mm_embedding_vlm_factory;
|
|
}
|
|
|
|
DiTModelFactory ModelRegistry::get_dit_model_factory(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
return instance->model_registry_[name].dit_model_factory;
|
|
}
|
|
|
|
InputProcessorFactory ModelRegistry::get_input_processor_factory(
|
|
const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].input_processor_factory;
|
|
}
|
|
|
|
ImageProcessorFactory ModelRegistry::get_image_processor_factory(
|
|
const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].image_processor_factory;
|
|
}
|
|
|
|
ModelArgsLoader ModelRegistry::get_model_args_loader(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].model_args_loader;
|
|
}
|
|
|
|
QuantArgsLoader ModelRegistry::get_quant_args_loader(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].quant_args_loader;
|
|
}
|
|
|
|
TokenizerArgsLoader ModelRegistry::get_tokenizer_args_loader(
|
|
const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
|
|
return instance->model_registry_[name].tokenizer_args_loader;
|
|
}
|
|
|
|
bool ModelRegistry::has_dit_model_factory(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
return (instance->model_registry_.find(name) !=
|
|
instance->model_registry_.end());
|
|
}
|
|
|
|
std::string ModelRegistry::get_model_backend(const std::string& name) {
|
|
ModelRegistry* instance = get_instance();
|
|
return instance->model_backend_[name];
|
|
}
|
|
|
|
std::unique_ptr<CausalLM> create_llm_model(const ModelContext& context) {
|
|
std::string resolved_name;
|
|
std::string error_message;
|
|
if (!resolve_model_registration_name(context.get_model_args().model_type(),
|
|
&resolved_name,
|
|
&error_message)) {
|
|
LOG(ERROR) << error_message;
|
|
return nullptr;
|
|
}
|
|
|
|
auto factory = ModelRegistry::get_causallm_factory(resolved_name);
|
|
if (factory) {
|
|
return factory(context);
|
|
}
|
|
|
|
LOG(ERROR) << "Unsupported model type: "
|
|
<< context.get_model_args().model_type();
|
|
|
|
return nullptr;
|
|
}
|
|
|
|
std::unique_ptr<CausalLM> create_rec_model(const ModelContext& context) {
|
|
std::string resolved_name;
|
|
std::string error_message;
|
|
if (!resolve_model_registration_name(context.get_model_args().model_type(),
|
|
&resolved_name,
|
|
&error_message)) {
|
|
LOG(ERROR) << error_message;
|
|
return nullptr;
|
|
}
|
|
|
|
auto factory = ModelRegistry::get_rec_model_factory(resolved_name);
|
|
if (factory) {
|
|
return factory(context);
|
|
}
|
|
|
|
LOG(ERROR) << "Unsupported rec model type: "
|
|
<< context.get_model_args().model_type();
|
|
|
|
return nullptr;
|
|
}
|
|
|
|
std::unique_ptr<CausalVLM> create_vlm_model(const ModelContext& context) {
|
|
std::string resolved_name;
|
|
std::string error_message;
|
|
if (!resolve_model_registration_name(context.get_model_args().model_type(),
|
|
&resolved_name,
|
|
&error_message)) {
|
|
LOG(ERROR) << error_message;
|
|
return nullptr;
|
|
}
|
|
|
|
auto factory = ModelRegistry::get_causalvlm_factory(resolved_name);
|
|
if (factory) {
|
|
return factory(context);
|
|
}
|
|
|
|
LOG(ERROR) << "Unsupported model type: "
|
|
<< context.get_model_args().model_type();
|
|
|
|
return nullptr;
|
|
}
|
|
|
|
std::unique_ptr<MMEmbeddingVLM> create_vlm_mm_embedding_model(
|
|
const ModelContext& context) {
|
|
std::string resolved_name;
|
|
std::string error_message;
|
|
if (!resolve_model_registration_name(context.get_model_args().model_type(),
|
|
&resolved_name,
|
|
&error_message)) {
|
|
LOG(ERROR) << error_message;
|
|
return nullptr;
|
|
}
|
|
|
|
auto factory = ModelRegistry::get_mm_embedding_vlm_factory(resolved_name);
|
|
if (factory) {
|
|
return factory(context);
|
|
}
|
|
|
|
LOG(ERROR) << "Unsupported model type: "
|
|
<< context.get_model_args().model_type();
|
|
|
|
return nullptr;
|
|
}
|
|
|
|
std::unique_ptr<DiTModel> create_dit_model(const DiTModelContext& context) {
|
|
// get the factory function for the model type from model registry
|
|
auto factory = ModelRegistry::get_dit_model_factory(context.model_type());
|
|
if (factory) {
|
|
return factory(context);
|
|
}
|
|
LOG(INFO) << "DiT Model type: " << context.model_type();
|
|
LOG(ERROR) << "Unsupported model type: " << context.model_type();
|
|
|
|
return nullptr;
|
|
}
|
|
|
|
} // namespace xllm
|