Files
project_6/upstream_ref/xllm/xllm/models/model_registry.cpp
EX Engine 002f9879b2 ref(upstream): FULL TREE — Deep-Spark xllm (1470) + ds_vllm csrc/models (703)
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
2026-08-10 02:54:03 +00:00

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 &registry;
}
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