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
This commit is contained in:
EX Engine
2026-08-10 02:53:54 +00:00
parent 9e4fb3712f
commit 002f9879b2
2179 changed files with 494021 additions and 79 deletions

View File

@@ -0,0 +1,133 @@
include(cc_binary)
include(cc_shared_library)
include_directories(.)
# Generate build-time version header from version.txt.
file(READ "${CMAKE_SOURCE_DIR}/version.txt" XLLM_BUILD_VERSION_RAW)
string(STRIP "${XLLM_BUILD_VERSION_RAW}" XLLM_BUILD_VERSION)
if(XLLM_BUILD_VERSION MATCHES "^v(.+)$")
set(XLLM_BUILD_VERSION "${CMAKE_MATCH_1}")
endif()
if(XLLM_BUILD_VERSION STREQUAL "")
set(XLLM_BUILD_VERSION "unknown")
endif()
file(MAKE_DIRECTORY "${CMAKE_CURRENT_BINARY_DIR}/core/common")
file(WRITE "${CMAKE_CURRENT_BINARY_DIR}/core/common/xllm_build_info.h"
"#pragma once\n\n#define XLLM_BUILD_VERSION \"${XLLM_BUILD_VERSION}\"\n")
include_directories("${CMAKE_CURRENT_BINARY_DIR}")
# Set warning-as-error for xllm code (third-party headers are marked as SYSTEM)
if(CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang")
add_compile_options(
$<$<COMPILE_LANGUAGE:CXX>:-Werror>
$<$<COMPILE_LANGUAGE:C>:-Werror>
)
if(CMAKE_CXX_COMPILER_ID MATCHES "Clang")
add_compile_options(
$<$<COMPILE_LANGUAGE:CXX>:-Wno-macro-redefined>
$<$<COMPILE_LANGUAGE:C>:-Wno-macro-redefined>
)
endif()
endif()
option(GENERATE_SO "Enable compile to generate .so" OFF)
add_subdirectory(api_service)
add_subdirectory(core)
add_subdirectory(function_call)
add_subdirectory(models)
add_subdirectory(parser)
add_subdirectory(processors)
add_subdirectory(proto)
add_subdirectory(pybind)
add_subdirectory(server)
if(GENERATE_SO)
message(STATUS "build libxllm.so and install")
cc_shared_library(
NAME
xllm
HDRS
c_api/llm.h
c_api/rec.h
c_api/default.h
c_api/types.h
c_api/internal/helper.h
SRCS
c_api/internal/llm.cpp
c_api/internal/rec.cpp
c_api/internal/helper.cpp
DEPS
:flags
:master
absl::strings
Boost::serialization
gflags::gflags
glog::glog
Folly::folly
nlohmann_json::nlohmann_json
)
else()
message(STATUS "build xllm binary and install")
cc_binary(
NAME
xllm
SRCS
xllm.cpp
DEPS
:flags
:xllm_server
:master
absl::strings
Boost::serialization
gflags::gflags
glog::glog
Folly::folly
nlohmann_json::nlohmann_json
)
endif()
# link brpc
target_link_libraries(xllm PRIVATE glog::glog brpc leveldb::leveldb protobuf::libprotobuf ${OpenCV_LIBS})
add_dependencies(xllm brpc-static)
if(USE_NPU)
set(COMMON_LIBS torch_npu torch_python Python::Python ascendcl atb_customize hccl c_sec nnopbase ms_tools_ext)
elseif(USE_MLU)
set(COMMON_LIBS Python::Python)
endif()
if(USE_MSPTI)
list(APPEND COMMON_LIBS mspti)
endif()
target_link_libraries(xllm PUBLIC ${COMMON_LIBS})
if (USE_MUSA)
target_link_libraries(xllm PUBLIC atomic musa_python torch_cpu c10)
endif()
# install xllm
install(TARGETS xllm RUNTIME DESTINATION ${CMAKE_INSTALL_BINDIR})
# Install all dependencies for xllm
install(CODE [[
file(GET_RUNTIME_DEPENDENCIES
RESOLVED_DEPENDENCIES_VAR DEPENDENCIES
UNRESOLVED_DEPENDENCIES_VAR UNRESOLVED_DEPENDENCIES
EXECUTABLES $<TARGET_FILE:xllm>)
file(INSTALL
DESTINATION "${CMAKE_INSTALL_PREFIX}/lib"
FILES ${DEPENDENCIES}
FOLLOW_SYMLINK_CHAIN)
# This should not be possible, but error out when a dependency cannot
# be resolved.
list(LENGTH UNRESOLVED_DEPENDENCIES UNRESOLVED_LENGTH)
if(${UNRESOLVED_LENGTH} GREATER 0)
message(FATAL_ERROR "Unresolved dependencies: ${UNRESOLVED_DEPENDENCIES}")
endif()
]])

View File

@@ -0,0 +1,86 @@
import importlib.util
import os
import sys
import sysconfig
def _get_python_version_tag() -> str:
# returns "310", "311", ...
return sysconfig.get_python_version().replace(".", "")
def _find_export_so_path() -> str:
pkg_dir = os.path.dirname(__file__)
pyver = _get_python_version_tag()
# Preferred, exact tags we build for today.
candidates = [
os.path.join(pkg_dir, f"xllm_export.cpython-{pyver}-x86_64-linux-gnu.so"),
os.path.join(pkg_dir, f"xllm_export.cpython-{pyver}-aarch64-linux-gnu.so"),
]
for p in candidates:
if os.path.exists(p):
return os.path.abspath(p)
# Fallback: accept any xllm_export*.so that got packaged (tag may differ).
for fname in os.listdir(pkg_dir):
if fname.startswith("xllm_export") and fname.endswith(".so"):
return os.path.abspath(os.path.join(pkg_dir, fname))
raise ImportError(
f"cannot find xllm_export shared library under {pkg_dir!r}. "
f"Expected one of: {candidates!r}"
)
_export_so_path = _find_export_so_path()
_spec = importlib.util.spec_from_file_location("xllm_export", _export_so_path)
if _spec is None or _spec.loader is None:
raise ImportError(f"failed to create import spec for xllm_export: {_export_so_path}")
# Make `import xllm_export` work for submodules (pybind/*) by loading and
# registering it before importing any modules that depend on it.
xllm_export = importlib.util.module_from_spec(_spec)
sys.modules["xllm_export"] = xllm_export
_spec.loader.exec_module(xllm_export)
from xllm.pybind.embedding import Embedding
from xllm.pybind.llm import LLM
try:
from xllm.pybind.vlm import VLM
except Exception:
VLM = None
from xllm.pybind.args import ArgumentParser
from xllm.pybind.params import SamplingParams, BeamSearchParams, PoolingParams
from xllm_export import (
LLMMaster,
VLMMaster,
Options,
RequestParams,
RequestOutput,
Usage,
SequenceOutput,
Status,
StatusCode,
MMType,
MMData,
)
__all__ = [
"ArgumentParser",
"Embedding",
"LLM",
"LLMMaster",
"VLM",
"VLMMaster",
"Options",
"SamplingParams",
"BeamSearchParams",
"PoolingParams",
"RequestParams",
"RequestOutput",
"Usage",
"SequenceOutput",
"Status",
"StatusCode",
]

View File

@@ -0,0 +1,58 @@
include(cc_library)
cc_library(
NAME
api_service
HDRS
api_service.h
api_service_impl.h
call.h
chat_json_parser.h
completion_service_impl.h
rec_completion_service_impl.h
chat_service_impl.h
anthropic_service_impl.h
sample_service_impl.h
embedding_service_impl.h
image_generation_service_impl.h
rerank_service_impl.h
qwen3_rerank_service_impl.h
non_stream_call.h
service_impl_factory.h
serving_mode.h
stream_call.h
models_service_impl.h
stream_output_parser.h
mm_service_utils.h
embedding_output_builder.h
utils.h
SRCS
api_service.cpp
chat_json_parser.cpp
service_impl_factory.cpp
call.cpp
completion_service_impl.cpp
rec_completion_service_impl.cpp
chat_service_impl.cpp
anthropic_service_impl.cpp
sample_service_impl.cpp
embedding_service_impl.cpp
image_generation_service_impl.cpp
models_service_impl.cpp
rerank_service_impl.cpp
stream_output_parser.cpp
qwen3_rerank_service_impl.cpp
embedding_output_builder.cpp
DEPS
:master
:chat_template
:util
glog::glog
proto::xllm_proto
absl::flat_hash_set
absl::random_random
:function_call
:reasoning
torch
$<$<BOOL:${USE_NPU}>:torch_npu>
)

View File

@@ -0,0 +1,764 @@
/* Copyright 2026 The xLLM 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 "anthropic_service_impl.h"
#include <absl/time/clock.h>
#include <absl/time/time.h>
#include <glog/logging.h>
#include <google/protobuf/util/json_util.h>
#include <cstdint>
#include <string>
#include <unordered_set>
#include "api_service/stream_output_parser.h"
#include "api_service/utils.h"
#include "core/common/types.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/framework/request/request_params.h"
#include "core/util/uuid.h"
#include "function_call/function_call.h"
namespace xllm {
namespace {
struct FunctionCallInfo {
std::string id = "";
std::string name = "";
std::string arguments = "";
};
struct ContentBlockInfo {
std::string normal_text = "";
std::vector<FunctionCallInfo> function_calls;
};
std::string convert_finish_reason_to_anthropic(
const std::string& finish_reason) {
if (finish_reason == "stop") {
return "end_turn";
} else if (finish_reason == "length") {
return "max_tokens";
} else if (finish_reason == "function_call") {
return "tool_use";
}
return "end_turn";
}
// Build messages from Anthropic protobuf request
std::vector<Message> build_messages(
const proto::AnthropicMessagesRequest& request) {
std::vector<Message> messages;
// Add system message if provided
if (request.has_system_string()) {
messages.emplace_back("system", request.system_string());
} else if (request.has_system_blocks()) {
std::string system_text;
for (const auto& block : request.system_blocks().blocks()) {
if (block.type() == "text" && block.has_text()) {
system_text += block.text();
}
}
if (!system_text.empty()) {
messages.emplace_back("system", system_text);
}
}
// Convert Anthropic messages to internal format
for (const auto& msg : request.messages()) {
const std::string& role = msg.role();
// Handle content - can be string or array of content blocks (oneof)
switch (msg.message_content_case()) {
case proto::AnthropicMessage::kContentString:
// Simple string content
messages.emplace_back(role, msg.content_string());
break;
case proto::AnthropicMessage::kContentBlocks: {
// Handle complex content blocks
std::vector<MMContent> content_parts;
Message::ToolCallVec tool_calls;
for (const auto& block : msg.content_blocks().blocks()) {
if (block.type() == "text" && block.has_text()) {
// Text content block
content_parts.emplace_back("text", block.text());
} else if (block.type() == "image" && block.has_source()) {
// Image content block - convert source to image_url
std::string image_url;
auto source_json = api_service::struct_to_json(block.source());
if (source_json.contains("data")) {
image_url = source_json["data"].get<std::string>();
}
content_parts.emplace_back("image_url", ImageURL{image_url});
} else if (block.type() == "tool_use") {
// Tool use block - convert to function call format
Message::ToolCall tool_call;
tool_call.id =
block.has_id()
? block.id()
: ("call_" +
std::to_string(absl::ToUnixSeconds(absl::Now())));
tool_call.type = "function";
tool_call.function.name = block.has_name() ? block.name() : "";
if (block.has_input()) {
tool_call.function.arguments =
api_service::struct_to_json(block.input()).dump();
} else {
tool_call.function.arguments = "{}";
}
tool_calls.emplace_back(std::move(tool_call));
} else if (block.type() == "tool_result") {
// Tool result block
if (role == "user") {
// User's tool result becomes a separate tool message
Message tool_msg("tool", "");
tool_msg.tool_call_id = block.has_id() ? block.id() : "";
if (block.has_content_string()) {
tool_msg.content = block.content_string();
}
messages.emplace_back(std::move(tool_msg));
} else {
// Assistant tool result becomes regular text
std::string tool_result_text = "Tool result: ";
if (block.has_content_string()) {
tool_result_text += block.content_string();
}
content_parts.emplace_back("text", tool_result_text);
}
}
}
if (!tool_calls.empty() || !content_parts.empty()) {
Message new_msg(role, "");
if (!tool_calls.empty()) {
new_msg.tool_calls = std::move(tool_calls);
}
if (!content_parts.empty()) {
if (content_parts.size() == 1 && content_parts[0].type == "text") {
// Single text content - use string directly
new_msg.content = content_parts[0].text;
} else {
// Multiple parts or non-text - use MMContentVec
new_msg.content = std::move(content_parts);
}
}
messages.emplace_back(std::move(new_msg));
}
break;
}
default:
break;
}
}
return messages;
}
// for non-streaming,
// generate chat response first and then convert to anthropic protobuf response.
void generate_chat_response(proto::ChatResponse& response,
const std::string& request_id,
const std::string& model,
const RequestOutput& req_output,
const std::string& tool_call_parser_format = "",
const std::string& reasoning_parser_format = "",
bool is_force_reasoning = false,
const std::vector<xllm::JsonTool>& tools = {}) {
response.set_object("chat.completion");
response.set_id(request_id);
response.set_model(model);
response.mutable_choices()->Reserve(req_output.outputs.size());
for (const auto& output : req_output.outputs) {
auto* choice = response.add_choices();
choice->set_index(output.index);
auto* message = choice->mutable_message();
message->set_role("assistant");
// 1) handle reasoning output
std::string cur_text = output.text;
if (!reasoning_parser_format.empty()) {
auto reasoning_parser = std::make_unique<ReasoningParser>(
reasoning_parser_format, false, is_force_reasoning);
auto result = reasoning_parser->parse_non_stream(cur_text);
if (result.normal_text.has_value()) {
cur_text = result.normal_text.value();
} else {
cur_text = "";
}
// set reasoning output
if (result.reasoning_text.has_value()) {
message->set_reasoning_content(result.reasoning_text.value());
}
}
// 2) handle tool call output
if (!tools.empty() && !tool_call_parser_format.empty() &&
!cur_text.empty()) {
auto* arena = response.GetArena();
auto result =
api_service::process_tool_calls(cur_text,
tools,
tool_call_parser_format,
output.finish_reason.value_or(""),
arena);
// set tool call output
message->mutable_content()->swap(result.text);
// set tool calls
if (result.tool_calls) {
auto& source_tool_calls = *result.tool_calls;
message->mutable_tool_calls()->Swap(&source_tool_calls);
}
// set finish reason
if (!result.finish_reason.empty()) {
choice->mutable_finish_reason()->swap(result.finish_reason);
}
} else {
// 3) handle text output
message->set_content(cur_text);
if (output.finish_reason.has_value()) {
choice->set_finish_reason(output.finish_reason.value());
}
}
}
// set usage
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
}
}
// for non-streaming,
// convert chat response to anthropic protobuf response
template <typename AnthropicCall>
bool send_result_to_client(std::shared_ptr<AnthropicCall> call,
const proto::ChatResponse& chat_response) {
auto& anthropic_response = call->response();
// Set basic fields
anthropic_response.set_id(chat_response.id());
anthropic_response.set_type("message");
anthropic_response.set_role("assistant");
anthropic_response.set_model(chat_response.model());
// Set usage
if (chat_response.has_usage()) {
auto* usage = anthropic_response.mutable_usage();
usage->set_input_tokens(chat_response.usage().prompt_tokens());
usage->set_output_tokens(chat_response.usage().completion_tokens());
}
// Process first choice
if (chat_response.choices_size() > 0) {
const auto& choice = chat_response.choices(0);
// set stop_reason
if (choice.has_finish_reason()) {
anthropic_response.set_stop_reason(std::move(
convert_finish_reason_to_anthropic(choice.finish_reason())));
}
// Add text content block
auto* text_block = anthropic_response.add_content();
text_block->set_type("text");
if (choice.has_message() && choice.message().has_content()) {
text_block->set_text(choice.message().content());
} else {
text_block->set_text("");
}
// Add tool_use blocks for each tool call
if (choice.has_message()) {
const auto& message = choice.message();
for (const auto& tool_call : message.tool_calls()) {
auto* tool_block = anthropic_response.add_content();
tool_block->set_type("tool_use");
tool_block->set_id(tool_call.id());
tool_block->set_name(tool_call.function().name());
// Parse arguments JSON string to Struct
if (!tool_call.function().arguments().empty()) {
google::protobuf::util::JsonStringToMessage(
tool_call.function().arguments(), tool_block->mutable_input());
}
}
}
}
return call->write_and_finish(anthropic_response);
}
// create a new content block like
// `<content_block_start>` ... `<content_block_stop>`
bool start_new_content_block(std::shared_ptr<AnthropicCall> call,
std::string& last_content_block_type,
const std::string& curr_content_block_type,
const ContentBlockInfo& content_block_info,
int& content_block_index) {
// if not the first content block,
// we need to create a content_block_stop
if (!last_content_block_type.empty()) {
proto::AnthropicStreamEvent stop_chunk;
stop_chunk.set_index(content_block_index);
stop_chunk.set_type("content_block_stop");
if (!call->write(stop_chunk.type(), stop_chunk)) {
LOG(ERROR) << "Failed to send content_block_stop event";
return false;
}
}
// update last_content_block_type
last_content_block_type = curr_content_block_type;
// create a new content block
proto::AnthropicStreamEvent content_start_event;
content_start_event.set_type("content_block_start");
content_start_event.set_index(++content_block_index);
auto* content_block = content_start_event.mutable_content_block();
content_block->set_type(curr_content_block_type);
if (curr_content_block_type == "text") {
content_block->set_text("");
} else if (curr_content_block_type == "tool_use") {
content_block->set_id(content_block_info.function_calls[0].id);
content_block->set_name(content_block_info.function_calls[0].name);
} else {
LOG(FATAL) << "Unknown content block type: " << curr_content_block_type;
}
if (!call->write(content_start_event.type(), content_start_event)) {
LOG(ERROR) << "Failed to send content_block_start event";
return false;
}
return true;
}
// send a content block delta content back
bool send_content_block_delta(std::shared_ptr<AnthropicCall> call,
std::string& last_content_block_type,
const std::string& curr_content_block_type,
const std::string& delta_type,
const ContentBlockInfo& content_block_info,
int& content_block_index) {
// counter new block or tool function call, we need a new content block
// <content_block_start> ... <content_block_stop>
if (last_content_block_type != curr_content_block_type ||
(delta_type == "tool_use_delta" &&
!content_block_info.function_calls[0].name.empty())) {
// try to create new content block
if (!start_new_content_block(call,
last_content_block_type,
curr_content_block_type,
content_block_info,
content_block_index)) {
return false;
}
}
proto::AnthropicStreamEvent chunk;
chunk.set_index(content_block_index);
chunk.set_type("content_block_delta");
auto* delta = chunk.mutable_delta();
if (delta_type == "text_delta") {
delta->set_type("text_delta");
delta->set_text(content_block_info.normal_text);
} else if (delta_type == "tool_use_delta") {
delta->set_type("input_json_delta");
if (!content_block_info.function_calls.empty() &&
!content_block_info.function_calls[0].arguments.empty()) {
delta->set_partial_json(content_block_info.function_calls[0].arguments);
}
} else {
LOG(FATAL) << "Unknown delta type: " << delta_type;
}
if (!call->write(chunk.type(), chunk)) {
LOG(ERROR) << "Failed to send content_block_delta event";
return false;
}
return true;
}
// for streaming,
// process tool call stream and send content block delta back
bool process_tool_call_stream(std::shared_ptr<AnthropicCall> call,
std::string& last_content_block_type,
int& content_block_index,
std::shared_ptr<StreamOutputParser> stream_parser,
size_t index,
const std::string& delta) {
auto* parser = stream_parser->get_tool_call_parser(index);
if (!parser) {
return true;
}
auto parse_result = parser->parse_streaming_increment(delta);
if (!parse_result.normal_text.empty()) {
ContentBlockInfo content_block_info;
content_block_info.normal_text = parse_result.normal_text;
if (!send_content_block_delta(call,
last_content_block_type,
"text",
"text_delta",
content_block_info,
content_block_index)) {
return false;
}
}
for (const auto& call_item : parse_result.calls) {
stream_parser->set_has_tool_call(index, true);
std::string tool_call_id;
std::string function_name;
if (call_item.name.has_value()) {
tool_call_id = function_call::utils::generate_tool_call_id();
function_name = call_item.name.value();
}
ContentBlockInfo content_block_info;
content_block_info.function_calls.emplace_back(FunctionCallInfo{
.id = tool_call_id,
.name = function_name,
.arguments = call_item.parameters,
});
if (!send_content_block_delta(call,
last_content_block_type,
"tool_use",
"tool_use_delta",
content_block_info,
content_block_index)) {
return false;
}
}
return true;
}
// for streaming,
// send stream delta content to client
bool send_delta_to_client(
std::shared_ptr<AnthropicCall> call,
ContentBlockInfo& content_block_info,
bool& content_block_started,
int& content_block_index,
std::string& last_content_block_type,
const std::string& request_id,
const std::string& model,
const RequestOutput& output,
std::shared_ptr<StreamOutputParser> stream_parser = nullptr) {
if (stream_parser && output.outputs.size() > 0) {
stream_parser->check_resize_for_index(output.outputs.size() - 1);
}
std::string finish_reason = "";
for (const auto& seq_output : output.outputs) {
const auto& index = seq_output.index;
std::string cur_text = seq_output.text;
// 1) Handle reasoning text
if (!cur_text.empty() && stream_parser && stream_parser->is_reasoning()) {
auto parser = stream_parser->get_reasoning_parser(index);
auto result = parser->parse_stream_chunk(cur_text);
if (result.normal_text.has_value()) {
cur_text = result.normal_text.value();
} else {
cur_text = "";
}
if (result.reasoning_text.has_value()) {
ContentBlockInfo content_block_info;
content_block_info.normal_text = result.reasoning_text.value();
if (!send_content_block_delta(call,
last_content_block_type,
"text",
"text_delta",
content_block_info,
content_block_index)) {
return false;
}
}
}
if (!cur_text.empty()) {
// 2) Handle tool call: text or tool_use
if (stream_parser && stream_parser->is_tool_call()) {
if (!process_tool_call_stream(call,
last_content_block_type,
content_block_index,
stream_parser,
index,
cur_text)) {
return false;
}
} else {
// 3) Handle text output
ContentBlockInfo content_block_info;
content_block_info.normal_text = cur_text;
if (!send_content_block_delta(call,
last_content_block_type,
"text",
"text_delta",
content_block_info,
content_block_index)) {
return false;
}
}
}
// Handle finish reason
if (seq_output.finish_reason.has_value()) {
// Check for unstreamed tool args before sending finish reason
if (stream_parser && stream_parser->get_has_tool_call(index)) {
auto send_func = [&](const std::string& arguments,
int tool_index) -> bool {
ContentBlockInfo content_block_info;
content_block_info.function_calls.push_back(FunctionCallInfo{
.arguments = arguments,
});
return send_content_block_delta(call,
last_content_block_type,
"tool_use",
"tool_use_delta",
content_block_info,
content_block_index);
};
if (!api_service::check_for_unstreamed_tool_args(
stream_parser, index, send_func)) {
return false;
}
}
finish_reason = seq_output.finish_reason.value();
}
}
// 4) finish request, we need to send the
// last `content_block_stop` and `message_delta` event
if (output.finished || output.cancelled) {
if (output.finished) {
finish_reason = convert_finish_reason_to_anthropic(finish_reason);
} else {
finish_reason = "stop";
}
// if content_block_index < 0, means no content block started
// so we don't need to send content_block_stop event
if (content_block_index >= 0) {
proto::AnthropicStreamEvent stop_chunk;
stop_chunk.set_index(content_block_index);
stop_chunk.set_type("content_block_stop");
if (!call->write(stop_chunk.type(), stop_chunk)) {
LOG(ERROR) << "Failed to send content_block_stop event";
return false;
}
}
// send message_delta event for the last message
proto::AnthropicStreamEvent message_delta;
message_delta.set_type("message_delta");
auto* delta = message_delta.mutable_delta();
delta->set_stop_reason(finish_reason);
// Set usage information
if (output.usage.has_value()) {
auto* usage = message_delta.mutable_usage();
usage->set_input_tokens(output.usage.value().num_prompt_tokens);
usage->set_output_tokens(output.usage.value().num_generated_tokens);
} else {
auto* usage = message_delta.mutable_usage();
usage->set_input_tokens(0);
usage->set_output_tokens(0);
}
if (!call->write(message_delta.type(), message_delta)) {
LOG(ERROR) << "Failed to send message_delta event";
return false;
}
// send message_stop event
proto::AnthropicStreamEvent stop_message;
stop_message.set_type("message_stop");
if (!call->write(stop_message.type(), stop_message)) {
LOG(ERROR) << "Failed to send message_stop event";
return false;
}
return call->finish();
}
return true;
}
} // namespace
AnthropicServiceImpl::AnthropicServiceImpl(
LLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models),
master_(master),
tool_call_parser_format_(
master_->options().tool_call_parser().value_or("")),
reasoning_parser_format_(
master_->options().reasoning_parser().value_or("")) {
CHECK(master_ != nullptr);
}
void AnthropicServiceImpl::process_async_impl(
std::shared_ptr<AnthropicCall> call) {
const auto& rpc_request = call->request();
const auto& model = rpc_request.model();
// Check if model is supported
if (!models_.contains(model)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
CHECK(master_ != nullptr);
// Check rate limit
if (master_->get_rate_limiter()->is_limited()) {
call->finish_with_error(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.");
return;
}
// Build request parameters
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
// Build messages
std::vector<Message> messages = build_messages(rpc_request);
// Create stream parser if needed
std::shared_ptr<StreamOutputParser> stream_parser;
if (request_params.streaming && (!tool_call_parser_format_.empty() ||
!reasoning_parser_format_.empty())) {
stream_parser =
std::make_shared<StreamOutputParser>(request_params.tools,
tool_call_parser_format_,
reasoning_parser_format_,
false /*is_force_reasoning_*/);
CHECK(stream_parser != nullptr) << "create StreamOutputParser failed!";
}
auto saved_streaming = request_params.streaming;
auto message_id = request_params.request_id;
auto saved_tools = request_params.tools;
// Handle request
master_->handle_request(
std::move(messages),
std::nullopt,
std::move(request_params),
call.get(),
[call,
model,
master = master_,
stream = saved_streaming,
message_id = std::move(message_id),
message_started = false,
content_block_started = false,
content_block_index = -1,
last_content_block_type = std::string{},
tools = std::move(saved_tools),
tool_call_parser_format = tool_call_parser_format_,
reasoning_parser_format = reasoning_parser_format_,
stream_parser =
stream_parser](const RequestOutput& req_output) mutable -> bool {
// Handle errors
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
master->get_rate_limiter()->decrease_one_request();
return call->finish_with_error(status.code(), status.message());
}
}
// Decrease rate limiter on completion
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
// Anthropic format:
//
// event: message_start
// event: content_block_start
// event: content_block_delta (may multiple times)
// event: content_block_stop
// event: message_delta (only once, at the end)
// event: message_stop
// data: [DONE]
if (stream) {
// 1. Send `message_start` event
if (!message_started) {
message_started = true;
proto::AnthropicStreamEvent start_event;
start_event.set_type("message_start");
auto* start_message = start_event.mutable_message();
start_message->set_id(message_id);
start_message->set_type("message");
start_message->set_role("assistant");
start_message->set_model(model);
auto* usage = start_message->mutable_usage();
usage->set_input_tokens(0);
usage->set_output_tokens(0);
if (!call->write(start_event.type(), start_event)) {
return false;
}
}
ContentBlockInfo content_block_info;
return send_delta_to_client(call,
content_block_info,
content_block_started,
content_block_index,
last_content_block_type,
message_id,
model,
req_output,
stream_parser);
}
// handle non-streaming response
proto::ChatResponse chat_response;
generate_chat_response(chat_response,
message_id,
model,
req_output,
tool_call_parser_format,
reasoning_parser_format,
false /*is_force_reasoning_*/,
tools);
return send_result_to_client(call, chat_response);
});
}
} // namespace xllm

View File

@@ -0,0 +1,45 @@
/* Copyright 2026 The xLLM 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.
==============================================================================*/
#pragma once
#include <type_traits>
#include "anthropic.pb.h"
#include "api_service/api_service_impl.h"
#include "api_service/stream_call.h"
namespace xllm {
// Specialize is_stream_call for AnthropicCall to recognize it as a stream call
template <>
struct is_stream_call<AnthropicCall> : std::true_type {};
class AnthropicServiceImpl final : public APIServiceImpl<AnthropicCall> {
public:
AnthropicServiceImpl(LLMMaster* master,
const std::vector<std::string>& models);
void process_async_impl(std::shared_ptr<AnthropicCall> call) override;
private:
DISALLOW_COPY_AND_ASSIGN(AnthropicServiceImpl);
LLMMaster* master_ = nullptr;
const std::string tool_call_parser_format_;
const std::string reasoning_parser_format_;
};
} // namespace xllm

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,211 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <functional>
#include <shared_mutex>
#include <string>
#include <unordered_map>
#include "anthropic_service_impl.h"
#include "chat_service_impl.h"
#include "completion_service_impl.h"
#include "embedding_service_impl.h"
#include "image_generation_service_impl.h"
#include "models_service_impl.h"
#include "qwen3_rerank_service_impl.h"
#include "rec_completion_service_impl.h"
#include "rerank_service_impl.h"
#include "sample_service_impl.h"
#include "xllm_service.pb.h"
namespace xllm {
class ClosureGuard;
class ServiceImplFactory;
class APIService : public proto::XllmAPIService {
friend class ServiceImplFactory;
public:
APIService(Master* master,
const std::vector<std::string>& model_names,
const std::vector<std::string>& model_versions);
~APIService() = default;
void Completions(::google::protobuf::RpcController* controller,
const proto::CompletionRequest* request,
proto::CompletionResponse* response,
::google::protobuf::Closure* done) override;
void CompletionsHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void Sample(::google::protobuf::RpcController* controller,
const proto::SampleRequest* request,
proto::SampleResponse* response,
::google::protobuf::Closure* done) override;
void SampleHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void ChatCompletions(::google::protobuf::RpcController* controller,
const proto::ChatRequest* request,
proto::ChatResponse* response,
::google::protobuf::Closure* done) override;
void ChatCompletionsHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void Embeddings(::google::protobuf::RpcController* controller,
const proto::EmbeddingRequest* request,
proto::EmbeddingResponse* response,
::google::protobuf::Closure* done) override;
void EmbeddingsHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void ImageGeneration(::google::protobuf::RpcController* controller,
const proto::ImageGenerationRequest* request,
proto::ImageGenerationResponse* response,
::google::protobuf::Closure* done) override;
void ImageGenerationHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void Rerank(::google::protobuf::RpcController* controller,
const proto::RerankRequest* request,
proto::RerankResponse* response,
::google::protobuf::Closure* done) override;
void RerankHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void Models(::google::protobuf::RpcController* controller,
const proto::ModelListRequest* request,
proto::ModelListResponse* response,
::google::protobuf::Closure* done) override;
void ModelsHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void ModelVersionsHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void AnthropicMessagesHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void ForkMaster(::google::protobuf::RpcController* controller,
const proto::MasterInfos* request,
proto::Status* response,
::google::protobuf::Closure* done) override;
void ForkMasterHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void Sleep(::google::protobuf::RpcController* controller,
const proto::MasterInfos* request,
proto::Status* response,
::google::protobuf::Closure* done) override;
void SleepHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void Wakeup(::google::protobuf::RpcController* controller,
const proto::MasterInfos* request,
proto::Status* response,
::google::protobuf::Closure* done) override;
void WakeupHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void LinkD2D(::google::protobuf::RpcController* controller,
const proto::D2DLinkRequest* request,
proto::Status* response,
::google::protobuf::Closure* done) override;
void LinkD2DHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
void UnlinkD2D(::google::protobuf::RpcController* controller,
const proto::D2DLinkRequest* request,
proto::Status* response,
::google::protobuf::Closure* done) override;
void UnlinkD2DHttp(::google::protobuf::RpcController* controller,
const proto::HttpRequest* request,
proto::HttpResponse* response,
::google::protobuf::Closure* done) override;
private:
using ChatHttpHandler = std::function<void(ClosureGuard&,
brpc::Controller*,
const proto::HttpRequest*,
proto::HttpResponse*)>;
void register_chat_completions_handler();
bool ParseForkMasterRequest(const proto::MasterInfos* request,
Options& options);
void set_model_master(const std::string& model_id, Master* master);
bool has_model_master(const std::string& model_id) const;
bool add_model_master_if_absent(const std::string& model_id, Master* master);
Master* get_model_master(const std::string& model_id) const;
Master* master_;
ChatHttpHandler chat_completions_handler_;
mutable std::shared_mutex masters_mutex_;
std::unordered_map<std::string, Master*> masters_;
std::unique_ptr<AnthropicServiceImpl> anthropic_service_impl_;
std::unique_ptr<CompletionServiceImpl> completion_service_impl_;
std::unique_ptr<SampleServiceImpl> sample_service_impl_;
std::unique_ptr<ChatServiceImpl> chat_service_impl_;
std::unique_ptr<MMChatServiceImpl> mm_chat_service_impl_;
std::unique_ptr<EmbeddingServiceImpl> embedding_service_impl_;
std::unique_ptr<MMEmbeddingServiceImpl> mm_embedding_service_impl_;
std::unique_ptr<ModelsServiceImpl> models_service_impl_;
std::unique_ptr<ImageGenerationServiceImpl> image_generation_service_impl_;
std::unique_ptr<RerankServiceImpl> rerank_service_impl_;
std::unique_ptr<RecCompletionServiceImpl> rec_completion_service_impl_;
};
} // namespace xllm

View File

@@ -0,0 +1,56 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <absl/container/flat_hash_set.h>
#include <memory>
#include "call.h"
#include "chat.pb.h"
#include "completion.pb.h"
#include "core/common/macros.h"
#include "core/distributed_runtime/llm_master.h"
namespace xllm {
template <typename CallType>
class APIServiceImpl {
public:
using RequestType = typename CallType::ReqType;
using ResponseType = typename CallType::ResType;
APIServiceImpl(const std::vector<std::string>& models)
: models_(models.begin(), models.end()) {
CHECK(!models_.empty());
}
virtual ~APIServiceImpl() = default;
void process_async(std::shared_ptr<Call> call) {
std::shared_ptr<CallType> call_cast =
std::dynamic_pointer_cast<CallType>(call);
process_async_impl(call_cast);
}
virtual void process_async_impl(std::shared_ptr<CallType> call) = 0;
virtual void process_async_rpc_impl(const RequestType* request) {
NOT_IMPLEMENTED();
}
protected:
absl::flat_hash_set<std::string> models_;
};
} // namespace xllm

View File

@@ -0,0 +1,64 @@
/* Copyright 2025 The xLLM 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 "call.h"
#include "core/common/constants.h"
namespace xllm {
Call::Call(brpc::Controller* controller) : controller_(controller) { init(); }
void Call::init() {
if (controller_->http_request().GetHeader("x-request-id")) {
x_request_id_ = *controller_->http_request().GetHeader("x-request-id");
} else if (controller_->http_request().GetHeader("x-ms-client-request-id")) {
x_request_id_ =
*controller_->http_request().GetHeader("x-ms-client-request-id");
}
if (controller_->http_request().GetHeader("x-request-time")) {
x_request_time_ = *controller_->http_request().GetHeader("x-request-time");
} else if (controller_->http_request().GetHeader("x-request-timems")) {
x_request_time_ =
*controller_->http_request().GetHeader("x-request-timems");
}
init_request_payload();
}
void Call::init_request_payload() {
const auto infer_content_len =
controller_->http_request().GetHeader(kInferContentLength);
const auto content_len =
controller_->http_request().GetHeader(kContentLength);
if (infer_content_len == nullptr || content_len == nullptr) return;
auto infer_len = std::stoul(*infer_content_len);
auto len = std::stoul(*content_len);
if (infer_len > len) {
LOG(ERROR) << " content length is invalid:"
<< " infer content len is " << infer_len
<< " , content length is " << len;
return;
}
controller_->request_attachment().copy_to(
&request_payload_, len - infer_len, infer_len);
}
} // namespace xllm

View File

@@ -0,0 +1,49 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <brpc/controller.h>
#include <string>
namespace xllm {
class Call {
public:
Call(brpc::Controller* controller);
virtual ~Call() = default;
std::string get_x_request_id() { return x_request_id_; }
std::string get_x_request_time() { return x_request_time_; }
std::string take_request_payload() { return std::move(request_payload_); }
void init_request_payload();
virtual bool is_disconnected() const = 0;
protected:
void init();
protected:
brpc::Controller* controller_;
std::string x_request_id_;
std::string x_request_time_;
std::string request_payload_;
};
} // namespace xllm

View File

@@ -0,0 +1,164 @@
/* Copyright 2026 The xLLM 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 "api_service/chat_json_parser.h"
#include <glog/logging.h>
#include <nlohmann/json.hpp>
namespace xllm {
const ChatJsonParser& ChatJsonParser::get(ServingMode mode) {
if (mode == ServingMode::VLM) {
static const VlmChatJsonParser k_vlm_parser;
return k_vlm_parser;
}
static const LlmChatJsonParser k_llm_parser;
return k_llm_parser;
}
const ChatJsonParser& ChatJsonParser::anthropic() {
static const AnthropicChatJsonParser k_anthropic_parser;
return k_anthropic_parser;
}
std::pair<Status, std::string> VlmChatJsonParser::preprocess(
std::string json_str) const {
return {Status(), std::move(json_str)};
}
std::pair<Status, std::string> LlmChatJsonParser::preprocess(
std::string json_str) const {
try {
auto json = nlohmann::json::parse(json_str);
if (!json.contains("messages") || !json["messages"].is_array()) {
return {Status(), std::move(json_str)};
}
bool modified = false;
for (auto& msg : json["messages"]) {
if (!msg.is_object()) {
return {Status(StatusCode::INVALID_ARGUMENT,
"Message in 'messages' array must be an object."),
""};
}
if (msg.contains("content") && msg["content"].is_array()) {
for (const auto& item : msg["content"]) {
if (!item.is_object()) {
return {Status(StatusCode::INVALID_ARGUMENT,
"Content array item must be an object."),
""};
}
if (!item.contains("type") || item["type"] != "text") {
return {Status(StatusCode::INVALID_ARGUMENT,
"Non-text content (e.g., image_url) requires "
"multimodal backend (-backend vlm)"),
""};
}
if (!item.contains("text") || !item["text"].is_string()) {
return {Status(StatusCode::INVALID_ARGUMENT,
"Missing or invalid 'text' field in content item."),
""};
}
}
size_t total_size = 0;
size_t num_items = msg["content"].size();
for (const auto& item : msg["content"]) {
total_size += item["text"].get_ref<const std::string&>().size();
}
if (num_items > 1) {
total_size += num_items - 1;
}
std::string combined_text;
combined_text.reserve(total_size);
bool first = true;
for (const auto& item : msg["content"]) {
if (!first) {
combined_text += '\n';
}
combined_text += item["text"].get_ref<const std::string&>();
first = false;
}
msg["content"] = combined_text;
modified = true;
}
}
return modified ? std::make_pair(Status(), json.dump())
: std::make_pair(Status(), std::move(json_str));
} catch (const nlohmann::json::exception& e) {
return {Status(StatusCode::INVALID_ARGUMENT,
"Invalid JSON format: " + std::string(e.what())),
""};
} catch (const std::exception& e) {
LOG(ERROR) << "Exception during JSON preprocessing: " << e.what();
return {Status(StatusCode::UNKNOWN,
"Internal server error during JSON processing."),
""};
}
}
std::pair<Status, std::string> AnthropicChatJsonParser::preprocess(
std::string json_str) const {
try {
auto j = nlohmann::json::parse(json_str);
if (j.contains("messages") && j["messages"].is_array()) {
for (auto& msg : j["messages"]) {
if (!msg.contains("content")) {
continue;
}
auto& content = msg["content"];
if (content.is_string()) {
msg["content_string"] = content.get<std::string>();
msg.erase("content");
} else if (content.is_array()) {
nlohmann::json content_blocks;
content_blocks["blocks"] = content;
msg["content_blocks"] = content_blocks;
msg.erase("content");
}
}
}
if (j.contains("system")) {
auto& system = j["system"];
if (system.is_string()) {
j["system_string"] = system.get<std::string>();
j.erase("system");
} else if (system.is_array()) {
nlohmann::json system_blocks;
system_blocks["blocks"] = system;
j["system_blocks"] = system_blocks;
j.erase("system");
}
}
return {Status(), j.dump()};
} catch (const nlohmann::json::exception& e) {
return {Status(StatusCode::INVALID_ARGUMENT,
"Invalid JSON format: " + std::string(e.what())),
""};
} catch (const std::exception& e) {
LOG(ERROR) << "Exception during Anthropic JSON preprocessing: " << e.what();
return {Status(StatusCode::UNKNOWN,
"Internal server error during JSON processing."),
""};
}
}
} // namespace xllm

View File

@@ -0,0 +1,68 @@
/* Copyright 2026 The xLLM 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.
==============================================================================*/
#pragma once
#include <string>
#include <utility>
#include "api_service/serving_mode.h"
#include "core/common/types.h"
namespace xllm {
// Normalizes OpenAI-style chat JSON before protobuf parsing. LLM backends
// collapse text-only content arrays into a single string; VLM backends pass
// JSON through for downstream multimodal handling.
class ChatJsonParser {
public:
virtual ~ChatJsonParser() = default;
[[nodiscard]] virtual std::pair<Status, std::string> preprocess(
std::string json_str) const = 0;
// Returns the singleton parser for the given serving mode.
// LLM/REC → LlmChatJsonParser, VLM → VlmChatJsonParser.
static const ChatJsonParser& get(ServingMode mode);
// Returns the Anthropic protocol parser (separate from serving mode).
static const ChatJsonParser& anthropic();
};
// Text-only backend: combines array content items of type "text" into one
// string; rejects non-text parts (e.g. image_url).
class LlmChatJsonParser final : public ChatJsonParser {
public:
std::pair<Status, std::string> preprocess(
std::string json_str) const override;
};
// Multimodal backend: no preprocessing; array content stays as-is.
class VlmChatJsonParser final : public ChatJsonParser {
public:
std::pair<Status, std::string> preprocess(
std::string json_str) const override;
};
// Anthropic Messages API: remaps "content" (string|array) to
// "content_string"/"content_blocks" and "system" (string|array) to
// "system_string"/"system_blocks" for protobuf compatibility.
class AnthropicChatJsonParser final : public ChatJsonParser {
public:
std::pair<Status, std::string> preprocess(
std::string json_str) const override;
};
} // namespace xllm

View File

@@ -0,0 +1,979 @@
/* 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 "chat_service_impl.h"
#include <absl/strings/escaping.h>
#include <absl/time/clock.h>
#include <absl/time/time.h>
#include <glog/logging.h>
#include <google/protobuf/util/json_util.h>
#include <torch/torch.h>
#include <boost/algorithm/string.hpp>
#include <cstdint>
#include <functional>
#include <string>
#include <unordered_set>
#include "api_service/stream_output_parser.h"
#include "api_service/utils.h"
#include "core/common/instance_name.h"
#include "core/common/types.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/distributed_runtime/rec_master.h"
#include "core/distributed_runtime/vlm_master.h"
#include "core/framework/request/rec_type.h"
#include "core/framework/request/request_params.h"
#include "core/util/utils.h"
#include "core/util/uuid.h"
#include "mm_service_utils.h"
namespace xllm {
namespace {
void set_logprobs(proto::ChatChoice* choice,
const std::optional<std::vector<LogProb>>& logprobs) {
if (!logprobs.has_value() || logprobs.value().empty()) {
return;
}
auto* proto_logprobs = choice->mutable_logprobs();
proto_logprobs->mutable_content()->Reserve(logprobs.value().size());
for (const auto& logprob : logprobs.value()) {
auto* logprob_proto = proto_logprobs->add_content();
logprob_proto->set_token(logprob.token);
logprob_proto->set_token_id(logprob.token_id);
logprob_proto->set_logprob(logprob.logprob);
if (logprob.top_logprobs.has_value()) {
for (const auto& top_logprob : logprob.top_logprobs.value()) {
auto* top_logprob_proto = logprob_proto->add_top_logprobs();
top_logprob_proto->set_token(top_logprob.token);
top_logprob_proto->set_token_id(top_logprob.token_id);
top_logprob_proto->set_logprob(top_logprob.logprob);
}
}
}
}
template <typename ChatCall>
bool send_tool_call_chunk(std::shared_ptr<ChatCall> call,
size_t index,
const std::string& tool_call_id,
const std::string& function_name,
const std::string& arguments,
int tool_index,
const std::string& request_id,
int64_t created_time,
const std::string& model) {
auto& response = call->response();
response.Clear();
response.set_object("chat.completion.chunk");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(index);
auto* delta = choice->mutable_delta();
auto* tool_call = delta->add_tool_calls();
if (!tool_call_id.empty()) {
tool_call->set_id(tool_call_id);
}
tool_call->set_index(tool_index);
tool_call->set_type("function");
auto* function = tool_call->mutable_function();
if (!function_name.empty()) {
function->set_name(function_name);
}
if (!arguments.empty()) {
function->set_arguments(arguments);
}
return call->write(response);
}
template <typename ChatCall>
bool send_normal_text_chunk(std::shared_ptr<ChatCall> call,
size_t index,
const std::string& content,
const std::string& request_id,
int64_t created_time,
const std::string& model) {
auto& response = call->response();
response.Clear();
response.set_object("chat.completion.chunk");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(index);
auto* delta = choice->mutable_delta();
delta->set_content(content);
return call->write(response);
}
template <typename ChatCall>
bool send_reasoning_text_chunk(std::shared_ptr<ChatCall> call,
size_t index,
const std::string& content,
const std::string& request_id,
int64_t created_time,
const std::string& model) {
auto& response = call->response();
response.Clear();
response.set_object("chat.completion.chunk");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(index);
auto* delta = choice->mutable_delta();
delta->set_reasoning_content(content);
return call->write(response);
}
template <typename ChatCall>
bool process_tool_call_stream(std::shared_ptr<ChatCall> call,
std::shared_ptr<StreamOutputParser> stream_parser,
size_t index,
const std::string& delta,
const std::string& request_id,
int64_t created_time,
const std::string& model) {
auto* parser = stream_parser->get_tool_call_parser(index);
if (!parser) {
return true;
}
auto parse_result = parser->parse_streaming_increment(delta);
if (!parse_result.normal_text.empty()) {
if (!send_normal_text_chunk(call,
index,
parse_result.normal_text,
request_id,
created_time,
model)) {
return false;
}
}
for (const auto& call_item : parse_result.calls) {
stream_parser->set_has_tool_call(index, true);
std::string tool_call_id;
std::string function_name;
if (call_item.name.has_value()) {
tool_call_id = function_call::utils::generate_tool_call_id();
function_name = call_item.name.value();
}
if (!send_tool_call_chunk(call,
index,
tool_call_id,
function_name,
call_item.parameters,
call_item.tool_index,
request_id,
created_time,
model)) {
return false;
}
}
return true;
}
bool get_enable_thinking_from_request(
const nlohmann::json& chat_template_kwargs,
const std::string& reasoning_parser_format) {
// Default to true if reasoning_parser is configured
// This matches chat template defaults for models like glm47 that enable
// thinking by default
bool default_value = !reasoning_parser_format.empty();
if (chat_template_kwargs.empty()) {
return default_value;
}
// Check for explicit enable_thinking or thinking setting
// qwen3 and glm45/glm47 use enable_thinking and deepseek-v3 uses thinking
if (chat_template_kwargs.contains("enable_thinking") ||
chat_template_kwargs.contains("thinking")) {
auto get_bool_val = [&](const char* key) {
auto it = chat_template_kwargs.find(key);
if (it != chat_template_kwargs.end() && it->is_boolean()) {
return it->get<bool>();
}
return false;
};
return get_bool_val("enable_thinking") || get_bool_val("thinking");
}
return default_value;
}
template <typename ChatCall>
bool send_delta_to_client_brpc(
std::shared_ptr<ChatCall> call,
bool include_usage,
std::unordered_set<size_t>* first_message_sent,
const std::string& request_id,
int64_t created_time,
const std::string& model,
const RequestOutput& output,
std::shared_ptr<StreamOutputParser> stream_parser = nullptr) {
auto& response = call->response();
if (stream_parser && output.outputs.size() > 0) {
stream_parser->check_resize_for_index(output.outputs.size() - 1);
}
// send delta to client
for (const auto& seq_output : output.outputs) {
const auto& index = seq_output.index;
std::string cur_text = seq_output.text;
if (first_message_sent->find(index) == first_message_sent->end()) {
response.Clear();
response.set_object("chat.completion.chunk");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(index);
auto* message = choice->mutable_delta();
message->set_role("assistant");
message->set_content("");
first_message_sent->insert(index);
if (!call->write(response)) {
return false;
}
}
// Handle reasoning text
if (!cur_text.empty()) {
if (stream_parser && stream_parser->is_reasoning()) {
auto parser = stream_parser->get_reasoning_parser(index);
auto result = parser->parse_stream_chunk(cur_text);
if (result.normal_text.has_value()) {
cur_text = result.normal_text.value();
} else {
cur_text = "";
}
if (result.reasoning_text.has_value()) {
send_reasoning_text_chunk(call,
index,
result.reasoning_text.value(),
request_id,
created_time,
model);
}
}
}
if (!cur_text.empty()) {
// Handle tool call text
if (stream_parser && stream_parser->is_tool_call()) {
if (!process_tool_call_stream(call,
stream_parser,
index,
cur_text,
request_id,
created_time,
model)) {
return false;
}
} else {
response.Clear();
response.set_object("chat.completion.chunk");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(index);
set_logprobs(choice, seq_output.logprobs);
auto* message = choice->mutable_delta();
message->set_content(cur_text);
if (!call->write(response)) {
return false;
}
}
}
// Handle finish reason
if (seq_output.finish_reason.has_value()) {
// Check for unstreamed tool args before sending finish reason
if (stream_parser && stream_parser->get_has_tool_call(index)) {
auto send_func = [&](const std::string& arguments, int tool_index) {
return send_tool_call_chunk(call,
index,
"",
"",
arguments,
tool_index,
request_id,
created_time,
model);
};
if (!api_service::check_for_unstreamed_tool_args(
stream_parser, index, send_func)) {
return false;
}
}
response.Clear();
response.set_object("chat.completion.chunk");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(index);
choice->mutable_delta();
if (stream_parser && stream_parser->get_has_tool_call(index) &&
seq_output.finish_reason.value() == "stop") {
choice->set_finish_reason("tool_calls");
} else {
choice->set_finish_reason(std::move(seq_output.finish_reason.value()));
}
if (!call->write(response)) {
return false;
}
}
}
if (include_usage && output.usage.has_value()) {
response.Clear();
const auto& usage = output.usage.value();
response.set_object("chat.completion.chunk");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
if (!call->write(response)) {
return false;
}
}
if (output.finished || output.cancelled) {
response.Clear();
return call->finish();
}
return true;
}
template <typename ChatCall>
bool send_result_to_client_brpc(std::shared_ptr<ChatCall> call,
const std::string& request_id,
int64_t created_time,
const std::string& model,
const RequestOutput& req_output,
const std::string& tool_call_parser_format = "",
const std::string& reasoning_parser_format = "",
bool is_force_reasoning = false,
const std::vector<xllm::JsonTool>& tools = {}) {
auto& response = call->response();
response.set_object("chat.completion");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
response.mutable_choices()->Reserve(req_output.outputs.size());
for (const auto& output : req_output.outputs) {
auto* choice = response.add_choices();
choice->set_index(output.index);
set_logprobs(choice, output.logprobs);
auto* message = choice->mutable_message();
message->set_role("assistant");
// handle reasoning output
std::string cur_text = output.text;
if (!reasoning_parser_format.empty()) {
auto reasoning_parser = std::make_unique<ReasoningParser>(
reasoning_parser_format, false, is_force_reasoning);
auto result = reasoning_parser->parse_non_stream(cur_text);
if (result.normal_text.has_value()) {
cur_text = result.normal_text.value();
} else {
cur_text = "";
}
if (result.reasoning_text.has_value()) {
message->set_reasoning_content(result.reasoning_text.value());
}
}
// handle tool call output
if (!tools.empty() && !tool_call_parser_format.empty() &&
!cur_text.empty()) {
auto* arena = response.GetArena();
auto result =
api_service::process_tool_calls(cur_text,
tools,
tool_call_parser_format,
output.finish_reason.value_or(""),
arena);
message->mutable_content()->swap(result.text);
if (result.tool_calls) {
auto& source_tool_calls = *result.tool_calls;
message->mutable_tool_calls()->Swap(&source_tool_calls);
}
if (!result.finish_reason.empty()) {
choice->mutable_finish_reason()->swap(result.finish_reason);
}
} else {
message->set_content(cur_text);
if (output.finish_reason.has_value()) {
choice->set_finish_reason(output.finish_reason.value());
}
}
}
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
}
return call->write_and_finish(response);
}
} // namespace
ChatServiceImpl::ChatServiceImpl(LLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models),
master_(master),
tool_call_parser_format_(
master_->options().tool_call_parser().value_or("")),
reasoning_parser_format_(
master_->options().reasoning_parser().value_or("")) {
CHECK(master_ != nullptr);
add_model_master(models[0], master);
}
ChatServiceImpl::ChatServiceImpl(RecMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models),
rec_master_(master),
// RecMaster does not expose tool_call_parser/reasoning_parser options,
// and the current Rec scenario does not require these features.
tool_call_parser_format_(""),
reasoning_parser_format_("") {
CHECK(rec_master_ != nullptr);
}
void ChatServiceImpl::add_model_master(const std::string& model,
LLMMaster* master) {
CHECK(master != nullptr);
std::unique_lock<std::shared_mutex> lock(llm_model_to_master_mutex_);
llm_model_to_master_.insert_or_assign(model, master);
models_.insert(model);
}
LLMMaster* ChatServiceImpl::get_model_master(const std::string& model) const {
std::shared_lock<std::shared_mutex> lock(llm_model_to_master_mutex_);
auto it = llm_model_to_master_.find(model);
if (it == llm_model_to_master_.end()) {
return nullptr;
}
return it->second;
}
void ChatServiceImpl::process_rec_chat_request(std::shared_ptr<ChatCall> call) {
CHECK(rec_master_ != nullptr);
const auto& rpc_request = call->request();
const auto& model = rpc_request.model();
if (rec_master_->rec_type() != RecType::kLlmRec) {
call->finish_with_error(StatusCode::INVALID_ARGUMENT,
"Chat is only supported for LLMRec models");
return;
}
if (rec_master_->get_rate_limiter()->is_limited()) {
call->finish_with_error(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.");
return;
}
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
// Build messages from rpc request (inline code, not extracted to helper)
std::vector<Message> messages;
messages.reserve(rpc_request.messages_size());
for (const auto& message : rpc_request.messages()) {
messages.emplace_back(message.role(), message.content());
auto& msg = messages.back();
if (message.has_tool_call_id()) {
msg.tool_call_id = message.tool_call_id();
}
if (message.has_reasoning_content()) {
msg.reasoning_content = message.reasoning_content();
}
if (message.tool_calls_size() > 0) {
Message::ToolCallVec tool_calls;
tool_calls.reserve(message.tool_calls_size());
for (const auto& tool_call : message.tool_calls()) {
tool_calls.emplace_back();
auto& tc = tool_calls.back();
tc.id = tool_call.id();
tc.type = tool_call.type();
tc.function.name = tool_call.function().name();
tc.function.arguments = tool_call.function().arguments();
}
msg.tool_calls = std::move(tool_calls);
}
}
bool include_usage = false;
if (rpc_request.has_stream_options()) {
include_usage = rpc_request.stream_options().include_usage();
}
// Parse prompt tokens from routing (inline code, not extracted to helper)
std::optional<std::vector<int>> prompt_tokens = std::nullopt;
if (rpc_request.has_routing()) {
prompt_tokens = std::vector<int>{};
prompt_tokens->reserve(rpc_request.token_ids_size());
for (int i = 0; i < rpc_request.token_ids_size(); i++) {
prompt_tokens->emplace_back(rpc_request.token_ids(i));
}
request_params.decode_address = rpc_request.routing().decode_name();
}
auto saved_streaming = request_params.streaming;
auto saved_request_id = request_params.request_id;
rec_master_->handle_request(
std::move(messages),
std::move(prompt_tokens),
std::nullopt,
std::move(request_params),
[call,
model,
master = rec_master_,
stream = std::move(saved_streaming),
include_usage = include_usage,
first_message_sent = std::unordered_set<size_t>(),
request_id = std::move(saved_request_id),
created_time = absl::ToUnixSeconds(absl::Now())](
const RequestOutput& req_output) mutable -> bool {
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
master->get_rate_limiter()->decrease_one_request();
return call->finish_with_error(status.code(), status.message());
}
}
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
if (stream) {
return send_delta_to_client_brpc(call,
include_usage,
&first_message_sent,
request_id,
created_time,
model,
req_output);
}
return send_result_to_client_brpc(
call, request_id, created_time, model, req_output);
});
}
// chat_async for brpc with xllm_serice
void ChatServiceImpl::process_async_rpc_impl(
const proto::ChatRequest* request) {
const auto& service_request_id = request->service_request_id();
const auto& target_xservice_addr = request->source_xservice_addr();
auto callback = [master = master_](const RequestOutput& req_output) -> bool {
req_output.log_request_status();
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
// Reduce the number of concurrent requests when a request is
// finished with error.
master->get_rate_limiter()->decrease_one_request();
return master->handle_rpc_response(req_output);
}
}
// Reduce the number of concurrent requests when a request is finished
// or canceled.
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
return master->handle_rpc_response(req_output);
};
// LLMMaster path (existing logic)
// Check if the request is being rate-limited.
CHECK(master_ != nullptr);
if (master_->get_rate_limiter()->is_limited()) {
CALLBACK_WITH_ERROR(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.",
service_request_id,
target_xservice_addr);
return;
}
// check if model is supported
const auto& rpc_request = *request;
const auto& model = rpc_request.model();
if (unlikely(!models_.contains(model))) {
CALLBACK_WITH_ERROR(StatusCode::UNKNOWN,
"Model not supported",
service_request_id,
target_xservice_addr);
return;
}
RequestParams request_params(rpc_request, "", "");
std::vector<Message> messages;
messages.reserve(rpc_request.messages_size());
for (const auto& message : rpc_request.messages()) {
messages.emplace_back(message.role(), message.content());
auto& msg = messages.back();
if (message.has_tool_call_id()) {
msg.tool_call_id = message.tool_call_id();
}
if (message.has_reasoning_content()) {
msg.reasoning_content = message.reasoning_content();
}
if (message.tool_calls_size() > 0) {
Message::ToolCallVec tool_calls;
tool_calls.reserve(message.tool_calls_size());
for (const auto& tool_call : message.tool_calls()) {
tool_calls.emplace_back();
auto& tc = tool_calls.back();
tc.id = tool_call.id();
tc.type = tool_call.type();
tc.function.name = tool_call.function().name();
tc.function.arguments = tool_call.function().arguments();
}
msg.tool_calls = std::move(tool_calls);
}
}
std::optional<std::vector<int>> prompt_tokens = std::nullopt;
if (rpc_request.has_routing()) {
prompt_tokens = std::vector<int>{};
prompt_tokens->reserve(rpc_request.token_ids_size());
for (int i = 0; i < rpc_request.token_ids_size(); i++) {
prompt_tokens->emplace_back(rpc_request.token_ids(i));
}
request_params.decode_address = rpc_request.routing().decode_name();
}
// Preserve parser-relevant special tokens in decoded output
// so tool call detectors can match their control markers.
// Aligns with vLLM parser adjust_request behavior.
if (!tool_call_parser_format_.empty() && !request_params.tools.empty()) {
request_params.skip_special_tokens = false;
}
master_->handle_request(std::move(messages),
std::move(prompt_tokens),
std::move(request_params),
std::nullopt,
callback);
}
// chat_async for brpc
void ChatServiceImpl::process_async_impl(std::shared_ptr<ChatCall> call) {
const auto& rpc_request = call->request();
const auto& model = rpc_request.model();
// Route to RecMaster if configured
if (rec_master_) {
if (unlikely(!models_.contains(model))) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
process_rec_chat_request(call);
return;
}
LLMMaster* master = get_model_master(model);
if (unlikely(master == nullptr)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
// LLMMaster path (existing logic)
// Check if the request is being rate-limited or model is sleeping.
// is_limited() returns true if sleeping or rate-limited.
if (unlikely(master->get_rate_limiter()->is_limited())) {
if (master->get_rate_limiter()->is_sleeping()) {
call->finish_with_error(StatusCode::UNAVAILABLE,
"Model is currently in sleep state.");
} else {
call->finish_with_error(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.");
}
return;
}
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
std::vector<Message> messages;
messages.reserve(rpc_request.messages_size());
for (const auto& message : rpc_request.messages()) {
messages.emplace_back(message.role(), message.content());
auto& msg = messages.back();
if (message.has_tool_call_id()) {
msg.tool_call_id = message.tool_call_id();
}
if (message.has_reasoning_content()) {
msg.reasoning_content = message.reasoning_content();
}
if (message.tool_calls_size() > 0) {
Message::ToolCallVec tool_calls;
tool_calls.reserve(message.tool_calls_size());
for (const auto& tool_call : message.tool_calls()) {
tool_calls.emplace_back();
auto& tc = tool_calls.back();
tc.id = tool_call.id();
tc.type = tool_call.type();
tc.function.name = tool_call.function().name();
tc.function.arguments = tool_call.function().arguments();
}
msg.tool_calls = std::move(tool_calls);
}
}
bool include_usage = false;
if (rpc_request.has_stream_options()) {
include_usage = rpc_request.stream_options().include_usage();
}
std::optional<std::vector<int>> prompt_tokens = std::nullopt;
if (rpc_request.has_routing()) {
prompt_tokens = std::vector<int>{};
prompt_tokens->reserve(rpc_request.token_ids_size());
for (int i = 0; i < rpc_request.token_ids_size(); i++) {
prompt_tokens->emplace_back(rpc_request.token_ids(i));
}
request_params.decode_address = rpc_request.routing().decode_name();
}
if (!tool_call_parser_format_.empty() && !request_params.tools.empty()) {
request_params.skip_special_tokens = false;
}
const bool is_force_reasoning = get_enable_thinking_from_request(
request_params.chat_template_kwargs, reasoning_parser_format_);
std::shared_ptr<StreamOutputParser> stream_parser;
if (request_params.streaming && (!tool_call_parser_format_.empty() ||
!reasoning_parser_format_.empty())) {
stream_parser =
std::make_shared<StreamOutputParser>(request_params.tools,
tool_call_parser_format_,
reasoning_parser_format_,
is_force_reasoning);
CHECK(stream_parser != nullptr) << "create StreamOutputParser failed!";
}
auto saved_tools = request_params.tools;
auto saved_streaming = request_params.streaming;
auto saved_request_id = request_params.request_id;
master->handle_request(
std::move(messages),
std::move(prompt_tokens),
std::move(request_params),
call.get(),
[call,
model,
master = master,
stream = std::move(saved_streaming),
include_usage = include_usage,
first_message_sent = std::unordered_set<size_t>(),
request_id = std::move(saved_request_id),
created_time = absl::ToUnixSeconds(absl::Now()),
json_tools = std::move(saved_tools),
tool_call_parser_format = tool_call_parser_format_,
reasoning_parser_format = reasoning_parser_format_,
is_force_reasoning = is_force_reasoning,
stream_parser =
stream_parser](const RequestOutput& req_output) mutable -> bool {
req_output.log_request_status();
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
// Reduce the number of concurrent requests when a
// request is finished with error.
master->get_rate_limiter()->decrease_one_request();
return call->finish_with_error(status.code(), status.message());
}
}
// Reduce the number of concurrent requests when a request
// is finished or canceled.
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
if (stream) {
return send_delta_to_client_brpc(call,
include_usage,
&first_message_sent,
request_id,
created_time,
model,
req_output,
stream_parser);
}
return send_result_to_client_brpc(call,
request_id,
created_time,
model,
req_output,
tool_call_parser_format,
reasoning_parser_format,
is_force_reasoning,
json_tools);
});
}
MMChatServiceImpl::MMChatServiceImpl(VLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_(master) {
CHECK(master != nullptr);
}
void MMChatServiceImpl::process_async_impl(std::shared_ptr<MMChatCall> call) {
const auto& rpc_request = call->request();
const auto& req_messages = rpc_request.messages();
const auto& model = rpc_request.model();
if (!models_.contains(model)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
// Check if the request is being rate-limited.
if (master_->get_rate_limiter()->is_limited()) {
call->finish_with_error(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.");
return;
}
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
std::vector<Message> messages;
if (!mm_service_utils::build_messages<MMChatCall>(
req_messages, messages, call, master_->get_image_limit())) {
return;
}
bool include_usage = false;
if (rpc_request.has_stream_options()) {
include_usage = rpc_request.stream_options().include_usage();
}
auto saved_streaming = request_params.streaming;
auto saved_request_id = request_params.request_id;
auto payload = call->take_request_payload();
// schedule the request
master_->handle_request(
std::move(messages),
std::move(request_params),
std::move(payload),
[call,
model,
master = master_,
stream = std::move(saved_streaming),
include_usage = include_usage,
first_message_sent = std::unordered_set<size_t>(),
request_id = std::move(saved_request_id),
created_time = absl::ToUnixSeconds(absl::Now())](
const RequestOutput& req_output) mutable -> bool {
req_output.log_request_status();
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
// Reduce the number of concurrent requests when a request is
// finished with error.
master->get_rate_limiter()->decrease_one_request();
return call->finish_with_error(status.code(), status.message());
}
}
// Reduce the number of concurrent requests when a request is finished
// or canceled.
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
if (stream) {
// send delta to client
return send_delta_to_client_brpc(call,
include_usage,
&first_message_sent,
request_id,
created_time,
model,
req_output);
}
return send_result_to_client_brpc(
call, request_id, created_time, model, req_output);
});
}
} // namespace xllm

View File

@@ -0,0 +1,80 @@
/* 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.
==============================================================================*/
#pragma once
#include <shared_mutex>
#include <unordered_map>
#include "api_service/api_service_impl.h"
#include "api_service/stream_call.h"
#include "chat.pb.h"
#include "multimodal.pb.h"
namespace xllm {
class RecMaster;
using ChatCall = StreamCall<proto::ChatRequest, proto::ChatResponse>;
// a class to handle completion requests
class ChatServiceImpl final : public APIServiceImpl<ChatCall> {
public:
// Constructor for LLM backend
ChatServiceImpl(LLMMaster* master, const std::vector<std::string>& models);
// Constructor for Rec backend (LlmRec only, e.g., Qwen3)
ChatServiceImpl(RecMaster* master, const std::vector<std::string>& models);
// brpc call_data needs to use shared_ptr
void process_async_impl(std::shared_ptr<ChatCall> call);
void process_async_rpc_impl(const proto::ChatRequest* request);
void add_model_master(const std::string& model, LLMMaster* master);
private:
void process_rec_chat_request(std::shared_ptr<ChatCall> call);
LLMMaster* get_model_master(const std::string& model) const;
DISALLOW_COPY_AND_ASSIGN(ChatServiceImpl);
LLMMaster* master_ = nullptr;
RecMaster* rec_master_ = nullptr;
mutable std::shared_mutex llm_model_to_master_mutex_;
std::unordered_map<std::string, LLMMaster*> llm_model_to_master_;
const std::string tool_call_parser_format_;
const std::string reasoning_parser_format_;
};
class VLMMaster;
using MMChatCall = StreamCall<proto::MMChatRequest, proto::ChatResponse>;
// a class to handle mm chat completion requests
class MMChatServiceImpl : public APIServiceImpl<MMChatCall> {
public:
MMChatServiceImpl(VLMMaster* master, const std::vector<std::string>& models);
// brpc call_data needs to use shared_ptr
void process_async_impl(std::shared_ptr<MMChatCall> call);
private:
DISALLOW_COPY_AND_ASSIGN(MMChatServiceImpl);
VLMMaster* master_ = nullptr;
};
} // namespace xllm

View File

@@ -0,0 +1,340 @@
/* 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 "completion_service_impl.h"
#include <absl/time/clock.h>
#include <absl/time/time.h>
#include <glog/logging.h>
#include <torch/torch.h>
#include <cstdint>
#include <string>
#include "common/instance_name.h"
#include "completion.pb.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/framework/request/request_output.h"
#include "core/util/utils.h"
#ifdef likely
#undef likely
#endif
#define likely(x) __builtin_expect(!!(x), 1)
#ifdef unlikely
#undef unlikely
#endif
#define unlikely(x) __builtin_expect(!!(x), 0)
namespace xllm {
namespace {
void set_logprobs(proto::Choice* choice,
const std::optional<std::vector<LogProb>>& logprobs) {
if (!logprobs.has_value() || logprobs.value().empty()) {
return;
}
auto* proto_logprobs = choice->mutable_logprobs();
for (const auto& logprob : logprobs.value()) {
proto_logprobs->add_tokens(logprob.token);
proto_logprobs->add_token_ids(logprob.token_id);
proto_logprobs->add_token_logprobs(logprob.logprob);
}
}
bool send_delta_to_client_brpc(std::shared_ptr<CompletionCall> call,
bool include_usage,
const std::string& request_id,
int64_t created_time,
const std::string& model,
const RequestOutput& output) {
auto& response = call->response();
for (const auto& seq_output : output.outputs) {
if (!seq_output.text.empty()) {
response.Clear();
response.set_object("text_completion");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(seq_output.index);
choice->set_text(seq_output.text);
set_logprobs(choice, seq_output.logprobs);
if (!call->write(response)) {
return false;
}
}
if (seq_output.finish_reason.has_value()) {
response.Clear();
response.set_object("text_completion");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* choice = response.add_choices();
choice->set_index(seq_output.index);
choice->set_text("");
choice->set_finish_reason(seq_output.finish_reason.value());
if (!call->write(response)) {
return false;
}
}
}
if (include_usage && output.usage.has_value()) {
const auto& usage = output.usage.value();
response.Clear();
response.set_object("text_completion");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
response.mutable_choices();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
if (!call->write(response)) {
return false;
}
}
if (output.finished || output.cancelled) {
response.Clear();
return call->finish();
}
return true;
}
bool send_result_to_client_brpc(std::shared_ptr<CompletionCall> call,
const std::string& request_id,
int64_t created_time,
const std::string& model,
const RequestOutput& req_output) {
auto& response = call->response();
response.set_object("text_completion");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
response.mutable_choices()->Reserve(req_output.outputs.size());
for (const auto& output : req_output.outputs) {
auto* choice = response.add_choices();
choice->set_index(output.index);
choice->set_text(output.text);
set_logprobs(choice, output.logprobs);
if (output.finish_reason.has_value()) {
choice->set_finish_reason(output.finish_reason.value());
}
}
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
}
return call->write_and_finish(response);
}
} // namespace
CompletionServiceImpl::CompletionServiceImpl(
LLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_(master) {
CHECK(master_ != nullptr);
add_model_master(models[0], master);
}
void CompletionServiceImpl::add_model_master(const std::string& model,
LLMMaster* master) {
CHECK(master != nullptr);
std::unique_lock<std::shared_mutex> lock(llm_model_to_master_mutex_);
llm_model_to_master_.insert_or_assign(model, master);
models_.insert(model);
}
LLMMaster* CompletionServiceImpl::get_model_master(
const std::string& model) const {
std::shared_lock<std::shared_mutex> lock(llm_model_to_master_mutex_);
auto it = llm_model_to_master_.find(model);
if (it == llm_model_to_master_.end()) {
return nullptr;
}
return it->second;
}
// complete_async for brpc from xllm_service
void CompletionServiceImpl::process_async_rpc_impl(
const proto::CompletionRequest* request) {
const auto& service_request_id = request->service_request_id();
const auto& target_xservice_addr = request->source_xservice_addr();
auto callback = [master = master_](const RequestOutput& req_output) -> bool {
req_output.log_request_status();
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
// Reduce the number of concurrent requests when a request is
// finished with error.
master->get_rate_limiter()->decrease_one_request();
return master->handle_rpc_response(req_output);
}
}
// Reduce the number of concurrent requests when a request is finished
// or canceled.
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
return master->handle_rpc_response(req_output);
};
// Check if the request is being rate-limited.
if (unlikely(master_->get_rate_limiter()->is_limited())) {
CALLBACK_WITH_ERROR(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.",
service_request_id,
target_xservice_addr);
return;
}
// check if model is supported
const auto& rpc_request = *request;
const auto& model = rpc_request.model();
if (unlikely(!models_.contains(model))) {
CALLBACK_WITH_ERROR(StatusCode::UNKNOWN,
"Model not supported",
service_request_id,
target_xservice_addr);
return;
}
RequestParams request_params(rpc_request, "", "");
std::optional<std::vector<int>> prompt_tokens = std::nullopt;
if (rpc_request.has_routing()) {
prompt_tokens = std::vector<int>{};
prompt_tokens->reserve(rpc_request.token_ids_size());
for (int i = 0; i < rpc_request.token_ids_size(); i++) {
prompt_tokens->emplace_back(rpc_request.token_ids(i));
}
request_params.decode_address = rpc_request.routing().decode_name();
}
// schedule the request
master_->handle_request(std::move(rpc_request.prompt()),
std::move(prompt_tokens),
std::move(request_params),
std::nullopt,
callback);
}
// complete_async for brpc
void CompletionServiceImpl::process_async_impl(
std::shared_ptr<CompletionCall> call) {
const auto& rpc_request = call->request();
const auto& model = rpc_request.model();
LLMMaster* master = get_model_master(model);
if (unlikely(master == nullptr)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
// Check if the request is being rate-limited or model is sleeping.
// is_limited() returns true if sleeping or rate-limited.
if (unlikely(master->get_rate_limiter()->is_limited())) {
if (master->get_rate_limiter()->is_sleeping()) {
call->finish_with_error(StatusCode::UNAVAILABLE,
"Model is currently in sleep state.");
} else {
call->finish_with_error(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.");
}
return;
}
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
bool include_usage = false;
if (rpc_request.has_stream_options()) {
include_usage = rpc_request.stream_options().include_usage();
}
std::optional<std::vector<int>> prompt_tokens = std::nullopt;
if (rpc_request.has_routing()) {
prompt_tokens = std::vector<int>{};
prompt_tokens->reserve(rpc_request.token_ids_size());
for (int i = 0; i < rpc_request.token_ids_size(); i++) {
prompt_tokens->emplace_back(rpc_request.token_ids(i));
}
request_params.decode_address = rpc_request.routing().decode_name();
}
auto saved_streaming = request_params.streaming;
auto saved_request_id = request_params.request_id;
// schedule the request
master->handle_request(
std::move(rpc_request.prompt()),
std::move(prompt_tokens),
std::move(request_params),
call.get(),
[call,
model,
master = master,
stream = std::move(saved_streaming),
include_usage = include_usage,
request_id = std::move(saved_request_id),
created_time = absl::ToUnixSeconds(absl::Now())](
const RequestOutput& req_output) -> bool {
req_output.log_request_status();
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
// Reduce the number of concurrent requests when a request is
// finished with error.
master->get_rate_limiter()->decrease_one_request();
return call->finish_with_error(status.code(), status.message());
}
}
// Reduce the number of concurrent requests when a request is finished
// or canceled.
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
if (stream) {
return send_delta_to_client_brpc(
call, include_usage, request_id, created_time, model, req_output);
}
// NOTE: maybe need to refactor along with service, currently for
// non-stream request in prefill instance, we send a virtual response
return send_result_to_client_brpc(
call, request_id, created_time, model, req_output);
});
}
} // namespace xllm

View File

@@ -0,0 +1,52 @@
/* 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.
==============================================================================*/
#pragma once
#include <shared_mutex>
#include <unordered_map>
#include "api_service_impl.h"
#include "completion.pb.h"
#include "stream_call.h"
namespace xllm {
using CompletionCall =
StreamCall<proto::CompletionRequest, proto::CompletionResponse>;
// a class to handle completion requests
class CompletionServiceImpl final : public APIServiceImpl<CompletionCall> {
public:
CompletionServiceImpl(LLMMaster* master,
const std::vector<std::string>& models);
// brpc call_data needs to use shared_ptr
void process_async_impl(std::shared_ptr<CompletionCall> call);
void process_async_rpc_impl(const proto::CompletionRequest* request);
void add_model_master(const std::string& model, LLMMaster* master);
private:
LLMMaster* get_model_master(const std::string& model) const;
DISALLOW_COPY_AND_ASSIGN(CompletionServiceImpl);
LLMMaster* master_ = nullptr;
mutable std::shared_mutex llm_model_to_master_mutex_;
std::unordered_map<std::string, LLMMaster*> llm_model_to_master_;
};
} // namespace xllm

View File

@@ -0,0 +1,179 @@
/* Copyright 2025 The xLLM 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 "embedding_output_builder.h"
namespace xllm {
TensorProtoBuilder::TensorProtoBuilder(bool use_binary_encoding)
: use_binary_encoding_(use_binary_encoding) {};
bool TensorProtoBuilder::build_repeated_tensor(
const std::vector<torch::Tensor>& in_tensors,
google::protobuf::RepeatedPtrField<xllm::proto::Tensor>& out_tensors,
std::string& binary_payload) {
for (const auto& in_tensor : in_tensors) {
CHECK(in_tensor.is_contiguous())
<< "Internal Error: only support contiguous mm_embedding";
xllm::proto::Tensor* out_tensor = out_tensors.Add();
if (!build_tensor(in_tensor, *out_tensor, binary_payload)) {
return false;
}
}
return true;
}
bool TensorProtoBuilder::build_tensor(const torch::Tensor& in_tensor,
xllm::proto::Tensor& out_tensor,
std::string& binary_payload) {
if (use_binary_encoding_) {
// build tensor in binary format
out_tensor.set_datatype(
util::torch_datatype_to_proto(in_tensor.scalar_type()));
for (auto dim : in_tensor.sizes()) {
out_tensor.add_shape(static_cast<int32_t>(dim));
}
auto numel = in_tensor.numel();
auto byte_len = numel * in_tensor.element_size();
size_t offset = binary_payload.size();
auto* params = out_tensor.mutable_parameters();
(*params)["offset"].set_int64_param(offset);
(*params)["len"].set_int64_param(byte_len);
(*params)["is_binary"].set_bool_param(true);
binary_payload.append(reinterpret_cast<const char*>(in_tensor.data_ptr()),
byte_len);
} else {
// build tensor in json format
util::torch_to_proto(in_tensor, &out_tensor);
}
return true;
}
bool TensorProtoBuilder::build_tensor(const xllm::proto::Tensor& in_tensor,
const std::string& binary_payload,
torch::Tensor& out_tensor) {
const auto& params = in_tensor.parameters();
auto it = params.find("is_binary");
bool is_binary = (it != params.end() && it->second.bool_param());
if (!is_binary) {
out_tensor = util::proto_to_torch(in_tensor);
return true;
}
auto offset_it = params.find("offset");
auto len_it = params.find("len");
if (offset_it == params.end() || len_it == params.end()) {
return false;
}
const int64_t offset = offset_it->second.int64_param();
const int64_t byte_len = len_it->second.int64_param();
if (offset < 0 || byte_len <= 0 ||
static_cast<size_t>(offset + byte_len) > binary_payload.size()) {
return false;
}
auto dtype = util::datatype_proto_to_torch(in_tensor.datatype());
std::vector<int64_t> sizes;
sizes.reserve(in_tensor.shape_size());
for (auto dim : in_tensor.shape()) {
sizes.push_back(dim);
}
const char* src = binary_payload.data() + offset;
out_tensor = torch::from_blob(
const_cast<char*>(src), sizes, torch::TensorOptions().dtype(dtype));
out_tensor = out_tensor.clone();
return true;
}
EmbeddingOutputBuilder::EmbeddingOutputBuilder(
bool embedding_use_binary_encoding,
bool metadata_use_binary_encoding)
: embedding_use_binary_encoding_(embedding_use_binary_encoding),
metadata_use_binary_encoding_(metadata_use_binary_encoding) {};
EmbeddingOutputBuilder::~EmbeddingOutputBuilder() {};
bool EmbeddingOutputBuilder::build_repeated_embedding_output(
const std::vector<EmbeddingOutput>& in_embeddings,
google::protobuf::RepeatedPtrField<xllm::proto::Embedding>& out_embeddings,
std::string& binary_payload) {
for (const auto& in_embedding : in_embeddings) {
xllm::proto::Embedding* out_embedding = out_embeddings.Add();
if (!build_embedding_output(in_embedding, *out_embedding, binary_payload)) {
return false;
}
}
return true;
}
bool EmbeddingOutputBuilder::build_embedding_output(
const EmbeddingOutput& in_embedding,
xllm::proto::Embedding& out_embedding,
std::string& binary_payload) {
TensorProtoBuilder embedding_output_builder(embedding_use_binary_encoding_);
embedding_output_builder.build_tensor(in_embedding.embedding,
*out_embedding.mutable_embedding(),
binary_payload);
auto* meta_map = out_embedding.mutable_metadata();
TensorProtoBuilder meta_output_builder(metadata_use_binary_encoding_);
for (const auto& [key, value] : in_embedding.metadata) {
xllm::proto::Tensor metadata_tensor;
meta_output_builder.build_tensor(
in_embedding.metadata.at(key), metadata_tensor, binary_payload);
(*meta_map)[key] = std::move(metadata_tensor);
}
return true;
};
bool EmbeddingOutputBuilder::build_embedding_output(
const xllm::proto::Embedding& in_embedding,
std::string& binary_payload,
EmbeddingOutput& out_embedding) {
TensorProtoBuilder embedding_tensor_builder(embedding_use_binary_encoding_);
if (!embedding_tensor_builder.build_tensor(
in_embedding.embedding(), binary_payload, out_embedding.embedding)) {
return false;
}
out_embedding.metadata.clear();
TensorProtoBuilder meta_builder(metadata_use_binary_encoding_);
for (const auto& [key, proto_tensor] : in_embedding.metadata()) {
torch::Tensor tensor;
if (!meta_builder.build_tensor(proto_tensor, binary_payload, tensor)) {
return false;
}
out_embedding.metadata.emplace(key, std::move(tensor));
}
return true;
}
}; // namespace xllm

View File

@@ -0,0 +1,66 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include "core/common/message.h"
#include "core/common/types.h"
#include "core/framework//request/request_output.h"
#include "core/util/utils.h"
#include "embedding.pb.h"
#include "tensor.pb.h"
namespace xllm {
class TensorProtoBuilder {
public:
TensorProtoBuilder(bool use_binary_encoding);
~TensorProtoBuilder() = default;
bool build_repeated_tensor(
const std::vector<torch::Tensor>& in_tensors,
google::protobuf::RepeatedPtrField<xllm::proto::Tensor>& out_tensors,
std::string& binary_payload);
bool build_tensor(const torch::Tensor& in_tensor,
xllm::proto::Tensor& out_tensor,
std::string& binary_payload);
bool build_tensor(const xllm::proto::Tensor& in_tensor,
const std::string& binary_payload,
torch::Tensor& out_tensor);
private:
bool use_binary_encoding_;
};
class EmbeddingOutputBuilder {
public:
EmbeddingOutputBuilder(bool embedding_use_binary_encoding,
bool metadata_use_binary_encoding);
~EmbeddingOutputBuilder();
bool build_repeated_embedding_output(
const std::vector<EmbeddingOutput>& in_embeddings,
google::protobuf::RepeatedPtrField<xllm::proto::Embedding>&
out_embeddings,
std::string& binary_payload);
bool build_embedding_output(const EmbeddingOutput& in_embedding,
xllm::proto::Embedding& out_embedding,
std::string& binary_payload);
bool build_embedding_output(const xllm::proto::Embedding& in_embedding,
std::string& binary_payload,
EmbeddingOutput& out_embedding);
private:
bool embedding_use_binary_encoding_;
bool metadata_use_binary_encoding_;
};
} // namespace xllm

View File

@@ -0,0 +1,187 @@
/* Copyright 2025 The xLLM 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 "embedding_service_impl.h"
#include <glog/logging.h>
#include <string>
#include "common/instance_name.h"
#include "distributed_runtime/llm_master.h"
#include "embedding_output_builder.h"
#include "framework/request/request_params.h"
#include "mm_service_utils.h"
#include "util/utils.h"
#include "util/uuid.h"
namespace xllm {
namespace {
template <typename EmbeddingCall>
bool send_result_to_client_brpc(std::shared_ptr<EmbeddingCall> call,
const std::string& request_id,
int64_t created_time,
const std::string& model,
const RequestOutput& req_output) {
auto& response = call->response();
response.set_object("list");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
response.mutable_data()->Reserve(req_output.outputs.size());
std::string encoding_format = call->request().encoding_format();
bool use_binary_format = encoding_format == "binary";
EmbeddingOutputBuilder mm_embeddings_output_builder(use_binary_format, false);
std::string binary_payload;
for (const auto& output : req_output.outputs) {
// add data into response
auto* data = response.add_data();
data->set_index(output.index);
data->set_object("embedding");
if (output.embeddings.has_value()) {
data->mutable_embedding()->Add(
output.embeddings->data(),
output.embeddings->data() + output.embeddings->size());
}
if (output.mm_embeddings.has_value()) {
call->set_bytes_to_base64(true);
mm_embeddings_output_builder.build_repeated_embedding_output(
*output.mm_embeddings,
*(data->mutable_mm_embeddings()),
binary_payload);
}
}
// add usage statistics
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
}
return call->write_and_finish(response, binary_payload);
}
} // namespace
EmbeddingServiceImpl::EmbeddingServiceImpl(
LLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_(master) {
CHECK(master_ != nullptr);
}
// embedding_async for brpc
void EmbeddingServiceImpl::process_async_impl(
std::shared_ptr<EmbeddingCall> call) {
const auto& rpc_request = call->request();
// check if model is supported
const auto& model = rpc_request.model();
if (!models_.contains(model)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
// create RequestParams for embeddings request
// set is_embeddings and max_tokens = 1 to control engine step once.
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
// TODO only support input_str for now
auto& input = rpc_request.input();
auto saved_request_id = request_params.request_id;
// schedule the request
master_->handle_request(
std::move(input),
std::nullopt,
std::move(request_params),
call.get(),
[call,
model,
request_id = std::move(saved_request_id),
created_time = absl::ToUnixSeconds(absl::Now())](
const RequestOutput& req_output) -> bool {
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
return call->finish_with_error(status.code(), status.message());
}
}
return send_result_to_client_brpc<EmbeddingCall>(
call, request_id, created_time, model, req_output);
});
}
MMEmbeddingServiceImpl::MMEmbeddingServiceImpl(
VLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_(master) {
CHECK(master_ != nullptr);
}
void MMEmbeddingServiceImpl::process_async_impl(
std::shared_ptr<MMEmbeddingCall> call) {
const auto& rpc_request = call->request();
// check if model is supported
const auto& model = rpc_request.model();
if (!models_.contains(model)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
// create RequestParams for embeddings request
// set is_embeddings and max_tokens = 1 to control engine step once.
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
auto& req_messages = rpc_request.messages();
std::vector<Message> messages;
if (!mm_service_utils::build_messages<MMEmbeddingCall>(
req_messages, messages, call, master_->get_image_limit())) {
return;
}
auto request_id = request_params.request_id;
auto payload = call->take_request_payload();
// schedule the request
master_->handle_request(
std::move(messages),
std::move(request_params),
std::move(payload),
[call,
model,
request_id = request_id,
created_time = absl::ToUnixSeconds(absl::Now())](
const RequestOutput& req_output) -> bool {
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
return call->finish_with_error(status.code(), status.message());
}
}
return send_result_to_client_brpc<MMEmbeddingCall>(
call, request_id, created_time, model, req_output);
});
}
} // namespace xllm

View File

@@ -0,0 +1,58 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <absl/container/flat_hash_set.h>
#include "api_service/api_service_impl.h"
#include "api_service/call.h"
#include "api_service/non_stream_call.h"
#include "core/distributed_runtime/vlm_master.h"
#include "embedding.pb.h"
namespace xllm {
using EmbeddingCall =
NonStreamCall<proto::EmbeddingRequest, proto::EmbeddingResponse>;
// a class to handle completion requests
class EmbeddingServiceImpl final : public APIServiceImpl<EmbeddingCall> {
public:
EmbeddingServiceImpl(LLMMaster* master,
const std::vector<std::string>& models);
// brpc call_data needs to use shared_ptr
void process_async_impl(std::shared_ptr<EmbeddingCall> call);
private:
DISALLOW_COPY_AND_ASSIGN(EmbeddingServiceImpl);
LLMMaster* master_ = nullptr;
};
using MMEmbeddingCall =
NonStreamCall<proto::MMEmbeddingRequest, proto::EmbeddingResponse>;
class MMEmbeddingServiceImpl : public APIServiceImpl<MMEmbeddingCall> {
public:
MMEmbeddingServiceImpl(VLMMaster* master,
const std::vector<std::string>& models);
// brpc call_data needs to use shared_ptr
void process_async_impl(std::shared_ptr<MMEmbeddingCall> call);
private:
DISALLOW_COPY_AND_ASSIGN(MMEmbeddingServiceImpl);
VLMMaster* master_ = nullptr;
};
} // namespace xllm

View File

@@ -0,0 +1,108 @@
/* Copyright 2025 The xLLM 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 "image_generation_service_impl.h"
#include <glog/logging.h>
#include <string>
#include "butil/base64.h"
#include "common/instance_name.h"
#include "distributed_runtime/dit_master.h"
#include "framework/request/dit_request_output.h"
#include "framework/request/dit_request_params.h"
#include "util/utils.h"
#include "util/uuid.h"
namespace xllm {
namespace {
bool send_result_to_client_brpc(std::shared_ptr<ImageGenerationCall> call,
const std::string& request_id,
int64_t created_time,
const std::string& model,
const DiTRequestOutput& req_output) {
auto& response = call->response();
response.set_object("list");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
auto* proto_output = response.mutable_output();
const std::vector<DiTGenerationOutput>& outputs = req_output.outputs;
proto_output->mutable_results()->Reserve(outputs.size());
std::string image;
for (const auto& output : outputs) {
auto* proto_result = proto_output->add_results();
image.clear();
butil::Base64Encode(output.image, &image);
proto_result->set_image(image);
proto_result->set_width(output.width);
proto_result->set_height(output.height);
proto_result->set_seed(output.seed);
}
return call->write_and_finish(response);
}
} // namespace
ImageGenerationServiceImpl::ImageGenerationServiceImpl(
DiTMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_{master} {
CHECK(master_ != nullptr);
}
// image_generation_async for brpc
void ImageGenerationServiceImpl::process_async_impl(
std::shared_ptr<ImageGenerationCall> call) {
const auto& rpc_request = call->request();
// check if model is supported
const auto& model = rpc_request.model();
if (!models_.contains(model)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
// create DiTRequestParams for image generation request
DiTRequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
auto saved_request_id = request_params.request_id;
// schedule the request
master_->handle_request(
std::move(request_params),
call.get(),
[call,
model,
request_id = std::move(saved_request_id),
created_time = absl::ToUnixSeconds(absl::Now())](
const DiTRequestOutput& req_output) -> bool {
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
return call->finish_with_error(status.code(), status.message());
}
}
return send_result_to_client_brpc(
call, request_id, created_time, model, req_output);
});
}
} // namespace xllm

View File

@@ -0,0 +1,42 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <absl/container/flat_hash_set.h>
#include "api_service/api_service_impl.h"
#include "api_service/non_stream_call.h"
#include "image_generation.pb.h"
namespace xllm {
using ImageGenerationCall = NonStreamCall<proto::ImageGenerationRequest,
proto::ImageGenerationResponse>;
class DiTMaster;
// a class to handle image generation requests
class ImageGenerationServiceImpl : public APIServiceImpl<ImageGenerationCall> {
public:
ImageGenerationServiceImpl(DiTMaster* master,
const std::vector<std::string>& models);
// brpc call_data needs to use shared_ptr
void process_async_impl(std::shared_ptr<ImageGenerationCall> call);
private:
DISALLOW_COPY_AND_ASSIGN(ImageGenerationServiceImpl);
DiTMaster* master_ = nullptr;
};
} // namespace xllm

View File

@@ -0,0 +1,86 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include "core/common/message.h"
#include "core/common/types.h"
#include "multimodal.pb.h"
namespace xllm {
namespace mm_service_utils {
template <typename Call>
bool build_messages(const google::protobuf::RepeatedPtrField<
xllm::proto::MMChatMessage>& req_messages,
std::vector<Message>& out_messages,
std::shared_ptr<Call> call,
int image_limit) {
out_messages.clear();
out_messages.reserve(req_messages.size());
for (const auto& req_message : req_messages) {
MMContentVec contents;
for (const auto& input : req_message.content()) {
auto& item = const_cast<::xllm::proto::MMInputData&>(input);
if (item.type() == "text") {
contents.emplace_back(item.type(), *item.release_text());
} else if (item.type() == "image_url") {
ImageURL image_url;
image_url.url = std::move(*item.mutable_image_url()->release_url());
contents.emplace_back(item.type(), image_url);
} else if (item.type() == "video_url") {
VideoURL video_url;
video_url.url = std::move(*item.mutable_video_url()->release_url());
contents.emplace_back(item.type(), video_url);
} else if (item.type() == "audio_url") {
AudioURL audio_url;
audio_url.url = std::move(*item.mutable_audio_url()->release_url());
contents.emplace_back(item.type(), audio_url);
} else if (item.type() == "image_embedding") {
contents.emplace_back("image_embedding", item.image_embedding());
} else if (item.type() == "video_embedding") {
contents.emplace_back("video_embedding", item.video_embedding());
} else if (item.type() == "audio_embedding") {
contents.emplace_back("audio_embedding", item.audio_embedding());
} else {
call->finish_with_error(StatusCode::INVALID_ARGUMENT,
"message content type is invalid.");
return false;
}
}
out_messages.emplace_back(req_message.role(), std::move(contents));
}
for (auto& msg : out_messages) {
if (msg.calc_count("image_url") > image_limit) {
call->finish_with_error(StatusCode::INVALID_ARGUMENT,
"Number of images in a single message exceeds "
"the allowed image limit.");
return false;
}
}
return true;
};
} // namespace mm_service_utils
} // namespace xllm

View File

@@ -0,0 +1,62 @@
/* 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 "models_service_impl.h"
#include <nlohmann/json.hpp>
#include <string>
#include "absl/time/clock.h"
#include "absl/time/time.h"
#include "models.pb.h"
namespace xllm {
ModelsServiceImpl::ModelsServiceImpl(
const std::vector<std::string>& model_names,
const std::vector<std::string>& model_versions)
: model_names_(model_names),
model_versions_(model_versions),
created_(absl::ToUnixSeconds(absl::Now())) {}
bool ModelsServiceImpl::list_models(const proto::ModelListRequest* request,
proto::ModelListResponse* response) {
for (const auto& model_id : model_names_) {
auto* model_card = response->add_data();
model_card->set_id(model_id);
model_card->set_created(created_);
model_card->set_object("model");
model_card->set_owned_by("xllm");
}
return true;
}
std::string ModelsServiceImpl::list_model_versions() {
nlohmann::json model_states_array = nlohmann::json::array();
for (size_t i = 0; i < model_names_.size(); ++i) {
nlohmann::json model_state;
model_state["name"] = model_names_[i];
model_state["version"] = model_versions_[i];
model_state["state"] = "READY";
model_state["reason"] = "normal";
model_states_array.push_back(model_state);
}
return model_states_array.dump();
}
} // namespace xllm

View File

@@ -0,0 +1,43 @@
/* 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.
==============================================================================*/
#pragma once
#include <string>
#include "core/common/macros.h"
#include "models.pb.h"
namespace xllm {
class ModelsServiceImpl final {
public:
ModelsServiceImpl(const std::vector<std::string>& model_names,
const std::vector<std::string>& model_versions);
bool list_models(const proto::ModelListRequest* request,
proto::ModelListResponse* response);
std::string list_model_versions();
private:
DISALLOW_COPY_AND_ASSIGN(ModelsServiceImpl);
std::vector<std::string> model_names_;
std::vector<std::string> model_versions_;
uint32_t created_;
};
} // namespace xllm

View File

@@ -0,0 +1,125 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <brpc/controller.h>
#include <butil/iobuf.h>
#include <glog/logging.h>
#include <json2pb/pb_to_json.h>
#include <atomic>
#include <memory>
#include <optional>
#include <string>
#include "call.h"
#include "core/common/types.h"
namespace xllm {
template <typename Request, typename Response>
class NonStreamCall : public Call {
public:
using ReqType = Request;
using ResType = Response;
NonStreamCall(brpc::Controller* controller,
::google::protobuf::Closure* done,
Request* request,
Response* response,
bool use_arena = false)
: Call(controller),
done_(done),
request_(request),
response_(response),
use_arena_(use_arena) {
controller_->http_response().set_content_type("application/json");
json_options_.bytes_to_base64 = false;
json_options_.jsonify_empty_array = true;
json_options_.always_print_primitive_fields = true;
}
~NonStreamCall() override {
done_->Run();
if (!use_arena_) {
delete request_;
delete response_;
}
}
void set_bytes_to_base64(bool bytes_to_base64) {
json_options_.bytes_to_base64 = bytes_to_base64;
}
// For non stream response
bool write_and_finish(Response& response) {
butil::IOBufAsZeroCopyOutputStream json_output(
&controller_->response_attachment());
std::string err_msg;
if (!json2pb::ProtoMessageToJson(
response, &json_output, json_options_, &err_msg)) {
return finish_with_error(StatusCode::UNKNOWN, err_msg);
}
return true;
}
// For non stream response with json and binary payload
bool write_and_finish(Response& response, const std::string& binary_payload) {
if (!write_and_finish(response)) {
return false;
}
if (binary_payload.size() == 0) {
return true;
}
size_t json_len = controller_->response_attachment().size();
// Overwrite the header to indicate that the response contains JSON + binary
controller_->http_response().SetHeader("Content-Type",
"application/octet-stream");
// Add additional header to indicate the json payload length
controller_->http_response().SetHeader("X-Json-Length",
std::to_string(json_len));
controller_->response_attachment().append(binary_payload);
return true;
}
// For non stream response
bool finish_with_error(const StatusCode& code,
const std::string& error_message) {
controller_->SetFailed(error_message);
return true;
}
bool is_disconnected() const override { return controller_->IsCanceled(); }
const Request& request() const { return *request_; }
Response& response() { return *response_; }
::google::protobuf::Closure* done() { return done_; }
private:
::google::protobuf::Closure* done_;
Request* request_ = nullptr;
Response* response_ = nullptr;
bool use_arena_ = false;
json2pb::Pb2JsonOptions json_options_;
};
}; // namespace xllm

View File

@@ -0,0 +1,92 @@
/* Copyright 2025 The xLLM 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 "api_service/qwen3_rerank_service_impl.h"
#include "distributed_runtime/llm_master.h"
#include "framework/request/request_params.h"
namespace xllm {
Qwen3RerankServiceImpl::Qwen3RerankServiceImpl(
LLMMaster* master,
const std::vector<std::string>& models)
: RerankServiceImpl(master, models) {}
void Qwen3RerankServiceImpl::process_async_impl(
std::shared_ptr<RerankCall> call) {
const auto& rpc_request = call->request();
const auto& model = rpc_request.model();
if (!models_.contains(model)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
auto query = rpc_request.query();
std::vector<std::string> documents;
if (rpc_request.documents_size() > 0) {
documents = std::vector<std::string>(rpc_request.documents().begin(),
rpc_request.documents().end());
}
std::vector<std::string> reqs;
reqs.reserve(documents.size());
for (size_t i = 0; i < documents.size(); ++i) {
reqs.emplace_back(query + documents[i]);
}
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
std::vector<RequestParams> sps(documents.size(), request_params);
auto request_id = request_params.request_id;
int32_t top_n = static_cast<int32_t>(documents.size());
if (rpc_request.has_top_n()) {
top_n = std::min(top_n, rpc_request.top_n());
}
// Logprobs-based score computer for Qwen3 rerank
auto compute_scores = [](const std::vector<std::string>& documents,
const std::vector<RequestOutput>& req_outputs)
-> std::vector<RerankRequestOutput> {
std::vector<RerankRequestOutput> rerank_outputs;
rerank_outputs.reserve(documents.size());
for (size_t i = 0; i < documents.size(); ++i) {
if (req_outputs[i].outputs[0].logprobs.has_value()) {
auto score = req_outputs[i].outputs[0].logprobs.value()[0].logprob;
rerank_outputs.emplace_back(i, documents[i], score);
}
}
return rerank_outputs;
};
auto ctx = std::make_shared<RerankContext>(call,
std::move(documents),
model,
request_id,
top_n,
sps.size(),
compute_scores);
auto batch_callback = [ctx](size_t index, RequestOutput output) -> bool {
ctx->on_complete(index, std::move(output));
return true;
};
master_->handle_batch_request(reqs, sps, batch_callback);
}
} // namespace xllm

View File

@@ -0,0 +1,33 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include "api_service/rerank_service_impl.h"
namespace xllm {
class Qwen3RerankServiceImpl final : public RerankServiceImpl {
public:
Qwen3RerankServiceImpl(LLMMaster* master,
const std::vector<std::string>& models);
void process_async_impl(std::shared_ptr<RerankCall> call) override;
private:
DISALLOW_COPY_AND_ASSIGN(Qwen3RerankServiceImpl);
};
} // namespace xllm

View File

@@ -0,0 +1,354 @@
/* Copyright 2025 The xLLM 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 "rec_completion_service_impl.h"
#include <absl/time/clock.h>
#include <absl/time/time.h>
#include <glog/logging.h>
#include <torch/torch.h>
#include <algorithm>
#include <cstdint>
#include <string>
#include <vector>
#include "common/global_flags.h"
#include "common/instance_name.h"
#include "completion.pb.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/distributed_runtime/rec_master.h"
#include "core/framework/request/request_output.h"
#ifdef likely
#undef likely
#endif
#define likely(x) __builtin_expect(!!(x), 1)
#ifdef unlikely
#undef unlikely
#endif
#define unlikely(x) __builtin_expect(!!(x), 0)
namespace xllm {
namespace {
struct RecEmitRecord {
int32_t output_index = 0;
int64_t item_id = 0;
std::optional<RecItemInfo> item_info;
};
void append_rec_logprobs(proto::InferTensorContents* logprobs_context,
const SequenceOutput& output,
int32_t expected_count) {
const auto& token_logprobs = output.token_ids_logprobs;
const int32_t actual_count = static_cast<int32_t>(token_logprobs.size());
for (int32_t i = 0; i < expected_count; ++i) {
if (i < actual_count && token_logprobs[i].has_value()) {
logprobs_context->mutable_fp32_contents()->Add(token_logprobs[i].value());
} else {
logprobs_context->mutable_fp32_contents()->Add(0.0f);
}
}
}
void set_logprobs(proto::Choice* choice,
const std::optional<std::vector<LogProb>>& logprobs) {
if (!logprobs.has_value() || logprobs.value().empty()) {
return;
}
auto* proto_logprobs = choice->mutable_logprobs();
for (const auto& logprob : logprobs.value()) {
proto_logprobs->add_tokens(logprob.token);
proto_logprobs->add_token_ids(logprob.token_id);
proto_logprobs->add_token_logprobs(logprob.logprob);
}
}
bool send_result_to_client_brpc_rec(std::shared_ptr<CompletionCall> call,
const std::string& request_id,
int64_t created_time,
const std::string& model,
const RequestOutput& req_output) {
auto& response = call->response();
response.set_object("text_completion");
response.set_id(request_id);
response.set_created(created_time);
response.set_model(model);
// add choices into response
response.mutable_choices()->Reserve(req_output.outputs.size());
for (const auto& output : req_output.outputs) {
auto* choice = response.add_choices();
choice->set_index(output.index);
choice->set_text(output.text);
set_logprobs(choice, output.logprobs);
if (output.finish_reason.has_value()) {
choice->set_finish_reason(output.finish_reason.value());
}
}
// add usage statistics
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
}
// Add rec specific output tensors
auto output_tensor = response.mutable_output_tensors()->Add();
output_tensor->set_name("rec_result");
proto::InferOutputTensor* logprobs_tensor = nullptr;
int32_t logprob_width = 0;
if (FLAGS_enable_output_sku_logprobs && !req_output.outputs.empty()) {
logprobs_tensor = response.mutable_output_tensors()->Add();
logprobs_tensor->set_name("sku_logprobs");
logprobs_tensor->set_datatype(proto::DataType::FLOAT);
logprob_width =
static_cast<int32_t>(req_output.outputs[0].token_ids_logprobs.size());
}
if (FLAGS_enable_convert_tokens_to_item) {
output_tensor->set_datatype(proto::DataType::INT64);
proto::InferOutputTensor* did_tensor = nullptr;
proto::InferOutputTensor* type_tensor = nullptr;
if (FLAGS_enable_extended_item_info) {
did_tensor = response.mutable_output_tensors()->Add();
did_tensor->set_name("item_did");
did_tensor->set_datatype(proto::DataType::STRING);
type_tensor = response.mutable_output_tensors()->Add();
type_tensor->set_name("item_type");
type_tensor->set_datatype(proto::DataType::STRING);
}
std::vector<RecEmitRecord> emitted_items;
emitted_items.reserve(req_output.outputs.size());
const int32_t total_threshold = FLAGS_total_conversion_threshold;
for (int32_t i = 0; i < static_cast<int32_t>(req_output.outputs.size());
++i) {
const auto& output = req_output.outputs[i];
if (!output.item_ids_list.empty()) {
const bool has_item_infos =
output.item_infos_list.size() == output.item_ids_list.size();
for (size_t item_idx = 0; item_idx < output.item_ids_list.size();
++item_idx) {
if (static_cast<int32_t>(emitted_items.size()) >= total_threshold) {
break;
}
std::optional<RecItemInfo> item_info;
if (has_item_infos) {
item_info = output.item_infos_list[item_idx];
}
RecEmitRecord emitted_item;
emitted_item.output_index = i;
emitted_item.item_id = output.item_ids_list[item_idx];
emitted_item.item_info = std::move(item_info);
emitted_items.emplace_back(std::move(emitted_item));
}
} else if (output.item_ids.has_value() &&
static_cast<int32_t>(emitted_items.size()) < total_threshold) {
RecEmitRecord emitted_item;
emitted_item.output_index = i;
emitted_item.item_id = output.item_ids.value();
emitted_item.item_info = output.item_info;
emitted_items.emplace_back(std::move(emitted_item));
}
if (static_cast<int32_t>(emitted_items.size()) >= total_threshold) {
break;
}
}
const int32_t emitted_count = static_cast<int32_t>(emitted_items.size());
output_tensor->mutable_shape()->Add(emitted_count);
if (logprobs_tensor != nullptr) {
logprobs_tensor->mutable_shape()->Add(emitted_count);
logprobs_tensor->mutable_shape()->Add(logprob_width);
}
if (did_tensor != nullptr && type_tensor != nullptr) {
did_tensor->mutable_shape()->Add(emitted_count);
type_tensor->mutable_shape()->Add(emitted_count);
}
auto* output_context = output_tensor->mutable_contents();
auto* logprobs_context = logprobs_tensor == nullptr
? nullptr
: logprobs_tensor->mutable_contents();
auto append_output_logprobs = [&](int32_t output_index) {
if (logprobs_context != nullptr) {
append_rec_logprobs(
logprobs_context, req_output.outputs[output_index], logprob_width);
}
};
for (const RecEmitRecord& emitted_item : emitted_items) {
output_context->mutable_int64_contents()->Add(emitted_item.item_id);
append_output_logprobs(emitted_item.output_index);
if (did_tensor != nullptr && type_tensor != nullptr) {
did_tensor->mutable_contents()->add_bytes_contents(
emitted_item.item_info.has_value() ? emitted_item.item_info->did
: "");
type_tensor->mutable_contents()->add_bytes_contents(
emitted_item.item_info.has_value() ? emitted_item.item_info->type
: "");
}
}
} else {
output_tensor->set_datatype(proto::DataType::INT32);
if (req_output.outputs.empty()) {
output_tensor->mutable_shape()->Add(0);
output_tensor->mutable_shape()->Add(0);
if (logprobs_tensor != nullptr) {
logprobs_tensor->mutable_shape()->Add(0);
logprobs_tensor->mutable_shape()->Add(0);
}
return call->write_and_finish(response);
}
const int32_t output_count =
static_cast<int32_t>(req_output.outputs.size());
output_tensor->mutable_shape()->Add(output_count);
output_tensor->mutable_shape()->Add(req_output.outputs[0].token_ids.size());
if (logprobs_tensor != nullptr) {
logprobs_tensor->mutable_shape()->Add(output_count);
logprobs_tensor->mutable_shape()->Add(logprob_width);
}
auto* context = output_tensor->mutable_contents();
auto* logprobs_context = logprobs_tensor == nullptr
? nullptr
: logprobs_tensor->mutable_contents();
auto append_output_logprobs = [&](int32_t output_index) {
if (logprobs_context != nullptr) {
append_rec_logprobs(
logprobs_context, req_output.outputs[output_index], logprob_width);
}
};
for (int32_t i = 0; i < output_count; ++i) {
// LOG(INFO) << req_output.outputs[i].token_ids;
context->mutable_int_contents()->Add(
req_output.outputs[i].token_ids.begin(),
req_output.outputs[i].token_ids.end());
append_output_logprobs(i);
}
}
return call->write_and_finish(response);
}
} // namespace
RecCompletionServiceImpl::RecCompletionServiceImpl(
RecMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_(master) {
CHECK(master_ != nullptr);
}
void RecCompletionServiceImpl::process_async_impl(
std::shared_ptr<CompletionCall> call) {
const auto& rpc_request = call->request();
// check if model is supported
const auto& model = rpc_request.model();
if (unlikely(!models_.contains(model))) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
// Check if the request is being rate-limited.
if (unlikely(master_->get_rate_limiter()->is_limited())) {
call->finish_with_error(
StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.");
return;
}
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
if (FLAGS_enable_output_sku_logprobs) {
request_params.logprobs = true;
}
bool include_usage = false;
if (rpc_request.has_stream_options()) {
include_usage = rpc_request.stream_options().include_usage();
}
std::optional<std::vector<int>> prompt_tokens = std::nullopt;
if (rpc_request.has_routing()) {
prompt_tokens = std::vector<int>{};
prompt_tokens->reserve(rpc_request.token_ids_size());
for (int i = 0; i < rpc_request.token_ids_size(); i++) {
prompt_tokens->emplace_back(rpc_request.token_ids(i));
}
request_params.decode_address = rpc_request.routing().decode_name();
}
const auto& rpc_request_ref = call->request();
std::optional<std::vector<proto::InferInputTensor>> input_tensors =
std::nullopt;
if (rpc_request_ref.input_tensors_size()) {
std::vector<proto::InferInputTensor> tensors;
tensors.reserve(rpc_request_ref.input_tensors_size());
for (int i = 0; i < rpc_request_ref.input_tensors_size(); ++i) {
tensors.push_back(rpc_request_ref.input_tensors(i));
}
input_tensors = std::move(tensors);
}
// schedule the request
auto saved_streaming = request_params.streaming;
auto saved_request_id = request_params.request_id;
master_->handle_request(
std::move(rpc_request_ref.prompt()),
std::move(prompt_tokens),
std::move(input_tensors),
std::move(request_params),
[call,
model,
master = master_,
stream = std::move(saved_streaming),
include_usage = include_usage,
request_id = saved_request_id,
created_time = absl::ToUnixSeconds(absl::Now())](
const RequestOutput& req_output) -> bool {
if (req_output.status.has_value()) {
const auto& status = req_output.status.value();
if (!status.ok()) {
// Reduce the number of concurrent requests when a request is
// finished with error.
master->get_rate_limiter()->decrease_one_request();
return call->finish_with_error(status.code(), status.message());
}
}
// Reduce the number of concurrent requests when a request is finished
// or canceled.
if (req_output.finished || req_output.cancelled) {
master->get_rate_limiter()->decrease_one_request();
}
return send_result_to_client_brpc_rec(
call, request_id, created_time, model, req_output);
});
}
} // namespace xllm

View File

@@ -0,0 +1,45 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <absl/container/flat_hash_set.h>
#include "api_service_impl.h"
#include "completion.pb.h"
#include "core/distributed_runtime/rec_master.h"
#include "rec.pb.h"
#include "stream_call.h"
namespace xllm {
using CompletionCall =
StreamCall<proto::CompletionRequest, proto::CompletionResponse>;
// a class to handle completion requests
class RecCompletionServiceImpl final : public APIServiceImpl<CompletionCall> {
public:
RecCompletionServiceImpl(RecMaster* master,
const std::vector<std::string>& models);
// brpc call_data needs to use shared_ptr
void process_async_impl(std::shared_ptr<CompletionCall> call);
private:
DISALLOW_COPY_AND_ASSIGN(RecCompletionServiceImpl);
RecMaster* master_ = nullptr;
};
} // namespace xllm

View File

@@ -0,0 +1,159 @@
/* Copyright 2025 The xLLM 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 "rerank_service_impl.h"
#include <torch/torch.h>
#include <string>
#include "common/instance_name.h"
#include "distributed_runtime/llm_master.h"
#include "framework/request/request_params.h"
#include "util/utils.h"
#include "util/uuid.h"
namespace xllm {
void RerankContext::finalize() {
auto rerank_outputs = compute_scores(documents, req_outputs);
if (rerank_outputs.empty()) {
call->finish_with_error(StatusCode::UNKNOWN, "Failed to compute scores");
return;
}
std::sort(rerank_outputs.begin(),
rerank_outputs.end(),
[](const RerankRequestOutput& a, const RerankRequestOutput& b) {
return a.score > b.score;
});
auto& response = call->response();
response.set_id(request_id);
response.set_model(model);
response.mutable_results()->Reserve(top_n);
for (int32_t i = 0;
i < top_n && i < static_cast<int32_t>(rerank_outputs.size());
++i) {
auto* result = response.add_results();
result->set_index(rerank_outputs[i].index);
result->mutable_document()->set_text(rerank_outputs[i].document);
result->set_relevance_score(rerank_outputs[i].score);
}
int32_t num_prompt_tokens = 0;
int32_t num_generated_tokens = 0;
int32_t num_total_tokens = 0;
for (const auto& req_output : req_outputs) {
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
num_prompt_tokens += usage.num_prompt_tokens;
num_generated_tokens += usage.num_generated_tokens;
num_total_tokens += usage.num_total_tokens;
}
}
if (num_total_tokens > 0) {
auto* proto_usage = response.mutable_usage();
proto_usage->set_prompt_tokens(num_prompt_tokens);
proto_usage->set_completion_tokens(num_generated_tokens);
proto_usage->set_total_tokens(num_total_tokens);
}
call->write_and_finish(response);
}
RerankServiceImpl::RerankServiceImpl(LLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_(master) {
CHECK(master_ != nullptr);
}
void RerankServiceImpl::process_async_impl(std::shared_ptr<RerankCall> call) {
const auto& rpc_request = call->request();
const auto& model = rpc_request.model();
if (!models_.contains(model)) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
std::vector<std::string> documents;
if (rpc_request.documents_size() > 0) {
documents = std::vector<std::string>(rpc_request.documents().begin(),
rpc_request.documents().end());
}
documents.emplace_back(rpc_request.query());
RequestParams request_params(
rpc_request, call->get_x_request_id(), call->get_x_request_time());
std::vector<RequestParams> sps(documents.size(), request_params);
auto request_id = request_params.request_id;
int32_t top_n = static_cast<int32_t>(documents.size() - 1);
if (rpc_request.has_top_n()) {
top_n = std::min(top_n, rpc_request.top_n());
}
// Cosine similarity score computer for embedding-based rerank
auto compute_scores = [](const std::vector<std::string>& documents,
const std::vector<RequestOutput>& req_outputs)
-> std::vector<RerankRequestOutput> {
size_t doc_size = documents.size() - 1;
auto& query_output = req_outputs[doc_size];
if (!query_output.outputs[0].embeddings.has_value()) {
return {};
}
auto query_embed = query_output.outputs[0].embeddings.value();
auto query_tensor =
torch::from_blob(query_embed.data(),
{static_cast<int64_t>(query_embed.size())},
torch::kFloat32);
std::vector<RerankRequestOutput> rerank_outputs;
rerank_outputs.reserve(doc_size);
for (size_t i = 0; i < doc_size; ++i) {
if (req_outputs[i].outputs[0].embeddings.has_value()) {
auto doc_embed = req_outputs[i].outputs[0].embeddings.value();
auto doc_tensor =
torch::from_blob(doc_embed.data(),
{static_cast<int64_t>(doc_embed.size())},
torch::kFloat32);
auto score =
torch::cosine_similarity(query_tensor, doc_tensor, 0).item<float>();
rerank_outputs.emplace_back(i, documents[i], score);
}
}
return rerank_outputs;
};
auto ctx = std::make_shared<RerankContext>(call,
std::move(documents),
model,
request_id,
top_n,
sps.size(),
compute_scores);
auto batch_callback = [ctx](size_t index, RequestOutput output) -> bool {
ctx->on_complete(index, std::move(output));
return true;
};
master_->handle_batch_request(ctx->documents, sps, batch_callback);
}
} // namespace xllm

View File

@@ -0,0 +1,100 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <absl/container/flat_hash_set.h>
#include <atomic>
#include <functional>
#include <vector>
#include "api_service/api_service_impl.h"
#include "api_service/call.h"
#include "api_service/non_stream_call.h"
#include "rerank.pb.h"
namespace xllm {
using RerankCall = NonStreamCall<proto::RerankRequest, proto::RerankResponse>;
struct RerankRequestOutput {
int32_t index = 0;
std::string document = "";
float score = 0.0f;
RerankRequestOutput(int32_t index, std::string document, float score)
: index(index), document(std::move(document)), score(score) {}
};
// Score computer function type: computes scores from request outputs
// Returns vector of RerankRequestOutput with computed scores
using ScoreComputer = std::function<std::vector<RerankRequestOutput>(
const std::vector<std::string>& documents,
const std::vector<RequestOutput>& req_outputs)>;
// Shared context for async aggregation of rerank sub-request results
// Template parameter allows different score computation strategies
struct RerankContext {
std::shared_ptr<RerankCall> call;
std::vector<std::string> documents;
std::string model;
std::string request_id;
int32_t top_n;
std::vector<RequestOutput> req_outputs;
std::atomic<size_t> pending_count;
ScoreComputer compute_scores;
RerankContext(std::shared_ptr<RerankCall> call,
std::vector<std::string> documents,
std::string model,
std::string request_id,
int32_t top_n,
size_t num_requests,
ScoreComputer compute_scores)
: call(std::move(call)),
documents(std::move(documents)),
model(std::move(model)),
request_id(std::move(request_id)),
top_n(top_n),
pending_count(num_requests),
compute_scores(std::move(compute_scores)) {
req_outputs.resize(num_requests);
}
void on_complete(size_t index, RequestOutput output) {
req_outputs[index] = std::move(output);
if (pending_count.fetch_sub(1, std::memory_order_acq_rel) == 1) {
finalize();
}
}
void finalize();
};
class RerankServiceImpl : public APIServiceImpl<RerankCall> {
public:
RerankServiceImpl(LLMMaster* master, const std::vector<std::string>& models);
virtual void process_async_impl(std::shared_ptr<RerankCall> call);
protected:
DISALLOW_COPY_AND_ASSIGN(RerankServiceImpl);
LLMMaster* master_ = nullptr;
};
} // namespace xllm

View File

@@ -0,0 +1,471 @@
/* Copyright 2026 The xLLM 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 "sample_service_impl.h"
#include <absl/time/clock.h>
#include <absl/time/time.h>
#include <glog/logging.h>
#include <algorithm>
#include <condition_variable>
#include <mutex>
#include "common/instance_name.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/framework/request/request_output.h"
#include "core/framework/request/request_params.h"
#include "core/framework/request/sample_slot.h"
#include "core/util/uuid.h"
namespace xllm {
namespace {
thread_local ShortUUID short_uuid;
const std::string kSelectorMatchFinishReason = "selector_match";
const std::string kEmptyLogprobsFinishReason = "empty_logprobs";
std::string generate_sample_request_id() {
return "sample-" + InstanceName::name()->get_name_hash() + "-" +
short_uuid.random();
}
void initialize_response(const std::string& request_id,
const std::string& model,
uint32_t created_time,
proto::SampleResponse* response) {
CHECK(response != nullptr);
response->Clear();
response->set_id(request_id);
response->set_object("sample_completion");
response->set_created(created_time);
response->set_model(model);
}
void set_choice_logprobs(proto::Choice* choice,
const std::optional<std::vector<LogProb>>& logprobs) {
if (choice == nullptr || !logprobs.has_value() || logprobs->empty()) {
return;
}
const auto& sampled_logprob = logprobs->front();
auto* proto_logprobs = choice->mutable_logprobs();
if (sampled_logprob.top_logprobs.has_value() &&
!sampled_logprob.top_logprobs->empty()) {
for (const auto& top_logprob : sampled_logprob.top_logprobs.value()) {
proto_logprobs->add_tokens(top_logprob.token);
proto_logprobs->add_token_ids(top_logprob.token_id);
proto_logprobs->add_token_logprobs(top_logprob.logprob);
}
return;
}
proto_logprobs->add_tokens(sampled_logprob.token);
proto_logprobs->add_token_ids(sampled_logprob.token_id);
proto_logprobs->add_token_logprobs(sampled_logprob.logprob);
}
bool has_choice_logprobs(const SequenceOutput& output) {
return output.logprobs.has_value() && !output.logprobs->empty();
}
std::string get_choice_text(const SequenceOutput& output) {
if (!has_choice_logprobs(output)) {
return output.text;
}
const auto& sampled_logprob = output.logprobs->front();
if (sampled_logprob.top_logprobs.has_value() &&
!sampled_logprob.top_logprobs->empty()) {
return sampled_logprob.top_logprobs->front().token;
}
return sampled_logprob.token;
}
std::string get_finish_reason(const SequenceOutput& output) {
if (output.finish_reason.has_value()) {
return output.finish_reason.value();
}
return has_choice_logprobs(output) ? kSelectorMatchFinishReason
: kEmptyLogprobsFinishReason;
}
uint32_t get_requested_logprobs(const proto::SampleRequest& request) {
return request.has_logprobs()
? request.logprobs()
: sample_service_internal::kDefaultSampleLogprobs;
}
Status get_rate_limit_status(LLMMaster* master) {
CHECK(master != nullptr);
if (!master->get_rate_limiter()->is_limited()) {
return Status();
}
if (master->get_rate_limiter()->is_sleeping()) {
return Status(StatusCode::UNAVAILABLE,
"Model is currently in sleep state.");
}
return Status(StatusCode::RESOURCE_EXHAUSTED,
"The number of concurrent requests has reached the limit.");
}
} // namespace
namespace sample_service_internal {
Status validate_request(const proto::SampleRequest& request) {
if (request.model().empty()) {
return Status(StatusCode::INVALID_ARGUMENT, "model is required");
}
if (request.prompt().empty()) {
return Status(StatusCode::INVALID_ARGUMENT, "prompt is required");
}
if (!request.has_selector()) {
return Status(StatusCode::INVALID_ARGUMENT, "selector is required");
}
if (request.selector().type() != "literal") {
return Status(StatusCode::INVALID_ARGUMENT,
"selector.type must be literal");
}
if (request.selector().value().empty()) {
return Status(StatusCode::INVALID_ARGUMENT, "selector.value is required");
}
if (request.has_logprobs() && (request.logprobs() < kMinSampleLogprobs ||
request.logprobs() > kMaxSampleLogprobs)) {
return Status(StatusCode::INVALID_ARGUMENT,
"logprobs must be between 1 and 5");
}
return Status();
}
Status validate_runtime_config(bool enable_schedule_overlap) {
if (enable_schedule_overlap) {
return Status(StatusCode::UNAVAILABLE,
"/v1/sample does not support async scheduling "
"(enable_schedule_overlap=true)");
}
return Status();
}
bool build_request_params(const proto::SampleRequest& request,
const Tokenizer& tokenizer,
RequestParams* request_params) {
if (request_params == nullptr) {
return false;
}
RequestParams params;
params.request_id = request.has_request_id() ? request.request_id()
: generate_sample_request_id();
params.logprobs = true;
params.top_logprobs =
request.has_logprobs() ? request.logprobs() : kDefaultSampleLogprobs;
params.max_tokens = 1;
params.n = 1;
params.best_of = 1;
params.add_special_tokens = true;
params.is_sample_request = true;
if (!build_sample_slots(params.request_id,
request.prompt(),
request.selector().value(),
tokenizer,
&params.sample_slots)) {
return false;
}
*request_params = std::move(params);
return true;
}
bool build_empty_response(const proto::SampleRequest& request,
const Tokenizer& tokenizer,
const std::string& request_id,
proto::SampleResponse* response) {
if (response == nullptr) {
return false;
}
std::vector<int32_t> prompt_tokens;
if (!tokenizer.encode(request.prompt(), &prompt_tokens, true)) {
return false;
}
initialize_response(request_id,
request.model(),
static_cast<uint32_t>(absl::ToUnixSeconds(absl::Now())),
response);
response->mutable_choices();
auto* usage = response->mutable_usage();
const int32_t prompt_tokens_count =
static_cast<int32_t>(prompt_tokens.size());
usage->set_prompt_tokens(prompt_tokens_count);
usage->set_completion_tokens(0);
usage->set_total_tokens(prompt_tokens_count);
return true;
}
bool build_response(const std::string& request_id,
const std::string& model,
uint32_t created_time,
const RequestOutput& req_output,
proto::SampleResponse* response) {
if (response == nullptr) {
return false;
}
initialize_response(request_id, model, created_time, response);
std::vector<const SequenceOutput*> sorted_outputs;
sorted_outputs.reserve(req_output.outputs.size());
for (const auto& output : req_output.outputs) {
sorted_outputs.push_back(&output);
}
std::stable_sort(
sorted_outputs.begin(),
sorted_outputs.end(),
[](const auto* lhs, const auto* rhs) { return lhs->index < rhs->index; });
response->mutable_choices()->Reserve(sorted_outputs.size());
for (const auto* output : sorted_outputs) {
auto* choice = response->add_choices();
choice->set_index(output->index);
choice->set_text(get_choice_text(*output));
set_choice_logprobs(choice, output->logprobs);
if (!has_choice_logprobs(*output)) {
choice->mutable_logprobs();
}
choice->set_finish_reason(get_finish_reason(*output));
}
if (req_output.usage.has_value()) {
const auto& usage = req_output.usage.value();
auto* proto_usage = response->mutable_usage();
proto_usage->set_prompt_tokens(usage.num_prompt_tokens);
proto_usage->set_completion_tokens(usage.num_generated_tokens);
proto_usage->set_total_tokens(usage.num_total_tokens);
}
return true;
}
} // namespace sample_service_internal
SampleServiceImpl::SampleServiceImpl(LLMMaster* master,
const std::vector<std::string>& models)
: APIServiceImpl(models), master_(master) {
CHECK(master_ != nullptr);
}
bool SampleServiceImpl::process_request(const proto::SampleRequest& request,
proto::SampleResponse* response,
Status* status) const {
CHECK(response != nullptr);
CHECK(status != nullptr);
response->Clear();
*status = sample_service_internal::validate_request(request);
if (!status->ok()) {
return false;
}
if (!models_.contains(request.model())) {
*status = Status(StatusCode::UNKNOWN, "Model not supported");
return false;
}
*status = sample_service_internal::validate_runtime_config(
master_->options().enable_schedule_overlap());
if (!status->ok()) {
return false;
}
RequestParams request_params;
if (!sample_service_internal::build_request_params(
request, master_->tokenizer(), &request_params)) {
*status = Status(StatusCode::UNKNOWN,
"Failed to build sample selector runtime mapping");
return false;
}
if (request_params.sample_slots.empty()) {
if (!sample_service_internal::build_empty_response(
request,
master_->tokenizer(),
request_params.request_id,
response)) {
*status = Status(StatusCode::UNKNOWN,
"Failed to build sample no-match response");
return false;
}
*status = Status();
return true;
}
*status = get_rate_limit_status(master_);
if (!status->ok()) {
return false;
}
RequestOutput final_output;
bool has_final_output = false;
std::mutex mu;
std::condition_variable cv;
const auto request_id = request_params.request_id;
const size_t match_count = request_params.sample_slots.size();
const auto created_time =
static_cast<uint32_t>(absl::ToUnixSeconds(absl::Now()));
master_->handle_request(
request.prompt(),
std::nullopt,
std::move(request_params),
std::nullopt,
[this, &mu, &cv, &final_output, &has_final_output](
const RequestOutput& req_output) -> bool {
req_output.log_request_status();
if (req_output.status.has_value() && !req_output.status->ok()) {
master_->get_rate_limiter()->decrease_one_request();
} else if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master_->get_rate_limiter()->decrease_one_request();
} else {
return true;
}
{
std::lock_guard<std::mutex> lock(mu);
if (!has_final_output) {
final_output = req_output;
has_final_output = true;
}
}
cv.notify_one();
return true;
});
{
std::unique_lock<std::mutex> lock(mu);
cv.wait(lock, [&has_final_output]() { return has_final_output; });
}
if (final_output.status.has_value() && !final_output.status->ok()) {
*status = final_output.status.value();
return false;
}
if (!sample_service_internal::build_response(
request_id, request.model(), created_time, final_output, response)) {
*status = Status(StatusCode::UNKNOWN, "Failed to build sample response");
return false;
}
*status = Status();
return true;
}
void SampleServiceImpl::process_async_impl(std::shared_ptr<SampleCall> call) {
const auto& request = call->request();
Status status = sample_service_internal::validate_request(request);
if (!status.ok()) {
call->finish_with_error(status.code(), status.message());
return;
}
if (!models_.contains(request.model())) {
call->finish_with_error(StatusCode::UNKNOWN, "Model not supported");
return;
}
status = sample_service_internal::validate_runtime_config(
master_->options().enable_schedule_overlap());
if (!status.ok()) {
call->finish_with_error(status.code(), status.message());
return;
}
RequestParams request_params;
if (!sample_service_internal::build_request_params(
request, master_->tokenizer(), &request_params)) {
call->finish_with_error(StatusCode::UNKNOWN,
"Failed to build sample selector runtime mapping");
return;
}
if (request_params.sample_slots.empty()) {
if (!sample_service_internal::build_empty_response(
request,
master_->tokenizer(),
request_params.request_id,
&call->response())) {
call->finish_with_error(StatusCode::UNKNOWN,
"Failed to build sample no-match response");
return;
}
call->write_and_finish(call->response());
return;
}
status = get_rate_limit_status(master_);
if (!status.ok()) {
call->finish_with_error(status.code(), status.message());
return;
}
const auto request_id = request_params.request_id;
const size_t match_count = request_params.sample_slots.size();
const auto created_time =
static_cast<uint32_t>(absl::ToUnixSeconds(absl::Now()));
master_->handle_request(
request.prompt(),
std::nullopt,
std::move(request_params),
call.get(),
[call,
master = master_,
model = request.model(),
request_id,
match_count,
created_time](const RequestOutput& req_output) -> bool {
req_output.log_request_status();
if (req_output.status.has_value()) {
const auto& output_status = req_output.status.value();
if (!output_status.ok()) {
master->get_rate_limiter()->decrease_one_request();
return call->finish_with_error(output_status.code(),
output_status.message());
}
}
if (req_output.finished || req_output.cancelled ||
req_output.finished_on_prefill_instance) {
master->get_rate_limiter()->decrease_one_request();
}
if (!sample_service_internal::build_response(request_id,
model,
created_time,
req_output,
&call->response())) {
return call->finish_with_error(StatusCode::UNKNOWN,
"Failed to build sample response");
}
return call->write_and_finish(call->response());
});
}
} // namespace xllm

View File

@@ -0,0 +1,74 @@
/* Copyright 2026 The xLLM 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.
==============================================================================*/
#pragma once
#include <cstdint>
#include <string>
#include "api_service/api_service_impl.h"
#include "api_service/non_stream_call.h"
#include "core/common/types.h"
#include "core/framework/request/request_params.h"
#include "sample.pb.h"
namespace xllm {
using SampleCall = NonStreamCall<proto::SampleRequest, proto::SampleResponse>;
namespace sample_service_internal {
inline constexpr uint32_t kDefaultSampleLogprobs = 5;
inline constexpr uint32_t kMinSampleLogprobs = 1;
inline constexpr uint32_t kMaxSampleLogprobs = 5;
Status validate_request(const proto::SampleRequest& request);
Status validate_runtime_config(bool enable_schedule_overlap);
bool build_request_params(const proto::SampleRequest& request,
const Tokenizer& tokenizer,
RequestParams* request_params);
bool build_empty_response(const proto::SampleRequest& request,
const Tokenizer& tokenizer,
const std::string& request_id,
proto::SampleResponse* response);
bool build_response(const std::string& request_id,
const std::string& model,
uint32_t created_time,
const RequestOutput& req_output,
proto::SampleResponse* response);
} // namespace sample_service_internal
class SampleServiceImpl final : public APIServiceImpl<SampleCall> {
public:
SampleServiceImpl(LLMMaster* master, const std::vector<std::string>& models);
bool process_request(const proto::SampleRequest& request,
proto::SampleResponse* response,
Status* status) const;
void process_async_impl(std::shared_ptr<SampleCall> call) override;
private:
DISALLOW_COPY_AND_ASSIGN(SampleServiceImpl);
LLMMaster* master_ = nullptr;
};
} // namespace xllm

View File

@@ -0,0 +1,124 @@
/* Copyright 2026 The xLLM 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 "api_service/service_impl_factory.h"
#include <glog/logging.h>
#include <functional>
#include <unordered_map>
#include "api_service.h"
#include "api_service/serving_mode.h"
#include "core/common/global_flags.h"
#include "core/distributed_runtime/dit_master.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/distributed_runtime/rec_master.h"
#include "core/distributed_runtime/vlm_master.h"
namespace xllm {
namespace {
template <typename T, typename MasterT>
std::unique_ptr<T> create_service_impl(
MasterT* master,
const std::vector<std::string>& model_names) {
return std::make_unique<T>(master, model_names);
}
} // namespace
void ServiceImplFactory::create(
APIService* service,
Master* master,
const std::vector<std::string>& model_names,
const std::vector<std::string>& model_versions) {
using InitFn = std::function<void(
APIService*, Master*, const std::vector<std::string>&)>;
static const std::unordered_map<int8_t, InitFn> kRegistry = {
{static_cast<int8_t>(ServingMode::LLM),
[](APIService* self,
Master* master,
const std::vector<std::string>& models) {
auto* llm_master = dynamic_cast<LLMMaster*>(master);
self->anthropic_service_impl_ =
std::make_unique<AnthropicServiceImpl>(llm_master, models);
self->completion_service_impl_ =
create_service_impl<CompletionServiceImpl>(llm_master, models);
self->sample_service_impl_ =
create_service_impl<SampleServiceImpl>(llm_master, models);
self->chat_service_impl_ =
create_service_impl<ChatServiceImpl>(llm_master, models);
self->embedding_service_impl_ =
create_service_impl<EmbeddingServiceImpl>(llm_master, models);
if (FLAGS_enable_qwen3_reranker) {
self->rerank_service_impl_ =
create_service_impl<Qwen3RerankServiceImpl>(llm_master, models);
} else {
self->rerank_service_impl_ =
create_service_impl<RerankServiceImpl>(llm_master, models);
}
}},
{static_cast<int8_t>(ServingMode::VLM),
[](APIService* self,
Master* master,
const std::vector<std::string>& models) {
auto* vlm_master = dynamic_cast<VLMMaster*>(master);
self->mm_chat_service_impl_ =
std::make_unique<MMChatServiceImpl>(vlm_master, models);
self->mm_embedding_service_impl_ =
std::make_unique<MMEmbeddingServiceImpl>(vlm_master, models);
}},
{static_cast<int8_t>(ServingMode::DIT),
[](APIService* self,
Master* master,
const std::vector<std::string>& models) {
self->image_generation_service_impl_ =
std::make_unique<ImageGenerationServiceImpl>(
dynamic_cast<DiTMaster*>(master), models);
}},
{static_cast<int8_t>(ServingMode::REC),
[](APIService* self,
Master* master,
const std::vector<std::string>& models) {
auto* rec_master = dynamic_cast<RecMaster*>(master);
self->rec_completion_service_impl_ =
std::make_unique<RecCompletionServiceImpl>(rec_master, models);
self->chat_service_impl_ =
std::make_unique<ChatServiceImpl>(rec_master, models);
}},
};
ServingMode mode = to_serving_mode(master->engine_type());
auto it = kRegistry.find(static_cast<int8_t>(mode));
if (it != kRegistry.end()) {
it->second(service, master, model_names);
} else {
LOG(FATAL) << "Unsupported serving mode for engine type: "
<< master->engine_type().to_string();
}
CHECK_EQ(model_names.size(), model_versions.size())
<< "Models and model_versions size mismatch: model_names.size()="
<< model_names.size()
<< ", model_versions.size()=" << model_versions.size();
service->models_service_impl_ =
std::make_unique<ModelsServiceImpl>(model_names, model_versions);
}
} // namespace xllm

View File

@@ -0,0 +1,37 @@
/* Copyright 2026 The xLLM 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.
==============================================================================*/
#pragma once
#include <string>
#include <vector>
namespace xllm {
class APIService;
class Master;
// Creates all service-impl instances that an APIService needs for the active
// engine type. Adding a new engine type only requires one new entry in the
// registry defined in service_impl_factory.cpp.
class ServiceImplFactory {
public:
static void create(APIService* service,
Master* master,
const std::vector<std::string>& model_names,
const std::vector<std::string>& model_versions);
};
} // namespace xllm

View File

@@ -0,0 +1,49 @@
/* Copyright 2026 The xLLM 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.
==============================================================================*/
#pragma once
#include <cstdint>
#include "core/common/types.h"
namespace xllm {
// Service-layer abstraction of the active serving pipeline.
// Unlike EngineType (which includes engine-internal variants such as SSM),
// ServingMode only exposes the distinctions that matter to the API surface.
enum class ServingMode : int8_t {
LLM = 0,
VLM = 1,
DIT = 2,
REC = 3,
};
// Maps an engine-layer EngineType to its corresponding ServingMode.
// SSM (speculative decoding) serves the same API as LLM.
inline ServingMode to_serving_mode(EngineType engine_type) {
switch (static_cast<EngineType::Value>(engine_type)) {
case EngineType::VLM:
return ServingMode::VLM;
case EngineType::DIT:
return ServingMode::DIT;
case EngineType::REC:
return ServingMode::REC;
default:
return ServingMode::LLM;
}
}
} // namespace xllm

View File

@@ -0,0 +1,223 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <brpc/controller.h>
#include <butil/iobuf.h>
#include <glog/logging.h>
#include <json2pb/pb_to_json.h>
#include <atomic>
#include <memory>
#include <optional>
#include <string>
#include "api_service/call.h"
#include "core/common/types.h"
namespace xllm {
template <typename Request, typename Response>
class StreamCall : public Call {
public:
using ReqType = Request;
using ResType = Response;
StreamCall(brpc::Controller* controller,
::google::protobuf::Closure* done,
Request* request,
Response* response,
bool use_arena = false)
: Call(controller),
done_(done),
request_(request),
response_(response),
use_arena_(use_arena) {
stream_ = request_->stream();
if (stream_) {
pa_ = controller_->CreateProgressiveAttachment();
// Send the first SSE response
controller_->http_response().set_content_type(
"text/event-stream; charset=utf-8");
controller_->http_response().set_status_code(200);
controller_->http_response().SetHeader("Connection", "keep-alive");
controller_->http_response().SetHeader("Cache-Control", "no-cache");
// Done Run first for steam response
done_->Run();
} else {
controller_->http_response().set_content_type("application/json");
}
json_options_.bytes_to_base64 = false;
json_options_.jsonify_empty_array = true;
}
~StreamCall() override {
// For non stream response, call brpc done Run
if (!stream_) {
done_->Run();
}
if (!use_arena_) {
delete request_;
delete response_;
}
}
bool write_and_finish(Response& response) {
butil::IOBufAsZeroCopyOutputStream json_output(
&controller_->response_attachment());
std::string err_msg;
if (!json2pb::ProtoMessageToJson(
response, &json_output, json_options_, &err_msg)) {
return finish_with_error(StatusCode::UNKNOWN, err_msg);
}
return true;
}
bool finish_with_error(const StatusCode& code,
const std::string& error_message) {
if (!stream_) {
controller_->SetFailed(error_message);
} else {
io_buf_.clear();
io_buf_.append(error_message);
pa_->Write(io_buf_);
}
return true;
}
// For stream response
bool write(Response& response) {
io_buf_.clear();
io_buf_.append("data: ");
butil::IOBufAsZeroCopyOutputStream json_output(&io_buf_);
std::string err_msg;
if (!json2pb::ProtoMessageToJson(
response, &json_output, json_options_, &err_msg)) {
LOG(ERROR) << "Failed to convert proto to json: " << err_msg;
return false;
}
io_buf_.append("\n\n");
connection_status_ |= pa_->Write(io_buf_);
return true;
}
// For stream response
bool finish() {
io_buf_.clear();
io_buf_.append("data: [DONE]\n\n");
pa_->Write(io_buf_);
return true;
}
bool is_disconnected() const override {
if (stream_) {
return connection_status_ != 0;
} else {
if (controller_) {
return controller_->IsCanceled();
}
return true;
}
}
const Request& request() const { return *request_; }
Response& response() { return *response_; }
::google::protobuf::Closure* done() { return done_; }
protected:
::google::protobuf::Closure* done_;
Request* request_ = nullptr;
Response* response_ = nullptr;
bool stream_ = false;
bool use_arena_ = false;
butil::intrusive_ptr<brpc::ProgressiveAttachment> pa_;
butil::IOBuf io_buf_;
json2pb::Pb2JsonOptions json_options_;
int connection_status_ = 0;
};
// Anthropic SSE stream call with custom event formatting
class AnthropicCall : public StreamCall<proto::AnthropicMessagesRequest,
proto::AnthropicMessagesResponse> {
public:
AnthropicCall(brpc::Controller* controller,
::google::protobuf::Closure* done,
proto::AnthropicMessagesRequest* request,
proto::AnthropicMessagesResponse* response,
bool use_arena = false)
: StreamCall<proto::AnthropicMessagesRequest,
proto::AnthropicMessagesResponse>(controller,
done,
request,
response,
use_arena) {}
~AnthropicCall() {}
// Write SSE event with Anthropic format: event: <type>\ndata: <json>\n\n
bool write(const std::string& event_type, const std::string& json_data) {
this->io_buf_.clear();
this->io_buf_.append("event: ");
this->io_buf_.append(event_type);
this->io_buf_.append("\ndata: ");
this->io_buf_.append(json_data);
this->io_buf_.append("\n\n");
this->connection_status_ |= this->pa_->Write(this->io_buf_);
return this->connection_status_ == 0;
}
// Write SSE event with proto message
template <typename ProtoMessage>
bool write(const std::string& event_type, const ProtoMessage& message) {
this->io_buf_.clear();
this->io_buf_.append("event: ");
this->io_buf_.append(event_type);
this->io_buf_.append("\ndata: ");
butil::IOBufAsZeroCopyOutputStream json_output(&this->io_buf_);
std::string err_msg;
if (!json2pb::ProtoMessageToJson(
message, &json_output, this->json_options_, &err_msg)) {
LOG(ERROR) << "Failed to convert proto to json: " << err_msg;
return false;
}
this->io_buf_.append("\n\n");
this->connection_status_ |= this->pa_->Write(this->io_buf_);
return this->connection_status_ == 0;
}
};
template <typename T>
struct is_stream_call : std::false_type {};
template <typename... Args>
struct is_stream_call<StreamCall<Args...>> : std::true_type {};
template <typename T>
inline constexpr bool is_stream_call_v = is_stream_call<T>::value;
} // namespace xllm

View File

@@ -0,0 +1,87 @@
/* 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 "stream_output_parser.h"
namespace xllm {
StreamOutputParser::StreamOutputParser(
const std::vector<function_call::JsonTool>& tools,
const std::string& tool_call_parser_format,
const std::string& reasoning_parser_format,
bool force_reasoning)
: tools_(tools),
tool_call_parser_format_(tool_call_parser_format),
reasoning_parser_format_(reasoning_parser_format),
force_reasoning_(force_reasoning) {
sequence_parsers_.resize(1);
if (is_tool_call()) {
sequence_parsers_[0].tool_call_parser =
std::make_unique<function_call::FunctionCallParser>(
tools_, tool_call_parser_format_);
}
if (is_reasoning()) {
sequence_parsers_[0].reasoning_parser_ = std::make_unique<ReasoningParser>(
reasoning_parser_format_, true, force_reasoning_);
}
}
bool StreamOutputParser::is_tool_call() {
return !tools_.empty() && !tool_call_parser_format_.empty();
}
bool StreamOutputParser::is_reasoning() {
return !reasoning_parser_format_.empty();
}
void StreamOutputParser::check_resize_for_index(size_t index) {
if (index >= sequence_parsers_.size()) {
sequence_parsers_.resize(index + 1);
}
}
function_call::FunctionCallParser* StreamOutputParser::get_tool_call_parser(
size_t index) {
if (!is_tool_call()) {
return nullptr;
}
check_resize_for_index(index);
if (!sequence_parsers_[index].tool_call_parser) {
sequence_parsers_[index].tool_call_parser =
std::make_unique<function_call::FunctionCallParser>(
tools_, tool_call_parser_format_);
}
return sequence_parsers_[index].tool_call_parser.get();
}
ReasoningParser* StreamOutputParser::get_reasoning_parser(size_t index) {
if (!is_reasoning()) {
return nullptr;
}
check_resize_for_index(index);
if (!sequence_parsers_[index].reasoning_parser_) {
sequence_parsers_[index].reasoning_parser_ =
std::make_unique<ReasoningParser>(
reasoning_parser_format_, true, force_reasoning_);
}
return sequence_parsers_[index].reasoning_parser_.get();
}
} // namespace xllm

View File

@@ -0,0 +1,69 @@
/* 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.
==============================================================================*/
#pragma once
#include "function_call/function_call.h"
#include "parser/reasoning_parser.h"
namespace xllm {
struct SequenceParser {
std::unique_ptr<function_call::FunctionCallParser> tool_call_parser;
bool has_tool_call = false;
std::unique_ptr<ReasoningParser> reasoning_parser_;
};
class StreamOutputParser {
public:
StreamOutputParser(const std::vector<function_call::JsonTool>& tools,
const std::string& tool_call_parser_format,
const std::string& reasoning_parser_format,
bool force_reasoning = false);
~StreamOutputParser() = default;
bool is_tool_call();
bool is_reasoning();
void check_resize_for_index(size_t index);
function_call::FunctionCallParser* get_tool_call_parser(size_t index);
ReasoningParser* get_reasoning_parser(size_t index);
bool get_has_tool_call(size_t index) {
check_resize_for_index(index);
return sequence_parsers_[index].has_tool_call;
}
void set_has_tool_call(size_t index, bool has_tool_call) {
check_resize_for_index(index);
sequence_parsers_[index].has_tool_call = has_tool_call;
}
private:
// candidate tools of requets
std::vector<function_call::JsonTool> tools_;
// list of parsers for each sequence
std::vector<SequenceParser> sequence_parsers_;
std::string tool_call_parser_format_;
std::string reasoning_parser_format_;
bool force_reasoning_;
};
} // namespace xllm

View File

@@ -0,0 +1,157 @@
/* Copyright 2026 The xLLM 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.
==============================================================================*/
#pragma once
#include <glog/logging.h>
#include <google/protobuf/util/json_util.h>
#include <functional>
#include <memory>
#include <nlohmann/json.hpp>
#include <string>
#include <tuple>
#include <vector>
#include "api_service/stream_output_parser.h"
#include "chat.pb.h"
#include "core/common/types.h"
#include "function_call/function_call.h"
namespace xllm {
namespace api_service {
// Check for unstreamed tool arguments and send them using the provided sender
// This is shared between Chat API and Anthropic API implementations
using SendFunc = std::function<bool(const std::string&, int)>;
inline bool check_for_unstreamed_tool_args(
std::shared_ptr<StreamOutputParser> stream_parser,
size_t index,
SendFunc send_func) {
auto* parser = stream_parser->get_tool_call_parser(index);
if (!parser) {
return true;
}
auto* detector = parser->get_detector();
if (!detector) {
return true;
}
if (!detector->prev_tool_call_arr_.empty() &&
!detector->streamed_args_for_tool_.empty()) {
size_t tool_index = detector->prev_tool_call_arr_.size() - 1;
if (tool_index < detector->streamed_args_for_tool_.size()) {
const auto& expected_args = detector->prev_tool_call_arr_[tool_index];
const std::string& actual_args =
detector->streamed_args_for_tool_[tool_index];
if (expected_args.find("arguments") != expected_args.end()) {
const std::string& expected_call = expected_args.at("arguments");
if (expected_call.length() > actual_args.length()) {
std::string remaining_call =
expected_call.substr(actual_args.length());
if (!remaining_call.empty()) {
return send_func(remaining_call, static_cast<int>(tool_index));
}
}
}
}
}
return true;
}
struct ToolCallResult {
std::optional<google::protobuf::RepeatedPtrField<proto::ToolCall>> tool_calls;
std::string text;
std::string finish_reason;
};
inline ToolCallResult process_tool_calls(
std::string text,
const std::vector<xllm::JsonTool>& tools,
const std::string& parser_format,
std::string finish_reason,
google::protobuf::Arena* arena = nullptr) {
ToolCallResult result;
function_call::FunctionCallParser parser(tools, parser_format);
if (!parser.has_tool_call(text)) {
result.text = std::move(text);
result.finish_reason = std::move(finish_reason);
return result;
}
if (finish_reason == "stop") {
result.finish_reason = "tool_calls";
} else {
result.finish_reason = std::move(finish_reason);
}
try {
auto [parsed_text, call_info_list] = parser.parse_non_stream(text);
result.text = std::move(parsed_text);
google::protobuf::RepeatedPtrField<proto::ToolCall> tool_calls;
for (const auto& call_info : call_info_list) {
proto::ToolCall* tool_call =
arena ? google::protobuf::Arena::CreateMessage<proto::ToolCall>(arena)
: new proto::ToolCall();
tool_call->set_id(function_call::utils::generate_tool_call_id());
tool_call->set_type("function");
auto* function = tool_call->mutable_function();
if (call_info.name) {
function->set_name(*call_info.name);
}
function->set_arguments(call_info.parameters);
tool_calls.AddAllocated(tool_call);
}
result.tool_calls = std::move(tool_calls);
} catch (const std::exception& e) {
LOG(ERROR) << "Tool call parsing error: " << e.what();
}
return result;
}
// Convert google::protobuf::Struct to nlohmann::json
inline nlohmann::json struct_to_json(
const google::protobuf::Struct& pb_struct) {
std::string json_str;
google::protobuf::util::JsonPrintOptions options;
options.preserve_proto_field_names = true;
auto status = google::protobuf::util::MessageToJsonString(
pb_struct, &json_str, options);
if (status.ok()) {
try {
return nlohmann::json::parse(json_str);
} catch (...) {
return nlohmann::json::object();
}
}
return nlohmann::json::object();
}
} // namespace api_service
} // namespace xllm

View File

@@ -0,0 +1,54 @@
/* Copyright 2025 The xLLM 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 <brpc/server.h>
#include <google/protobuf/service.h>
#include "core/common/metrics.h"
namespace xllm {
static void request_in_metric(void* context) {
COUNTER_INC(server_request_in_total);
}
static void request_out_metric(void* context) {
auto ctrl = reinterpret_cast<brpc::Controller*>(context);
if (ctrl == nullptr) {
LOG(ERROR) << "ctrl is nullptr";
return;
}
if (!ctrl->Failed()) {
COUNTER_INC(server_request_total_ok);
} else {
COUNTER_INC(server_request_total_fail);
if (ctrl->ErrorCode() == brpc::ELIMIT) {
COUNTER_INC(server_request_total_limit);
}
}
}
static void device_info_metric() {
// TODO: get cpu device info
GAUGE_SET(xllm_cpu_num, 0);
GAUGE_SET(xllm_cpu_utilization, 0);
// TODO: get gpu device info
GAUGE_SET(xllm_gpu_num, 0);
GAUGE_SET(xllm_gpu_utilization, 0);
}
} // namespace xllm

View File

@@ -0,0 +1,104 @@
### How to compile xllm dynamic library
Run the following command in root directory:
```
python setup.py build --generate-so true
```
If you want to debug, it needs to set DEBUG environment variable.
```
export DEBUG=1
```
### How to install dynamic library
Run installation script xllm/c_api/install.sh, headers and dynamic library will be installed in /usr/local/xllm directory.
```
cd xllm/c_api/tools
sh install.sh
```
You will see the following files in /usr/local/xllm directory:
```
[root@A03-R40-I189-101-4100046]# tree /usr/local/xllm
/usr/local/xllm
|-- include
| |-- llm.h
| |-- default.h
| |-- rec.h
| `-- types.h
`-- lib
`-- libxllm.so
3 directories, 5 files
```
### How to compile c_api examples
GPU builds and NPU builds use different link commands. Replace
`<example>.cpp` and `<example>` with the example source file and output binary
name you want to build.
#### GPU
```
cd xllm/c_api/examples
g++ <example>.cpp -o <example> \
-std=c++17 \
-DUSE_CUDA \
-I/usr/local/xllm/include \
-L/usr/local/xllm/lib \
-lxllm \
-Wl,-rpath=/usr/local/xllm/lib
```
#### NPU
Before compiling or running examples, source the Ascend environment first:
```
source /usr/local/Ascend/ascend-toolkit/set_env.sh
```
Then compile with the extra custom op library used by the NPU build:
```
cd xllm/c_api/examples
g++ <example>.cpp -o <example> \
-std=c++17 \
-DUSE_NPU \
-I/usr/local/xllm/include \
-L/usr/local/xllm/lib \
-L/usr/local/Ascend/ascend-toolkit/latest/opp/vendors/xllm/op_api/lib \
-lxllm \
-lcust_opapi \
-Wl,-rpath=/usr/local/xllm/lib \
-Wl,-rpath=/usr/local/Ascend/ascend-toolkit/latest/opp/vendors/xllm/op_api/lib
```
If `-lcust_opapi` is missing from the NPU link command, the linker may report
undefined references to symbols such as `aclnnBeamSearchGroup` and
`aclnnXAttention`.
### How to run c_api examples
Some examples, such as `simple_rec_completions`, support overriding the target
device from `argv[1]`.
#### NPU
```
./simple_rec_completions npu:14
```
#### GPU
```
./simple_rec_completions cuda:0
```

View File

@@ -0,0 +1,161 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#ifndef XLLM_LLM_DEFAULT_H
#define XLLM_LLM_DEFAULT_H
#ifdef __cplusplus
extern "C" {
#endif
#include "types.h"
const XLLM_InitOptions XLLM_INIT_LLM_OPTIONS_DEFAULT = {
.enable_chunked_prefill = false,
.enable_prefill_sp = false,
.enable_prefix_cache = false,
.enable_disagg_pd = false,
.enable_pd_ooc = false,
.enable_schedule_overlap = false,
.enable_shm = false,
.transfer_listen_port = 26000,
.nnodes = 1,
.node_rank = 0,
.dp_size = 1,
.ep_size = 1,
.block_size = 32,
.max_cache_size = 0,
.max_tokens_per_batch = 20480,
.max_seqs_per_batch = 256,
.max_tokens_per_chunk_for_prefill = 0,
.num_speculative_tokens = 0,
.num_request_handling_threads = 4,
.expert_parallel_degree = 0,
.server_idx = 0,
.max_memory_utilization = 0.9,
.task = "generate",
.communication_backend = "lccl",
.instance_role = "DEFAULT",
.device_ip = "",
.master_node_addr = "127.0.0.1:18899",
.xservice_addr = "",
.instance_name = "",
.kv_cache_transfer_mode = "PUSH",
.log_dir = "",
.draft_model = "",
.draft_devices = ""};
const XLLM_RequestParams XLLM_LLM_REQUEST_PARAMS_DEFAULT = {
.echo = false,
.offline = false,
.logprobs = false,
.ignore_eos = false,
.n = 1,
.max_tokens = 5120,
.best_of = 1,
.ttlt_slo_ms = INT32_MAX,
.ttft_slo_ms = INT32_MAX,
.tpot_slo_ms = INT32_MAX,
.beam_width = 0,
.num_return_sequences = 0,
.top_logprobs = 0,
.top_k = -1,
.top_p = 1.0,
.frequency_penalty = 0.0,
.presence_penalty = 0.0,
.repetition_penalty = 1.0,
.temperature = 0.0,
.request_id = ""};
const XLLM_InitOptions XLLM_INIT_REC_OPTIONS_DEFAULT = {
.enable_chunked_prefill = false,
.enable_prefill_sp = false,
.enable_prefix_cache = false,
.enable_disagg_pd = false,
.enable_pd_ooc = false,
.enable_schedule_overlap = false,
.enable_shm = false,
.enable_graph = true,
.enable_rec_fast_sampler = true,
.enable_prefill_piecewise_graph = true,
.enable_xattention_one_stage = false,
.enable_graph_mode_decode_no_padding = true,
.enable_block_copy_kernel = false,
.enable_topk_sorted = false,
.enable_rec_prefill_only = false,
.transfer_listen_port = 26000,
.nnodes = 1,
.node_rank = 0,
.dp_size = 1,
.ep_size = 1,
.block_size = 1,
.max_cache_size = 1000000,
.max_tokens_per_batch = 4096,
.max_seqs_per_batch = 4,
.max_tokens_per_chunk_for_prefill = 0,
.num_speculative_tokens = 0,
.num_request_handling_threads = 4,
.expert_parallel_degree = 0,
.server_idx = 0,
.beam_width = 128,
.max_decode_rounds = 3,
.max_token_per_req = 1000,
.max_memory_utilization = 0.55,
.rec_worker_max_concurrency = 2,
.task = "generate",
.communication_backend = "lccl",
.instance_role = "DEFAULT",
.device_ip = "",
.master_node_addr = "127.0.0.1:18899",
.xservice_addr = "",
.instance_name = "",
.kv_cache_transfer_mode = "PUSH",
.log_dir = "",
.draft_model = "",
.draft_devices = ""};
const XLLM_RequestParams XLLM_REC_REQUEST_PARAMS_DEFAULT = {
.echo = false,
.offline = false,
.logprobs = false,
.ignore_eos = false,
.n = 1,
.max_tokens = 5120,
.best_of = 1,
.ttlt_slo_ms = INT32_MAX,
.ttft_slo_ms = INT32_MAX,
.tpot_slo_ms = INT32_MAX,
.beam_width = 128,
.num_return_sequences = 0,
.top_logprobs = 0,
.top_k = -1,
.top_p = 1.0,
.frequency_penalty = 0.0,
.presence_penalty = 0.0,
.repetition_penalty = 1.0,
.temperature = 0.0,
.request_id = ""};
#ifdef __cplusplus
}
#endif
#endif // XLLM_LLM_DEFAULT_H

View File

@@ -0,0 +1,956 @@
/* Copyright 2026 The xLLM 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 <unistd.h>
#include <algorithm>
#include <atomic>
#include <chrono>
#include <cmath>
#include <condition_variable>
#include <cstdint>
#include <cstdlib>
#include <cstring>
#include <deque>
#include <filesystem>
#include <fstream>
#include <iomanip>
#include <iostream>
#include <limits>
#include <memory>
#include <mutex>
#include <optional>
#include <random>
#include <sstream>
#include <stdexcept>
#include <string>
#include <thread>
#include <utility>
#include <vector>
#include "rec.h"
namespace {
using Clock = std::chrono::steady_clock;
constexpr int32_t kFixedRequestMaxTokens = 3;
constexpr int32_t kFixedRequestBeamWidth = 128;
constexpr int32_t kFixedRequestTopK = 128;
constexpr int32_t kFixedRequestTopLogprobs = 128;
constexpr bool kFixedRequestLogprobs = true;
constexpr int32_t kFixedMmItemsPerRequest = 2;
constexpr int32_t kMaxClientParallelism = 128;
struct CliOptions {
std::string model_path;
std::string devices = "cuda:0";
std::string master_node_addr = "127.0.0.1:18899";
int32_t prompt_size = 128;
int32_t token_min_size = 1024;
int32_t token_max_size = 1024;
double qps = 1.0;
int32_t duration_s = 60;
int32_t client_threads = 0;
int32_t timeout_ms = 30000;
int32_t mm_min_span = 8;
int32_t mm_max_span = 64;
uint32_t seed = 20260410U;
};
struct ModelConfig {
int32_t hidden_size = 0;
int32_t vocab_size = 0;
int32_t max_position_embeddings = 0;
std::string model_type;
};
struct Metrics {
std::atomic<uint64_t> sent{0};
std::atomic<uint64_t> succeeded{0};
std::atomic<uint64_t> failed{0};
std::atomic<uint64_t> timeout{0};
std::atomic<uint64_t> invalid_request{0};
std::atomic<uint64_t> internal_error{0};
std::atomic<uint64_t> total_prompt_tokens{0};
std::atomic<uint64_t> total_completion_tokens{0};
std::atomic<uint64_t> total_latency_us{0};
std::atomic<uint64_t> max_latency_us{0};
std::atomic<uint64_t> current_in_flight{0};
std::atomic<uint64_t> max_in_flight{0};
};
struct RequestPayload {
std::vector<int32_t> token_ids;
};
struct ScheduledRequest {
uint64_t request_index = 0;
size_t pool_index = 0;
};
class EmbeddingMmDataBuilder {
public:
EmbeddingMmDataBuilder() = default;
const XLLM_MM_Data* Build(
const std::vector<std::pair<uint32_t, uint32_t>>& spans,
int32_t hidden_size,
uint64_t request_index,
std::mt19937& rng) {
Reset();
if (spans.empty()) {
return nullptr;
}
mm_data_.type_mask = static_cast<uint32_t>(XLLM_MM_TYPE_EMBEDDING);
mm_data_.is_dict = false;
items_.reserve(spans.size());
buffers_.reserve(spans.size());
for (size_t item_idx = 0; item_idx < spans.size(); ++item_idx) {
const auto [offset, length] = spans[item_idx];
if (length == 0) {
continue;
}
XLLM_MM_Item item{};
item.type = XLLM_MM_TYPE_EMBEDDING;
item.state.token_pos.offset = offset;
item.state.token_pos.length = length;
item.data.is_single_tensor = true;
item.data.data.tensor.dtype = XLLM_DTYPE_BFLOAT16;
item.data.data.tensor.dims.rank = 2;
item.data.data.tensor.dims.dim[0] = static_cast<int>(length);
item.data.data.tensor.dims.dim[1] = hidden_size;
auto buffer = std::make_unique<uint16_t[]>(
static_cast<size_t>(length) * static_cast<size_t>(hidden_size));
FillEmbeddingBuffer(
buffer.get(), length, hidden_size, request_index, item_idx, rng);
item.data.data.tensor.data = buffer.get();
buffers_.push_back(std::move(buffer));
items_.push_back(item);
}
if (items_.empty()) {
return nullptr;
}
mm_data_.data.items.entries = items_.data();
mm_data_.data.items.entries_size = items_.size();
return &mm_data_;
}
private:
static uint16_t FloatToBFloat16(float value) {
union {
float f32;
uint32_t u32;
} bits;
bits.f32 = value;
return static_cast<uint16_t>(bits.u32 >> 16);
}
void FillEmbeddingBuffer(uint16_t* dst,
uint32_t length,
int32_t hidden_size,
uint64_t request_index,
size_t item_idx,
std::mt19937& rng) {
std::uniform_real_distribution<float> dist(-1.0f, 1.0f);
const size_t element_count =
static_cast<size_t>(length) * static_cast<size_t>(hidden_size);
for (size_t element_idx = 0; element_idx < element_count; ++element_idx) {
const float noise = dist(rng) * 0.03125f;
const float base =
std::sin(static_cast<float>((request_index + 1) * 0.013) +
static_cast<float>(item_idx) * 0.17f +
static_cast<float>(element_idx % hidden_size) * 0.001f);
dst[element_idx] = FloatToBFloat16(base + noise);
}
}
void Reset() {
std::memset(&mm_data_, 0, sizeof(mm_data_));
items_.clear();
buffers_.clear();
}
XLLM_MM_Data mm_data_{};
std::vector<XLLM_MM_Item> items_;
std::vector<std::unique_ptr<uint16_t[]>> buffers_;
};
std::string Trim(const std::string& value) {
const auto begin = value.find_first_not_of(" \t\r\n");
if (begin == std::string::npos) {
return "";
}
const auto end = value.find_last_not_of(" \t\r\n");
return value.substr(begin, end - begin + 1);
}
std::string RemoveJsonComments(std::string content) {
std::string output;
output.reserve(content.size());
bool in_string = false;
bool escape = false;
for (size_t i = 0; i < content.size(); ++i) {
const char ch = content[i];
if (in_string) {
output.push_back(ch);
if (escape) {
escape = false;
} else if (ch == '\\') {
escape = true;
} else if (ch == '"') {
in_string = false;
}
continue;
}
if (ch == '"') {
in_string = true;
output.push_back(ch);
continue;
}
if (ch == '/' && i + 1 < content.size()) {
if (content[i + 1] == '/') {
i += 2;
while (i < content.size() && content[i] != '\n') {
++i;
}
if (i < content.size()) {
output.push_back('\n');
}
continue;
}
if (content[i + 1] == '*') {
i += 2;
while (i + 1 < content.size() &&
!(content[i] == '*' && content[i + 1] == '/')) {
++i;
}
++i;
continue;
}
}
output.push_back(ch);
}
return output;
}
const std::string* FindObjectRange(const std::string& content,
const std::string& key,
size_t* object_begin,
size_t* object_end) {
const std::string pattern = "\"" + key + "\"";
const size_t key_pos = content.find(pattern);
if (key_pos == std::string::npos) {
return nullptr;
}
size_t pos = content.find(':', key_pos + pattern.size());
if (pos == std::string::npos) {
return nullptr;
}
++pos;
while (pos < content.size() &&
std::isspace(static_cast<unsigned char>(content[pos]))) {
++pos;
}
if (pos >= content.size() || content[pos] != '{') {
return nullptr;
}
const size_t begin = pos;
int depth = 0;
bool in_string = false;
bool escape = false;
for (; pos < content.size(); ++pos) {
const char ch = content[pos];
if (in_string) {
if (escape) {
escape = false;
} else if (ch == '\\') {
escape = true;
} else if (ch == '"') {
in_string = false;
}
continue;
}
if (ch == '"') {
in_string = true;
continue;
}
if (ch == '{') {
++depth;
} else if (ch == '}') {
--depth;
if (depth == 0) {
*object_begin = begin;
*object_end = pos + 1;
return &content;
}
}
}
return nullptr;
}
std::optional<std::string> FindJsonStringValueInRange(
const std::string& content,
size_t begin,
size_t end,
const std::string& key) {
const std::string pattern = "\"" + key + "\"";
size_t pos = content.find(pattern, begin);
if (pos == std::string::npos || pos >= end) {
return std::nullopt;
}
pos = content.find(':', pos + pattern.size());
if (pos == std::string::npos || pos >= end) {
return std::nullopt;
}
++pos;
while (pos < end && std::isspace(static_cast<unsigned char>(content[pos]))) {
++pos;
}
if (pos >= end || content[pos] != '"') {
return std::nullopt;
}
++pos;
std::string value;
bool escape = false;
while (pos < end) {
const char ch = content[pos++];
if (escape) {
value.push_back(ch);
escape = false;
continue;
}
if (ch == '\\') {
escape = true;
continue;
}
if (ch == '"') {
return value;
}
value.push_back(ch);
}
return std::nullopt;
}
std::optional<int32_t> FindJsonIntValueInRange(const std::string& content,
size_t begin,
size_t end,
const std::string& key) {
const std::string pattern = "\"" + key + "\"";
size_t pos = content.find(pattern, begin);
if (pos == std::string::npos || pos >= end) {
return std::nullopt;
}
pos = content.find(':', pos + pattern.size());
if (pos == std::string::npos || pos >= end) {
return std::nullopt;
}
++pos;
while (pos < end && std::isspace(static_cast<unsigned char>(content[pos]))) {
++pos;
}
if (pos >= end) {
return std::nullopt;
}
size_t value_end = pos;
if (content[value_end] == '-') {
++value_end;
}
while (value_end < end &&
std::isdigit(static_cast<unsigned char>(content[value_end]))) {
++value_end;
}
if (value_end == pos || (value_end == pos + 1 && content[pos] == '-')) {
return std::nullopt;
}
return std::stoi(content.substr(pos, value_end - pos));
}
ModelConfig LoadModelConfig(const std::string& model_path) {
const std::filesystem::path config_path =
std::filesystem::path(model_path) / "config.json";
std::ifstream ifs(config_path);
if (!ifs.is_open()) {
throw std::runtime_error("failed to open model config: " +
config_path.string());
}
std::stringstream buffer;
buffer << ifs.rdbuf();
const std::string content = RemoveJsonComments(buffer.str());
ModelConfig cfg;
cfg.hidden_size =
FindJsonIntValueInRange(content, 0, content.size(), "hidden_size")
.value_or(0);
cfg.vocab_size =
FindJsonIntValueInRange(content, 0, content.size(), "vocab_size")
.value_or(0);
cfg.max_position_embeddings =
FindJsonIntValueInRange(
content, 0, content.size(), "max_position_embeddings")
.value_or(0);
cfg.model_type =
FindJsonStringValueInRange(content, 0, content.size(), "model_type")
.value_or(FindJsonStringValueInRange(
content, 0, content.size(), "model_name")
.value_or(""));
size_t text_begin = 0;
size_t text_end = 0;
if (FindObjectRange(content, "text_config", &text_begin, &text_end) !=
nullptr) {
if (cfg.hidden_size <= 0) {
cfg.hidden_size =
FindJsonIntValueInRange(content, text_begin, text_end, "hidden_size")
.value_or(0);
}
if (cfg.vocab_size <= 0) {
cfg.vocab_size =
FindJsonIntValueInRange(content, text_begin, text_end, "vocab_size")
.value_or(0);
}
if (cfg.max_position_embeddings <= 0) {
cfg.max_position_embeddings =
FindJsonIntValueInRange(
content, text_begin, text_end, "max_position_embeddings")
.value_or(0);
}
}
if (cfg.hidden_size <= 0) {
throw std::runtime_error(
"config.json missing hidden_size/text_config.hidden_size");
}
if (cfg.vocab_size <= 0) {
throw std::runtime_error(
"config.json missing vocab_size/text_config.vocab_size");
}
if (cfg.max_position_embeddings <= 0) {
throw std::runtime_error(
"config.json missing "
"max_position_embeddings/text_config.max_position_embeddings");
}
return cfg;
}
std::string ResolveModelId(const std::string& model_path) {
std::filesystem::path path =
std::filesystem::path(model_path).lexically_normal();
if (path.has_filename()) {
return path.filename().string();
}
return path.string();
}
void PrintUsage(const char* argv0) {
std::cerr
<< "Usage: " << argv0 << " --model_path PATH [options]\n"
<< "Options:\n"
<< " --master_node_addr STR default: 127.0.0.1:18899\n"
<< " --prompt_size N number of requests kept in pool, default: "
"128\n"
<< " --token_min_size N minimum prompt token length, default: "
"1024\n"
<< " --token_max_size N maximum prompt token length, default: "
"1024\n"
<< " --qps FLOAT target QPS, default: 1.0\n"
<< " --client_threads N concurrent client request threads; 0 "
"means auto (ceil(qps))\n"
<< " --duration_s N duration in seconds, default: 60; 0 "
"means run forever\n"
<< " --mm_min_span N default: 8\n"
<< " --mm_max_span N default: 64\n"
<< " --seed N default: 20260410\n";
}
CliOptions ParseArgs(int argc, char** argv) {
CliOptions options;
for (int i = 1; i < argc; ++i) {
const std::string arg = argv[i];
auto next = [&](const char* name) -> std::string {
if (i + 1 >= argc) {
throw std::runtime_error(std::string("missing value for ") + name);
}
return argv[++i];
};
if (arg == "--model_path") {
options.model_path = next("--model_path");
} else if (arg == "--master_node_addr") {
options.master_node_addr = next("--master_node_addr");
} else if (arg == "--prompt_size") {
options.prompt_size = std::stoi(next("--prompt_size"));
} else if (arg == "--token_min_size") {
options.token_min_size = std::stoi(next("--token_min_size"));
} else if (arg == "--token_max_size") {
options.token_max_size = std::stoi(next("--token_max_size"));
} else if (arg == "--qps") {
options.qps = std::stod(next("--qps"));
} else if (arg == "--client_threads") {
options.client_threads = std::stoi(next("--client_threads"));
} else if (arg == "--duration_s") {
options.duration_s = std::stoi(next("--duration_s"));
} else if (arg == "--mm_min_span") {
options.mm_min_span = std::stoi(next("--mm_min_span"));
} else if (arg == "--mm_max_span") {
options.mm_max_span = std::stoi(next("--mm_max_span"));
} else if (arg == "--seed") {
options.seed = static_cast<uint32_t>(std::stoul(next("--seed")));
} else if (arg == "--help" || arg == "-h") {
PrintUsage(argv[0]);
std::exit(0);
} else {
throw std::runtime_error("unknown argument: " + arg);
}
}
if (options.model_path.empty()) {
throw std::runtime_error("--model_path is required");
}
if (options.prompt_size <= 0) {
throw std::runtime_error("--prompt_size must be > 0");
}
if (options.token_min_size <= 0 || options.token_max_size <= 0) {
throw std::runtime_error(
"--token_min_size and --token_max_size must be > 0");
}
if (options.token_min_size > options.token_max_size) {
throw std::runtime_error("--token_min_size must be <= --token_max_size");
}
if (options.qps <= 0.0) {
throw std::runtime_error("--qps must be > 0");
}
if (options.client_threads < 0) {
throw std::runtime_error("--client_threads must be >= 0");
}
if (options.duration_s < 0) {
throw std::runtime_error("--duration_s must be >= 0");
}
if (options.mm_min_span <= 0 || options.mm_max_span <= 0 ||
options.mm_min_span > options.mm_max_span) {
throw std::runtime_error("invalid mm span range");
}
return options;
}
std::vector<int32_t> MakeRandomTokens(int32_t token_size,
int32_t vocab_size,
std::mt19937& rng) {
std::uniform_int_distribution<int32_t> dist(1, std::max(2, vocab_size - 1));
std::vector<int32_t> token_ids(token_size);
for (int32_t& token_id : token_ids) {
token_id = dist(rng);
}
return token_ids;
}
std::vector<std::pair<uint32_t, uint32_t>> MakeRandomMmSpans(
int32_t token_size,
const CliOptions& options,
std::mt19937& rng) {
std::vector<std::pair<uint32_t, uint32_t>> spans;
spans.reserve(kFixedMmItemsPerRequest);
const int32_t max_valid_end = token_size - 1;
int32_t cursor = 1;
for (int32_t item_idx = 0; item_idx < kFixedMmItemsPerRequest; ++item_idx) {
const int32_t remaining_items = kFixedMmItemsPerRequest - item_idx;
const int32_t remaining_tokens = max_valid_end - cursor;
if (remaining_tokens <= options.mm_min_span) {
break;
}
const int32_t max_span =
std::min(options.mm_max_span, remaining_tokens - remaining_items + 1);
if (max_span < options.mm_min_span) {
break;
}
std::uniform_int_distribution<int32_t> span_dist(options.mm_min_span,
max_span);
const int32_t length = span_dist(rng);
const int32_t max_offset =
max_valid_end - length - (remaining_items - 1) * options.mm_min_span;
if (max_offset < cursor) {
break;
}
std::uniform_int_distribution<int32_t> offset_dist(cursor, max_offset);
const int32_t offset = offset_dist(rng);
spans.emplace_back(static_cast<uint32_t>(offset),
static_cast<uint32_t>(length));
cursor = offset + length + 1;
}
if (spans.empty()) {
const uint32_t fallback_len =
static_cast<uint32_t>(std::min(options.mm_min_span, token_size - 2));
spans.emplace_back(1U, std::max<uint32_t>(1U, fallback_len));
}
return spans;
}
std::vector<RequestPayload> BuildRequestPool(const CliOptions& options,
const ModelConfig& config) {
std::mt19937 rng(options.seed);
std::vector<RequestPayload> pool;
pool.reserve(options.prompt_size);
const int32_t max_model_token_size = config.max_position_embeddings - 1;
const int32_t token_min_size =
std::min(options.token_min_size, max_model_token_size);
const int32_t token_max_size =
std::min(options.token_max_size, max_model_token_size);
if (token_min_size <= 1 || token_max_size <= 1) {
throw std::runtime_error(
"token_size range is too large for model max_position_embeddings");
}
if (token_min_size > token_max_size) {
throw std::runtime_error(
"token_size range becomes invalid after clamping to model limits");
}
std::uniform_int_distribution<int32_t> token_size_dist(token_min_size,
token_max_size);
for (int32_t i = 0; i < options.prompt_size; ++i) {
const int32_t token_size = token_size_dist(rng);
RequestPayload payload;
payload.token_ids = MakeRandomTokens(token_size, config.vocab_size, rng);
pool.push_back(std::move(payload));
}
return pool;
}
void UpdateMax(std::atomic<uint64_t>& target, uint64_t value) {
uint64_t prev = target.load(std::memory_order_relaxed);
while (
prev < value &&
!target.compare_exchange_weak(
prev, value, std::memory_order_relaxed, std::memory_order_relaxed)) {
}
}
void RecordResponseMetrics(const XLLM_Response* resp,
uint64_t latency_us,
Metrics* metrics) {
metrics->sent.fetch_add(1, std::memory_order_relaxed);
metrics->total_latency_us.fetch_add(latency_us, std::memory_order_relaxed);
UpdateMax(metrics->max_latency_us, latency_us);
if (resp == nullptr) {
metrics->failed.fetch_add(1, std::memory_order_relaxed);
return;
}
metrics->total_prompt_tokens.fetch_add(
static_cast<uint64_t>(std::max(resp->usage.prompt_tokens, 0)),
std::memory_order_relaxed);
metrics->total_completion_tokens.fetch_add(
static_cast<uint64_t>(std::max(resp->usage.completion_tokens, 0)),
std::memory_order_relaxed);
if (resp->status_code == kSuccess) {
metrics->succeeded.fetch_add(1, std::memory_order_relaxed);
return;
}
metrics->failed.fetch_add(1, std::memory_order_relaxed);
if (resp->status_code == kTimeout) {
metrics->timeout.fetch_add(1, std::memory_order_relaxed);
} else if (resp->status_code == kInvalidRequest) {
metrics->invalid_request.fetch_add(1, std::memory_order_relaxed);
} else if (resp->status_code == kInternalError) {
metrics->internal_error.fetch_add(1, std::memory_order_relaxed);
}
}
void IncrementInFlight(Metrics* metrics) {
const uint64_t current =
metrics->current_in_flight.fetch_add(1, std::memory_order_relaxed) + 1;
UpdateMax(metrics->max_in_flight, current);
}
void DecrementInFlight(Metrics* metrics) {
metrics->current_in_flight.fetch_sub(1, std::memory_order_relaxed);
}
void PrintSummary(const CliOptions& options,
const ModelConfig& config,
const Metrics& metrics,
double actual_duration_s) {
const uint64_t sent = metrics.sent.load(std::memory_order_relaxed);
const uint64_t succeeded = metrics.succeeded.load(std::memory_order_relaxed);
const uint64_t failed = metrics.failed.load(std::memory_order_relaxed);
const uint64_t total_latency_us =
metrics.total_latency_us.load(std::memory_order_relaxed);
const uint64_t max_latency_us =
metrics.max_latency_us.load(std::memory_order_relaxed);
const uint64_t max_in_flight =
metrics.max_in_flight.load(std::memory_order_relaxed);
const double avg_latency_ms = sent == 0
? 0.0
: static_cast<double>(total_latency_us) /
static_cast<double>(sent) / 1000.0;
const double actual_qps = actual_duration_s <= 0.0
? 0.0
: static_cast<double>(sent) / actual_duration_s;
std::cout << "=== stress_rec_multimodal_completions summary ===\n";
std::cout << "model_id=" << ResolveModelId(options.model_path) << "\n";
std::cout << "model_type=" << config.model_type << "\n";
std::cout << "devices=" << options.devices << "\n";
std::cout << "master_node_addr=" << options.master_node_addr << "\n";
std::cout << "hidden_size=" << config.hidden_size
<< ", vocab_size=" << config.vocab_size
<< ", max_position_embeddings=" << config.max_position_embeddings
<< "\n";
std::cout << "prompt_size=" << options.prompt_size << ", token_size_range=["
<< options.token_min_size << ", " << options.token_max_size << "]"
<< ", qps_target=" << options.qps
<< ", duration_s=" << options.duration_s << "\n";
std::cout << "client_threads="
<< (options.client_threads == 0
? static_cast<int32_t>(std::ceil(options.qps))
: options.client_threads)
<< "\n";
std::cout << "mm_span_range=[" << options.mm_min_span << ", "
<< options.mm_max_span
<< "], mm_items_per_request=" << kFixedMmItemsPerRequest << "\n";
std::cout << "sent=" << sent << ", succeeded=" << succeeded
<< ", failed=" << failed << "\n";
std::cout << "max_client_in_flight=" << max_in_flight << "\n";
std::cout << "timeout=" << metrics.timeout.load(std::memory_order_relaxed)
<< ", invalid_request="
<< metrics.invalid_request.load(std::memory_order_relaxed)
<< ", internal_error="
<< metrics.internal_error.load(std::memory_order_relaxed) << "\n";
std::cout << std::fixed << std::setprecision(3)
<< "actual_duration_s=" << actual_duration_s
<< ", actual_qps=" << actual_qps
<< ", avg_latency_ms=" << avg_latency_ms << ", max_latency_ms="
<< static_cast<double>(max_latency_us) / 1000.0 << "\n";
std::cout << "prompt_tokens_total="
<< metrics.total_prompt_tokens.load(std::memory_order_relaxed)
<< ", completion_tokens_total="
<< metrics.total_completion_tokens.load(std::memory_order_relaxed)
<< "\n";
}
} // namespace
int main(int argc, char** argv) {
try {
const CliOptions options = ParseArgs(argc, argv);
const ModelConfig config = LoadModelConfig(options.model_path);
auto request_pool = BuildRequestPool(options, config);
const std::string model_id = ResolveModelId(options.model_path);
std::cout << "Loaded model config from " << options.model_path << "\n";
std::cout << "hidden_size=" << config.hidden_size
<< ", vocab_size=" << config.vocab_size
<< ", max_position_embeddings=" << config.max_position_embeddings
<< ", model_type=" << config.model_type << "\n";
std::cout << "model_id=" << model_id << "\n";
XLLM_REC_Handler* rec_handler = xllm_rec_create();
if (rec_handler == nullptr) {
throw std::runtime_error("xllm_rec_create returned nullptr");
}
XLLM_InitOptions init_options;
xllm_rec_init_options_default(&init_options);
std::snprintf(init_options.master_node_addr,
sizeof(init_options.master_node_addr),
"%s",
options.master_node_addr.c_str());
const bool init_ok = xllm_rec_initialize(rec_handler,
options.model_path.c_str(),
options.devices.c_str(),
&init_options);
if (!init_ok) {
xllm_rec_destroy(rec_handler);
throw std::runtime_error("xllm_rec_initialize failed");
}
XLLM_RequestParams request_params;
xllm_rec_request_params_default(&request_params);
request_params.max_tokens = kFixedRequestMaxTokens;
request_params.beam_width = kFixedRequestBeamWidth;
request_params.logprobs = kFixedRequestLogprobs;
request_params.top_k = kFixedRequestTopK;
request_params.top_logprobs = kFixedRequestTopLogprobs;
std::cout << "Initialized REC with fixed request params: max_tokens="
<< request_params.max_tokens
<< ", beam_width=" << request_params.beam_width
<< ", logprobs=" << (request_params.logprobs ? 1 : 0)
<< ", top_k=" << request_params.top_k
<< ", top_logprobs=" << request_params.top_logprobs << "\n";
Metrics metrics;
const auto start_time = Clock::now();
const bool run_forever = options.duration_s == 0;
const auto stop_time =
run_forever ? Clock::time_point::max()
: start_time + std::chrono::seconds(options.duration_s);
const double interval_us = 1e6 / options.qps;
std::mutex queue_mutex;
std::condition_variable queue_cv;
std::deque<ScheduledRequest> queue;
bool producer_done = false;
const int32_t auto_worker_count =
std::max<int32_t>(1, static_cast<int32_t>(std::ceil(options.qps)));
const int32_t worker_count = std::max<int32_t>(
1,
std::min<int32_t>(kMaxClientParallelism,
options.client_threads > 0 ? options.client_threads
: auto_worker_count));
std::cout << "Using client_threads=" << worker_count
<< (options.client_threads > 0 ? " (explicit)" : " (auto)")
<< "\n";
auto worker_fn = [&](int32_t worker_id) {
std::mt19937 rng(options.seed ^ 0x9e3779b9U ^
static_cast<uint32_t>(worker_id * 0x85ebca6bU));
EmbeddingMmDataBuilder mm_builder;
while (true) {
ScheduledRequest scheduled;
{
std::unique_lock<std::mutex> lock(queue_mutex);
queue_cv.wait(lock,
[&]() { return producer_done || !queue.empty(); });
if (queue.empty()) {
if (producer_done) {
return;
}
continue;
}
scheduled = queue.front();
queue.pop_front();
}
const RequestPayload& payload = request_pool[scheduled.pool_index];
const auto mm_spans = MakeRandomMmSpans(
static_cast<int32_t>(payload.token_ids.size()), options, rng);
const XLLM_MM_Data* mm_data = mm_builder.Build(
mm_spans, config.hidden_size, scheduled.request_index, rng);
IncrementInFlight(&metrics);
const auto req_begin = Clock::now();
XLLM_Response* resp =
xllm_rec_multimodal_completions(rec_handler,
model_id.c_str(),
payload.token_ids.data(),
payload.token_ids.size(),
mm_data,
options.timeout_ms,
&request_params);
const auto req_end = Clock::now();
DecrementInFlight(&metrics);
const uint64_t latency_us = static_cast<uint64_t>(
std::chrono::duration_cast<std::chrono::microseconds>(req_end -
req_begin)
.count());
RecordResponseMetrics(resp, latency_us, &metrics);
if (resp != nullptr && resp->status_code != kSuccess) {
std::cerr << "request " << scheduled.request_index
<< " failed: status=" << resp->status_code
<< ", error=" << resp->error_info << "\n";
std::cerr << "request " << scheduled.request_index
<< " shape: token_size=" << payload.token_ids.size()
<< ", mm_items=" << mm_spans.size();
for (const auto& [offset, length] : mm_spans) {
std::cerr << " [" << offset << "," << length << "]";
}
std::cerr << "\n";
}
xllm_rec_free_response(resp);
}
};
std::vector<std::thread> workers;
workers.reserve(static_cast<size_t>(worker_count));
for (int32_t worker_id = 0; worker_id < worker_count; ++worker_id) {
workers.emplace_back(worker_fn, worker_id);
}
uint64_t request_index = 0;
size_t next_pool_index = 0;
while (Clock::now() < stop_time) {
const auto scheduled_time =
start_time + std::chrono::microseconds(
static_cast<int64_t>(request_index * interval_us));
std::this_thread::sleep_until(scheduled_time);
{
std::lock_guard<std::mutex> lock(queue_mutex);
queue.push_back(ScheduledRequest{request_index, next_pool_index});
}
queue_cv.notify_one();
next_pool_index = (next_pool_index + 1) % request_pool.size();
++request_index;
}
{
std::lock_guard<std::mutex> lock(queue_mutex);
producer_done = true;
}
queue_cv.notify_all();
for (auto& worker : workers) {
worker.join();
}
const double actual_duration_s =
std::chrono::duration_cast<std::chrono::duration<double>>(Clock::now() -
start_time)
.count();
PrintSummary(options, config, metrics, actual_duration_s);
xllm_rec_destroy(rec_handler);
return 0;
} catch (const std::exception& e) {
std::cerr << "fatal: " << e.what() << "\n";
return 1;
}
}

View File

@@ -0,0 +1,113 @@
/* Copyright 2025 The xLLM 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 <unistd.h>
#include <cstring>
#include <iostream>
#include "llm.h"
std::string devices = "cuda:4";
std::string model_name = "Qwen3-8B";
std::string model_path = "/export/home/models/Qwen3-8B";
XLLM_LLM_Handler* service_startup_hook() {
XLLM_LLM_Handler* llm_handler = xllm_llm_create();
// If there is no separate setting, init_options can be passed as nullptr, and
// the default value(XLLM_INIT_LLM_OPTIONS_DEFAULT) will be used
XLLM_InitLLMOptions init_options;
xllm_llm_init_options_default(&init_options);
snprintf(
init_options.log_dir, sizeof(init_options.log_dir), "/export/xllm/log");
bool ret = xllm_llm_initialize(
llm_handler, model_path.c_str(), devices.c_str(), &init_options);
if (!ret) {
std::cout << "LLM init failed" << std::endl;
xllm_llm_destroy(llm_handler);
return nullptr;
}
std::cout << "LLM init successfully" << std::endl;
return llm_handler;
}
void service_stop_hook(XLLM_LLM_Handler* llm_handler) {
xllm_llm_destroy(llm_handler);
std::cout << "LLM stop" << std::endl;
}
int main(int argc, char** argv) {
XLLM_LLM_Handler* llm_handler = service_startup_hook();
if (nullptr == llm_handler) {
return -1;
}
// If there is no separate setting, request_params can be passed as nullptr,
// and the default value(XLLM_REQUEST_PARAMS_DEFAULT) will be used
XLLM_RequestParams request_params;
xllm_llm_request_params_default(&request_params);
request_params.max_tokens = 300;
std::string content =
"You are an expert in e-commerce scenarios. The current scenario is an "
"e-commerce search engine with a comprehensive range of business "
"categories. Your task is to determine whether 'user query' and "
"'product title' are related in the e-commerce search engine. "
"Discrimination criteria: If the search for 'user query' returns' "
"product title 'that meets the user's needs, then the task is "
"relevant. Output requirement: Please provide the answer in the "
"'related' or 'unrelated' section, without mentioning any other "
"content. User query: 'Hotpot sauce'. Product title: 'Grassland Red "
"Sun Hotpot Base Dip Multi flavored Barbecue Sauce Tomato Sauce Leek "
"Flower Sauce Nightsnack Paired with [New] Spicy Barbecue Sauce 100g'";
XLLM_ChatMessage message = {0};
strncpy(message.role, "user", sizeof(message.role) - 1);
message.content = const_cast<char*>(content.c_str());
XLLM_Response* resp = xllm_llm_chat_completions(
llm_handler, model_name.c_str(), &message, 1, 10000, &request_params);
if (nullptr == resp) {
std::cout << "LLM completions failed, response is nullptr" << std::endl;
service_stop_hook(llm_handler);
return -1;
}
if (resp->status_code != XLLM_StatusCode::kSuccess) {
std::cout << "LLM completions failed, status code:" << resp->status_code
<< ", error info:" << resp->error_info << std::endl;
} else {
std::cout << "LLM completions successfully" << std::endl;
if (nullptr != resp->choices.entries) {
for (int i = 0; i < resp->choices.entries_size; ++i) {
XLLM_Choice& choice = resp->choices.entries[i];
std::cout << "xllm answer[" << choice.index
<< "]:" << "role:" << choice.message->role
<< ",content:" << choice.message->content << std::endl;
}
}
}
xllm_llm_free_response(resp);
service_stop_hook(llm_handler);
return 0;
}

View File

@@ -0,0 +1,97 @@
/* Copyright 2025 The xLLM 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 <unistd.h>
#include <cstring>
#include <iostream>
#include "llm.h"
std::string devices = "cuda:1";
std::string model_name = "Qwen3-8B";
std::string model_path = "/export/home/models/Qwen3-8B";
XLLM_LLM_Handler* service_startup_hook() {
XLLM_LLM_Handler* llm_handler = xllm_llm_create();
// If there is no separate setting, init_options can be passed as nullptr, and
// the default value(XLLM_INIT_LLM_OPTIONS_DEFAULT) will be used
XLLM_InitOptions init_options;
xllm_llm_init_options_default(&init_options);
snprintf(
init_options.log_dir, sizeof(init_options.log_dir), "/export/xllm/log");
bool ret = xllm_llm_initialize(
llm_handler, model_path.c_str(), devices.c_str(), &init_options);
if (!ret) {
std::cout << "LLM init failed" << std::endl;
xllm_llm_destroy(llm_handler);
return nullptr;
}
std::cout << "LLM init successfully" << std::endl;
return llm_handler;
}
void service_stop_hook(XLLM_LLM_Handler* llm_handler) {
xllm_llm_destroy(llm_handler);
std::cout << "LLM stop" << std::endl;
}
int main(int argc, char** argv) {
XLLM_LLM_Handler* llm_handler = service_startup_hook();
if (nullptr == llm_handler) {
return -1;
}
// If there is no separate setting, request_params can be passed as nullptr,
// and the default value(XLLM_REQUEST_PARAMS_DEFAULT) will be used
XLLM_RequestParams request_params;
xllm_llm_request_params_default(&request_params);
request_params.max_tokens = 300;
std::string prompt = "please briefly introduce XLLM for me";
XLLM_Response* resp = xllm_llm_completions(
llm_handler, model_name.c_str(), prompt.c_str(), 10000, &request_params);
if (nullptr == resp) {
std::cout << "LLM completions failed, response is nullptr" << std::endl;
service_stop_hook(llm_handler);
return -1;
}
if (resp->status_code != XLLM_StatusCode::kSuccess) {
std::cout << "LLM completions failed, status code:" << resp->status_code
<< ", error info:" << resp->error_info << std::endl;
} else {
std::cout << "LLM completions successfully" << std::endl;
if (nullptr != resp->choices.entries) {
for (int i = 0; i < resp->choices.entries_size; ++i) {
XLLM_Choice& choice = resp->choices.entries[i];
std::cout << "xllm answer[" << choice.index << "]:" << choice.text
<< std::endl;
}
}
}
xllm_llm_free_response(resp);
service_stop_hook(llm_handler);
return 0;
}

View File

@@ -0,0 +1,153 @@
/* Copyright 2025 The xLLM 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 <unistd.h>
#include <cstring>
#include <iostream>
#include <random>
#include "rec.h"
#if defined(USE_NPU)
std::string devices = "npu:14";
#elif defined(USE_CUDA)
std::string devices = "cuda:0";
#else
std::string devices = "npu:14";
#endif
std::string model_name = "Qwen3-0.6B";
std::string model_path = "/export/home/models/Qwen3-0.6B";
XLLM_REC_Handler* service_startup_hook() {
XLLM_REC_Handler* rec_handler = xllm_rec_create();
// If there is no separate setting, init_options can be passed as nullptr, and
// the default value(XLLM_INIT_REC_OPTIONS_DEFAULT) will be used
XLLM_InitOptions init_options;
xllm_rec_init_options_default(&init_options);
init_options.block_size = 1;
init_options.max_tokens_per_batch = 8192;
init_options.max_seqs_per_batch = 4;
init_options.max_memory_utilization = 0.8;
init_options.max_cache_size = 500000;
init_options.beam_width = 64;
init_options.max_decode_rounds = 3;
init_options.enable_chunked_prefill = false;
init_options.enable_prefix_cache = false;
#if defined(USE_NPU)
init_options.enable_graph = false;
init_options.enable_graph_mode_decode_no_padding = false;
init_options.enable_prefill_piecewise_graph = false;
init_options.rec_worker_max_concurrency = 1;
#endif
bool ret = xllm_rec_initialize(
rec_handler, model_path.c_str(), devices.c_str(), &init_options);
if (!ret) {
std::cout << "REC init failed" << std::endl;
xllm_rec_destroy(rec_handler);
return nullptr;
}
std::cout << "REC init successfully" << std::endl;
return rec_handler;
}
void service_stop_hook(XLLM_REC_Handler* rec_handler) {
xllm_rec_destroy(rec_handler);
std::cout << "REC stop" << std::endl;
}
int main(int argc, char** argv) {
if (argc > 1) {
devices = argv[1];
}
std::cout << "Using model path: " << model_path << std::endl;
std::cout << "Using devices: " << devices << std::endl;
XLLM_REC_Handler* rec_handler = service_startup_hook();
if (nullptr == rec_handler) {
return -1;
}
// If there is no separate setting, request_params can be passed as nullptr,
// and the default value(XLLM_REQUEST_PARAMS_DEFAULT) will be used
XLLM_RequestParams request_params;
xllm_rec_request_params_default(&request_params);
request_params.max_tokens = 3;
request_params.beam_width = 64;
request_params.logprobs = true;
// request_params.temperature = 1.0;
request_params.top_k = 64;
request_params.top_logprobs = 64;
// request_params.top_p = 1.0;
// request_params.repetition_penalty = 1.0;
// Qwen3-0.6B tokenizer ids for: "where is bejing?".
std::vector<int32_t> token_ids = {2870, 374, 387, 98168, 30};
size_t token_size = token_ids.size();
const int32_t* token_ids_ptr = token_ids.data();
XLLM_Response* resp = xllm_rec_token_completions(rec_handler,
model_name.c_str(),
token_ids_ptr,
token_size,
100000,
&request_params);
if (nullptr == resp) {
std::cout << "REC completions failed, response is nullptr" << std::endl;
service_stop_hook(rec_handler);
return -1;
}
if (resp->status_code != XLLM_StatusCode::kSuccess) {
std::cout << "REC completions failed, status code:" << resp->status_code
<< ", error info:" << resp->error_info << std::endl;
} else {
std::cout << "REC completions successfully, size:"
<< resp->choices.entries_size << std::endl;
if (nullptr != resp->choices.entries) {
for (int i = 0; i < resp->choices.entries_size; ++i) {
XLLM_Choice& choice = resp->choices.entries[i];
std::cout << "token size: " << choice.token_size
<< ",logprobs size:" << choice.logprobs.entries_size
<< std::endl;
for (int j = 0; j < choice.token_size; j++) {
std::cout << "xllm answer[" << choice.index
<< "]: token id=" << choice.token_ids[j] << std::endl;
}
for (int j = 0; j < choice.logprobs.entries_size; j++) {
XLLM_LogProb& logprob = choice.logprobs.entries[j];
std::cout << "xllm answer[" << choice.index
<< "]: token id=" << logprob.token_id
<< ", token logprob=" << logprob.logprob << std::endl;
}
}
}
}
xllm_rec_free_response(resp);
service_stop_hook(rec_handler);
return 0;
}

View File

@@ -0,0 +1,376 @@
/* Copyright 2026 The xLLM 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 <unistd.h>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <iostream>
#include <memory>
#include <random>
#include <stdexcept>
#include <utility>
#include <vector>
#include "rec.h"
std::string devices = "cuda:0";
std::string model_name = "homepage_qwen_06b_6_raw";
std::string model_path = "/export/home/models/homepage_qwen_06b_6_raw";
#define MODEL_WORD_EMBEDDING_SIZE 1024
class XLLM_MM_Data_Wrapper {
public:
XLLM_MM_Data_Wrapper() = default;
~XLLM_MM_Data_Wrapper() { reset(); }
XLLM_MM_Data_Wrapper(const XLLM_MM_Data_Wrapper&) = delete;
XLLM_MM_Data_Wrapper& operator=(const XLLM_MM_Data_Wrapper&) = delete;
bool build(
const std::vector<std::pair<uint32_t, uint32_t>>& token_positions) {
if (is_built_ || token_positions.empty()) {
fprintf(stderr,
"build() failed: already built or empty token positions\n");
return false;
}
mm_data_.type_mask = static_cast<uint32_t>(XLLM_MM_TYPE_EMBEDDING);
mm_data_.is_dict = false;
for (size_t i = 0; i < token_positions.size(); ++i) {
const auto& [offset, length] = token_positions[i];
if (length == 0) {
fprintf(stderr, "build() skipped item %zu: length cannot be 0\n", i);
continue;
}
items_.emplace_back(create_embedding_item(offset, length));
}
if (items_.empty()) {
fprintf(stderr, "build() failed: no valid embedding items created\n");
return false;
}
mm_data_.data.items.entries_size = items_.size();
mm_data_.data.items.entries = items_.data();
is_built_ = true;
return true;
}
void reset() {
memset(&mm_data_, 0, sizeof(mm_data_));
items_.clear();
tensor_buffers_.clear();
is_built_ = false;
}
const XLLM_MM_Data* get_data() const {
return is_built_ ? &mm_data_ : nullptr;
}
void validate() const {
if (!is_built_) {
fprintf(stderr,
"validate() failed: no data available (call build() first)\n");
return;
}
const size_t item_count = mm_data_.data.items.entries_size;
printf("=== Validating %zu Embedding Items ===\n\n", item_count);
for (size_t i = 0; i < item_count; ++i) {
const auto& item = mm_data_.data.items.entries[i];
printf("=== Embedding Item %zu ===\n", i + 1);
printf("Token Position: offset=%u, length=%u\n",
item.state.token_pos.offset,
item.state.token_pos.length);
printf("Data Type: (%d)\n", item.data.data.tensor.dtype);
printf("Tensor Shape: rank=%d, dim=[%d, %d]\n\n",
item.data.data.tensor.dims.rank,
item.data.data.tensor.dims.dim[0],
item.data.data.tensor.dims.dim[1]);
}
}
bool is_built() const { return is_built_; }
size_t get_item_count() const {
return is_built_ ? mm_data_.data.items.entries_size : 0;
}
private:
XLLM_MM_Data mm_data_{};
std::vector<XLLM_MM_Item> items_;
std::vector<std::unique_ptr<uint16_t[]>> tensor_buffers_;
bool is_built_ = false;
inline uint16_t float_to_bfloat16(float f) {
union {
float f32;
uint32_t u32;
} u;
u.f32 = f;
return static_cast<uint16_t>(u.u32 >> 16);
}
XLLM_MM_Item create_embedding_item(uint32_t offset, uint32_t length) {
XLLM_MM_Item item{};
item.type = XLLM_MM_TYPE_EMBEDDING;
item.state.token_pos.offset = offset;
item.state.token_pos.length = length;
item.data.is_single_tensor = true;
item.data.data.tensor.dtype = XLLM_DTYPE_BFLOAT16;
item.data.data.tensor.dims.rank = 2;
memset(item.data.data.tensor.dims.dim,
0,
sizeof(item.data.data.tensor.dims.dim));
item.data.data.tensor.dims.dim[0] = static_cast<int>(length);
item.data.data.tensor.dims.dim[1] = MODEL_WORD_EMBEDDING_SIZE;
const size_t element_count = length * MODEL_WORD_EMBEDDING_SIZE;
const size_t buffer_size_bytes = element_count * sizeof(uint16_t);
auto buffer = std::make_unique<uint16_t[]>(element_count);
for (size_t i = 0; i < length; ++i) {
for (size_t j = 0; j < MODEL_WORD_EMBEDDING_SIZE; ++j) {
float float_val =
static_cast<float>(i * MODEL_WORD_EMBEDDING_SIZE + j) /
static_cast<float>(element_count);
uint16_t bf16_val = float_to_bfloat16(float_val);
buffer[i * MODEL_WORD_EMBEDDING_SIZE + j] = bf16_val;
}
}
item.data.data.tensor.data = buffer.get();
tensor_buffers_.push_back(std::move(buffer));
return item;
}
};
XLLM_REC_Handler* service_startup_hook() {
XLLM_REC_Handler* rec_handler = xllm_rec_create();
// If there is no separate setting, init_options can be passed as nullptr, and
// the default value(XLLM_INIT_REC_OPTIONS_DEFAULT) will be used
XLLM_InitOptions init_options;
xllm_rec_init_options_default(&init_options);
// init_options.beam_width = 1;
// init_options.max_decode_rounds = 0;
snprintf(init_options.log_dir,
sizeof(init_options.log_dir),
"/export/home/huheng7/log");
bool ret = xllm_rec_initialize(
rec_handler, model_path.c_str(), devices.c_str(), &init_options);
if (!ret) {
std::cout << "REC init failed" << std::endl;
xllm_rec_destroy(rec_handler);
return nullptr;
}
std::cout << "REC init successfully" << std::endl;
return rec_handler;
}
void service_stop_hook(XLLM_REC_Handler* rec_handler) {
xllm_rec_destroy(rec_handler);
std::cout << "REC stop" << std::endl;
}
int generate_random_int(int min, int max) {
if (min > max) {
throw std::invalid_argument("min cannot be greater than max");
}
static std::random_device rd;
static std::mt19937 gen(rd());
std::uniform_int_distribution<int> dist(min, max);
return dist(gen);
}
int main(int argc, char** argv) {
XLLM_REC_Handler* rec_handler = service_startup_hook();
if (nullptr == rec_handler) {
return -1;
}
// If there is no separate setting, request_params can be passed as nullptr,
// and the default value(XLLM_REQUEST_PARAMS_DEFAULT) will be used
XLLM_RequestParams request_params;
xllm_rec_request_params_default(&request_params);
// request_params.beam_width = 128;
request_params.max_tokens = 3;
request_params.beam_width = 128;
request_params.logprobs = true;
// request_params.temperature = 1.0;
request_params.top_k = 128;
request_params.top_logprobs = 128;
// request_params.top_p = 1.0;
// request_params.repetition_penalty = 1.0;
std::vector<int32_t> token_ids = {
151644, 8948, 198, 56568, 101909, 101215, 104799, 101914, 101057,
3837, 103929, 100032, 44956, 15946, 55338, 45943, 104570, 11622,
105801, 72881, 64559, 307, 71817, 51463, 3837, 56568, 107618,
100345, 20002, 104754, 72651, 105565, 45943, 116951, 101034, 67949,
72651, 109348, 36407, 104538, 20002, 104326, 87267, 72651, 109348,
1773, 151645, 198, 151644, 872, 198, 20002, 21, 15,
35727, 31843, 36667, 59879, 20450, 99805, 32044, 72651, 105565,
45943, 32044, 113507, 153479, 155828, 160439, 11, 153479, 157177,
160439, 11, 153479, 155828, 160439, 11, 153479, 155828, 160439,
11, 153479, 155828, 160439, 11, 153479, 155828, 160439, 11,
155622, 158228, 160337, 11, 152907, 158228, 159858, 11, 153036,
158228, 160333, 11, 153258, 159797, 160105, 11, 153186, 157627,
160740, 11, 152907, 158228, 160680, 11, 154562, 157329, 160321,
11, 153326, 157680, 163928, 11, 153258, 159634, 160105, 11,
152847, 157129, 162841, 11, 152847, 157399, 162841, 11, 152847,
158228, 163388, 11, 153036, 159807, 162840, 11, 154562, 157329,
160321, 11, 154562, 156839, 160321, 11, 154562, 158181, 160321,
11, 153326, 158534, 163886, 11, 153326, 157177, 163041, 11,
155622, 158228, 163359, 11, 152569, 155800, 162738, 11, 153390,
158228, 160357, 11, 152663, 157649, 162738, 11, 155193, 158667,
162738, 11, 155622, 158228, 160706, 11, 151685, 158473, 162738,
11, 152907, 158228, 162653, 11, 151876, 158228, 159909, 11,
152907, 158228, 162407, 11, 152907, 158228, 163551, 11, 151685,
158473, 162738, 11, 152686, 155927, 162029, 11, 152663, 158228,
161841, 11, 152686, 155927, 162603, 11, 153516, 157280, 161980,
11, 153516, 159807, 160708, 11, 153516, 157900, 163856, 11,
153516, 155967, 161020, 11, 153516, 157280, 160838, 11, 153200,
157591, 162582, 11, 151924, 158696, 160358, 11, 154562, 159113,
160860, 11, 153386, 159086, 161519, 11, 154625, 159807, 160781,
11, 153479, 155828, 160439, 11, 153479, 155828, 160439, 11,
153479, 157177, 160439, 11, 153479, 155828, 160439, 11, 154213,
157866, 160523, 11, 153036, 156918, 163610, 11, 153036, 157351,
160974, 11, 153688, 158228, 160337, 11, 155507, 159807, 162736,
11, 155370, 159219, 161059, 11, 155002, 158118, 160019, 11,
155370, 159219, 161059, 11, 153792, 159022, 161003, 11, 155576,
155927, 161581, 11, 155576, 155927, 163189, 11, 155576, 159630,
162853, 11, 155576, 159630, 163527, 11, 155576, 159630, 162164,
11, 155576, 158048, 163339, 11, 155576, 157177, 163339, 11,
155576, 159630, 163527, 11, 155576, 157177, 163339, 11, 155576,
157680, 163339, 11, 155576, 159630, 160653, 11, 155576, 159630,
162153, 11, 155576, 159630, 161747, 11, 155576, 157505, 163339,
11, 153831, 158228, 160026, 11, 153390, 158228, 161841, 11,
153831, 156324, 162738, 11, 153390, 158228, 161491, 11, 153390,
159145, 162738, 11, 155507, 158473, 162738, 11, 153831, 157649,
162738, 11, 155507, 157770, 162738, 11, 153390, 158228, 161033,
11, 155507, 158473, 162738, 11, 153390, 158228, 160824, 11,
153479, 157649, 160439, 11, 153479, 157649, 160439, 11, 153479,
155828, 160439, 11, 153479, 157649, 160439, 11, 153479, 157649,
160439, 11, 153479, 157649, 160439, 11, 153849, 159380, 162841,
11, 152663, 158107, 162738, 11, 152271, 157371, 161110, 11,
152663, 157176, 160199, 11, 154936, 158966, 162841, 11, 153390,
158228, 161491, 11, 153036, 158228, 162840, 11, 155646, 158228,
162408, 11, 152663, 156814, 162738, 11, 152569, 158473, 162738,
11, 155646, 158228, 161308, 11, 152663, 158228, 163631, 11,
155370, 159786, 163029, 11, 153534, 159283, 161094, 11, 153534,
157756, 163778, 11, 151905, 156698, 163573, 11, 151905, 156698,
161534, 11, 151905, 156698, 162140, 11, 153534, 157931, 161817,
11, 153534, 157121, 161059, 11, 154826, 158585, 163433, 11,
154826, 158585, 160756, 11, 154826, 157666, 161504, 11, 154826,
157351, 161808, 11, 154826, 158585, 161062, 11, 154826, 157666,
161504, 11, 154826, 156537, 163635, 11, 155370, 159219, 161059,
11, 155370, 156903, 160381, 11, 155370, 156903, 160381, 11,
155370, 159219, 162223, 11, 155370, 159330, 162223, 11, 153464,
159219, 161059, 11, 154809, 156903, 160381, 11, 153464, 156878,
162223, 11, 154809, 157794, 162010, 11, 154809, 159219, 161059,
11, 151893, 159807, 162666, 11, 151893, 158534, 160890, 11,
153326, 157177, 163620, 11, 153326, 159462, 163041, 11, 152663,
156348, 162738, 11, 152663, 158473, 162736, 11, 152463, 156537,
160873, 11, 155507, 157176, 162738, 11, 155193, 158473, 162738,
11, 152663, 157649, 162738, 11, 152663, 158107, 162738, 11,
152663, 155780, 162738, 11, 152663, 158473, 162738, 11, 152663,
157649, 162738, 11, 152663, 157649, 162738, 11, 152663, 155828,
162738, 11, 152663, 158621, 162738, 11, 152663, 157176, 162738,
11, 155646, 158228, 160017, 11, 155682, 158228, 162859, 67949,
103969, 72651, 109348, 17714, 155646, 158228, 162234, 1773, 104210,
67949, 9370, 72651, 45943, 9370, 111450, 37945, 104538, 20002,
104326, 104309, 72651, 9370, 16, 15, 18947, 45943, 3837,
11622, 107463, 17992, 71817, 17177, 99859, 1773, 151645, 198,
151644, 77091, 198};
size_t token_size = token_ids.size();
const int32_t* token_ids_ptr = token_ids.data();
XLLM_MM_Data_Wrapper multimodal_data_wrapper;
std::vector<std::pair<uint32_t, uint32_t>> positions = {{100, 32}, {300, 64}};
multimodal_data_wrapper.build(positions);
multimodal_data_wrapper.validate();
// multimodal_data_wrapper.get_data(),
XLLM_Response* resp =
xllm_rec_multimodal_completions(rec_handler,
model_name.c_str(),
token_ids_ptr,
token_size,
multimodal_data_wrapper.get_data(),
10000,
&request_params);
if (nullptr == resp) {
std::cout << "REC completions failed, response is nullptr" << std::endl;
service_stop_hook(rec_handler);
return -1;
}
if (resp->status_code != XLLM_StatusCode::kSuccess) {
std::cout << "REC completions failed, status code:" << resp->status_code
<< ", error info:" << resp->error_info << std::endl;
} else {
std::cout << "REC completions successfully, size:"
<< resp->choices.entries_size << std::endl;
if (nullptr != resp->choices.entries) {
for (int i = 0; i < resp->choices.entries_size; ++i) {
XLLM_Choice& choice = resp->choices.entries[i];
std::cout << "token size: " << choice.token_size
<< ",logprobs size:" << choice.logprobs.entries_size
<< std::endl;
for (int j = 0; j < choice.token_size; j++) {
std::cout << "xllm answer[" << choice.index
<< "]: token id=" << choice.token_ids[j] << std::endl;
}
for (int j = 0; j < choice.logprobs.entries_size; j++) {
XLLM_LogProb& logprob = choice.logprobs.entries[j];
std::cout << "xllm answer[" << choice.index
<< "]: token id=" << logprob.token_id
<< ", token logprob=" << logprob.logprob << std::endl;
}
}
}
}
xllm_rec_free_response(resp);
service_stop_hook(rec_handler);
return 0;
}

View File

@@ -0,0 +1,751 @@
/* Copyright 2025 The xLLM 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 "helper.h"
#include <glog/logging.h>
#include <pthread.h>
#include <torch/torch.h>
#include <atomic>
#include <string>
#include "core/common/global_flags.h"
#include "core/util/env_var.h"
#include "core/util/rec_model_utils.h"
#include "core/util/uuid.h"
namespace xllm {
namespace helper {
namespace {
thread_local ShortUUID short_uuid;
static std::atomic<bool> g_glog_inited = false;
static pthread_mutex_t g_log_init_mutex = PTHREAD_MUTEX_INITIALIZER;
} // namespace
std::string generate_request_id() {
return "xllm-" + InstanceName::name()->get_name_hash() + "-" +
short_uuid.random();
}
void init_log(const std::string& log_dir) {
if (g_glog_inited.load(std::memory_order_acquire)) {
return;
}
pthread_mutex_lock(&g_log_init_mutex);
if (!g_glog_inited.load(std::memory_order_relaxed)) {
google::InitGoogleLogging("xllm");
std::string log_prefix = log_dir.empty() ? "./" : log_dir + "/";
google::SetLogDestination(google::INFO,
(log_prefix + "xllm.log.INFO.").c_str());
google::SetLogDestination(google::WARNING,
(log_prefix + "xllm.log.WARNING.").c_str());
google::SetLogDestination(google::ERROR,
(log_prefix + "xllm.log.ERROR.").c_str());
google::SetStderrLogging(google::FATAL);
g_glog_inited.store(true, std::memory_order_release);
}
pthread_mutex_unlock(&g_log_init_mutex);
}
void shutdown_log() {
if (!g_glog_inited.load(std::memory_order_acquire)) {
return;
}
pthread_mutex_lock(&g_log_init_mutex);
if (g_glog_inited.load(std::memory_order_relaxed)) {
google::ShutdownGoogleLogging();
g_glog_inited.store(false, std::memory_order_release);
}
pthread_mutex_unlock(&g_log_init_mutex);
}
void set_init_options(BackendType backend_type,
const XLLM_InitOptions* init_options,
XLLM_InitOptions* xllm_init_options) {
if (init_options == nullptr) {
if (backend_type == BackendType::LLM) {
memcpy(xllm_init_options,
&XLLM_INIT_LLM_OPTIONS_DEFAULT,
sizeof(XLLM_InitOptions));
} else if (backend_type == BackendType::REC) {
memcpy(xllm_init_options,
&XLLM_INIT_REC_OPTIONS_DEFAULT,
sizeof(XLLM_InitOptions));
}
} else {
memcpy(xllm_init_options, init_options, sizeof(XLLM_InitOptions));
}
return;
}
void transfer_request_params(InferenceType inference_type,
const XLLM_RequestParams* request_params,
xllm::RequestParams* xllm_request_params) {
XLLM_RequestParams final_request_params;
if (nullptr == request_params) {
if (inference_type == InferenceType::LLM_COMPLETIONS ||
inference_type == InferenceType::LLM_CHAT_COMPLETIONS) {
memcpy(&final_request_params,
&XLLM_LLM_REQUEST_PARAMS_DEFAULT,
sizeof(XLLM_RequestParams));
} else if (inference_type == InferenceType::REC_COMPLETIONS ||
inference_type == InferenceType::REC_CHAT_COMPLETIONS) {
memcpy(&final_request_params,
&XLLM_REC_REQUEST_PARAMS_DEFAULT,
sizeof(XLLM_RequestParams));
}
} else {
memcpy(&final_request_params, request_params, sizeof(XLLM_RequestParams));
}
xllm_request_params->echo = final_request_params.echo;
xllm_request_params->offline = final_request_params.offline;
xllm_request_params->logprobs = final_request_params.logprobs;
xllm_request_params->ignore_eos = final_request_params.ignore_eos;
xllm_request_params->best_of = final_request_params.best_of;
xllm_request_params->top_k = final_request_params.top_k;
xllm_request_params->top_p = final_request_params.top_p;
xllm_request_params->n = final_request_params.n;
xllm_request_params->max_tokens = final_request_params.max_tokens;
xllm_request_params->frequency_penalty =
final_request_params.frequency_penalty;
xllm_request_params->presence_penalty = final_request_params.presence_penalty;
xllm_request_params->repetition_penalty =
final_request_params.repetition_penalty;
xllm_request_params->beam_width = final_request_params.beam_width;
xllm_request_params->num_return_sequences =
final_request_params.num_return_sequences;
xllm_request_params->top_logprobs = final_request_params.top_logprobs;
xllm_request_params->temperature = final_request_params.temperature;
xllm_request_params->request_id = final_request_params.request_id;
xllm_request_params->ttlt_slo_ms = final_request_params.ttlt_slo_ms;
xllm_request_params->ttft_slo_ms = final_request_params.ttft_slo_ms;
xllm_request_params->tpot_slo_ms = final_request_params.tpot_slo_ms;
return;
}
XLLM_Response* build_error_response(const std::string& request_id,
XLLM_StatusCode status_code,
const std::string& error_info) {
XLLM_Response* response = new XLLM_Response();
CHECK(nullptr != response);
response->status_code = status_code;
strncpy(
response->error_info, error_info.c_str(), XLLM_ERROR_INFO_MAX_LEN - 1);
response->error_info[XLLM_ERROR_INFO_MAX_LEN - 1] = '\0';
XLLM_SET_META_STRING_FIELD(response->id, request_id);
LOG(ERROR) << "Request [" << request_id << "] error: " << error_info
<< " (code: " << static_cast<int>(response->status_code) << ")";
return response;
}
XLLM_Response* build_success_response(const InferenceType& inference_type,
const RequestOutput& output,
RecPipelineType rec_pipeline_type,
const std::string& request_id,
int64_t created_time,
const std::string& model) {
XLLM_Response* response = new XLLM_Response();
CHECK(nullptr != response);
response->status_code = XLLM_StatusCode::kSuccess;
response->created = created_time;
XLLM_SET_META_STRING_FIELD(response->id, request_id);
XLLM_SET_META_STRING_FIELD(response->model, model);
if (inference_type == InferenceType::LLM_COMPLETIONS ||
inference_type == InferenceType::REC_COMPLETIONS) {
snprintf(response->object, sizeof(response->object), "text_completion");
} else if (inference_type == InferenceType::LLM_CHAT_COMPLETIONS ||
inference_type == InferenceType::REC_CHAT_COMPLETIONS) {
snprintf(response->object, sizeof(response->object), "chat.completion");
}
response->choices.entries_size = output.outputs.size();
response->choices.entries = new XLLM_Choice[response->choices.entries_size]();
CHECK(nullptr != response->choices.entries);
const bool is_rec_inference =
inference_type == InferenceType::REC_COMPLETIONS ||
inference_type == InferenceType::REC_CHAT_COMPLETIONS;
const bool is_onerec_pipeline =
is_rec_inference && is_onerec_pipeline_type(rec_pipeline_type);
if (is_onerec_pipeline) {
response->rec_outputs.entries_size = output.outputs.size();
response->rec_outputs.entries =
new XLLM_RecOutput[response->rec_outputs.entries_size]();
CHECK(nullptr != response->rec_outputs.entries);
}
int32_t total_item_count = 0;
const int32_t total_threshold = FLAGS_total_conversion_threshold;
for (int i = 0; i < output.outputs.size(); i++) {
const auto& seq_output = output.outputs[i];
XLLM_Choice& choice = response->choices.entries[i];
choice.index = seq_output.index;
XLLM_RecOutput* rec_output = nullptr;
if (response->rec_outputs.entries != nullptr) {
rec_output = &response->rec_outputs.entries[i];
rec_output->index = seq_output.index;
}
if (inference_type == InferenceType::LLM_COMPLETIONS ||
inference_type == InferenceType::REC_COMPLETIONS) {
size_t text_len = seq_output.text.length();
choice.text = new char[text_len + 1];
CHECK(nullptr != choice.text);
strncpy(choice.text, seq_output.text.c_str(), text_len + 1);
choice.text[text_len] = '\0';
} else if (inference_type == InferenceType::LLM_CHAT_COMPLETIONS ||
inference_type == InferenceType::REC_CHAT_COMPLETIONS) {
choice.message = new XLLM_ChatMessage();
CHECK(nullptr != choice.message);
snprintf(choice.message->role, sizeof(choice.message->role), "assistant");
size_t text_len = seq_output.text.length();
choice.message->content = new char[text_len + 1];
CHECK(nullptr != choice.message->content);
strncpy(choice.message->content, seq_output.text.c_str(), text_len + 1);
choice.message->content[text_len] = '\0';
}
if (seq_output.finish_reason.has_value()) {
XLLM_SET_META_STRING_FIELD(choice.finish_reason,
seq_output.finish_reason.value());
}
if (seq_output.token_ids.size() > 0) {
choice.token_size = seq_output.token_ids.size();
choice.token_ids = new int32_t[choice.token_size];
CHECK(nullptr != choice.token_ids);
for (int j = 0; j < choice.token_size; j++) {
choice.token_ids[j] = seq_output.token_ids[j];
}
}
if (seq_output.logprobs.has_value()) {
choice.logprobs.entries_size = seq_output.logprobs.value().size();
choice.logprobs.entries =
new XLLM_LogProb[choice.logprobs.entries_size]();
CHECK(nullptr != choice.logprobs.entries);
for (int j = 0; j < seq_output.logprobs.value().size(); j++) {
const auto& logprob = seq_output.logprobs.value()[j];
XLLM_LogProb& xllm_logprob = choice.logprobs.entries[j];
xllm_logprob.token_id = logprob.token_id;
xllm_logprob.logprob = logprob.logprob;
}
}
if (is_onerec_pipeline && FLAGS_enable_convert_tokens_to_item &&
rec_output != nullptr) {
size_t copied_item_count = 0;
if (!seq_output.item_ids_list.empty()) {
copied_item_count =
std::min(seq_output.item_ids_list.size(),
static_cast<size_t>(
std::max(total_threshold - total_item_count, 0)));
if (copied_item_count > 0) {
rec_output->item_ids_size = copied_item_count;
rec_output->item_ids = new int64_t[copied_item_count];
CHECK(nullptr != rec_output->item_ids);
for (size_t j = 0; j < copied_item_count; ++j) {
rec_output->item_ids[j] = seq_output.item_ids_list[j];
}
total_item_count += static_cast<int32_t>(copied_item_count);
}
} else if (seq_output.item_ids.has_value() &&
total_item_count < total_threshold) {
rec_output->item_ids_size = 1;
rec_output->item_ids = new int64_t[1];
CHECK(nullptr != rec_output->item_ids);
rec_output->item_ids[0] = seq_output.item_ids.value();
++total_item_count;
}
}
if (is_onerec_pipeline && FLAGS_enable_output_sku_logprobs &&
!seq_output.token_ids_logprobs.empty() && rec_output != nullptr) {
rec_output->rec_token_logprobs_size =
seq_output.token_ids_logprobs.size();
rec_output->rec_token_logprobs =
new float[rec_output->rec_token_logprobs_size];
CHECK(nullptr != rec_output->rec_token_logprobs);
for (size_t j = 0; j < rec_output->rec_token_logprobs_size; ++j) {
if (seq_output.token_ids_logprobs[j].has_value()) {
rec_output->rec_token_logprobs[j] =
seq_output.token_ids_logprobs[j].value();
} else {
rec_output->rec_token_logprobs[j] = 0.0f;
}
}
}
}
if (output.usage.has_value()) {
const auto& usage = output.usage.value();
response->usage.prompt_tokens = usage.num_prompt_tokens;
response->usage.completion_tokens = usage.num_generated_tokens;
response->usage.total_tokens = usage.num_total_tokens;
}
return response;
}
template <typename HandlerType, typename InputType>
XLLM_Response* handle_inference_request(
HandlerType* handler,
InferenceType inference_type,
const std::string& model_id,
const InputType& input,
void* extra,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params) {
CHECK(nullptr != handler);
std::string request_id;
if (nullptr != request_params && strlen(request_params->request_id) > 0) {
request_id = request_params->request_id;
} else {
request_id = generate_request_id();
}
if (!handler->initialized) {
return build_error_response(
request_id, XLLM_StatusCode::kNotInitialized, "LLM is not initialized");
}
if (std::find(handler->model_ids.begin(),
handler->model_ids.end(),
model_id) == handler->model_ids.end()) {
return build_error_response(request_id,
XLLM_StatusCode::kModelNotFound,
"Specified model ID not loaded: " + model_id);
}
xllm::RequestParams xllm_request_params;
transfer_request_params(inference_type, request_params, &xllm_request_params);
xllm_request_params.request_id = request_id;
RecPipelineType rec_pipeline_type = RecPipelineType::kLlmRecDefault;
if constexpr (std::is_same_v<HandlerType, XLLM_REC_Handler>) {
rec_pipeline_type = handler->pipeline_type;
if (FLAGS_enable_output_sku_logprobs &&
is_onerec_pipeline_type(rec_pipeline_type)) {
xllm_request_params.logprobs = true;
}
}
const int64_t created_time = absl::ToUnixSeconds(absl::Now());
try {
auto promise_ptr = std::make_shared<folly::Promise<XLLM_Response*>>();
auto future = promise_ptr->getSemiFuture();
auto on_request_complete = [model_id,
request_id,
created_time,
inference_type,
rec_pipeline_type,
weak_promise = std::weak_ptr(promise_ptr)](
const RequestOutput& req_output) -> bool {
if (auto locked_promise = weak_promise.lock()) {
try {
if (req_output.status.has_value()) {
if (req_output.status.value().ok()) {
locked_promise->setValue(build_success_response(inference_type,
req_output,
rec_pipeline_type,
request_id,
created_time,
model_id));
} else {
locked_promise->setValue(build_error_response(
request_id,
XLLM_StatusCode::kInternalError,
"RequestOutput status is not ok, message: " +
req_output.status.value().message()));
}
} else {
locked_promise->setValue(
build_error_response(request_id,
XLLM_StatusCode::kInternalError,
"RequestOutput status has no value"));
}
return true;
} catch (const std::exception& e) {
LOG(ERROR) << "Build response failed: " << e.what();
locked_promise->setValue(build_error_response(
request_id,
XLLM_StatusCode::kInternalError,
"Build response failed: " + std::string(e.what())));
}
}
return false;
};
if constexpr (std::is_same_v<HandlerType, XLLM_LLM_Handler>) {
handler->master->handle_request(input,
std::nullopt,
xllm_request_params,
std::nullopt,
on_request_complete);
} else if constexpr (std::is_same_v<HandlerType, XLLM_REC_Handler>) {
if constexpr (std::is_same_v<InputType, std::vector<int>>) {
if (nullptr != extra) {
xllm::MMData* mm_data =
dynamic_cast<xllm::MMData*>(static_cast<xllm::MMData*>(extra));
CHECK(nullptr != mm_data);
std::optional<xllm::MMData> opt_mm_data = std::move(*mm_data);
handler->master->handle_request(
input, opt_mm_data, xllm_request_params, on_request_complete);
} else {
handler->master->handle_request("",
input,
std::nullopt,
xllm_request_params,
on_request_complete);
}
} else {
handler->master->handle_request(input,
std::nullopt,
std::nullopt,
xllm_request_params,
on_request_complete);
}
} else {
CHECK(false);
}
return std::move(future)
.via(handler->executor.get())
.within(std::chrono::milliseconds(timeout_ms))
.thenTry([request_id](
folly::Try<XLLM_Response*>&& result) -> XLLM_Response* {
if (result.hasValue()) return std::move(result).value();
std::string error_msg;
XLLM_StatusCode code = XLLM_StatusCode::kInternalError;
try {
result.throwUnlessValue();
} catch (const folly::FutureTimeout& e) {
error_msg = "Request timed out: " + std::string(e.what());
code = XLLM_StatusCode::kTimeout;
} catch (const std::exception& e) {
error_msg = "Inference failed: " + std::string(e.what());
} catch (...) {
error_msg = "Inference failed with unknown exception";
}
return build_error_response(request_id, code, error_msg);
})
.get();
} catch (...) {
return build_error_response(request_id,
XLLM_StatusCode::kInternalError,
"Critical error in inference pipeline");
}
}
void xllm_free_response(XLLM_Response* resp) {
if (nullptr == resp) {
return;
}
if (nullptr != resp->choices.entries) {
for (int i = 0; i < resp->choices.entries_size; ++i) {
XLLM_Choice& choice = resp->choices.entries[i];
if (nullptr != choice.text) {
delete[] choice.text;
choice.text = nullptr;
}
if (nullptr != choice.message) {
if (nullptr != choice.message->content) {
delete[] choice.message->content;
choice.message->content = nullptr;
}
delete choice.message;
choice.message = nullptr;
}
if (nullptr != choice.token_ids) {
delete[] choice.token_ids;
choice.token_ids = nullptr;
choice.token_size = 0;
}
if (nullptr != choice.logprobs.entries) {
delete[] choice.logprobs.entries;
choice.logprobs.entries = nullptr;
}
choice.logprobs.entries_size = 0;
}
delete[] resp->choices.entries;
resp->choices.entries = nullptr;
}
resp->choices.entries_size = 0;
if (nullptr != resp->rec_outputs.entries) {
for (size_t i = 0; i < resp->rec_outputs.entries_size; ++i) {
XLLM_RecOutput& rec_output = resp->rec_outputs.entries[i];
if (nullptr != rec_output.item_ids) {
delete[] rec_output.item_ids;
rec_output.item_ids = nullptr;
rec_output.item_ids_size = 0;
}
if (nullptr != rec_output.rec_token_logprobs) {
delete[] rec_output.rec_token_logprobs;
rec_output.rec_token_logprobs = nullptr;
rec_output.rec_token_logprobs_size = 0;
}
}
delete[] resp->rec_outputs.entries;
resp->rec_outputs.entries = nullptr;
}
resp->rec_outputs.entries_size = 0;
delete resp;
return;
}
torch::ScalarType xllm_dtype_to_torch_scalar_type(XLLM_DataType dtype) {
switch (dtype) {
case XLLM_DTYPE_UNDEFINED:
throw std::runtime_error(
"XLLM_DTYPE_UNDEFINED is not a valid dtype for tensor conversion");
case XLLM_DTYPE_FLOAT16:
return torch::kFloat16;
case XLLM_DTYPE_FLOAT32:
return torch::kFloat32;
case XLLM_DTYPE_FLOAT64:
return torch::kFloat64;
case XLLM_DTYPE_BFLOAT16:
return torch::kBFloat16;
case XLLM_DTYPE_INT8:
return torch::kInt8;
case XLLM_DTYPE_INT16:
return torch::kInt16;
case XLLM_DTYPE_INT32:
return torch::kInt32;
case XLLM_DTYPE_INT64:
return torch::kInt64;
case XLLM_DTYPE_BOOL:
return torch::kBool;
case XLLM_DTYPE_STRING:
throw std::runtime_error(
"String dtype is not supported for torch::Tensor");
default:
throw std::runtime_error("Unsupported XLLM_DataType: " +
std::to_string(dtype));
}
}
torch::Tensor convert_xllm_tensor_to_torch(const XLLM_Tensor& xllm_tensor) {
if (xllm_tensor.data == nullptr) {
throw std::runtime_error("XLLM_Tensor data pointer is null");
}
torch::ScalarType scalar_type =
xllm_dtype_to_torch_scalar_type(xllm_tensor.dtype);
std::vector<int64_t> shape;
for (int i = 0; i < xllm_tensor.dims.rank; ++i) {
int dim = xllm_tensor.dims.dim[i];
if (dim > 0) {
shape.push_back(dim);
}
}
if (shape.empty()) {
throw std::runtime_error("XLLM_Tensor all dimensions are invalid value");
}
torch::Tensor tensor =
torch::from_blob(const_cast<void*>(xllm_tensor.data), shape, scalar_type)
.clone();
return tensor;
}
xllm::MMDataItem convert_xllm_mm_item_to_internal(
const XLLM_MM_Item& xllm_item) {
uint32_t xllm_type_val = static_cast<uint32_t>(xllm_item.type);
xllm::MMType::Value internal_val = xllm::MMType::NONE;
switch (xllm_type_val) {
case XLLM_MM_TYPE_EMBEDDING:
internal_val = xllm::MMType::EMBEDDING;
break;
case XLLM_MM_TYPE_IMAGE:
internal_val = xllm::MMType::IMAGE;
break;
case XLLM_MM_TYPE_VIDEO:
internal_val = xllm::MMType::VIDEO;
break;
case XLLM_MM_TYPE_AUDIO:
internal_val = xllm::MMType::AUDIO;
break;
case XLLM_MM_TYPE_NONE:
internal_val = xllm::MMType::NONE;
break;
default:
throw std::runtime_error(std::string("Unsupported XLLM_MM_Type: ") +
std::to_string(xllm_type_val));
}
xllm::MMType item_type(internal_val);
xllm::MMDataItem internal_item(item_type);
xllm::MMItemState& state = internal_item.mutable_state();
xllm::MMItemState::TokenPos& token_pos = state.mutable_token_pos();
token_pos.offset = xllm_item.state.token_pos.offset;
token_pos.length = xllm_item.state.token_pos.length;
if (xllm_item.data.is_single_tensor) {
torch::Tensor tensor =
convert_xllm_tensor_to_torch(xllm_item.data.data.tensor);
internal_item.add("tensor", tensor);
} else {
std::vector<torch::Tensor> tensor_list;
const XLLM_Tensors& xllm_tensors = xllm_item.data.data.tensors;
for (size_t i = 0; i < xllm_tensors.entries_size; ++i) {
tensor_list.push_back(
convert_xllm_tensor_to_torch(xllm_tensors.entries[i]));
}
internal_item.add("tensor_list", tensor_list);
}
return internal_item;
}
bool convert_xllm_mm_data_to_internal(const XLLM_MM_Data* mm_data,
xllm::MMData& internal_mm_data) {
if (mm_data == nullptr || mm_data->type_mask == XLLM_MM_TYPE_NONE) {
return false;
}
xllm::MMType::Value internal_val =
static_cast<xllm::MMType::Value>(mm_data->type_mask);
xllm::MMType mm_type(internal_val);
if (mm_data->is_dict) {
const XLLM_MM_Dict& xllm_dict = mm_data->data.dict;
xllm::MMDict internal_dict;
for (size_t i = 0; i < xllm_dict.entries_size; ++i) {
const XLLM_MM_DictEntry& xllm_entry = xllm_dict.entries[i];
xllm::MMKey key(xllm_entry.key);
const XLLM_MM_Value& xllm_value = xllm_entry.value;
if (xllm_value.is_single_tensor) {
torch::Tensor tensor =
convert_xllm_tensor_to_torch(xllm_value.data.tensor);
internal_dict.insert({key, tensor});
} else {
std::vector<torch::Tensor> tensor_list;
const XLLM_Tensors& xllm_tensors = xllm_value.data.tensors;
for (size_t j = 0; j < xllm_tensors.entries_size; ++j) {
tensor_list.push_back(
convert_xllm_tensor_to_torch(xllm_tensors.entries[j]));
}
internal_dict.insert({key, tensor_list});
}
}
internal_mm_data.set<xllm::MMDict>(mm_type, internal_dict);
} else {
const XLLM_MM_Items& xllm_items = mm_data->data.items;
xllm::MMItemVec internal_item_vec;
for (size_t i = 0; i < xllm_items.entries_size; ++i) {
const XLLM_MM_Item& xllm_item = xllm_items.entries[i];
xllm::MMDataItem internal_item =
convert_xllm_mm_item_to_internal(xllm_item);
internal_item_vec.push_back(std::move(internal_item));
}
internal_mm_data.set<xllm::MMItemVec>(mm_type, internal_item_vec);
}
return true;
}
// 1. LLM Handler + const char* (text completions)
template XLLM_Response* handle_inference_request<XLLM_LLM_Handler, const char*>(
XLLM_LLM_Handler* handler,
InferenceType inference_type,
const std::string& model_id,
const char* const& input,
void* extra,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
// 2. LLM Handler + std::vector<xllm::Message> (chat completions)
template XLLM_Response*
handle_inference_request<XLLM_LLM_Handler, std::vector<xllm::Message>>(
XLLM_LLM_Handler* handler,
InferenceType inference_type,
const std::string& model_id,
const std::vector<xllm::Message>& input,
void* extra,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
// 3. REC Handler + const char* (REC completions)
template XLLM_Response* handle_inference_request<XLLM_REC_Handler, const char*>(
XLLM_REC_Handler* handler,
InferenceType inference_type,
const std::string& model_id,
const char* const& input,
void* extra,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
// 4. REC Handler + std::vector<xllm::Message> (REC chat completions)
template XLLM_Response*
handle_inference_request<XLLM_REC_Handler, std::vector<xllm::Message>>(
XLLM_REC_Handler* handler,
InferenceType inference_type,
const std::string& model_id,
const std::vector<xllm::Message>& input,
void* extra,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
// 5. REC Handler + std::vector<int> (chat completions)
template XLLM_Response*
handle_inference_request<XLLM_REC_Handler, std::vector<int>>(
XLLM_REC_Handler* handler,
InferenceType inference_type,
const std::string& model_id,
const std::vector<int>& input,
void* extra,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
} // namespace helper
} // namespace xllm

View File

@@ -0,0 +1,172 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <absl/time/clock.h>
#include <absl/time/time.h>
#include <folly/executors/CPUThreadPoolExecutor.h>
#include <folly/futures/Future.h>
#include <folly/futures/Promise.h>
#include <glog/logging.h>
#include <algorithm>
#include <memory>
#include <string>
#include <vector>
#include "c_api/default.h"
#include "c_api/types.h"
#include "core/common/instance_name.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/distributed_runtime/rec_master.h"
#include "core/framework/request/request_output.h"
#include "core/framework/request/request_params.h"
#include "core/util/rec_model_utils.h"
/**
* @brief Opaque handle for LLM inference instance
*/
struct XLLM_LLM_Handler {
/** Flag indicating if LLM instance is initialized and ready for inference */
bool initialized{false};
/** List of loaded model IDs (for model existence validation) */
std::vector<std::string> model_ids;
/** Core controller for LLM runtime management */
std::unique_ptr<xllm::LLMMaster> master;
/** Thread pool for asynchronous inference task scheduling */
std::unique_ptr<folly::CPUThreadPoolExecutor> executor;
};
/**
* @brief Opaque handle for REC (Recommendation) inference instance
*/
struct XLLM_REC_Handler {
/** Flag indicating if REC instance is initialized and ready for inference */
bool initialized{false};
/** Selected REC pipeline type for the loaded model */
xllm::RecPipelineType pipeline_type{xllm::RecPipelineType::kLlmRecDefault};
/** List of loaded recommendation model IDs */
std::vector<std::string> model_ids;
/** Core controller for REC runtime management */
std::unique_ptr<xllm::RecMaster> master;
/** Thread pool for asynchronous recommendation task scheduling */
std::unique_ptr<folly::CPUThreadPoolExecutor> executor;
};
namespace xllm {
namespace helper {
enum class BackendType { LLM = 0, VLM = 1, REC = 2 };
enum class InferenceType {
LLM_COMPLETIONS = 0,
LLM_CHAT_COMPLETIONS = 1,
REC_COMPLETIONS = 2,
REC_CHAT_COMPLETIONS = 3,
REC_TOKENID_COMPLETIONS = 4,
};
#define XLLM_SET_META_STRING_FIELD(DST, SRC_STR) \
do { \
static_assert(sizeof(DST) > 1, "Destination buffer is too small"); \
strncpy( \
(char*)(DST), (SRC_STR).c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1); \
(DST)[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0'; \
} while (0)
/**
* @brief Thread-safe glog initialization for xLLM framework
* @note This API is idempotent (multiple calls have same effect as single call)
* @note Thread-safe: protected by pthread mutex to prevent race condition
* @param log_dir Directory to store log files (empty = current directory)
*/
void init_log(const std::string& log_dir);
/**
* @brief Safely shutdown glog and release resources
* @note Call this function before program exit (optional but recommended)
*/
void shutdown_log();
/**
* @brief Set init options, merge default options
*/
void set_init_options(BackendType backend_type,
const XLLM_InitOptions* init_options,
XLLM_InitOptions* xllm_init_options);
/**
* @brief Transfer C API request params to xLLM internal request params
*/
void transfer_request_params(InferenceType inference_type,
const XLLM_RequestParams* request_params,
xllm::RequestParams* xllm_request_params);
/**
* @brief Build error response for failed inference requests
*/
XLLM_Response* build_error_response(const std::string& request_id,
XLLM_StatusCode status_code,
const std::string& error_info);
/**
* @brief Build success response for completed inference requests
*/
XLLM_Response* build_success_response(const InferenceType& inference_type,
const xllm::RequestOutput& output,
xllm::RecPipelineType rec_pipeline_type,
const std::string& request_id,
int64_t created_time,
const std::string& model);
/**
* @brief Generic inference request handler (template function)
*/
template <typename HandlerType, typename InputType>
XLLM_Response* handle_inference_request(
HandlerType* handler,
InferenceType inference_type,
const std::string& model_id,
const InputType& input,
void* extra,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
/**
* @brief Safely free all memory allocated in XLLM_Response
*/
void xllm_free_response(XLLM_Response* resp);
/**
* @brief Generate unique request ID for tracing
*/
std::string generate_request_id();
torch::ScalarType xllm_dtype_to_torch_scalar_type(XLLM_DataType dtype);
torch::Tensor convert_xllm_tensor_to_torch(const XLLM_Tensor& xllm_tensor);
xllm::MMDataItem convert_xllm_mm_item_to_internal(
const XLLM_MM_Item& xllm_item);
bool convert_xllm_mm_data_to_internal(const XLLM_MM_Data* mm_data,
xllm::MMData& internal_mm_data);
} // namespace helper
} // namespace xllm

View File

@@ -0,0 +1,221 @@
/* Copyright 2025 The xLLM 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 "c_api/llm.h"
#include <folly/Unit.h>
#include <folly/experimental/coro/Timeout.h>
#include <folly/futures/Future.h>
#include <glog/logging.h>
#include <pthread.h>
#include <atomic>
#include <cstring>
#include <exception>
#include <stdexcept>
#include "core/common/global_flags.h"
#include "helper.h"
XLLM_CAPI_EXPORT XLLM_LLM_Handler* xllm_llm_create(void) {
XLLM_LLM_Handler* handler = new XLLM_LLM_Handler();
CHECK(nullptr != handler);
handler->initialized = false;
return handler;
}
XLLM_CAPI_EXPORT void xllm_llm_destroy(XLLM_LLM_Handler* handler) {
if (!handler) return;
handler->master.reset();
handler->executor.reset();
handler->model_ids.clear();
handler->initialized = false;
delete handler;
}
XLLM_CAPI_EXPORT void xllm_llm_init_options_default(
XLLM_InitOptions* init_options) {
if (nullptr == init_options) return;
*init_options = XLLM_INIT_LLM_OPTIONS_DEFAULT;
}
XLLM_CAPI_EXPORT bool xllm_llm_initialize(
XLLM_LLM_Handler* handler,
const char* model_path,
const char* devices,
const XLLM_InitOptions* init_options) {
if (!handler || !model_path || !devices) return false;
try {
XLLM_InitOptions xllm_init_options;
xllm::helper::set_init_options(
xllm::helper::BackendType::LLM, init_options, &xllm_init_options);
std::string log_dir(xllm_init_options.log_dir);
if (!log_dir.empty()) {
xllm::helper::init_log(xllm_init_options.log_dir);
}
if (!std::filesystem::exists(model_path)) {
LOG(ERROR) << "model path[" << model_path << "] does not exist";
return false;
}
xllm::Options options;
options.model_path(model_path)
.task_type(xllm_init_options.task)
.devices(devices)
.draft_model_path(xllm_init_options.draft_model)
.draft_devices(xllm_init_options.draft_devices)
.backend("llm")
.block_size(xllm_init_options.block_size)
.max_cache_size(xllm_init_options.max_cache_size)
.max_memory_utilization(xllm_init_options.max_memory_utilization)
.enable_prefix_cache(xllm_init_options.enable_prefix_cache)
.max_tokens_per_batch(xllm_init_options.max_tokens_per_batch)
.max_seqs_per_batch(xllm_init_options.max_seqs_per_batch)
.max_tokens_per_chunk_for_prefill(
xllm_init_options.max_tokens_per_chunk_for_prefill)
.num_speculative_tokens(xllm_init_options.num_speculative_tokens)
.num_request_handling_threads(
xllm_init_options.num_request_handling_threads)
.communication_backend(xllm_init_options.communication_backend)
.expert_parallel_degree(xllm_init_options.expert_parallel_degree)
.enable_chunked_prefill(xllm_init_options.enable_chunked_prefill)
.enable_prefill_sp(xllm_init_options.enable_prefill_sp)
.master_node_addr(xllm_init_options.master_node_addr)
.device_ip(xllm_init_options.device_ip)
.transfer_listen_port(xllm_init_options.transfer_listen_port)
.nnodes(xllm_init_options.nnodes)
.node_rank(xllm_init_options.node_rank)
.dp_size(xllm_init_options.dp_size)
.ep_size(xllm_init_options.ep_size)
.instance_name(xllm_init_options.instance_name)
.enable_disagg_pd(xllm_init_options.enable_disagg_pd)
.enable_schedule_overlap(xllm_init_options.enable_schedule_overlap)
.enable_pd_ooc(xllm_init_options.enable_pd_ooc)
.kv_cache_transfer_mode(xllm_init_options.kv_cache_transfer_mode)
.enable_shm(xllm_init_options.enable_shm)
.is_local(true)
.server_idx(xllm_init_options.server_idx);
options.enable_graph(FLAGS_enable_graph);
#if !defined(USE_NPU) && !defined(USE_CUDA)
FLAGS_enable_block_copy_kernel = false;
#endif
handler->master = std::make_unique<xllm::LLMMaster>(options);
handler->master->run();
size_t cpu_cores = std::thread::hardware_concurrency();
size_t thread_num = std::clamp((cpu_cores == 0) ? 8 : cpu_cores / 2,
static_cast<size_t>(4),
static_cast<size_t>(16));
handler->executor =
std::make_unique<folly::CPUThreadPoolExecutor>(thread_num);
std::filesystem::path model_path_fs =
std::filesystem::path(model_path).lexically_normal();
std::string model_id;
if (model_path_fs.has_filename()) {
model_id = model_path_fs.filename().string();
} else if (!model_path_fs.empty()) {
model_id = model_path_fs.string();
} else {
model_id = "default";
}
handler->model_ids.clear();
handler->model_ids.emplace_back(model_id);
handler->initialized = true;
return true;
} catch (const std::exception& e) {
LOG(ERROR) << "LLM initialization failed: " << e.what();
}
handler->master.reset();
handler->executor.reset();
handler->model_ids.clear();
handler->initialized = false;
return false;
}
XLLM_CAPI_EXPORT void xllm_llm_request_params_default(
XLLM_RequestParams* request_params) {
if (nullptr == request_params) return;
*request_params = XLLM_LLM_REQUEST_PARAMS_DEFAULT;
}
XLLM_CAPI_EXPORT XLLM_Response* xllm_llm_completions(
XLLM_LLM_Handler* handler,
const char* model_id,
const char* prompt,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params) {
if (!handler || !model_id || *model_id == '\0' || !prompt ||
*prompt == '\0') {
return xllm::helper::build_error_response(
"", XLLM_StatusCode::kInvalidRequest, "Invalid input parameters");
}
return xllm::helper::handle_inference_request(
handler,
xllm::helper::InferenceType::LLM_COMPLETIONS,
model_id,
prompt,
nullptr,
timeout_ms,
request_params);
}
XLLM_CAPI_EXPORT XLLM_Response* xllm_llm_chat_completions(
XLLM_LLM_Handler* handler,
const char* model_id,
const XLLM_ChatMessage* messages,
size_t messages_count,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params) {
if (!handler || !model_id || *model_id == '\0' || !messages ||
messages_count == 0) {
return xllm::helper::build_error_response(
"", XLLM_StatusCode::kInvalidRequest, "Invalid input parameters");
}
std::vector<xllm::Message> xllm_messages;
xllm_messages.reserve(messages_count);
for (int i = 0; i < messages_count; i++) {
xllm_messages.emplace_back(messages[i].role, messages[i].content);
}
return xllm::helper::handle_inference_request(
handler,
xllm::helper::InferenceType::LLM_CHAT_COMPLETIONS,
model_id,
xllm_messages,
nullptr,
timeout_ms,
request_params);
}
XLLM_CAPI_EXPORT void xllm_llm_free_response(XLLM_Response* resp) {
return xllm::helper::xllm_free_response(resp);
}

View File

@@ -0,0 +1,432 @@
/* Copyright 2025 The xLLM 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 "c_api/rec.h"
#include <folly/Unit.h>
#include <folly/experimental/coro/Timeout.h>
#include <folly/futures/Future.h>
#include <glog/logging.h>
#include <pthread.h>
#include <atomic>
#include <cstring>
#include <exception>
#include <stdexcept>
#include "core/framework/model_loader.h"
#include "core/util/rec_model_utils.h"
#include "helper.h"
namespace {
const char* get_rec_pipeline_name(xllm::RecPipelineType pipeline_type) {
switch (pipeline_type) {
case xllm::RecPipelineType::kLlmRecDefault:
return "LlmRecEnginePipeline";
case xllm::RecPipelineType::kLlmRecWithMmData:
return "LlmRecWithMmData";
case xllm::RecPipelineType::kLlmRecMultiRoundPipeline:
return "RecMultiRoundEnginePipeline";
case xllm::RecPipelineType::kOneRecDefault:
return "OneRecPrefillOnlyEnginePipeline";
case xllm::RecPipelineType::kOneRecXAttentionPipeline:
return "OneRecXAttentionEnginePipeline";
default:
return "UnknownRecPipeline";
}
}
void reset_pipeline_runtime_toggles() {
FLAGS_enable_rec_fast_sampler = false;
FLAGS_enable_prefill_piecewise_graph = false;
FLAGS_enable_xattention_one_stage = false;
FLAGS_enable_graph_mode_decode_no_padding = false;
FLAGS_enable_rec_prefill_only = false;
FLAGS_enable_constrained_decoding = false;
FLAGS_enable_topk_sorted = false;
}
void apply_multi_round_pipeline_toggles() {
FLAGS_enable_rec_fast_sampler = true;
FLAGS_enable_prefill_piecewise_graph = true;
FLAGS_enable_xattention_one_stage = false;
FLAGS_enable_graph_mode_decode_no_padding = true;
FLAGS_enable_topk_sorted = false;
}
void apply_onerec_pipeline_toggles(xllm::Options* options) {
const bool enable_onerec_xattention = FLAGS_max_decode_rounds > 0;
FLAGS_enable_rec_prefill_only = !enable_onerec_xattention;
FLAGS_enable_constrained_decoding = true;
FLAGS_enable_prefix_cache = false;
FLAGS_enable_schedule_overlap = false;
FLAGS_enable_chunked_prefill = false;
options->enable_prefix_cache(false)
.enable_schedule_overlap(false)
.enable_chunked_prefill(false);
if (!enable_onerec_xattention) {
// Legacy OneRec keeps the historical fixed decode-step behavior.
FLAGS_max_decode_rounds = 0;
}
}
} // namespace
XLLM_CAPI_EXPORT XLLM_REC_Handler* xllm_rec_create(void) {
XLLM_REC_Handler* handler = new XLLM_REC_Handler();
CHECK(nullptr != handler);
handler->initialized = false;
handler->pipeline_type = xllm::RecPipelineType::kLlmRecDefault;
return handler;
}
XLLM_CAPI_EXPORT void xllm_rec_destroy(XLLM_REC_Handler* handler) {
if (!handler) return;
handler->master.reset();
handler->executor.reset();
handler->model_ids.clear();
handler->pipeline_type = xllm::RecPipelineType::kLlmRecDefault;
handler->initialized = false;
delete handler;
}
XLLM_CAPI_EXPORT void xllm_rec_init_options_default(
XLLM_InitOptions* init_options) {
if (nullptr == init_options) return;
*init_options = XLLM_INIT_REC_OPTIONS_DEFAULT;
}
XLLM_CAPI_EXPORT bool xllm_rec_initialize(
XLLM_REC_Handler* handler,
const char* model_path,
const char* devices,
const XLLM_InitOptions* init_options) {
if (!handler || !model_path || !devices) return false;
try {
XLLM_InitOptions xllm_init_options;
xllm::helper::set_init_options(
xllm::helper::BackendType::REC, init_options, &xllm_init_options);
std::string log_dir(xllm_init_options.log_dir);
if (!log_dir.empty()) {
xllm::helper::init_log(xllm_init_options.log_dir);
}
if (!std::filesystem::exists(model_path)) {
LOG(ERROR) << "model path[" << model_path << "] does not exist";
return false;
}
xllm::Options options;
options.model_path(model_path)
.task_type(xllm_init_options.task)
.devices(devices)
.draft_model_path(xllm_init_options.draft_model)
.draft_devices(xllm_init_options.draft_devices)
.backend("rec")
.block_size(xllm_init_options.block_size)
.max_cache_size(xllm_init_options.max_cache_size)
.max_memory_utilization(xllm_init_options.max_memory_utilization)
.enable_prefix_cache(xllm_init_options.enable_prefix_cache)
.max_tokens_per_batch(xllm_init_options.max_tokens_per_batch)
.max_seqs_per_batch(xllm_init_options.max_seqs_per_batch)
.max_tokens_per_chunk_for_prefill(
xllm_init_options.max_tokens_per_chunk_for_prefill)
.num_speculative_tokens(xllm_init_options.num_speculative_tokens)
.num_request_handling_threads(
xllm_init_options.num_request_handling_threads)
.communication_backend(xllm_init_options.communication_backend)
.expert_parallel_degree(xllm_init_options.expert_parallel_degree)
.enable_chunked_prefill(xllm_init_options.enable_chunked_prefill)
.master_node_addr(xllm_init_options.master_node_addr)
.device_ip(xllm_init_options.device_ip)
.transfer_listen_port(xllm_init_options.transfer_listen_port)
.nnodes(xllm_init_options.nnodes)
.node_rank(xllm_init_options.node_rank)
.dp_size(xllm_init_options.dp_size)
.ep_size(xllm_init_options.ep_size)
.instance_name(xllm_init_options.instance_name)
.enable_disagg_pd(xllm_init_options.enable_disagg_pd)
.enable_schedule_overlap(xllm_init_options.enable_schedule_overlap)
.enable_pd_ooc(xllm_init_options.enable_pd_ooc)
.kv_cache_transfer_mode(xllm_init_options.kv_cache_transfer_mode)
.enable_shm(xllm_init_options.enable_shm)
.is_local(true)
.server_idx(xllm_init_options.server_idx);
// @TODO: Currently, gflags are configured through hard coding, which needs
// to be improved in the future. For example, a separate gflags
// configuration file can be provided to the so for setting gflags.
//
// REC so still has two configuration paths:
// - some request/runtime code reads FLAGS_* directly
// - master/worker construction reads xllm::Options
//
// The fields copied from init options below are read from FLAGS_* today.
// beam_width/block_size/max_tokens/max_seqs are also represented in
// Options, so duplicated values must stay aligned.
FLAGS_beam_width = xllm_init_options.beam_width;
FLAGS_max_decode_rounds = xllm_init_options.max_decode_rounds;
FLAGS_max_seqs_per_batch = xllm_init_options.max_seqs_per_batch;
FLAGS_max_tokens_per_batch = xllm_init_options.max_tokens_per_batch;
FLAGS_block_size = xllm_init_options.block_size;
FLAGS_enable_rec_prefill_only = xllm_init_options.enable_rec_prefill_only;
FLAGS_enable_prefix_cache = xllm_init_options.enable_prefix_cache;
FLAGS_enable_schedule_overlap = xllm_init_options.enable_schedule_overlap;
FLAGS_enable_chunked_prefill = xllm_init_options.enable_chunked_prefill;
FLAGS_enable_graph = xllm_init_options.enable_graph;
FLAGS_rec_worker_max_concurrency =
xllm_init_options.rec_worker_max_concurrency;
FLAGS_enable_block_copy_kernel = xllm_init_options.enable_block_copy_kernel;
auto model_loader = xllm::ModelLoader::create(model_path);
if (model_loader == nullptr) {
LOG(ERROR) << "Failed to create model loader for path: " << model_path;
return false;
}
const auto& model_args = model_loader->model_args();
const xllm::RecModelKind rec_model_kind =
xllm::get_rec_model_kind(model_args.model_type());
if (rec_model_kind == xllm::RecModelKind::kNone) {
LOG(ERROR) << "Unsupported rec model_type: " << model_args.model_type();
return false;
}
const xllm::RecPipelineType pipeline_type =
xllm::get_rec_pipeline_type(rec_model_kind);
// Pipeline-specific runtime toggles in the REC so path.
reset_pipeline_runtime_toggles();
switch (pipeline_type) {
case xllm::RecPipelineType::kLlmRecMultiRoundPipeline:
apply_multi_round_pipeline_toggles();
break;
case xllm::RecPipelineType::kOneRecDefault:
case xllm::RecPipelineType::kOneRecXAttentionPipeline:
apply_onerec_pipeline_toggles(&options);
break;
case xllm::RecPipelineType::kLlmRecDefault:
case xllm::RecPipelineType::kLlmRecWithMmData:
break;
default:
LOG(ERROR) << "Unsupported rec pipeline type: "
<< static_cast<int32_t>(pipeline_type);
return false;
}
// Keep dual-source settings aligned with the FLAGS_* values above.
options.enable_graph(FLAGS_enable_graph)
.beam_width(FLAGS_beam_width)
.rec_worker_max_concurrency(FLAGS_rec_worker_max_concurrency);
LOG(INFO) << "REC C API selected pipeline="
<< get_rec_pipeline_name(pipeline_type)
<< ", model_type=" << model_args.model_type()
<< ", enable_rec_prefill_only=" << FLAGS_enable_rec_prefill_only
<< ", enable_constrained_decoding="
<< FLAGS_enable_constrained_decoding
<< ", enable_prefix_cache=" << FLAGS_enable_prefix_cache
<< ", enable_schedule_overlap=" << FLAGS_enable_schedule_overlap
<< ", enable_chunked_prefill=" << FLAGS_enable_chunked_prefill
<< ", enable_rec_fast_sampler=" << FLAGS_enable_rec_fast_sampler
<< ", max_decode_rounds=" << FLAGS_max_decode_rounds;
#if !defined(USE_NPU) && !defined(USE_CUDA)
FLAGS_enable_block_copy_kernel = false;
#endif
handler->master = std::make_unique<xllm::RecMaster>(options);
handler->master->run();
size_t cpu_cores = std::thread::hardware_concurrency();
size_t thread_num = std::clamp((cpu_cores == 0) ? 8 : cpu_cores / 2,
static_cast<size_t>(4),
static_cast<size_t>(16));
handler->executor =
std::make_unique<folly::CPUThreadPoolExecutor>(thread_num);
std::filesystem::path model_path_fs =
std::filesystem::path(model_path).lexically_normal();
std::string model_id;
if (model_path_fs.has_filename()) {
model_id = model_path_fs.filename().string();
} else if (!model_path_fs.empty()) {
model_id = model_path_fs.string();
} else {
model_id = "default";
}
handler->model_ids.clear();
handler->model_ids.emplace_back(model_id);
handler->pipeline_type = pipeline_type;
handler->initialized = true;
return true;
} catch (const std::exception& e) {
LOG(ERROR) << "LLM initialization failed: " << e.what();
}
handler->master.reset();
handler->executor.reset();
handler->model_ids.clear();
handler->pipeline_type = xllm::RecPipelineType::kLlmRecDefault;
handler->initialized = false;
return false;
}
XLLM_CAPI_EXPORT void xllm_rec_request_params_default(
XLLM_RequestParams* request_params) {
if (nullptr == request_params) return;
*request_params = XLLM_REC_REQUEST_PARAMS_DEFAULT;
}
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_text_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const char* prompt,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params) {
if (!handler || !model_id || *model_id == '\0' || !prompt ||
*prompt == '\0') {
return xllm::helper::build_error_response(
"", XLLM_StatusCode::kInvalidRequest, "Invalid input parameters");
}
return xllm::helper::handle_inference_request(
handler,
xllm::helper::InferenceType::REC_COMPLETIONS,
model_id,
prompt,
nullptr,
timeout_ms,
request_params);
}
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_token_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const int32_t* token_ids,
size_t token_size,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params) {
if (!handler || !model_id || *model_id == '\0' || !token_ids ||
token_size == 0) {
return xllm::helper::build_error_response(
"", XLLM_StatusCode::kInvalidRequest, "Invalid input parameters");
}
std::vector<int> token_ids_vec;
for (int i = 0; i < token_size; i++) {
token_ids_vec.push_back(token_ids[i]);
}
return xllm::helper::handle_inference_request(
handler,
xllm::helper::InferenceType::REC_COMPLETIONS,
model_id,
token_ids_vec,
nullptr,
timeout_ms,
request_params);
}
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_multimodal_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const int32_t* token_ids,
size_t token_size,
const XLLM_MM_Data* mm_data,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params) {
if (!handler || !model_id || *model_id == '\0' || !token_ids ||
token_size == 0) {
return xllm::helper::build_error_response(
"", XLLM_StatusCode::kInvalidRequest, "Invalid input parameters");
}
if (!mm_data) {
return xllm_rec_token_completions(
handler, model_id, token_ids, token_size, timeout_ms, request_params);
}
xllm::MMData internal_mm_data;
try {
bool ret = xllm::helper::convert_xllm_mm_data_to_internal(mm_data,
internal_mm_data);
if (!ret) {
return xllm::helper::build_error_response(
"", XLLM_StatusCode::kInternalError, "Fail in mm_data conversion");
}
} catch (const std::exception& e) {
return xllm::helper::build_error_response(
"",
XLLM_StatusCode::kInternalError,
"Critical error in mm_data conversion: " + std::string(e.what()));
}
std::vector<int> token_ids_vec;
for (int i = 0; i < token_size; i++) {
token_ids_vec.push_back(token_ids[i]);
}
return xllm::helper::handle_inference_request(
handler,
xllm::helper::InferenceType::REC_COMPLETIONS,
model_id,
token_ids_vec,
static_cast<void*>(&internal_mm_data),
timeout_ms,
request_params);
}
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_chat_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const XLLM_ChatMessage* messages,
size_t messages_count,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params) {
if (!handler || !model_id || *model_id == '\0' || !messages ||
messages_count == 0) {
return xllm::helper::build_error_response(
"", XLLM_StatusCode::kInvalidRequest, "Invalid input parameters");
}
std::vector<xllm::Message> xllm_messages;
xllm_messages.reserve(messages_count);
for (int i = 0; i < messages_count; i++) {
xllm_messages.emplace_back(messages[i].role, messages[i].content);
}
return xllm::helper::handle_inference_request(
handler,
xllm::helper::InferenceType::REC_CHAT_COMPLETIONS,
model_id,
xllm_messages,
nullptr,
timeout_ms,
request_params);
}
XLLM_CAPI_EXPORT void xllm_rec_free_response(XLLM_Response* resp) {
return xllm::helper::xllm_free_response(resp);
}

View File

@@ -0,0 +1,227 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#ifndef XLLM_LLM_API_H
#define XLLM_LLM_API_H
#ifdef __cplusplus
extern "C" {
#endif
#include <stdbool.h>
#include <stddef.h>
#include <stdint.h>
#include "types.h"
/**
* @brief Opaque handle to an LLM inference instance
*
* This handle encapsulates all internal state of an LLM inference runtime,
* including model weights, device context, and generation cache.
* The handle MUST be created via xllm_llm_create() and destroyed via
* xllm_llm_destroy() to prevent memory/device resource leaks.
*/
typedef struct XLLM_LLM_Handler XLLM_LLM_Handler;
/**
* @brief Create a new LLM inference instance handle
*
* Allocates memory and initializes a new LLM handler with default internal
* state (empty model, uninitialized device context). This is the first function
* that must be called before using any other LLM APIs.
*
* @return Valid XLLM_LLM_Handler* on success; NULL if memory allocation fails
* @see xllm_llm_destroy
*/
XLLM_CAPI_EXPORT XLLM_LLM_Handler* xllm_llm_create(void);
/**
* @brief Destroy an LLM instance handle and release all associated resources
*
* Frees all memory allocated for the LLM instance, including:
* - Model weights (host/device memory)
* - Runtime context (CUDA/NPU streams, compute graphs)
* - Generation cache and temporary buffers
* - Device resources (contexts, queues)
*
* This function is idempotent—calling with NULL has no effect.
*
* @param handler LLM instance handle (NULL = no operation)
* @note Mandatory: Must be called to avoid memory/device resource leaks
* @see xllm_llm_create
*/
XLLM_CAPI_EXPORT void xllm_llm_destroy(XLLM_LLM_Handler* handler);
/**
* @brief Initialize XLLM_InitOptions with canonical default values
*
* Populates the XLLM_InitOptions struct with industry-standard default values
*
* @param init_options Pointer to XLLM_InitOptions to initialize (NULL = no-op)
* @see xllm_llm_initialize, XLLM_INIT_LLM_OPTIONS_DEFAULT
*/
XLLM_CAPI_EXPORT void xllm_llm_init_options_default(
XLLM_InitOptions* init_options);
/**
* @brief Initialize the LLM model and runtime environment
*
* Loads model weights from the specified path, configures target devices,
* initializes compute contexts, and prepares the inference runtime.
* Must be called exactly once per handler before using completion/chat APIs.
*
* If init_options is NULL, this function automatically uses the default values
* from XLLM_INIT_LLM_OPTIONS_DEFAULT (via xllm_llm_init_options_default()).
*
* @param handler Valid LLM instance handle (must not be NULL)
* @param model_path Null-terminated string of the model directory/file path
* (supports .bin/.pth/.safetensors formats)
* @param devices Null-terminated string specifying target devices (format:
* "npu:0,1" (specific NPUs), "cuda:0" (single GPU), "auto"
* (automatic selection))
* @param init_options Advanced initialization options (NULL = use defaults)
*
* @return true if initialization succeeds; false on failure (see failure causes
* below)
*
* @failure_causes
* - Invalid handler (NULL or already destroyed)
* - Invalid model_path (non-existent, corrupted, or unsupported format)
* - Invalid devices string (malformed format or unavailable devices)
* - Model load error (mismatched model architecture or weight corruption)
* - Device initialization failure (out of memory, driver error)
*
* @see xllm_llm_init_options_default, XLLM_INIT_LLM_OPTIONS_DEFAULT,
* xllm_llm_create
*/
XLLM_CAPI_EXPORT bool xllm_llm_initialize(XLLM_LLM_Handler* handler,
const char* model_path,
const char* devices,
const XLLM_InitOptions* init_options);
/**
* @brief Initialize XLLM_RequestParams with canonical generation defaults
*
* Populates the XLLM_RequestParams struct with safe default generation values
*
* @param request_params Pointer to XLLM_RequestParams to initialize (NULL =
* no-op)
* @see xllm_llm_completions, xllm_llm_chat_completions,
* XLLM_LLM_REQUEST_PARAMS_DEFAULT
*/
XLLM_CAPI_EXPORT void xllm_llm_request_params_default(
XLLM_RequestParams* request_params);
/**
* @brief Generate text completions for a single prompt
*
* Generates continuation text for the input prompt using the initialized LLM
* model. Returns a dynamically allocated response struct that MUST be freed
* with xllm_llm_free_response() to avoid memory leaks.
*
* If request_params is NULL, this function automatically uses the default
* values from XLLM_LLM_REQUEST_PARAMS_DEFAULT (via
* xllm_llm_request_params_default()).
*
* @param handler Valid, initialized LLM instance handle (must not be NULL)
* @param model_id Null-terminated string of the loaded model ID (must match
* model_path)
* @param prompt Null-terminated string of input text to complete (non-empty)
* @param timeout_ms Timeout in milliseconds (0 = no timeout, wait indefinitely)
* @param request_params Generation parameters (NULL = use defaults)
*
* @return Pointer to XLLM_Response on success; NULL ONLY if memory allocation
* fails (response->status indicates the actual result status)
*
* @response_status_codes
* - kSuccess: Valid response generated (check response->choices for results)
* - kNotInitialized: Handler not initialized with xllm_llm_initialize()
* - kInvalidRequest: Invalid prompt (empty/NULL) or model_id (mismatch)
* - kTimeout: Generation exceeded timeout_ms (partial results may be available)
*
* @warning Mandatory: Call xllm_llm_free_response() to release response memory
* @see xllm_llm_request_params_default, XLLM_LLM_REQUEST_PARAMS_DEFAULT,
* xllm_llm_free_response
*/
XLLM_CAPI_EXPORT XLLM_Response* xllm_llm_completions(
XLLM_LLM_Handler* handler,
const char* model_id,
const char* prompt,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
/**
* @brief Generate chat completions from a conversation history
*
* Generates model responses for a multi-turn conversation using chat-formatted
* message history (user/assistant/system roles). Returns a dynamically
* allocated response struct that MUST be freed with xllm_llm_free_response().
*
* If request_params is NULL, this function automatically uses the default
* values from XLLM_LLM_REQUEST_PARAMS_DEFAULT (via
* xllm_llm_request_params_default()).
*
* @param handler Valid, initialized LLM instance handle (must not be NULL)
* @param model_id Null-terminated string of the loaded model ID
* @param messages Array of XLLM_ChatMessage structs (conversation history)
* @param messages_count Number of messages in the messages array (must be ≥ 0)
* @param timeout_ms Timeout in milliseconds (0 = no timeout)
* @param request_params Generation parameters (NULL = use defaults)
*
* @return Pointer to XLLM_Response on success; NULL ONLY if memory allocation
* fails (response->status indicates the actual result status)
*
* @response_status_codes
* - kSuccess: Valid chat response generated (check
* response->choices[0].message)
* - kNotInitialized: Handler not initialized
* - kInvalidRequest: Invalid messages (NULL with count>0, empty role/content)
* - kTimeout: Generation exceeded timeout_ms
*
* @warning Mandatory: Call xllm_llm_free_response() to release response memory
* @see xllm_llm_request_params_default, XLLM_LLM_REQUEST_PARAMS_DEFAULT,
* xllm_llm_free_response
*/
XLLM_CAPI_EXPORT XLLM_Response* xllm_llm_chat_completions(
XLLM_LLM_Handler* handler,
const char* model_id,
const XLLM_ChatMessage* messages,
size_t messages_count,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
/**
* @brief Free all dynamically allocated memory in an XLLM_Response
*
* Releases all heap memory used by the response struct
*
* After freeing, all fields are reset to safe defaults (NULL/0) to prevent
* use-after-free.
*
* @param resp Pointer to XLLM_Response to free (NULL = no operation)
*
* @note Idempotent: Safe to call multiple times on the same response
* @warning Mandatory: Must be called after using completions/chat completions
* responses
* @see xllm_llm_completions, xllm_llm_chat_completions
*/
XLLM_CAPI_EXPORT void xllm_llm_free_response(XLLM_Response* resp);
#ifdef __cplusplus
}
#endif
#endif // XLLM_LLM_API_H

View File

@@ -0,0 +1,339 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#ifndef XLLM_REC_API_H
#define XLLM_REC_API_H
#ifdef __cplusplus
extern "C" {
#endif
#include <stdbool.h>
#include <stddef.h>
#include <stdint.h>
#include "types.h"
/**
* @brief Opaque handle to a Generative Recommendation (REC) inference instance
* This handle encapsulates all internal state of a REC-specialized runtime,
* including:
* - Generative recommendation model weights (item embedding, ranking head)
* - Device context (CUDA/NPU streams for batch inference)
* - Generation cache (user behavior context, item candidate pool)
* - Runtime config (recommendation-specific decoding strategy)
* The handle MUST be created via xllm_rec_create() and destroyed via
* xllm_rec_destroy() to prevent memory/device resource leaks.
*/
typedef struct XLLM_REC_Handler XLLM_REC_Handler;
/**
* @brief Create a new Generative Recommendation (REC) inference instance handle
* This is the first function that must be called before using any other REC
* APIs.
* @return Valid XLLM_REC_Handler* on success; NULL if memory allocation fails
* @see xllm_rec_destroy
*/
XLLM_CAPI_EXPORT XLLM_REC_Handler* xllm_rec_create(void);
/**
* @brief Destroy a Generative Recommendation (REC) inference instance handle
* and release resources Frees all memory allocated for the REC instance,
* including:
* - Model weights (host/device memory for item embedding and ranking head)
* - Runtime context (CUDA/NPU streams, compute graphs for batch recommendation)
* - Generation cache (user behavior sequence, item candidate pool, attention
* cache)
* - Device resources (contexts, queues, memory pools for batch inference)
* This function is idempotent—calling with NULL has no effect.
* @param handler REC inference instance handle (NULL = no operation)
* @note Mandatory: Must be called to avoid memory/device resource leaks
* @see xllm_rec_create
*/
XLLM_CAPI_EXPORT void xllm_rec_destroy(XLLM_REC_Handler* handler);
/**
* @brief Helper to initialize XLLM_InitOptions with REC default values
* Copies the predefined XLLM_INIT_REC_OPTIONS_DEFAULT values into the target
* init_options struct. Convenient alternative to manually setting each field,
* ensuring consistency with REC best practices.
* @param init_options Pointer to XLLM_InitOptions to initialize (NULL = no-op)
* @see XLLM_INIT_REC_OPTIONS_DEFAULT, xllm_rec_initialize
*/
XLLM_CAPI_EXPORT void xllm_rec_init_options_default(
XLLM_InitOptions* init_options);
/**
* @brief Initialize the Generative Recommendation (REC) model and runtime
* environment Loads generative recommendation model weights from the specified
* path, configures target devices, initializes compute contexts, and prepares
* the recommendation inference runtime
* @param handler Valid REC inference instance handle (must not be NULL)
* @param model_path Null-terminated string of the REC model directory/file path
* (supports .bin/.pth/.safetensors formats with ranking head)
* @param devices Null-terminated string specifying target devices (format:
* "npu:0,1" (specific NPUs), "cuda:0" (single GPU), "auto"
* (automatic selection))
* @param init_options Advanced initialization options (NULL = use REC defaults)
* @return true if initialization succeeds; false on failure (see failure causes
* below)
* @par Failure Causes
* - Invalid handler (NULL or already destroyed)
* - Invalid model_path (non-existent, corrupted, or missing ranking head
* weights)
* - Invalid devices string (malformed format or unavailable devices)
* - Model load error (mismatched REC model architecture or embedding table
* corruption)
* - Device initialization failure (out of memory, driver error, insufficient
* batch size)
* @see xllm_rec_init_options_default, XLLM_INIT_REC_OPTIONS_DEFAULT,
* xllm_rec_create
*/
XLLM_CAPI_EXPORT bool xllm_rec_initialize(XLLM_REC_Handler* handler,
const char* model_path,
const char* devices,
const XLLM_InitOptions* init_options);
/**
* @brief Helper to initialize XLLM_RequestParams with REC default values
* Copies the predefined XLLM_REC_REQUEST_PARAMS_DEFAULT values into the target
* request_params struct.
* @param request_params Pointer to XLLM_RequestParams to initialize (NULL =
* no-op)
* @see XLLM_REC_REQUEST_PARAMS_DEFAULT, xllm_rec_text_completions,
* xllm_rec_token_completions, xllm_rec_chat_completions
*/
XLLM_CAPI_EXPORT void xllm_rec_request_params_default(
XLLM_RequestParams* request_params);
/**
* @brief Generate generative recommendation text completions for a user prompt
* Generates recommendation-focused continuation text for the input user prompt
* using the initialized REC model
* @param handler Valid, initialized REC inference instance handle (must not be
* NULL)
* @param model_id Null-terminated string of the loaded REC model ID (must match
* model_path)
* @param prompt Null-terminated string of user input prompt (non-empty,
* recommendation-focused)
* @param timeout_ms Timeout in milliseconds (0 = no timeout, wait indefinitely)
* @param request_params Generation parameters (NULL = use REC defaults)
* @return Pointer to XLLM_Response on success; NULL ONLY if memory allocation
* fails (response->status indicates the actual result status)
* @par Response Status Codes
* - kSuccess: Valid recommendation response generated (check response->choices
* for item list + explanations)
* - kNotInitialized: Handler not initialized with xllm_rec_initialize()
* - kInvalidRequest: Invalid prompt (empty/NULL) or model_id (mismatch)
* - kTimeout: Generation exceeded timeout_ms (partial recommendation results
* may be available)
* @warning Mandatory: Call xllm_rec_free_response() to release response memory
* @see xllm_rec_request_params_default, XLLM_REC_REQUEST_PARAMS_DEFAULT,
* xllm_rec_free_response
*/
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_text_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const char* prompt,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
/**
* @brief Generate generative recommendation completions for tokenized input
* (TOKEN ID INPUT) Generates recommendation results from pre-tokenized user
* input (bypasses the REC model's tokenizer)
*
* @param handler Valid, initialized REC inference instance handle (must not be
* NULL) Created via xllm_rec_create() and initialized via xllm_rec_initialize()
* @param model_id Null-terminated string of the loaded REC model ID (must match
* the model_path used in xllm_rec_initialize())
* @param token_ids Pointer to int32_t array of pre-tokenized input IDs (NULL
* only if token_size = 0) Token IDs must be compatible with the REC model's
* tokenizer vocabulary (e.g., GPT-2/BERT token IDs for text-based REC models)
* @param token_size Number of tokens in the token_ids array (must be ≥ 0)
* Valid ranges: 1 ≤ token_size ≤ xxx (model-dependent max input
* length) token_size = 0 will return kInvalidRequest status
* @param timeout_ms Timeout in milliseconds (0 = no timeout, wait indefinitely)
* @param request_params Generation parameters (NULL = use REC defaults)
*
* @return Pointer to XLLM_Response on success; NULL ONLY if memory allocation
* fails (response->status indicates the actual result status, even if non-NULL)
*
* @par Response Status Codes (XLLM_StatusCode)
* - kSuccess: Valid recommendation response generated
* Check response->choices for recommended item list and explanation
* text
* - kNotInitialized: Handler not initialized with xllm_rec_initialize()
* - kModelNotFound: model_id does not match any loaded REC model
* - kInvalidRequest:
* - token_ids = NULL and token_size > 0 (invalid null pointer with non-zero
* size)
* - token_size = 0 (empty token input)
* - token_ids contain invalid IDs (out of vocabulary range)
* - model_id is NULL/empty/mismatch
* - kTimeout: Generation exceeded timeout_ms
* - kInternalError: Internal REC runtime error (e.g., token embedding failure,
* item retrieval error)
* @warning Mandatory: Call xllm_rec_free_response() to release response memory
* @note 1. Token IDs must be generated using the SAME tokenizer as the REC
* model (e.g., same vocab.txt)
* 2. Invalid token IDs (e.g., < 0 or > vocab_size) will trigger
* kInvalidRequest or kInternalError
* 3. For token_size > model's max input length, the input will be
* truncated to max length
* @see xllm_rec_request_params_default, XLLM_REC_REQUEST_PARAMS_DEFAULT,
* xllm_rec_free_response
*/
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_token_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const int32_t* token_ids,
size_t token_size,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
/**
* @brief Generate generative recommendation completions for multimodal input
* (TOKEN ID + MULTIMODAL DATA INPUT)
* @details Generates recommendation results from pre-tokenized text input
* (MANDATORY) supplemented with multimodal data that replaces/augments
* information for specific tokens in the token_ids array. This API extends
* xllm_rec_token_completions to support multi-modal recommendation scenarios
* where partial text tokens are enriched with image/audio/video/embedding
* features (e.g., replacing product text tokens with image embeddings).
*
* @param handler Valid, initialized REC inference instance handle (must not be
* NULL) Created via xllm_rec_create() and initialized via xllm_rec_initialize()
* @param model_id Null-terminated string of the loaded REC model ID (must match
* the model_path used in xllm_rec_initialize())
* Must be a multi-modal REC model (text-only models will return
* kInvalidRequest)
* @param token_ids Pointer to int32_t array of pre-tokenized text input IDs
* (MUST NOT be NULL) Token IDs must be compatible with the REC model's
* tokenizer vocabulary This is the core input and cannot be empty (token_size >
* 0 required)
* @param token_size Number of tokens in the token_ids array (MUST be ≥ 1)
* Valid ranges: 1 ≤ token_size ≤ model-dependent max input
* length token_size = 0 will return kInvalidRequest status (core text input
* required)
* @param mm_data Pointer to multi-modal data container (XLLM_MM_Data) (NULL =
* no multimodal augmentation) Used to replace/augment information for specific
* tokens in token_ids (via XLLM_MM_TokenPos) Supports
* image/audio/video/embedding modalities (see XLLM_MM_Type) Must be valid
* (mm_data->type_mask != XLLM_MM_TYPE_NONE) if non-NULL, and token positions in
* mm_data must be within [0, token_size-1] (out-of-range positions trigger
* kInvalidRequest)
* @param timeout_ms Timeout in milliseconds (0 = no timeout, wait indefinitely)
* @param request_params Generation parameters (NULL = use REC defaults)
* See XLLM_RequestParams for configurable options (e.g.,
* top_k, top_p)
* @return Pointer to XLLM_Response on success; NULL ONLY if memory allocation
* fails (response->status indicates the actual result status, even if
* non-NULL)
* @par Response Status Codes (XLLM_StatusCode)
* - kSuccess: Valid multi-modal recommendation response generated
* Check response->choices for recommended item list and explanation
* text Multimodal data has been applied to augment/replace specified tokens
* - kNotInitialized: Handler not initialized with xllm_rec_initialize()
* - kModelNotFound: model_id does not match any loaded REC model
* - kInvalidRequest:
* - token_ids = NULL (core text input is mandatory)
* - token_size = 0 (empty core text input)
* - token_ids contain invalid IDs (out of vocabulary range)
* - model_id is NULL/empty/mismatch or is a text-only model
* - mm_data is non-NULL but invalid:
* - mm_data->type_mask = XLLM_MM_TYPE_NONE (empty multimodal data)
* - token positions in mm_data (XLLM_MM_TokenPos) are out of [0,
* token_size-1] range
* - mismatched tensor types/shape in mm_data (e.g., embedding dim mismatch)
* - kTimeout: Generation exceeded timeout_ms
* - kInternalError: Internal REC runtime error (e.g., multimodal embedding
* fusion failure, token augmentation/replacement error, item retrieval error)
* @warning Mandatory: Call xllm_rec_free_response() to release response memory
* Failing to free will cause memory leaks
* @note 1. Token IDs must be generated using the SAME tokenizer as the REC
* model (e.g., same vocab.txt)
* 2. Invalid token IDs (e.g., < 0 or > vocab_size) will trigger
* kInvalidRequest or kInternalError
* 3. For token_size > model's max input length, the input will be
* truncated to max length
* 4. mm_data is used to replace/augment specific tokens (via
* XLLM_MM_TokenPos.offset/length):
* - offset: start index of tokens in token_ids to be
* augmented/replaced
* - length: number of consecutive tokens to apply multimodal data to
* 5. If mm_data is NULL, this API behaves identically to
* xllm_rec_token_completions (text-only inference)
* 6. Multimodal data must be aligned with token positions (offset +
* length ≤ token_size)
* @see xllm_rec_token_completions, xllm_rec_request_params_default,
* XLLM_REC_REQUEST_PARAMS_DEFAULT, xllm_rec_free_response, XLLM_MM_Data,
* XLLM_MM_TokenPos
*/
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_multimodal_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const int32_t* token_ids,
size_t token_size,
const XLLM_MM_Data* mm_data,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
/**
* @brief Generate generative recommendation chat completions from multi-turn
* conversation history Generates personalized recommendation responses for a
* multi-turn user-assistant conversation
* @param handler Valid, initialized REC inference instance handle (must not be
* NULL)
* @param model_id Null-terminated string of the loaded REC model ID
* @param messages Array of XLLM_ChatMessage structs (recommendation-focused
* conversation history)
* @param messages_count Number of messages in the messages array (must be ≥ 0)
* @param timeout_ms Timeout in milliseconds (0 = no timeout, wait indefinitely)
* @param request_params Generation parameters (NULL = use REC defaults)
* @return Pointer to XLLM_Response on success; NULL ONLY if memory allocation
* fails (response->status indicates the actual result status)
* @par Response Status Codes
* - kSuccess: Valid chat recommendation response generated (check
* response->choices[0].message for item list)
* - kNotInitialized: Handler not initialized with xllm_rec_initialize()
* - kInvalidRequest: Invalid messages (NULL with count>0, empty role/content,
* non-recommendation context)
* - kTimeout: Generation exceeded timeout_ms
* @warning Mandatory: Call xllm_rec_free_response() to release response memory
* @see xllm_rec_request_params_default, XLLM_REC_REQUEST_PARAMS_DEFAULT,
* xllm_rec_free_response
*/
XLLM_CAPI_EXPORT XLLM_Response* xllm_rec_chat_completions(
XLLM_REC_Handler* handler,
const char* model_id,
const XLLM_ChatMessage* messages,
size_t messages_count,
uint32_t timeout_ms,
const XLLM_RequestParams* request_params);
/**
* @brief Free all dynamically allocated memory in a generative recommendation
* XLLM_Response Releases all heap memory used by the REC response struct
* @param resp Pointer to XLLM_Response to free (NULL = no operation)
* @warning Mandatory: Must be called after using REC completions/chat
* completions responses
* @see xllm_rec_text_completions, xllm_rec_token_completions,
* xllm_rec_chat_completions
*/
XLLM_CAPI_EXPORT void xllm_rec_free_response(XLLM_Response* resp);
#ifdef __cplusplus
}
#endif
#endif // XLLM_REC_API_H

View File

@@ -0,0 +1,280 @@
cmake_minimum_required(VERSION 3.10)
project(xllm_capi_test LANGUAGES CXX)
set(CMAKE_CXX_STANDARD 14)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
# -----------------------------------------------------------------------------
# Dependencies — Protobuf (vcpkg uses protobuf-config.cmake; package name is "protobuf")
# -----------------------------------------------------------------------------
# Recommended (matches main project / vcpkg.json):
# cmake -B build \
# -DCMAKE_TOOLCHAIN_FILE=$VCPKG_ROOT/scripts/buildsystems/vcpkg.cmake \
# -DVCPKG_MANIFEST_DIR=<path/to/xllm/repo_root>
# To reuse an existing xLLM build without re-installing deps:
# -DCMAKE_PREFIX_PATH=<xllm_build>/vcpkg_installed/<triplet>
#
# Fallback: system FindProtobuf / pkg-config / manual paths (no vcpkg).
find_package(protobuf CONFIG QUIET)
if(NOT protobuf_FOUND)
find_package(Protobuf CONFIG QUIET)
endif()
# vcpkg protobuf::protoc may not set Protobuf_PROTOC_EXECUTABLE; resolve from target.
if((protobuf_FOUND OR Protobuf_FOUND) AND TARGET protobuf::libprotobuf)
if(NOT Protobuf_PROTOC_EXECUTABLE AND TARGET protobuf::protoc)
get_target_property(Protobuf_PROTOC_EXECUTABLE protobuf::protoc IMPORTED_LOCATION)
if(NOT Protobuf_PROTOC_EXECUTABLE)
get_target_property(Protobuf_PROTOC_EXECUTABLE protobuf::protoc IMPORTED_LOCATION_RELEASE)
endif()
if(NOT Protobuf_PROTOC_EXECUTABLE)
get_target_property(Protobuf_PROTOC_EXECUTABLE protobuf::protoc IMPORTED_LOCATION_DEBUG)
endif()
endif()
if(NOT Protobuf_PROTOC_EXECUTABLE)
message(FATAL_ERROR
"Protobuf CONFIG found but could not resolve protoc (protobuf::protoc / Protobuf_PROTOC_EXECUTABLE).")
endif()
message(STATUS "Found Protobuf (CONFIG / vcpkg): protoc=${Protobuf_PROTOC_EXECUTABLE}")
set(_PROTOBUF_LIB protobuf::libprotobuf)
else()
set(Protobuf_FOUND FALSE)
set(protobuf_FOUND FALSE)
find_package(Protobuf QUIET)
if(Protobuf_FOUND)
message(STATUS "Found Protobuf (module): protoc=${Protobuf_PROTOC_EXECUTABLE}")
endif()
if(NOT Protobuf_FOUND)
find_package(PkgConfig QUIET)
if(PkgConfig_FOUND)
pkg_check_modules(PC_PROTOBUF QUIET protobuf)
endif()
if(PC_PROTOBUF_FOUND)
set(Protobuf_INCLUDE_DIRS ${PC_PROTOBUF_INCLUDE_DIRS})
set(Protobuf_LIBRARIES ${PC_PROTOBUF_LIBRARIES})
find_program(Protobuf_PROTOC_EXECUTABLE protoc)
if(NOT Protobuf_PROTOC_EXECUTABLE)
message(FATAL_ERROR
"Found libprotobuf via pkg-config, but protoc was not found in PATH. "
"Please install protoc or set Protobuf_PROTOC_EXECUTABLE.")
endif()
set(_PROTOBUF_LIB ${Protobuf_LIBRARIES})
endif()
endif()
if(NOT Protobuf_FOUND AND NOT PC_PROTOBUF_FOUND)
find_path(Protobuf_INCLUDE_DIR
NAMES google/protobuf/message.h
PATHS /usr/include /usr/local/include
)
find_library(Protobuf_LIBRARY
NAMES protobuf libprotobuf
PATHS /usr/lib /usr/lib64 /usr/local/lib /usr/local/lib64
)
find_program(Protobuf_PROTOC_EXECUTABLE
NAMES protoc
PATHS /usr/bin /usr/local/bin
)
if(Protobuf_INCLUDE_DIR AND Protobuf_LIBRARY AND Protobuf_PROTOC_EXECUTABLE)
set(Protobuf_INCLUDE_DIRS ${Protobuf_INCLUDE_DIR})
set(Protobuf_LIBRARIES ${Protobuf_LIBRARY})
set(_PROTOBUF_LIB ${Protobuf_LIBRARY})
message(STATUS "Found Protobuf manually:"
" include=${Protobuf_INCLUDE_DIR}"
" lib=${Protobuf_LIBRARY}"
" protoc=${Protobuf_PROTOC_EXECUTABLE}")
else()
message(FATAL_ERROR
"Could NOT find Protobuf.\n"
"Preferred: configure with vcpkg like the main xLLM build, e.g.\n"
" -DCMAKE_TOOLCHAIN_FILE=<vcpkg>/scripts/buildsystems/vcpkg.cmake\n"
" -DVCPKG_MANIFEST_DIR=<xllm_repo_root>\n"
"Or install system protobuf + protoc, or set:\n"
" -DProtobuf_INCLUDE_DIR=/path/to/include\n"
" -DProtobuf_LIBRARY=/path/to/libprotobuf.so\n"
" -DProtobuf_PROTOC_EXECUTABLE=/path/to/protoc\n")
endif()
endif()
if(TARGET protobuf::libprotobuf)
set(_PROTOBUF_LIB protobuf::libprotobuf)
endif()
endif()
# -----------------------------------------------------------------------------
# Generate protobuf sources from xllm_test.proto
# -----------------------------------------------------------------------------
set(XLLM_TEST_PROTO ${CMAKE_CURRENT_LIST_DIR}/xllm_test.proto)
set(XLLM_TEST_PB_CC ${CMAKE_CURRENT_BINARY_DIR}/xllm_test.pb.cc)
set(XLLM_TEST_PB_H ${CMAKE_CURRENT_BINARY_DIR}/xllm_test.pb.h)
add_custom_command(
OUTPUT ${XLLM_TEST_PB_CC} ${XLLM_TEST_PB_H}
COMMAND ${Protobuf_PROTOC_EXECUTABLE}
--cpp_out=${CMAKE_CURRENT_BINARY_DIR}
-I ${CMAKE_CURRENT_LIST_DIR}
${XLLM_TEST_PROTO}
DEPENDS ${XLLM_TEST_PROTO}
COMMENT "Generating xllm_test.pb.cc/h from xllm_test.proto"
VERBATIM
)
add_library(c_api_test_proto STATIC ${XLLM_TEST_PB_CC})
target_include_directories(c_api_test_proto PUBLIC ${CMAKE_CURRENT_BINARY_DIR})
if(TARGET protobuf::libprotobuf)
target_include_directories(c_api_test_proto PUBLIC
$<TARGET_PROPERTY:protobuf::libprotobuf,INTERFACE_INCLUDE_DIRECTORIES>)
elseif(Protobuf_INCLUDE_DIRS)
target_include_directories(c_api_test_proto PUBLIC ${Protobuf_INCLUDE_DIRS})
endif()
# PRIVATE: xllm_test controls link order (brpc before protobuf) for static archives.
target_link_libraries(c_api_test_proto PRIVATE ${_PROTOBUF_LIB})
# -----------------------------------------------------------------------------
# brpc + gflags + glog (server binary; brpc uses glog — link glog after libbrpc.a)
# -----------------------------------------------------------------------------
find_package(gflags CONFIG REQUIRED)
find_package(glog CONFIG REQUIRED)
find_package(Threads REQUIRED)
find_package(OpenSSL REQUIRED)
# brpc: main xLLM may use <repo>/build/third_party/brpc/output or
# <repo>/build/<toolchain>/third_party/brpc/output (e.g. cmake.linux-aarch64-*).
# Headers: output/include next to the library, else third_party/brpc/src.
get_filename_component(_XLLM_REPO_ROOT "${CMAKE_CURRENT_LIST_DIR}/../../.."
ABSOLUTE)
if(NOT BRPC_ROOT)
if(DEFINED ENV{BRPC_ROOT} AND NOT "$ENV{BRPC_ROOT}" STREQUAL "")
set(BRPC_ROOT "$ENV{BRPC_ROOT}")
else()
set(BRPC_ROOT
"${_XLLM_REPO_ROOT}/build/third_party/brpc/output")
endif()
endif()
set(BRPC_ROOT "${BRPC_ROOT}" CACHE PATH
"brpc prefix: include/brpc/server.h and lib/libbrpc.a (override if needed)")
set(BRPC_LIB "")
foreach(_brpc_lib_cand
"${BRPC_ROOT}/lib/libbrpc.a"
"${BRPC_ROOT}/lib/libbrpc.so"
"${_XLLM_REPO_ROOT}/build/third_party/brpc/output/lib/libbrpc.a"
"${_XLLM_REPO_ROOT}/build/third_party/brpc/output/lib/libbrpc.so")
if(EXISTS "${_brpc_lib_cand}")
set(BRPC_LIB "${_brpc_lib_cand}")
break()
endif()
endforeach()
# e.g. build/cmake.linux-aarch64-cpython-311/third_party/brpc/output/lib/...
if(NOT BRPC_LIB)
file(GLOB _brpc_lib_glob LIST_DIRECTORIES false
"${_XLLM_REPO_ROOT}/build/*/third_party/brpc/output/lib/libbrpc.a"
"${_XLLM_REPO_ROOT}/build/*/third_party/brpc/output/lib/libbrpc.so")
if(_brpc_lib_glob)
list(SORT _brpc_lib_glob)
list(GET _brpc_lib_glob 0 BRPC_LIB)
endif()
endif()
if(NOT BRPC_LIB)
find_library(BRPC_LIB NAMES brpc libbrpc.a
PATHS "${BRPC_ROOT}/lib"
"${CMAKE_PREFIX_PATH}/lib"
/usr/local/lib
/usr/lib64
/usr/lib)
endif()
set(BRPC_INCLUDE_DIR "")
if(BRPC_LIB)
get_filename_component(_brpc_out "${BRPC_LIB}" DIRECTORY)
get_filename_component(_brpc_out "${_brpc_out}" DIRECTORY)
if(EXISTS "${_brpc_out}/include/brpc/server.h")
set(BRPC_INCLUDE_DIR "${_brpc_out}/include")
endif()
endif()
if(NOT BRPC_INCLUDE_DIR AND EXISTS "${BRPC_ROOT}/include/brpc/server.h")
set(BRPC_INCLUDE_DIR "${BRPC_ROOT}/include")
endif()
if(NOT BRPC_INCLUDE_DIR
AND EXISTS "${_XLLM_REPO_ROOT}/third_party/brpc/src/brpc/server.h")
set(BRPC_INCLUDE_DIR "${_XLLM_REPO_ROOT}/third_party/brpc/src")
endif()
if(NOT BRPC_INCLUDE_DIR)
message(FATAL_ERROR
"brpc headers not found (BRPC_ROOT='${BRPC_ROOT}').\n"
" Expected include/brpc/server.h next to libbrpc, or submodule\n"
" '${_XLLM_REPO_ROOT}/third_party/brpc/src/brpc/server.h'.\n"
" Set -DBRPC_ROOT=... to a prefix with include/brpc/server.h, or init\n"
" git submodules so third_party/brpc exists.")
endif()
if(NOT BRPC_LIB)
message(FATAL_ERROR
"libbrpc not found (BRPC_ROOT='${BRPC_ROOT}').\n"
" xllm_test links brpc from the main xLLM build, e.g.\n"
" ${_XLLM_REPO_ROOT}/build/third_party/brpc/output/lib/libbrpc.a\n"
" ${_XLLM_REPO_ROOT}/build/*/third_party/brpc/output/lib/libbrpc.a\n"
" Build the main project first, or pass -DBRPC_ROOT=/path/to/.../output\n"
" (directory with include/ and lib/libbrpc.a).")
endif()
find_package(leveldb CONFIG REQUIRED)
find_package(ZLIB REQUIRED)
# -----------------------------------------------------------------------------
# Build xllm_test (links external /usr/local/xllm/lib/libxllm.so)
# -----------------------------------------------------------------------------
add_executable(xllm_test
${CMAKE_CURRENT_LIST_DIR}/xllm_test.cpp
${CMAKE_CURRENT_LIST_DIR}/utils.cpp
)
# rec.h / types.h: installed C API under /usr/local/xllm/include (run xllm/c_api/install.sh).
# Local test headers (utils.h): ${CMAKE_CURRENT_LIST_DIR}
target_include_directories(xllm_test
PRIVATE
${CMAKE_CURRENT_BINARY_DIR}
/usr/local/xllm/include
${CMAKE_CURRENT_LIST_DIR}
${BRPC_INCLUDE_DIR}
)
target_link_directories(xllm_test PRIVATE /usr/local/xllm/lib)
# Static libbrpc.a pulls protobuf gzip + glog symbols; libprotobuf must appear
# after brpc (or use --start-group) so GzipOutputStream etc. resolve; ZLIB for gzip.
# --start-group/--end-group: GNU ld needs them so libprotobuf resolves symbols
# referenced from libbrpc.a (e.g. GzipOutputStream) when linking static archives.
if(CMAKE_CXX_COMPILER_ID MATCHES "GNU|Clang" AND NOT MSVC)
set(_BRPC_LINK_GROUP_START -Wl,--start-group)
set(_BRPC_LINK_GROUP_END -Wl,--end-group)
else()
set(_BRPC_LINK_GROUP_START "")
set(_BRPC_LINK_GROUP_END "")
endif()
target_link_libraries(xllm_test
PRIVATE
c_api_test_proto
gflags::gflags
${_BRPC_LINK_GROUP_START}
${BRPC_LIB}
${_PROTOBUF_LIB}
glog::glog
ZLIB::ZLIB
${_BRPC_LINK_GROUP_END}
leveldb::leveldb
OpenSSL::SSL
OpenSSL::Crypto
Threads::Threads
dl
xllm
)
# Keep runtime able to locate libxllm.so without setting LD_LIBRARY_PATH.
set_target_properties(xllm_test PROPERTIES
BUILD_RPATH "/usr/local/xllm/lib"
INSTALL_RPATH "/usr/local/xllm/lib"
)

View File

@@ -0,0 +1,109 @@
# c_api/test — `xllm_test` 说明
## 作用
`xllm_test` 是一个基于 **brpc** 的小型服务进程,用于在 **RPC 层** 验证 xLLM 的 **C API**(`xllm/c_api/llm.h` 或 `rec.h`):
- 对客户端暴露 **一个** RPC:`Inference(XLLM_Request) -> XLLM_Response`(定义见 `xllm_test.proto`)。
- 根据请求里的 **`call_function`** 字符串,转发到对应的 C API(例如 `xllm_llm_completions`、`xllm_rec_text_completions` 等)。
- 请求/响应中的结构与 `types.h` 对齐,由 `utils.cpp` 在 **Protobuf** 与 **C 结构体** 之间做转换。
**注意**:一次进程只加载 **一种** 后端,由 **`--backend`** 决定:
| `--backend` | 使用的 C API | 仅有效的 `call_function` 前缀 |
|-------------|--------------|--------------------------------|
| `llm` | `llm.h` | `xllm_llm_*` |
| `rec` | `rec.h` | `xllm_rec_*` |
若后端与 `call_function` 不匹配(例如在 `rec` 模式下调用 `xllm_llm_completions`),会返回错误(例如 handler 为空)。
---
## 依赖与前置条件
1. **已安装的 C API 头文件与 `libxllm.so`**
默认按 **`/usr/local/xllm/include`** 与 **`/usr/local/xllm/lib`** 查找(与 `CMakeLists.txt` 一致)。
若尚未安装,可在仓库内执行 `xllm/c_api/install.sh`(或你们环境约定的安装方式)。
2. **主工程已构建出的 brpc**
`libbrpc.a`(及头文件)通常位于仓库根目录下类似路径:
`build/third_party/brpc/output/` 或
`build/<toolchain>/third_party/brpc/output/`(例如 `cmake.linux-aarch64-cpython-311`)。
CMake 会自动在 `build/*/third_party/brpc/output` 下搜索;若仍找不到,可设置:
`-DBRPC_ROOT=/path/to/.../third_party/brpc/output` 或环境变量 **`BRPC_ROOT`**。
3. **Protobuf、gflags、glog、leveldb、OpenSSL、Zlib**
推荐与主工程一致,使用 **vcpkg**(仓库根目录 `vcpkg.json`),配置时传入
`-DCMAKE_TOOLCHAIN_FILE=$VCPKG_ROOT/scripts/buildsystems/vcpkg.cmake`、
`-DVCPKG_MANIFEST_DIR=<xllm 仓库根目录>`。
---
## 编译
在 **`xllm/c_api/test`** 目录下新建构建目录并配置、编译(请将占位路径换成你本机路径):
```bash
cd xllm/c_api/test
cmake -B build \
-DCMAKE_BUILD_TYPE=Release \
-DCMAKE_TOOLCHAIN_FILE=$VCPKG_ROOT/scripts/buildsystems/vcpkg.cmake \
-DVCPKG_MANIFEST_DIR=/path/to/xllm
cmake --build build -j$(nproc)
```
生成可执行文件:**`build/xllm_test`**(具体路径以 CMake 生成位置为准)。
若 vcpkg 依赖已安装在主工程构建目录中,也可通过 **`CMAKE_PREFIX_PATH`** 指向
`<主工程 build>/vcpkg_installed/<triplet>`,避免重复安装。
---
## 运行
1. 编辑示例 flags:**`xllm_test.flags`**(至少设置 **`--model_path`**、**`--devices`**,并按需设置 **`--backend=llm`** 或 **`--backend=rec`**)。
2. 启动服务:
```bash
/path/to/build/xllm_test --flagfile=/path/to/xllm/c_api/test/xllm_test.flags
```
或在命令行直接传参,例如:
```bash
./build/xllm_test \
--backend=rec \
--model_path=/path/to/model \
--devices=auto \
--port=8000
```
3. **监听地址**
- 默认使用 **`--port`**(如 `8000`)在 `0.0.0.0` 上监听。
- 若设置 **`--listen_addr=host:port`**,则优先使用该地址(与 `xllm_test.flags` 中注释一致)。
4. **调用方式**
任意支持 **brpc + 同一套 `xllm_test.proto`** 的客户端,向上述地址发起 **`XllmRecCapiService/Inference`**,在 **`XLLM_Request.call_function`** 中填入与当前 **`--backend`** 一致的 API 名称即可。
---
## 目录内主要文件
| 文件 | 说明 |
|------|------|
| `xllm_test.cpp` | brpc 服务入口、`Inference` 分发逻辑 |
| `xllm_test.proto` | RPC 与消息定义 |
| `utils.cpp` / `utils.h` | Protobuf ↔ C API 类型转换、gflags 定义 |
| `xllm_test.flags` | 示例运行参数 |
| `CMakeLists.txt` | 构建配置 |
---
## 常见问题
- **`brpc` / `libbrpc.a` 找不到**:先在仓库根目录完整配置并编译主工程,使 `third_party/brpc` 产物出现;或使用 `-DBRPC_ROOT`。
- **链接或运行找不到 `libxllm.so`**:确认已安装到 **`/usr/local/xllm/lib`**,或自行修改 `CMakeLists.txt` 中的 include/lib 路径并设置 **`LD_LIBRARY_PATH`**。
- **与主进程 `127.0.0.1:18899` 相关日志**:那是 xLLM **分布式 engine/worker** 的地址,与 `xllm_test` 的 **`--port` / `--listen_addr`** 无关;需按主工程文档单独启动 engine。

View File

@@ -0,0 +1,647 @@
/* Copyright 2026 The xLLM 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 "utils.h"
#include <gflags/gflags.h>
#include <algorithm>
#include <cstring>
#include <string>
// --- Server + XLLM_InitOptions gflags (defaults aligned with REC defaults) ---
DEFINE_string(model_path, "", "Path to REC model weights");
DEFINE_string(devices, "auto", "Devices string, e.g. npu:0 or auto");
DEFINE_int32(port, 8000, "brpc TCP port");
DEFINE_string(listen_addr,
"",
"If non-empty, brpc listen endpoint (host:port), overrides port");
DEFINE_int32(idle_timeout_s,
-1,
"brpc connection idle timeout in seconds; -1 = no limit");
DEFINE_string(backend,
"rec",
"C API mode for xllm_test: \"llm\" (c_api/llm.h) or \"rec\" "
"(c_api/rec.h); only one is loaded");
DEFINE_bool(enable_chunked_prefill, false, "");
DEFINE_bool(enable_prefill_sp, false, "");
DEFINE_bool(enable_prefix_cache, false, "");
DEFINE_bool(enable_disagg_pd, false, "");
DEFINE_bool(enable_pd_ooc, false, "");
DEFINE_bool(enable_schedule_overlap, false, "");
DEFINE_bool(enable_shm, false, "");
DEFINE_uint32(transfer_listen_port, 26000, "");
DEFINE_uint32(nnodes, 1, "");
DEFINE_uint32(node_rank, 0, "");
DEFINE_uint32(dp_size, 1, "");
DEFINE_uint32(ep_size, 1, "");
DEFINE_uint32(block_size, 1, "");
DEFINE_uint32(max_cache_size, 1000000, "");
DEFINE_uint32(max_tokens_per_batch, 4096, "");
DEFINE_uint32(max_seqs_per_batch, 4, "");
DEFINE_uint32(max_tokens_per_chunk_for_prefill, 0, "");
DEFINE_uint32(num_speculative_tokens, 0, "");
DEFINE_uint32(num_request_handling_threads, 4, "");
DEFINE_uint32(expert_parallel_degree, 0, "");
DEFINE_uint32(server_idx, 0, "");
DEFINE_uint32(beam_width, 128, "");
DEFINE_uint32(max_decode_rounds, 3, "");
DEFINE_uint32(max_token_per_req, 1000, "");
DEFINE_double(max_memory_utilization, 0.55, "");
DEFINE_string(init_task, "generate", "XLLM_InitOptions.task");
DEFINE_string(communication_backend, "lccl", "");
DEFINE_string(instance_role, "DEFAULT", "");
DEFINE_string(device_ip, "", "");
DEFINE_string(master_node_addr, "127.0.0.1:18899", "");
DEFINE_string(xservice_addr, "", "");
DEFINE_string(instance_name, "", "");
DEFINE_string(kv_cache_transfer_mode, "PUSH", "");
// Not named "log_dir": glog already registers FLAGS_log_dir.
DEFINE_string(xllm_init_log_dir, "", "");
DEFINE_string(draft_model, "", "");
DEFINE_string(draft_devices, "", "");
namespace xllm_capi_test {
namespace {
void CopyToFixed(char* dst, const std::string& s, size_t cap) {
if (cap == 0) {
return;
}
std::strncpy(dst, s.c_str(), cap - 1);
dst[cap - 1] = '\0';
}
std::unique_ptr<char[]> CopyCStr(const std::string& s) {
auto p = std::make_unique<char[]>(s.size() + 1);
if (!s.empty()) {
std::memcpy(p.get(), s.data(), s.size());
}
p[s.size()] = '\0';
return p;
}
size_t TensorNumElements(const XLLM_Dims& d) {
if (d.rank <= 0) {
return 0;
}
size_t n = 1;
for (int i = 0; i < d.rank && i < 8; ++i) {
if (d.dim[i] <= 0) {
return 0;
}
n *= static_cast<size_t>(d.dim[i]);
}
return n;
}
size_t DTypeSize(XLLM_DataType dt) {
switch (dt) {
case XLLM_DTYPE_FLOAT16:
case XLLM_DTYPE_BFLOAT16:
return 2;
case XLLM_DTYPE_FLOAT32:
return 4;
case XLLM_DTYPE_FLOAT64:
return 8;
case XLLM_DTYPE_INT8:
case XLLM_DTYPE_UINT8:
return 1;
case XLLM_DTYPE_INT16:
case XLLM_DTYPE_UINT16:
return 2;
case XLLM_DTYPE_INT32:
case XLLM_DTYPE_UINT32:
return 4;
case XLLM_DTYPE_INT64:
case XLLM_DTYPE_UINT64:
return 8;
case XLLM_DTYPE_BOOL:
return 1;
case XLLM_DTYPE_STRING:
case XLLM_DTYPE_UNDEFINED:
default:
return 0;
}
}
} // namespace
void ApplyGflagsToXllmInitOptions(XLLM_InitOptions* o) {
o->enable_chunked_prefill = FLAGS_enable_chunked_prefill;
o->enable_prefill_sp = FLAGS_enable_prefill_sp;
o->enable_prefix_cache = FLAGS_enable_prefix_cache;
o->enable_disagg_pd = FLAGS_enable_disagg_pd;
o->enable_pd_ooc = FLAGS_enable_pd_ooc;
o->enable_schedule_overlap = FLAGS_enable_schedule_overlap;
o->enable_shm = FLAGS_enable_shm;
o->transfer_listen_port = FLAGS_transfer_listen_port;
o->nnodes = FLAGS_nnodes;
o->node_rank = FLAGS_node_rank;
o->dp_size = FLAGS_dp_size;
o->ep_size = FLAGS_ep_size;
o->block_size = FLAGS_block_size;
o->max_cache_size = FLAGS_max_cache_size;
o->max_tokens_per_batch = FLAGS_max_tokens_per_batch;
o->max_seqs_per_batch = FLAGS_max_seqs_per_batch;
o->max_tokens_per_chunk_for_prefill = FLAGS_max_tokens_per_chunk_for_prefill;
o->num_speculative_tokens = FLAGS_num_speculative_tokens;
o->num_request_handling_threads = FLAGS_num_request_handling_threads;
o->expert_parallel_degree = FLAGS_expert_parallel_degree;
o->server_idx = FLAGS_server_idx;
o->beam_width = FLAGS_beam_width;
o->max_decode_rounds = FLAGS_max_decode_rounds;
o->max_token_per_req = FLAGS_max_token_per_req;
o->max_memory_utilization = static_cast<float>(FLAGS_max_memory_utilization);
CopyToFixed(o->task, FLAGS_init_task, XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(o->communication_backend,
FLAGS_communication_backend,
XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(
o->instance_role, FLAGS_instance_role, XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(o->device_ip, FLAGS_device_ip, XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(o->master_node_addr,
FLAGS_master_node_addr,
XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(
o->xservice_addr, FLAGS_xservice_addr, XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(
o->instance_name, FLAGS_instance_name, XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(o->kv_cache_transfer_mode,
FLAGS_kv_cache_transfer_mode,
XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(
o->log_dir, FLAGS_xllm_init_log_dir, XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(
o->draft_model, FLAGS_draft_model, XLLM_META_STRING_FIELD_MAX_LEN);
CopyToFixed(
o->draft_devices, FLAGS_draft_devices, XLLM_META_STRING_FIELD_MAX_LEN);
}
void PbToXllmDims(const c_api_test::XLLM_Dims& pb, XLLM_Dims* out) {
std::memset(out->dim, 0, sizeof(out->dim));
out->rank = pb.rank();
const int n = std::min(8, pb.dim_size());
for (int i = 0; i < n; ++i) {
out->dim[i] = pb.dim(i);
}
}
void XllmDimsToPb(const XLLM_Dims& in, c_api_test::XLLM_Dims* pb) {
pb->set_rank(in.rank);
pb->clear_dim();
const int n = std::min(8, in.rank);
for (int i = 0; i < n; ++i) {
pb->add_dim(in.dim[i]);
}
}
void PbToXllmTensor(const c_api_test::XLLM_Tensor& pb,
XLLM_Tensor* out,
MmDataOwned* owned) {
out->dtype = static_cast<XLLM_DataType>(pb.dtype());
PbToXllmDims(pb.dims(), &out->dims);
owned->tensor_byte_buffers.emplace_back(pb.data().begin(), pb.data().end());
std::vector<uint8_t>& buf = owned->tensor_byte_buffers.back();
out->data = buf.empty() ? nullptr : static_cast<const void*>(buf.data());
}
void XllmTensorToPb(const XLLM_Tensor& in, c_api_test::XLLM_Tensor* pb) {
pb->set_dtype(static_cast<c_api_test::XLLM_DataType>(in.dtype));
XllmDimsToPb(in.dims, pb->mutable_dims());
const size_t n = TensorNumElements(in.dims);
const size_t es = DTypeSize(in.dtype);
if (in.data != nullptr && n > 0 && es > 0) {
pb->set_data(static_cast<const char*>(in.data), n * es);
} else {
pb->clear_data();
}
}
void PbToXllmTensors(const c_api_test::XLLM_Tensors& pb,
XLLM_Tensors* out,
MmDataOwned* owned) {
owned->tensor_lists.emplace_back();
std::vector<XLLM_Tensor>& row = owned->tensor_lists.back();
row.reserve(static_cast<size_t>(pb.entries_size()));
for (int i = 0; i < pb.entries_size(); ++i) {
XLLM_Tensor t{};
PbToXllmTensor(pb.entries(i), &t, owned);
row.push_back(t);
}
out->entries = row.data();
out->entries_size = row.size();
}
void XllmTensorsToPb(const XLLM_Tensors& in, c_api_test::XLLM_Tensors* pb) {
pb->clear_entries();
for (size_t i = 0; i < in.entries_size; ++i) {
XllmTensorToPb(in.entries[i], pb->add_entries());
}
}
void PbToXllmMmValue(const c_api_test::XLLM_MM_Value& pb,
XLLM_MM_Value* out,
MmDataOwned* owned) {
std::memset(out, 0, sizeof(*out));
out->is_single_tensor = pb.is_single_tensor();
switch (pb.data_case()) {
case c_api_test::XLLM_MM_Value::kTensor:
out->is_single_tensor = true;
PbToXllmTensor(pb.tensor(), &out->data.tensor, owned);
break;
case c_api_test::XLLM_MM_Value::kTensors:
out->is_single_tensor = false;
PbToXllmTensors(pb.tensors(), &out->data.tensors, owned);
break;
default:
break;
}
}
void XllmMmValueToPb(const XLLM_MM_Value& in, c_api_test::XLLM_MM_Value* pb) {
pb->set_is_single_tensor(in.is_single_tensor);
if (in.is_single_tensor) {
XllmTensorToPb(in.data.tensor, pb->mutable_tensor());
} else {
XllmTensorsToPb(in.data.tensors, pb->mutable_tensors());
}
}
void PbToXllmMmDict(const c_api_test::XLLM_MM_Dict& pb,
XLLM_MM_Dict* out,
MmDataOwned* owned) {
owned->mm_dict_entries.clear();
owned->mm_dict_entries.reserve(static_cast<size_t>(pb.entries_size()));
for (int i = 0; i < pb.entries_size(); ++i) {
owned->mm_dict_entries.emplace_back();
XLLM_MM_DictEntry& e = owned->mm_dict_entries.back();
std::memset(e.key, 0, sizeof(e.key));
const std::string& k = pb.entries(i).key();
std::strncpy(e.key, k.c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1);
PbToXllmMmValue(pb.entries(i).value(), &e.value, owned);
}
out->entries = owned->mm_dict_entries.data();
out->entries_size = owned->mm_dict_entries.size();
}
void XllmMmDictToPb(const XLLM_MM_Dict& in, c_api_test::XLLM_MM_Dict* pb) {
pb->clear_entries();
for (size_t i = 0; i < in.entries_size; ++i) {
c_api_test::XLLM_MM_DictEntry* e = pb->add_entries();
e->set_key(in.entries[i].key);
XllmMmValueToPb(in.entries[i].value, e->mutable_value());
}
}
void PbToXllmMmItems(const c_api_test::XLLM_MM_Items& pb,
XLLM_MM_Items* out,
MmDataOwned* owned) {
owned->mm_items.clear();
owned->mm_items.reserve(static_cast<size_t>(pb.entries_size()));
for (int i = 0; i < pb.entries_size(); ++i) {
owned->mm_items.emplace_back();
PbToXllmMmItem(pb.entries(i), &owned->mm_items.back(), owned);
}
out->entries = owned->mm_items.data();
out->entries_size = owned->mm_items.size();
}
void XllmMmItemsToPb(const XLLM_MM_Items& in, c_api_test::XLLM_MM_Items* pb) {
pb->clear_entries();
for (size_t i = 0; i < in.entries_size; ++i) {
XllmMmItemToPb(in.entries[i], pb->add_entries());
}
}
void PbToXllmMmState(const c_api_test::XLLM_MM_State& pb, XLLM_MM_State* out) {
out->token_pos.offset = pb.token_pos().offset();
out->token_pos.length = pb.token_pos().length();
}
void XllmMmStateToPb(const XLLM_MM_State& in, c_api_test::XLLM_MM_State* pb) {
pb->mutable_token_pos()->set_offset(in.token_pos.offset);
pb->mutable_token_pos()->set_length(in.token_pos.length);
}
void PbToXllmMmItem(const c_api_test::XLLM_MM_Item& pb,
XLLM_MM_Item* out,
MmDataOwned* owned) {
std::memset(out, 0, sizeof(*out));
out->type = static_cast<XLLM_MM_Type>(pb.type());
PbToXllmMmValue(pb.data(), &out->data, owned);
PbToXllmMmState(pb.state(), &out->state);
}
void XllmMmItemToPb(const XLLM_MM_Item& in, c_api_test::XLLM_MM_Item* pb) {
pb->set_type(static_cast<uint32_t>(in.type));
XllmMmValueToPb(in.data, pb->mutable_data());
XllmMmStateToPb(in.state, pb->mutable_state());
}
bool PbToXllmMmData(const c_api_test::XLLM_MM_Data& pb,
XLLM_MM_Data* out,
MmDataOwned* owned) {
std::memset(out, 0, sizeof(*out));
out->type_mask = pb.type_mask();
out->is_dict = pb.is_dict();
switch (pb.storage_case()) {
case c_api_test::XLLM_MM_Data::kDict:
out->is_dict = true;
PbToXllmMmDict(pb.dict(), &out->data.dict, owned);
return true;
case c_api_test::XLLM_MM_Data::kItems:
out->is_dict = false;
PbToXllmMmItems(pb.items(), &out->data.items, owned);
return true;
default:
return false;
}
}
void XllmMmDataToPb(const XLLM_MM_Data& in, c_api_test::XLLM_MM_Data* pb) {
pb->set_type_mask(in.type_mask);
pb->set_is_dict(in.is_dict);
if (in.is_dict) {
XllmMmDictToPb(in.data.dict, pb->mutable_dict());
} else {
XllmMmItemsToPb(in.data.items, pb->mutable_items());
}
}
void PbToXllmRequestParams(const c_api_test::XLLM_RequestParams& pb,
XLLM_RequestParams* out) {
out->echo = pb.echo();
out->offline = pb.offline();
out->logprobs = pb.logprobs();
out->ignore_eos = pb.ignore_eos();
out->n = pb.n();
out->max_tokens = pb.max_tokens();
out->best_of = pb.best_of();
out->ttlt_slo_ms = pb.ttlt_slo_ms();
out->ttft_slo_ms = pb.ttft_slo_ms();
out->tpot_slo_ms = pb.tpot_slo_ms();
out->beam_width = pb.beam_width();
out->top_logprobs = pb.top_logprobs();
out->top_k = pb.top_k();
out->top_p = pb.top_p();
out->frequency_penalty = pb.frequency_penalty();
out->presence_penalty = pb.presence_penalty();
out->repetition_penalty = pb.repetition_penalty();
out->temperature = pb.temperature();
std::strncpy(out->request_id,
pb.request_id().c_str(),
XLLM_META_STRING_FIELD_MAX_LEN - 1);
out->request_id[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
}
void XllmRequestParamsToPb(const XLLM_RequestParams& in,
c_api_test::XLLM_RequestParams* pb) {
pb->set_echo(in.echo);
pb->set_offline(in.offline);
pb->set_logprobs(in.logprobs);
pb->set_ignore_eos(in.ignore_eos);
pb->set_n(in.n);
pb->set_max_tokens(in.max_tokens);
pb->set_best_of(in.best_of);
pb->set_ttlt_slo_ms(in.ttlt_slo_ms);
pb->set_ttft_slo_ms(in.ttft_slo_ms);
pb->set_tpot_slo_ms(in.tpot_slo_ms);
pb->set_beam_width(in.beam_width);
pb->set_top_logprobs(in.top_logprobs);
pb->set_top_k(in.top_k);
pb->set_top_p(in.top_p);
pb->set_frequency_penalty(in.frequency_penalty);
pb->set_presence_penalty(in.presence_penalty);
pb->set_repetition_penalty(in.repetition_penalty);
pb->set_temperature(in.temperature);
pb->set_request_id(in.request_id);
}
void PbToXllmChatMessage(const c_api_test::XLLM_ChatMessage& pb,
XLLM_ChatMessage* out) {
std::memset(out->role, 0, sizeof(out->role));
std::strncpy(
out->role, pb.role().c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1);
out->role[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
if (!pb.content().empty()) {
out->content = new char[pb.content().size() + 1];
std::memcpy(out->content, pb.content().data(), pb.content().size());
out->content[pb.content().size()] = '\0';
} else {
out->content = nullptr;
}
}
void FreeXllmChatMessageContent(XLLM_ChatMessage* out) {
delete[] out->content;
out->content = nullptr;
}
void XllmChatMessageToPb(const XLLM_ChatMessage* in,
c_api_test::XLLM_ChatMessage* pb) {
if (!in) {
return;
}
pb->set_role(in->role);
if (in->content != nullptr) {
pb->set_content(in->content);
} else {
pb->clear_content();
}
}
void PbToXllmUsage(const c_api_test::XLLM_Usage& pb, XLLM_Usage* out) {
out->prompt_tokens = pb.prompt_tokens();
out->completion_tokens = pb.completion_tokens();
out->total_tokens = pb.total_tokens();
}
void XllmUsageToPb(const XLLM_Usage& in, c_api_test::XLLM_Usage* pb) {
pb->set_prompt_tokens(in.prompt_tokens);
pb->set_completion_tokens(in.completion_tokens);
pb->set_total_tokens(in.total_tokens);
}
void PbToXllmLogProbs(const c_api_test::XLLM_LogProbs& pb,
XLLM_LogProbs* out,
std::vector<XLLM_LogProb>* storage) {
storage->clear();
storage->reserve(static_cast<size_t>(pb.entries_size()));
for (int i = 0; i < pb.entries_size(); ++i) {
XLLM_LogProb e{};
e.token_id = pb.entries(i).token_id();
e.logprob = pb.entries(i).logprob();
storage->push_back(e);
}
out->entries = storage->empty() ? nullptr : storage->data();
out->entries_size = storage->size();
}
void XllmLogProbsToPb(const XLLM_LogProbs& in, c_api_test::XLLM_LogProbs* pb) {
pb->clear_entries();
if (in.entries == nullptr || in.entries_size == 0) {
return;
}
for (size_t i = 0; i < in.entries_size; ++i) {
c_api_test::XLLM_LogProb* e = pb->add_entries();
e->set_token_id(in.entries[i].token_id);
e->set_logprob(in.entries[i].logprob);
}
}
void PbToXllmChoice(const c_api_test::XLLM_Choice& pb,
XLLM_Choice* out,
ResponseOwned* ro) {
std::memset(out, 0, sizeof(*out));
out->index = pb.index();
if (!pb.text().empty()) {
ro->choice_text_bufs.push_back(CopyCStr(pb.text()));
out->text = ro->choice_text_bufs.back().get();
}
if (pb.has_chat_message()) {
ro->chat_messages.emplace_back();
XLLM_ChatMessage& cm = ro->chat_messages.back();
std::memset(cm.role, 0, sizeof(cm.role));
std::strncpy(cm.role,
pb.chat_message().role().c_str(),
XLLM_META_STRING_FIELD_MAX_LEN - 1);
cm.role[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
if (!pb.chat_message().content().empty()) {
ro->chat_message_contents.push_back(
CopyCStr(pb.chat_message().content()));
cm.content = ro->chat_message_contents.back().get();
} else {
cm.content = nullptr;
}
out->message = &ro->chat_messages.back();
}
ro->token_ids_vecs.emplace_back();
ro->token_ids_vecs.back().reserve(static_cast<size_t>(pb.token_ids_size()));
for (int i = 0; i < pb.token_ids_size(); ++i) {
ro->token_ids_vecs.back().push_back(pb.token_ids(i));
}
out->token_ids = ro->token_ids_vecs.back().data();
out->token_size = ro->token_ids_vecs.back().size();
ro->logprob_vecs.emplace_back();
std::vector<XLLM_LogProb>& le = ro->logprob_vecs.back();
le.reserve(static_cast<size_t>(pb.logprobs().entries_size()));
for (int i = 0; i < pb.logprobs().entries_size(); ++i) {
XLLM_LogProb e{};
e.token_id = pb.logprobs().entries(i).token_id();
e.logprob = pb.logprobs().entries(i).logprob();
le.push_back(e);
}
out->logprobs.entries = le.data();
out->logprobs.entries_size = le.size();
std::strncpy(out->finish_reason,
pb.finish_reason().c_str(),
XLLM_META_STRING_FIELD_MAX_LEN - 1);
out->finish_reason[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
}
void XllmChoiceToPb(const XLLM_Choice& in, c_api_test::XLLM_Choice* pb) {
pb->set_index(in.index);
if (in.text != nullptr) {
pb->set_text(in.text);
} else {
pb->clear_text();
}
if (in.message != nullptr) {
XllmChatMessageToPb(in.message, pb->mutable_chat_message());
} else {
pb->clear_chat_message();
}
pb->clear_token_ids();
if (in.token_ids != nullptr) {
for (size_t i = 0; i < in.token_size; ++i) {
pb->add_token_ids(in.token_ids[i]);
}
}
XllmLogProbsToPb(in.logprobs, pb->mutable_logprobs());
pb->set_finish_reason(in.finish_reason);
}
void PbToXllmChoices(const c_api_test::XLLM_Choices& pb,
XLLM_Choices* out,
ResponseOwned* ro) {
ro->choices.clear();
ro->choices.reserve(static_cast<size_t>(pb.entries_size()));
for (int i = 0; i < pb.entries_size(); ++i) {
ro->choices.emplace_back();
PbToXllmChoice(pb.entries(i), &ro->choices.back(), ro);
}
out->entries = ro->choices.data();
out->entries_size = ro->choices.size();
}
void XllmChoicesToPb(const XLLM_Choices& in, c_api_test::XLLM_Choices* pb) {
pb->clear_entries();
if (in.entries == nullptr) {
return;
}
for (size_t i = 0; i < in.entries_size; ++i) {
XllmChoiceToPb(in.entries[i], pb->add_entries());
}
}
void PbToXllmResponse(const c_api_test::XLLM_Response& pb,
XLLM_Response* out,
ResponseOwned* owned) {
std::memset(out, 0, sizeof(*out));
out->status_code = static_cast<XLLM_StatusCode>(pb.status_code());
std::strncpy(
out->error_info, pb.error_info().c_str(), XLLM_ERROR_INFO_MAX_LEN - 1);
out->error_info[XLLM_ERROR_INFO_MAX_LEN - 1] = '\0';
std::strncpy(out->id, pb.id().c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1);
out->id[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
std::strncpy(
out->object, pb.object().c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1);
out->object[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
out->created = pb.created();
std::strncpy(
out->model, pb.model().c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1);
out->model[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
PbToXllmUsage(pb.usage(), &out->usage);
PbToXllmChoices(pb.choices(), &out->choices, owned);
}
void XllmResponseToPb(const XLLM_Response* in, c_api_test::XLLM_Response* pb) {
if (!in) {
pb->Clear();
return;
}
pb->set_status_code(
static_cast<c_api_test::XLLM_StatusCode>(in->status_code));
pb->set_error_info(in->error_info);
pb->set_id(in->id);
pb->set_object(in->object);
pb->set_created(in->created);
pb->set_model(in->model);
XllmUsageToPb(in->usage, pb->mutable_usage());
XllmChoicesToPb(in->choices, pb->mutable_choices());
}
} // namespace xllm_capi_test

View File

@@ -0,0 +1,146 @@
/* Copyright 2026 The xLLM 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.
==============================================================================*/
#ifndef XLLM_C_API_TEST_UTILS_H_
#define XLLM_C_API_TEST_UTILS_H_
#include <gflags/gflags.h>
#include <cstddef>
#include <memory>
#include <vector>
#include "types.h"
#include "xllm_test.pb.h"
DECLARE_string(model_path);
DECLARE_string(devices);
DECLARE_int32(port);
DECLARE_string(listen_addr);
DECLARE_int32(idle_timeout_s);
DECLARE_string(backend);
namespace xllm_capi_test {
// Applies gflags (after ParseCommandLineFlags) into XLLM_InitOptions.
void ApplyGflagsToXllmInitOptions(XLLM_InitOptions* opt);
// Owns buffers referenced by XLLM_MM_Data after PbToXllmMmData.
struct MmDataOwned {
std::vector<std::vector<uint8_t>> tensor_byte_buffers;
std::vector<std::vector<XLLM_Tensor>> tensor_lists;
std::vector<XLLM_MM_Item> mm_items;
std::vector<XLLM_MM_DictEntry> mm_dict_entries;
};
// Owns heap data for XLLM_Response filled by PbToXllmResponse (text, message,
// token_ids, logprobs arrays).
struct ResponseOwned {
std::vector<std::unique_ptr<char[]>> choice_text_bufs;
std::vector<XLLM_ChatMessage> chat_messages;
std::vector<std::unique_ptr<char[]>> chat_message_contents;
std::vector<std::vector<int32_t>> token_ids_vecs;
std::vector<std::vector<XLLM_LogProb>> logprob_vecs;
std::vector<XLLM_Choice> choices;
};
void PbToXllmRequestParams(const c_api_test::XLLM_RequestParams& pb,
XLLM_RequestParams* out);
void XllmRequestParamsToPb(const XLLM_RequestParams& in,
c_api_test::XLLM_RequestParams* pb);
void PbToXllmChatMessage(const c_api_test::XLLM_ChatMessage& pb,
XLLM_ChatMessage* out);
void FreeXllmChatMessageContent(XLLM_ChatMessage* out);
void XllmChatMessageToPb(const XLLM_ChatMessage* in,
c_api_test::XLLM_ChatMessage* pb);
bool PbToXllmMmData(const c_api_test::XLLM_MM_Data& pb,
XLLM_MM_Data* out,
MmDataOwned* owned);
void XllmMmDataToPb(const XLLM_MM_Data& in, c_api_test::XLLM_MM_Data* pb);
void XllmResponseToPb(const XLLM_Response* in, c_api_test::XLLM_Response* pb);
void PbToXllmResponse(const c_api_test::XLLM_Response& pb,
XLLM_Response* out,
ResponseOwned* owned);
// --- Lower-level (types.h <-> pb) ---
void PbToXllmDims(const c_api_test::XLLM_Dims& pb, XLLM_Dims* out);
void XllmDimsToPb(const XLLM_Dims& in, c_api_test::XLLM_Dims* pb);
void PbToXllmTensor(const c_api_test::XLLM_Tensor& pb,
XLLM_Tensor* out,
MmDataOwned* owned);
void XllmTensorToPb(const XLLM_Tensor& in, c_api_test::XLLM_Tensor* pb);
void PbToXllmTensors(const c_api_test::XLLM_Tensors& pb,
XLLM_Tensors* out,
MmDataOwned* owned);
void XllmTensorsToPb(const XLLM_Tensors& in, c_api_test::XLLM_Tensors* pb);
void PbToXllmMmValue(const c_api_test::XLLM_MM_Value& pb,
XLLM_MM_Value* out,
MmDataOwned* owned);
void XllmMmValueToPb(const XLLM_MM_Value& in, c_api_test::XLLM_MM_Value* pb);
void PbToXllmMmDict(const c_api_test::XLLM_MM_Dict& pb,
XLLM_MM_Dict* out,
MmDataOwned* owned);
void XllmMmDictToPb(const XLLM_MM_Dict& in, c_api_test::XLLM_MM_Dict* pb);
void PbToXllmMmItems(const c_api_test::XLLM_MM_Items& pb,
XLLM_MM_Items* out,
MmDataOwned* owned);
void XllmMmItemsToPb(const XLLM_MM_Items& in, c_api_test::XLLM_MM_Items* pb);
void PbToXllmMmState(const c_api_test::XLLM_MM_State& pb, XLLM_MM_State* out);
void XllmMmStateToPb(const XLLM_MM_State& in, c_api_test::XLLM_MM_State* pb);
void PbToXllmMmItem(const c_api_test::XLLM_MM_Item& pb,
XLLM_MM_Item* out,
MmDataOwned* owned);
void XllmMmItemToPb(const XLLM_MM_Item& in, c_api_test::XLLM_MM_Item* pb);
void PbToXllmUsage(const c_api_test::XLLM_Usage& pb, XLLM_Usage* out);
void XllmUsageToPb(const XLLM_Usage& in, c_api_test::XLLM_Usage* pb);
void PbToXllmLogProbs(const c_api_test::XLLM_LogProbs& pb,
XLLM_LogProbs* out,
std::vector<XLLM_LogProb>* storage);
void XllmLogProbsToPb(const XLLM_LogProbs& in, c_api_test::XLLM_LogProbs* pb);
void PbToXllmChoice(const c_api_test::XLLM_Choice& pb,
XLLM_Choice* out,
ResponseOwned* owned);
void XllmChoiceToPb(const XLLM_Choice& in, c_api_test::XLLM_Choice* pb);
void PbToXllmChoices(const c_api_test::XLLM_Choices& pb,
XLLM_Choices* out,
ResponseOwned* owned);
void XllmChoicesToPb(const XLLM_Choices& in, c_api_test::XLLM_Choices* pb);
} // namespace xllm_capi_test
#endif // XLLM_C_API_TEST_UTILS_H_

View File

@@ -0,0 +1,320 @@
/* Copyright 2026 The xLLM 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 <brpc/controller.h>
#include <brpc/server.h>
#include <butil/endpoint.h>
#include <butil/logging.h>
#include <gflags/gflags.h>
#include <cctype>
#include <cstddef>
#include <cstdint>
#include <cstring>
#include <memory>
#include <mutex>
#include <string>
#include <vector>
#include "llm.h"
#include "rec.h"
#include "utils.h"
#include "xllm_test.pb.h"
namespace xllm_capi_test {
namespace {
std::unique_ptr<char[]> CopyContent(const std::string& s) {
if (s.empty()) {
return nullptr;
}
auto p = std::make_unique<char[]>(s.size() + 1);
std::memcpy(p.get(), s.data(), s.size());
p[s.size()] = '\0';
return p;
}
void SetErrorResponse(c_api_test::XLLM_Response* res,
c_api_test::XLLM_StatusCode code,
const std::string& msg) {
res->Clear();
res->set_status_code(code);
res->set_error_info(msg);
}
} // namespace
class XllmRecCapiServiceImpl : public c_api_test::XllmRecCapiService {
public:
XllmRecCapiServiceImpl(XLLM_REC_Handler* rec_handler,
XLLM_LLM_Handler* llm_handler)
: rec_handler_(rec_handler), llm_handler_(llm_handler) {}
void Inference(google::protobuf::RpcController* cntl_base,
const c_api_test::XLLM_Request* request,
c_api_test::XLLM_Response* response,
google::protobuf::Closure* done) override {
brpc::ClosureGuard done_guard(done);
(void)cntl_base;
std::lock_guard<std::mutex> lock(mu_);
const std::string& fn = request->call_function();
const bool is_llm_fn =
(fn == "xllm_llm_completions" || fn == "xllm_llm_chat_completions");
const bool is_rec_fn = (fn == "xllm_rec_text_completions" ||
fn == "xllm_rec_token_completions" ||
fn == "xllm_rec_multimodal_completions" ||
fn == "xllm_rec_chat_completions");
if (is_llm_fn && !llm_handler_) {
SetErrorResponse(response,
c_api_test::XLLM_STATUS_INTERNAL_ERROR,
"LLM handler is null");
return;
}
if (is_rec_fn && !rec_handler_) {
SetErrorResponse(response,
c_api_test::XLLM_STATUS_INTERNAL_ERROR,
"REC handler is null");
return;
}
XLLM_RequestParams params{};
if (is_llm_fn) {
xllm_llm_request_params_default(&params);
} else {
xllm_rec_request_params_default(&params);
}
if (request->params().ByteSizeLong() > 0) {
PbToXllmRequestParams(request->params(), &params);
}
const char* model_id =
request->model_id().empty() ? "" : request->model_id().c_str();
const uint32_t timeout_ms = request->timeout_ms();
XLLM_Response* raw = nullptr;
if (fn == "xllm_llm_completions") {
raw = xllm_llm_completions(llm_handler_,
model_id,
request->prompt().c_str(),
timeout_ms,
&params);
} else if (fn == "xllm_llm_chat_completions") {
std::vector<XLLM_ChatMessage> cms;
std::vector<std::unique_ptr<char[]>> contents;
cms.reserve(request->messages_size());
for (int i = 0; i < request->messages_size(); ++i) {
const auto& m = request->messages(i);
XLLM_ChatMessage cm{};
std::memset(cm.role, 0, sizeof(cm.role));
std::strncpy(
cm.role, m.role().c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1);
cm.role[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
contents.push_back(CopyContent(m.content()));
cm.content = contents.back() ? contents.back().get() : nullptr;
cms.push_back(cm);
}
raw = xllm_llm_chat_completions(llm_handler_,
model_id,
cms.empty() ? nullptr : cms.data(),
cms.size(),
timeout_ms,
&params);
} else if (fn == "xllm_rec_text_completions") {
raw = xllm_rec_text_completions(rec_handler_,
model_id,
request->prompt().c_str(),
timeout_ms,
&params);
} else if (fn == "xllm_rec_token_completions") {
std::vector<int32_t> token_ids;
token_ids.reserve(request->token_ids_size());
for (int i = 0; i < request->token_ids_size(); ++i) {
token_ids.push_back(request->token_ids(i));
}
raw = xllm_rec_token_completions(
rec_handler_,
model_id,
token_ids.empty() ? nullptr : token_ids.data(),
token_ids.size(),
timeout_ms,
&params);
} else if (fn == "xllm_rec_multimodal_completions") {
std::vector<int32_t> token_ids;
token_ids.reserve(request->token_ids_size());
for (int i = 0; i < request->token_ids_size(); ++i) {
token_ids.push_back(request->token_ids(i));
}
XLLM_MM_Data mm{};
MmDataOwned mm_owned;
if (!PbToXllmMmData(request->mm_data(), &mm, &mm_owned)) {
SetErrorResponse(response,
c_api_test::XLLM_STATUS_INVALID_REQUEST,
"invalid or empty mm_data");
return;
}
raw = xllm_rec_multimodal_completions(
rec_handler_,
model_id,
token_ids.empty() ? nullptr : token_ids.data(),
token_ids.size(),
&mm,
timeout_ms,
&params);
} else if (fn == "xllm_rec_chat_completions") {
std::vector<XLLM_ChatMessage> cms;
std::vector<std::unique_ptr<char[]>> contents;
cms.reserve(request->messages_size());
for (int i = 0; i < request->messages_size(); ++i) {
const auto& m = request->messages(i);
XLLM_ChatMessage cm{};
std::memset(cm.role, 0, sizeof(cm.role));
std::strncpy(
cm.role, m.role().c_str(), XLLM_META_STRING_FIELD_MAX_LEN - 1);
cm.role[XLLM_META_STRING_FIELD_MAX_LEN - 1] = '\0';
contents.push_back(CopyContent(m.content()));
cm.content = contents.back() ? contents.back().get() : nullptr;
cms.push_back(cm);
}
raw = xllm_rec_chat_completions(rec_handler_,
model_id,
cms.empty() ? nullptr : cms.data(),
cms.size(),
timeout_ms,
&params);
} else {
SetErrorResponse(response,
c_api_test::XLLM_STATUS_INVALID_REQUEST,
"unsupported call_function: " + fn);
return;
}
if (raw == nullptr) {
SetErrorResponse(response,
c_api_test::XLLM_STATUS_INTERNAL_ERROR,
"C API returned null response");
return;
}
XllmResponseToPb(raw, response);
if (is_llm_fn) {
xllm_llm_free_response(raw);
} else {
xllm_rec_free_response(raw);
}
}
private:
XLLM_REC_Handler* rec_handler_;
XLLM_LLM_Handler* llm_handler_;
std::mutex mu_;
};
} // namespace xllm_capi_test
int main(int argc, char* argv[]) {
GFLAGS_NAMESPACE::ParseCommandLineFlags(&argc, &argv, true);
if (FLAGS_model_path.empty()) {
LOG(ERROR) << "Missing --model_path (set in gflags file or command line)";
return -1;
}
std::string backend = FLAGS_backend;
for (char& c : backend) {
c = static_cast<char>(std::tolower(static_cast<unsigned char>(c)));
}
if (backend != "llm" && backend != "rec") {
LOG(ERROR) << "Invalid --backend=\"" << FLAGS_backend
<< "\" (expected llm or rec)";
return -2;
}
XLLM_REC_Handler* rec_ptr = nullptr;
XLLM_LLM_Handler* llm_ptr = nullptr;
std::unique_ptr<XLLM_REC_Handler, void (*)(XLLM_REC_Handler*)> rec_holder(
nullptr, xllm_rec_destroy);
std::unique_ptr<XLLM_LLM_Handler, void (*)(XLLM_LLM_Handler*)> llm_holder(
nullptr, xllm_llm_destroy);
if (backend == "rec") {
rec_holder.reset(xllm_rec_create());
if (!rec_holder) {
LOG(ERROR) << "xllm_rec_create failed";
return -3;
}
XLLM_InitOptions init{};
xllm_rec_init_options_default(&init);
xllm_capi_test::ApplyGflagsToXllmInitOptions(&init);
if (!xllm_rec_initialize(rec_holder.get(),
FLAGS_model_path.c_str(),
FLAGS_devices.c_str(),
&init)) {
LOG(ERROR) << "xllm_rec_initialize failed model_path=" << FLAGS_model_path
<< " devices=" << FLAGS_devices;
return -4;
}
rec_ptr = rec_holder.get();
} else {
llm_holder.reset(xllm_llm_create());
if (!llm_holder) {
LOG(ERROR) << "xllm_llm_create failed";
return -5;
}
XLLM_InitOptions init{};
xllm_llm_init_options_default(&init);
xllm_capi_test::ApplyGflagsToXllmInitOptions(&init);
if (!xllm_llm_initialize(llm_holder.get(),
FLAGS_model_path.c_str(),
FLAGS_devices.c_str(),
&init)) {
LOG(ERROR) << "xllm_llm_initialize failed model_path=" << FLAGS_model_path
<< " devices=" << FLAGS_devices;
return -6;
}
llm_ptr = llm_holder.get();
}
xllm_capi_test::XllmRecCapiServiceImpl svc(rec_ptr, llm_ptr);
brpc::Server server;
if (server.AddService(&svc, brpc::SERVER_DOESNT_OWN_SERVICE) != 0) {
LOG(ERROR) << "Fail to add XllmRecCapiService";
return -7;
}
butil::EndPoint point;
if (!FLAGS_listen_addr.empty()) {
if (butil::str2endpoint(FLAGS_listen_addr.c_str(), &point) < 0) {
LOG(ERROR) << "Invalid --listen_addr=" << FLAGS_listen_addr;
return -8;
}
} else {
point = butil::EndPoint(butil::IP_ANY, FLAGS_port);
}
brpc::ServerOptions options;
options.idle_timeout_sec = FLAGS_idle_timeout_s;
if (server.Start(point, &options) != 0) {
LOG(ERROR) << "Fail to start brpc server";
return -9;
}
LOG(INFO) << "xllm_test C API brpc server backend=" << backend
<< " listening on " << butil::endpoint2str(point).c_str();
server.RunUntilAskedToQuit();
return 0;
}

View File

@@ -0,0 +1,24 @@
# Example gflags for xllm_test brpc C API server.
# Start: xllm_test --flagfile=xllm/c_api/test/xllm_test.flags
#
# --- Backend: only one of llm (c_api/llm.h) or rec (c_api/rec.h) is loaded ---
--backend=rec
# --backend=llm
#
# --- Required for xllm_*_initialize ---
--model_path=/export/home/models/Qwen3-8B
--devices=npu:4
#
# --- brpc listen (optional) ---
--port=8000
# --listen_addr=0.0.0.0:8000
# --idle_timeout_s=-1
#
# --- XLLM_InitOptions (override as needed; defaults match REC) ---
# --enable_chunked_prefill=false
# --transfer_listen_port=26000
# --max_memory_utilization=0.55
# --init_task=generate
# --communication_backend=lccl
# --master_node_addr=127.0.0.1:18899
# --xllm_init_log_dir= # maps to XLLM_InitOptions.log_dir (not glog's --log_dir)

View File

@@ -0,0 +1,224 @@
syntax = "proto3";
package c_api_test;
option cc_generic_services = true;
// =============================================================================
// Mirrors xllm/c_api/types.h (all public structs/enums) except:
// - XLLM_InitOptions / XLLM_InitLLMOptions are intentionally omitted.
// =============================================================================
// --- XLLM_DataType ---
enum XLLM_DataType {
XLLM_DTYPE_UNDEFINED = 0;
XLLM_DTYPE_FLOAT16 = 1;
XLLM_DTYPE_FLOAT32 = 2;
XLLM_DTYPE_FLOAT64 = 3;
XLLM_DTYPE_BFLOAT16 = 4;
XLLM_DTYPE_INT8 = 5;
XLLM_DTYPE_INT16 = 6;
XLLM_DTYPE_INT32 = 7;
XLLM_DTYPE_INT64 = 8;
XLLM_DTYPE_UINT8 = 9;
XLLM_DTYPE_UINT16 = 10;
XLLM_DTYPE_UINT32 = 11;
XLLM_DTYPE_UINT64 = 12;
XLLM_DTYPE_BOOL = 13;
XLLM_DTYPE_STRING = 14;
}
// --- XLLM_StatusCode ---
enum XLLM_StatusCode {
XLLM_STATUS_SUCCESS = 0;
XLLM_STATUS_NOT_INITIALIZED = 1;
XLLM_STATUS_MODEL_NOT_FOUND = 2;
XLLM_STATUS_TIMEOUT = 3;
XLLM_STATUS_INVALID_REQUEST = 4;
XLLM_STATUS_INTERNAL_ERROR = 5;
}
// --- XLLM_MM_Type (same numeric values as C bitmask enum) ---
enum XLLM_MM_Type {
XLLM_MM_TYPE_NONE = 0;
XLLM_MM_TYPE_IMAGE = 1;
XLLM_MM_TYPE_AUDIO = 2;
XLLM_MM_TYPE_VIDEO = 4;
XLLM_MM_TYPE_TEXT = 8;
XLLM_MM_TYPE_EMBEDDING = 16;
}
// --- XLLM_Dims ---
message XLLM_Dims {
int32 rank = 1;
repeated int32 dim = 2;
}
// --- XLLM_Tensor ---
message XLLM_Tensor {
XLLM_DataType dtype = 1;
XLLM_Dims dims = 2;
bytes data = 3;
}
// --- XLLM_Tensors ---
message XLLM_Tensors {
repeated XLLM_Tensor entries = 1;
}
// --- XLLM_MM_Value ---
message XLLM_MM_Value {
bool is_single_tensor = 1;
oneof data {
XLLM_Tensor tensor = 2;
XLLM_Tensors tensors = 3;
}
}
// --- XLLM_MM_Meta (placeholder, matches empty struct in types.h) ---
message XLLM_MM_Meta {}
// --- XLLM_MM_TokenPos ---
message XLLM_MM_TokenPos {
uint32 offset = 1;
uint32 length = 2;
}
// --- XLLM_MM_State ---
message XLLM_MM_State {
XLLM_MM_TokenPos token_pos = 1;
}
// --- XLLM_MM_DictEntry ---
message XLLM_MM_DictEntry {
string key = 1;
XLLM_MM_Value value = 2;
}
// --- XLLM_MM_Dict ---
message XLLM_MM_Dict {
repeated XLLM_MM_DictEntry entries = 1;
}
// --- XLLM_MM_Item ---
message XLLM_MM_Item {
uint32 type = 1;
XLLM_MM_Value data = 2;
XLLM_MM_Meta meta = 3;
XLLM_MM_State state = 4;
}
// --- XLLM_MM_Items ---
message XLLM_MM_Items {
repeated XLLM_MM_Item entries = 1;
}
// --- XLLM_MM_Data ---
message XLLM_MM_Data {
uint32 type_mask = 1;
bool is_dict = 2;
oneof storage {
XLLM_MM_Dict dict = 3;
XLLM_MM_Items items = 4;
}
}
// --- XLLM_ChatMessage ---
message XLLM_ChatMessage {
string role = 1;
string content = 2;
}
// --- XLLM_RequestParams ---
message XLLM_RequestParams {
bool echo = 1;
bool offline = 2;
bool logprobs = 3;
bool ignore_eos = 4;
uint32 n = 5;
uint32 max_tokens = 6;
uint32 best_of = 7;
int32 ttlt_slo_ms = 8;
int32 ttft_slo_ms = 9;
int32 tpot_slo_ms = 10;
uint32 beam_width = 11;
int64 top_logprobs = 12;
int64 top_k = 13;
float top_p = 14;
float frequency_penalty = 15;
float presence_penalty = 16;
float repetition_penalty = 17;
float temperature = 18;
string request_id = 19;
}
// --- XLLM_Usage ---
message XLLM_Usage {
int32 prompt_tokens = 1;
int32 completion_tokens = 2;
int32 total_tokens = 3;
}
// --- XLLM_LogProb / XLLM_LogProbs ---
message XLLM_LogProb {
uint32 token_id = 1;
float logprob = 2;
}
message XLLM_LogProbs {
repeated XLLM_LogProb entries = 1;
}
// --- XLLM_Choice / XLLM_Choices ---
message XLLM_Choice {
uint32 index = 1;
string text = 2;
XLLM_ChatMessage chat_message = 3;
repeated int32 token_ids = 4;
XLLM_LogProbs logprobs = 5;
string finish_reason = 6;
}
message XLLM_Choices {
repeated XLLM_Choice entries = 1;
}
// --- XLLM_Response ---
message XLLM_Response {
XLLM_StatusCode status_code = 1;
string error_info = 2;
string id = 3;
string object = 4;
int64 created = 5;
string model = 6;
XLLM_Choices choices = 7;
XLLM_Usage usage = 8;
}
message XLLM_Request {
string call_function = 1;
uint32 timeout_ms = 2;
string prompt = 3;
repeated XLLM_ChatMessage messages = 4;
repeated int32 token_ids = 5;
XLLM_MM_Data mm_data = 6;
XLLM_RequestParams params = 7;
string model_id = 8;
}
// One dump file record: request + optional response (xllm_dump).
message XLLM_DumpRecord {
XLLM_Request request = 1;
XLLM_Response response = 2;
}
// brpc service: dispatch by XLLM_Request.call_function (must match xllm_test
// --backend: rec -> xllm_rec_*; llm -> xllm_llm_*).
service XllmRecCapiService {
rpc Inference(XLLM_Request) returns (XLLM_Response);
}

View File

@@ -0,0 +1,118 @@
#!/bin/bash
set -euo pipefail
SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)
BIN_DIR="${SCRIPT_DIR}/../../../bin"
TEMP_DIR="xllm"
INCLUDE_DIR="${TEMP_DIR}/include"
LIB_DIR="${TEMP_DIR}/lib"
VERSION_FILE="${SCRIPT_DIR}/../../../version.txt"
TAR_BASE_NAME="xllm"
LOCAL_INSTALL_DIR="/usr/local"
LOCAL_TARGET_DIR="${LOCAL_INSTALL_DIR}/xllm"
HEADERS=("${SCRIPT_DIR}/../llm.h" "${SCRIPT_DIR}/../rec.h" "${SCRIPT_DIR}/../default.h" "${SCRIPT_DIR}/../types.h")
SO_FILES=(
"${SCRIPT_DIR}/../../../build/xllm/core/server/libxllm.so"
)
error_exit() {
echo -e "\033[31merror: $1\033[0m" >&2
exit 1
}
cd_bin_dir() {
if [ ! -d "${BIN_DIR}" ]; then
mkdir -p "${BIN_DIR}" || error_exit "failed to create bin directory: ${BIN_DIR}"
fi
cd "${BIN_DIR}" || error_exit "failed to enter bin directory: ${BIN_DIR}"
}
read_version() {
if [ ! -f "${VERSION_FILE}" ]; then
error_exit "${VERSION_FILE} is not existed"
fi
VERSION=$(cat "${VERSION_FILE}" | tr -d '[:space:]')
if [ -z "${VERSION}" ]; then
error_exit "version content is empty"
fi
TAR_FILE="${TAR_BASE_NAME}_${VERSION}.tar.gz"
}
check_files() {
for header in "${HEADERS[@]}"; do
if [ ! -f "${header}" ]; then
error_exit "${header} is not existed"
fi
done
for so_file in "${SO_FILES[@]}"; do
if [ ! -f "${so_file}" ]; then
error_exit "${so_file} is not existed"
fi
done
}
create_dirs() {
mkdir -p "${INCLUDE_DIR}" || error_exit "create include directory failed"
mkdir -p "${LIB_DIR}" || error_exit "create lib directory failed"
}
copy_headers() {
for header in "${HEADERS[@]}"; do
cp -f "${header}" "${INCLUDE_DIR}/" || error_exit "copy ${header} failed"
done
}
copy_so() {
for so_file in "${SO_FILES[@]}"; do
cp -f "${so_file}" "${LIB_DIR}/" || error_exit "copy ${so_file} failed"
done
}
package_tar() {
tar -czf "${TAR_FILE}" "${TEMP_DIR}" || error_exit "tar failed"
}
cleanup_temp() {
rm -rf "${TEMP_DIR}" || error_exit "rm temp directory failed"
}
extract_to_local() {
if [ ! -f "${TAR_FILE}" ]; then
error_exit "${TAR_FILE} is not existed"
fi
if [ ! -d "${LOCAL_INSTALL_DIR}" ]; then
error_exit "local install directory is not existed"
fi
if [ -d "${LOCAL_TARGET_DIR}" ]; then
rm -rf "${LOCAL_TARGET_DIR}" || error_exit "rm old xllm directory failed"
fi
tar -xzf "${TAR_FILE}" -C "${LOCAL_INSTALL_DIR}" || error_exit "extract failed"
}
main() {
cd_bin_dir
read_version
check_files
create_dirs
copy_headers
copy_so
package_tar
cleanup_temp
extract_to_local
echo -e "install file: \033[33m${TAR_FILE}\033[0m"
echo -e "install path: \033[33m/usr/local/${TEMP_DIR}\033[0m"
}
main

View File

@@ -0,0 +1,627 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#ifndef XLLM_C_TYPES_H
#define XLLM_C_TYPES_H
#ifdef __cplusplus
extern "C" {
#endif
#include <stdbool.h>
#include <stddef.h>
#include <stdint.h>
// Maximum length for meta string fields (includes '\0' terminator)
#define XLLM_META_STRING_FIELD_MAX_LEN 128
// Export Macro Definition
#ifndef XLLM_CAPI_EXPORT
#define XLLM_CAPI_EXPORT __attribute__((visibility("default")))
#endif
// Core Struct & Enum Definitions
/**
* @brief Configuration options for initializing an LLM instance
* @note All string fields are fixed-length arrays. Default values are defined
* in macros. Empty string indicates disable/use default value.
*/
typedef struct XLLM_CAPI_EXPORT XLLM_InitOptions {
/** Whether to enable chunked prefill for inference */
bool enable_chunked_prefill;
/** Whether to enable prefill-only sequence parallel */
bool enable_prefill_sp;
/** Whether to enable prefix cache optimization */
bool enable_prefix_cache;
/** Whether to enable disaggregated prefill and decode execution */
bool enable_disagg_pd;
/** Whether to enable online-offline co-location in disaggregated PD mode */
bool enable_pd_ooc;
/** Whether to enable schedule overlap for parallel execution */
bool enable_schedule_overlap;
/** Whether to enable shared memory for model execution */
bool enable_shm;
/** Whether to enable graph execution for REC */
bool enable_graph;
/** Whether to enable REC fast sampler */
bool enable_rec_fast_sampler;
/** Whether to enable prefill piecewise graph for REC */
bool enable_prefill_piecewise_graph;
/** Whether to enable xattention one-stage execution for REC */
bool enable_xattention_one_stage;
/** Whether to enable graph-mode decode without padding for REC */
bool enable_graph_mode_decode_no_padding;
/** Whether to enable block copy kernel */
bool enable_block_copy_kernel;
/** Whether to keep REC top-k outputs sorted */
bool enable_topk_sorted;
/** Whether to enable rec prefill only */
bool enable_rec_prefill_only;
/** KVCache transfer listen port */
uint32_t transfer_listen_port;
/** Number of multi-nodes in distributed deployment */
uint32_t nnodes;
/** Node rank in distributed deployment */
uint32_t node_rank;
/** Data parallel size for MLA attention */
uint32_t dp_size;
/** Expert parallel size for MoE model */
uint32_t ep_size;
/** Number of slots per kv cache block */
uint32_t block_size;
/** Max GPU memory size for kv cache (0 = auto-calculate available memory) */
uint32_t max_cache_size;
/** Max number of tokens per batch */
uint32_t max_tokens_per_batch;
/** Max number of sequences per batch */
uint32_t max_seqs_per_batch;
/** Max number of token per chunk in prefill stage */
uint32_t max_tokens_per_chunk_for_prefill;
/** Number of speculative tokens for speculative decoding */
uint32_t num_speculative_tokens;
/** Number of threads for handling input requests */
uint32_t num_request_handling_threads;
/** Expert parallel degree for MoE model */
uint32_t expert_parallel_degree;
/** Index ID for internal server ID (unique for multiple models/versions) */
uint32_t server_idx;
/** Beam width for beam search decoding (1 for greedy search) */
uint32_t beam_width;
/** Maximum number of decode rounds for each inference request */
uint32_t max_decode_rounds;
/** Maximum number of tokens allowed per inference request */
uint32_t max_token_per_req;
/** Maximum GPU memory utilization ratio for model inference */
float max_memory_utilization;
/** Maximum REC worker pipeline concurrency */
uint32_t rec_worker_max_concurrency;
/** Model task type (generate/embed) */
char task[XLLM_META_STRING_FIELD_MAX_LEN];
/** NPU communication backend (lccl/hccl). Use hccl when dp is enabled */
char communication_backend[XLLM_META_STRING_FIELD_MAX_LEN];
/** Instance role (DEFAULT/PREFILL/DECODE/MIX) */
char instance_role[XLLM_META_STRING_FIELD_MAX_LEN];
/** Device IP address for NPU communication */
char device_ip[XLLM_META_STRING_FIELD_MAX_LEN];
/** Master address for multi-node distributed serving (e.g. 10.18.1.1:9999) */
char master_node_addr[XLLM_META_STRING_FIELD_MAX_LEN];
/** XService server address (empty string = disable XService) */
char xservice_addr[XLLM_META_STRING_FIELD_MAX_LEN];
/** Unique instance name for identification */
char instance_name[XLLM_META_STRING_FIELD_MAX_LEN];
/** KV cache transfer mode (PUSH/PULL) */
char kv_cache_transfer_mode[XLLM_META_STRING_FIELD_MAX_LEN];
/** Log directory path (empty string = disable logging) */
char log_dir[XLLM_META_STRING_FIELD_MAX_LEN];
/** Draft hf model path (empty string = no draft model) */
char draft_model[XLLM_META_STRING_FIELD_MAX_LEN];
/**
* Devices to run the draft model on (e.g. npu:0, npu:0,npu:1).
* Empty string = use the same devices as main model
*/
char draft_devices[XLLM_META_STRING_FIELD_MAX_LEN];
} XLLM_InitLLMOptions;
/**
* @brief Chat message structure (for ChatCompletions)
*/
typedef struct XLLM_CAPI_EXPORT XLLM_ChatMessage {
/** Message role (system/user/assistant) */
char role[XLLM_META_STRING_FIELD_MAX_LEN];
/** Message content (NULL for function call messages) */
char* content;
} XLLM_ChatMessage;
/**
* @brief Inference request parameters
* @note All numeric fields are fixed-width integers with value ranges defined
* in macros;
*/
typedef struct XLLM_CAPI_EXPORT XLLM_RequestParams {
/** Whether to include original prompt in response */
bool echo;
/** Whether it's an offline request */
bool offline;
/** Whether to return token log probabilities */
bool logprobs;
/** Whether to ignore EOS token */
bool ignore_eos;
/** Number of completions to return per prompt */
uint32_t n;
/** Maximum number of tokens to generate. Must be <= model context length */
uint32_t max_tokens;
/** Number of sequences to generate per prompt for top-n selection */
uint32_t best_of;
/** SLO timeout in milliseconds (0 = unlimited) */
int32_t ttlt_slo_ms;
int32_t ttft_slo_ms;
int32_t tpot_slo_ms;
/** Beam search width (0 = disable beam search) */
uint32_t beam_width;
/** Final number of beam search results to return (0 = use beam_width) */
uint32_t num_return_sequences;
/** Number of top log probabilities to return */
int64_t top_logprobs;
/** Top-K sampling cutoff (-1 = 0xFFFFFFFF means disabled) */
int64_t top_k;
/** Top-P sampling cutoff (range: [0.0, 1.0]) */
float top_p;
/** Frequency penalty (range: [0.0, 2.0]) */
float frequency_penalty;
/** Presence penalty (range: [-2.0, 2.0]) */
float presence_penalty;
/** Repetition penalty. >1.0 encourages new tokens, <1.0 encourages repetition
*/
float repetition_penalty;
/** Sampling temperature (range: [0.0, 2.0]) */
float temperature;
/** Request id */
char request_id[XLLM_META_STRING_FIELD_MAX_LEN];
} XLLM_RequestParams;
/**
* @brief API response status codes
*/
typedef enum XLLM_CAPI_EXPORT XLLM_StatusCode {
/** Request succeeded */
kSuccess = 0,
/** LLM instance not initialized */
kNotInitialized = 1,
/** Specified model ID not loaded */
kModelNotFound = 2,
/** Request timed out */
kTimeout = 3,
/** Invalid input parameters */
kInvalidRequest = 4,
/** Internal system error */
kInternalError = 5
} XLLM_StatusCode;
/**
* @brief Token usage statistics for inference request
*/
typedef struct XLLM_CAPI_EXPORT XLLM_Usage {
/** Number of tokens in the prompt */
int32_t prompt_tokens;
/** Number of tokens in the generated completion */
int32_t completion_tokens;
/** Total tokens used (prompt + completion) */
int32_t total_tokens;
} XLLM_Usage;
/**
* @brief Token log probability structure
*/
typedef struct XLLM_CAPI_EXPORT XLLM_LogProb {
/** Token ID */
uint32_t token_id;
/** Log probability of the token */
float logprob;
} XLLM_LogProb;
/**
* @brief List of token log probabilities
*/
typedef struct XLLM_CAPI_EXPORT XLLM_LogProbs {
/** Pointer to array of log probability entries */
XLLM_LogProb* entries;
/** Number of entries in the logprobs array */
size_t entries_size;
} XLLM_LogProbs;
/**
* @brief Inference result candidate
*/
typedef struct XLLM_CAPI_EXPORT XLLM_Choice {
/** Index of the generated completion candidate */
uint32_t index;
/** Generated text for completions inference (NULL for Chat mode) */
char* text;
/** Generated message for chatcompletions inference (NULL for Completion mode)
*/
XLLM_ChatMessage* message;
/** Generated token ids */
int32_t* token_ids;
/** Generated token ids size */
size_t token_size;
/** Token log probabilities */
XLLM_LogProbs logprobs;
/** Reason generation stopped (stop/length/function_call) */
char finish_reason[XLLM_META_STRING_FIELD_MAX_LEN];
} XLLM_Choice;
/**
* @brief List of inference result candidates
*/
typedef struct XLLM_CAPI_EXPORT XLLM_Choices {
/** Pointer to array of completion choice entries */
XLLM_Choice* entries;
/** Number of entries in the choices array */
size_t entries_size;
} XLLM_Choices;
/**
* @brief REC/OneRec specific output extension aligned by choice index
*/
typedef struct XLLM_CAPI_EXPORT XLLM_RecOutput {
/** Choice index this REC extension belongs to */
uint32_t index;
/** Selected REC item ids for this choice */
int64_t* item_ids;
/** Number of item ids in the item_ids array */
size_t item_ids_size;
/** Token-aligned REC/OneRec logprobs for this choice */
float* rec_token_logprobs;
/** Number of entries in rec_token_logprobs */
size_t rec_token_logprobs_size;
} XLLM_RecOutput;
/**
* @brief List of REC/OneRec specific output extensions
*/
typedef struct XLLM_CAPI_EXPORT XLLM_RecOutputs {
/** Pointer to array of REC output entries */
XLLM_RecOutput* entries;
/** Number of entries in the REC output array */
size_t entries_size;
} XLLM_RecOutputs;
#define XLLM_ERROR_INFO_MAX_LEN 512
/**
* @brief Inference response structure
*/
typedef struct XLLM_CAPI_EXPORT XLLM_Response {
/** Response status code (0 = success, non-zero = error) */
XLLM_StatusCode status_code;
/** Error details (NULL = no error) */
char error_info[XLLM_ERROR_INFO_MAX_LEN];
/** Unique ID for the completion request (fixed-length string) */
char id[XLLM_META_STRING_FIELD_MAX_LEN];
/** Object type (fixed to "text_completion") */
char object[XLLM_META_STRING_FIELD_MAX_LEN];
/** Unix timestamp (seconds) of when the completion was created */
int64_t created;
/** Model name used for the completion */
char model[XLLM_META_STRING_FIELD_MAX_LEN];
/** List of generated completion candidates */
XLLM_Choices choices;
/** Token usage statistics for the request */
XLLM_Usage usage;
/** REC/OneRec specific response extensions */
XLLM_RecOutputs rec_outputs;
} XLLM_Response;
/**
* @brief Enumeration of tensor data types
*/
typedef enum XLLM_CAPI_EXPORT XLLM_DataType {
XLLM_DTYPE_UNDEFINED = 0,
XLLM_DTYPE_FLOAT16 = 1,
XLLM_DTYPE_FLOAT32 = 2,
XLLM_DTYPE_FLOAT64 = 3,
XLLM_DTYPE_BFLOAT16 = 4,
XLLM_DTYPE_INT8 = 5,
XLLM_DTYPE_INT16 = 6,
XLLM_DTYPE_INT32 = 7,
XLLM_DTYPE_INT64 = 8,
XLLM_DTYPE_UINT8 = 9,
XLLM_DTYPE_UINT16 = 10,
XLLM_DTYPE_UINT32 = 11,
XLLM_DTYPE_UINT64 = 12,
XLLM_DTYPE_BOOL = 13,
XLLM_DTYPE_STRING = 14
} XLLM_DataType;
/**
* @brief Structure representing tensor dimensions (shape)
* @note Max supported rank is 8 (matches dim array length)
*/
typedef struct XLLM_CAPI_EXPORT XLLM_Dims {
/** Number of dimensions (0=scalar, 1=vector, ..., 8) */
int rank;
/** Size of each dimension (unused dims must be 0) */
int dim[8];
} XLLM_Dims;
/**
* @brief Core tensor structure for numerical computation
* @warning 1. data pointer is read-only, managed by external caller
* 2. dtype must match the actual type of data buffer
* 3. dims.rank must not exceed 8
*/
typedef struct XLLM_CAPI_EXPORT XLLM_Tensor {
/** Data type of tensor elements */
XLLM_DataType dtype;
/** Dimension information (shape) of the tensor */
XLLM_Dims dims;
/** Read-only pointer to tensor data buffer */
const void* data;
} XLLM_Tensor;
/**
* @brief Dynamic list of tensors (replaces C++ std::vector<XLLM_Tensor>)
*/
typedef struct XLLM_CAPI_EXPORT XLLM_Tensors {
XLLM_Tensor* entries;
size_t entries_size;
} XLLM_Tensors;
/**
* @brief Enumeration of multimodal data types (bitmask compatible)
* @note Each type is a bit flag (supports multiple types via type_mask)
*/
typedef enum XLLM_CAPI_EXPORT XLLM_MM_Type {
/** No multimodal type (invalid state) */
XLLM_MM_TYPE_NONE = 0,
/** Image modality (JPG/PNG/BMP) */
XLLM_MM_TYPE_IMAGE = 1 << 0,
/** Audio modality (WAV/MP3) */
XLLM_MM_TYPE_AUDIO = 1 << 1,
/** Video modality (H264/H265) */
XLLM_MM_TYPE_VIDEO = 1 << 2,
/** Text modality (tokenized text) */
XLLM_MM_TYPE_TEXT = 1 << 3,
/** Embedding modality (token embeddings) */
XLLM_MM_TYPE_EMBEDDING = 1 << 4
} XLLM_MM_Type;
/**
* @brief Multimodal value (variant type: single tensor or tensor list)
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_Value {
/** Type flag: true=single tensor, false=tensor list */
bool is_single_tensor;
union {
/** Single tensor (valid if is_single_tensor=true) */
XLLM_Tensor tensor;
/** Tensor list (valid if is_single_tensor=false) */
XLLM_Tensors tensors;
} data;
} XLLM_MM_Value;
/**
* @brief Single entry in multimodal dictionary (key-value pair)
* @note 1. Key is fixed-length string (null-terminated if shorter than max len)
* 2. Key must be unique within a dictionary
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_DictEntry {
/** Fixed-length key */
char key[XLLM_META_STRING_FIELD_MAX_LEN];
/** Value associated with the key */
XLLM_MM_Value value;
} XLLM_MM_DictEntry;
/**
* @brief Multimodal dictionary (array of key-value entries)
* @note 1. entries is a heap-allocated array (must be freed by caller)
* 2. entries_size = number of valid entries (no empty slots)
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_Dict {
XLLM_MM_DictEntry* entries;
size_t entries_size;
} XLLM_MM_Dict;
/**
* @brief Token position information (offset + length) for multimodal data
* @note Used to map multimodal data to token positions in sequence
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_TokenPos {
/** Start offset of tokens (0-based) */
uint32_t offset;
/** Number of tokens (must be >0 for valid position) */
uint32_t length;
} XLLM_MM_TokenPos;
/**
* @brief Base struct for multimodal metadata (to be extended by specific
* modalities)
* @note This is a placeholder for modality-specific metadata (e.g., image size,
* audio sample rate) Extend with union for image/audio/video metadata in
* production use
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_Meta {
// Placeholder for future extension,
// e.g., XLLM_ImageMeta|XLLM_AudioMeta|XLLM_VideoMeta
} XLLM_MM_Meta;
/**
* @brief State information for a single multimodal item
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_State {
/** Token position for multimodal data alignment */
XLLM_MM_TokenPos token_pos;
} XLLM_MM_State;
/**
* @brief Single multimodal data item (core unit of multimodal data)
* @note 1. type = modality type (image/audio/video/text/embedding)
* 2. data = numerical content (tensor/tensor list)
* 3. meta = modality-specific metadata (empty in base version)
* 4. state = token position and processing state
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_Item {
/** Modality type (e.g., XLLM_MM_TYPE_EMBEDDING) */
XLLM_MM_Type type;
/** Core data (tensor/tensor list) */
XLLM_MM_Value data;
/** Modality-specific metadata (extendable) */
XLLM_MM_Meta meta;
/** Processing state and token position */
XLLM_MM_State state;
} XLLM_MM_Item;
/**
* @brief List of multimodal items (replaces C++ std::vector<XLLM_MM_Item>)
* @note 1. entries is a heap-allocated array (must be freed by caller)
* 2. entries_size = number of valid items (no empty slots)
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_Items {
XLLM_MM_Item* entries;
size_t entries_size;
} XLLM_MM_Items;
/**
* @brief Core multimodal data container (supports list/dict storage)
*/
typedef struct XLLM_CAPI_EXPORT XLLM_MM_Data {
/** Bitmask of multimodal types (e.g., IMAGE | EMBEDDING) */
uint32_t type_mask;
/** Storage type: true=XLLM_MM_Dict, false=XLLM_MM_Items */
bool is_dict;
union {
/** Dict storage (valid if is_dict=true) */
XLLM_MM_Dict dict;
/** List storage (valid if is_dict=false) */
XLLM_MM_Items items;
} data;
} XLLM_MM_Data;
#ifdef __cplusplus
}
#endif
#endif // XLLM_C_TYPES_H

View File

@@ -0,0 +1,44 @@
### How to compile xllm dynamic library
Run the following command in root directory:
```
python setup.py build --generate-so true
```
If you want to debug, it needs to set DEBUG environment variable.
```
export DEBUG=1
```
### How to install dynamic library
Run installation script xllm/cc_api/install.sh, headers and dynamic library will be installed in /usr/local/xllm directory.
```
cd xllm/cc_api
sh install.sh
```
You will see the following files in /usr/local/xllm directory:
```
[root@A03-R40-I189-101-4100046 cc_api]# tree /usr/local/xllm
/usr/local/xllm
|-- include
| |-- llm.h
| |-- macros.h
| `-- types.h
`-- lib
|-- libcust_opapi.so
`-- libxllm.so
3 directories, 5 files
```
### How to run cc_api examples
It provides two examples which use cc_api to create xllm instance and run inference. The single_llm_instance.cpp creates one instance which is used in most LLM scenes. The multiple_llm_instances.cpp creates two instances which is used in multiple-models scene or one model with multiple versions.
You can follow the commands to compile and run these examples:
```
cd examples && mkdir build
cd build && cmake .. && make && cd ..
sh start-llm-instance.sh
```

View File

@@ -0,0 +1,97 @@
cmake_minimum_required(VERSION 3.26)
project(llm_examples_project)
set(CMAKE_CXX_STANDARD 20)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
set(CMAKE_CXX_EXTENSIONS ON)
link_directories(/usr/local/xllm/lib)
### build single_llm_instance
add_executable(single_llm_instance single_llm_instance.cpp service_request.cpp)
target_include_directories(single_llm_instance
PRIVATE /usr/local/xllm/include
)
target_link_libraries(single_llm_instance
PRIVATE xllm
)
### get cxx11 ABI config from torch api.
set(CHECK_TORCH_ABI_SCRIPT
"
import torch
print(0 if not torch.compiled_with_cxx11_abi() else 1)
"
)
execute_process(
COMMAND python3 -c "${CHECK_TORCH_ABI_SCRIPT}"
OUTPUT_VARIABLE TORCH_CXX11_ABI_VALUE
OUTPUT_STRIP_TRAILING_WHITESPACE
RESULT_VARIABLE TORCH_ABI_CHECK_RESULT
)
if(NOT TORCH_ABI_CHECK_RESULT EQUAL 0)
message(WARNING "Failed to detect PyTorch CXX11 ABI! "
"Reason: ${TORCH_ABI_CHECK_RESULT}\n"
"Fallback to _GLIBCXX_USE_CXX11_ABI=0 (default)")
set(TORCH_CXX11_ABI_VALUE 0)
else()
message(STATUS "Detected PyTorch CXX11 ABI: ${TORCH_CXX11_ABI_VALUE} "
"(0=OFF, 1=ON)")
endif()
if(TORCH_CXX11_ABI_VALUE EQUAL 0)
target_compile_definitions(single_llm_instance
PRIVATE
_GLIBCXX_USE_CXX11_ABI=0
USE_CXX11_ABI=OFF
)
else()
target_compile_definitions(single_llm_instance
PRIVATE
_GLIBCXX_USE_CXX11_ABI=1
USE_CXX11_ABI=ON
)
endif()
set_target_properties(single_llm_instance PROPERTIES
BUILD_WITH_INSTALL_RPATH ON
INSTALL_RPATH "/usr/local/xllm/lib"
RPATH "/usr/local/xllm/lib"
)
### build multiple_llm_instances
add_executable(multiple_llm_instances multiple_llm_instances.cpp service_request.cpp)
target_include_directories(multiple_llm_instances
PRIVATE /usr/local/xllm/include
)
target_link_libraries(multiple_llm_instances
PRIVATE xllm
)
if(TORCH_CXX11_ABI_VALUE EQUAL 0)
target_compile_definitions(multiple_llm_instances
PRIVATE
_GLIBCXX_USE_CXX11_ABI=0
USE_CXX11_ABI=OFF
)
else()
target_compile_definitions(multiple_llm_instances
PRIVATE
_GLIBCXX_USE_CXX11_ABI=1
USE_CXX11_ABI=ON
)
endif()
set_target_properties(multiple_llm_instances PROPERTIES
BUILD_WITH_INSTALL_RPATH ON
INSTALL_RPATH "/usr/local/xllm/lib"
RPATH "/usr/local/xllm/lib"
)

View File

@@ -0,0 +1,78 @@
/* Copyright 2025 The xLLM 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 <unistd.h>
#include <iostream>
#include <memory>
#include "llm.h"
#include "service_request.h"
/**
* In some scenes, such as, one model needs to support multiple versions, or
* multiple models needs to be supported, then you can follow this example.
*/
std::string devices = "npu:0";
std::string model_path = "/export/home/models/Qwen3-4B";
int main(int argc, char** argv) {
std::cout << "Start to bootup the first LLM instance." << std::endl;
std::shared_ptr<xllm::LLM> llm_instance_01 = std::make_shared<xllm::LLM>();
xllm::XLLM_InitLLMOptions options_01;
options_01.max_memory_utilization = 0.45;
bool ret = llm_instance_01->Initialize(model_path, devices, options_01);
if (!ret) {
std::cout << "LLM instance init failed." << std::endl;
return -1;
}
std::cout << "LLM init succefully." << std::endl;
const std::string model_name = "Qwen3-4B";
xllm::cc_api_test::run_completion_request(model_name, llm_instance_01.get());
xllm::cc_api_test::run_chat_completion_request(model_name,
llm_instance_01.get());
std::cout << "Start to bootup the second LLM instance." << std::endl;
std::shared_ptr<xllm::LLM> llm_instance_02 = std::make_shared<xllm::LLM>();
xllm::XLLM_InitLLMOptions options_02;
// The following options must be set to create different internal servers.
options_02.master_node_addr = "127.0.0.1:28899";
options_02.transfer_listen_port = 27000;
options_02.server_idx = 1;
options_02.max_memory_utilization = 0.9;
ret = llm_instance_02->Initialize(model_path, devices, options_02);
if (!ret) {
std::cout << "LLM instance init failed." << std::endl;
return -1;
}
std::cout << "LLM init succefully." << std::endl;
xllm::cc_api_test::run_completion_request(model_name, llm_instance_02.get());
xllm::cc_api_test::run_chat_completion_request(model_name,
llm_instance_02.get());
sleep(10);
return 0;
}

View File

@@ -0,0 +1,97 @@
/* Copyright 2025 The xLLM 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 "service_request.h"
#include <iostream>
namespace xllm {
namespace cc_api_test {
void run_completion_request(const std::string& model_name,
xllm::LLM* llm_instance) {
std::cout << "*** llm completions start ***" << std::endl;
std::string prompt =
"recommend 3 cheap and easy-to-use electric shavers, briefly "
"describe the product name, price, and features";
xllm::XLLM_RequestParams params;
params.max_tokens = 500;
xllm::XLLM_Response response =
llm_instance->Completions(model_name, prompt, 20000, params);
if (response.status_code != xllm::XLLM_StatusCode::kSuccess) {
std::cout << "LLM completions failed, error info: " << response.error_info
<< std::endl;
return;
} else {
for (auto choice : response.choices) {
if (choice.text.has_value()) {
std::cout << "LLM completions output: " << choice.text.value().c_str()
<< std::endl;
}
}
}
std::cout << "*** llm completions end ***" << std::endl;
}
void run_chat_completion_request(const std::string& model_name,
xllm::LLM* llm_instance) {
std::cout << "*** llm chat completions start ***" << std::endl;
xllm::XLLM_ChatMessage message;
message.role = "user";
message.content =
"You are an expert in e-commerce scenarios. The current scenario is an "
"e-commerce search engine with a comprehensive range of business "
"categories. Your task is to determine whether 'user query' and "
"'product title' are related in the e-commerce search engine. "
"Discrimination criteria: If the search for 'user query' returns' "
"product title 'that meets the user's needs, then the task is "
"relevant. Output requirement: Please provide the answer in the "
"'related' or 'unrelated' section, without mentioning any other "
"content. User query: 'Hotpot sauce'. Product title: 'Grassland Red "
"Sun Hotpot Base Dip Multi flavored Barbecue Sauce Tomato Sauce Leek "
"Flower Sauce Nightsnack Paired with [New] Spicy Barbecue Sauce 100g'";
std::vector<xllm::XLLM_ChatMessage> messages;
messages.emplace_back(message);
xllm::XLLM_RequestParams params;
params.max_tokens = 100;
xllm::XLLM_Response response =
llm_instance->ChatCompletions(model_name, messages, 20000, params);
if (response.status_code != xllm::XLLM_StatusCode::kSuccess) {
std::cout << "LLM completions failed, error info: " << response.error_info
<< std::endl;
return;
} else {
for (auto choice : response.choices) {
if (choice.message.has_value()) {
std::cout << "LLM completions output: role: "
<< choice.message.value().role.c_str()
<< ",content: " << choice.message.value().content.c_str()
<< std::endl;
}
}
}
std::cout << "*** llm chat completions end ***" << std::endl;
}
} // namespace cc_api_test
} // namespace xllm

View File

@@ -0,0 +1,28 @@
/* Copyright 2025 The xLLM 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 "llm.h"
namespace xllm {
namespace cc_api_test {
// Send Completion request and print the inference result.
void run_completion_request(const std::string& model_name,
xllm::LLM* llm_instance);
// Send ChatCompletion request and print the inference result.
void run_chat_completion_request(const std::string& model_name,
xllm::LLM* llm_instance);
} // namespace cc_api_test
} // namespace xllm

View File

@@ -0,0 +1,51 @@
/* Copyright 2025 The xLLM 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 <unistd.h>
#include <iostream>
#include "llm.h"
#include "service_request.h"
/**
* In most scenes, you can follow this example to integrate xllm as an internal
* inference engine.
*/
std::string devices = "npu:0";
std::string model_path = "/export/home/models/Qwen3-4B";
int main(int argc, char** argv) {
xllm::LLM llm_instance;
xllm::XLLM_InitLLMOptions options;
bool ret = llm_instance.Initialize(model_path, devices, options);
if (!ret) {
std::cout << "LLM init failed." << std::endl;
return -1;
}
std::cout << "LLM init succefully." << std::endl;
const std::string model_name = "Qwen3-4B";
xllm::cc_api_test::run_completion_request(model_name, &llm_instance);
xllm::cc_api_test::run_chat_completion_request(model_name, &llm_instance);
sleep(10);
return 0;
}

View File

@@ -0,0 +1,14 @@
#!/bin/bash
clear
\rm -rf core.*
cd build && make && cd ..
# export ASDOPS_LOG_LEVEL=DEBUG
# export ASDOPS_LOG_TO_STDOUT=1
export ASCEND_RT_VISIBLE_DEVICES=12
python3 -c "import torch; import torch_npu; torch_npu.npu.set_device('npu:0')"
# build/single_llm_instance
build/multiple_llm_instances

View File

@@ -0,0 +1,119 @@
#!/bin/bash
set -euo pipefail
SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)
BIN_DIR="${SCRIPT_DIR}/../../bin"
TEMP_DIR="xllm"
INCLUDE_DIR="${TEMP_DIR}/include"
LIB_DIR="${TEMP_DIR}/lib"
VERSION_FILE="${SCRIPT_DIR}/../../version.txt"
TAR_BASE_NAME="xllm"
LOCAL_INSTALL_DIR="/usr/local"
LOCAL_TARGET_DIR="${LOCAL_INSTALL_DIR}/xllm"
HEADERS=("${SCRIPT_DIR}/llm.h" "${SCRIPT_DIR}/macros.h" "${SCRIPT_DIR}/types.h")
SO_FILES=(
"${SCRIPT_DIR}/../../build/lib.linux-aarch64-cpython-311/xllm/libxllm.so"
"/usr/local/Ascend/ascend-toolkit/8.2.RC1/opp/vendors/xllm/op_api/lib/libcust_opapi.so"
)
error_exit() {
echo -e "\033[31merror: $1\033[0m" >&2
exit 1
}
cd_bin_dir() {
if [ ! -d "${BIN_DIR}" ]; then
mkdir -p "${BIN_DIR}" || error_exit "failed to create bin directory: ${BIN_DIR}"
fi
cd "${BIN_DIR}" || error_exit "failed to enter bin directory: ${BIN_DIR}"
}
read_version() {
if [ ! -f "${VERSION_FILE}" ]; then
error_exit "${VERSION_FILE} is not existed"
fi
VERSION=$(cat "${VERSION_FILE}" | tr -d '[:space:]')
if [ -z "${VERSION}" ]; then
error_exit "version content is empty"
fi
TAR_FILE="${TAR_BASE_NAME}_${VERSION}.tar.gz"
}
check_files() {
for header in "${HEADERS[@]}"; do
if [ ! -f "${header}" ]; then
error_exit "${header} is not existed"
fi
done
for so_file in "${SO_FILES[@]}"; do
if [ ! -f "${so_file}" ]; then
error_exit "${so_file} is not existed"
fi
done
}
create_dirs() {
mkdir -p "${INCLUDE_DIR}" || error_exit "create include directory failed"
mkdir -p "${LIB_DIR}" || error_exit "create lib directory failed"
}
copy_headers() {
for header in "${HEADERS[@]}"; do
cp -f "${header}" "${INCLUDE_DIR}/" || error_exit "copy ${header} failed"
done
}
copy_so() {
for so_file in "${SO_FILES[@]}"; do
cp -f "${so_file}" "${LIB_DIR}/" || error_exit "copy ${so_file} failed"
done
}
package_tar() {
tar -czf "${TAR_FILE}" "${TEMP_DIR}" || error_exit "tar failed"
}
cleanup_temp() {
rm -rf "${TEMP_DIR}" || error_exit "rm temp directory failed"
}
extract_to_local() {
if [ ! -f "${TAR_FILE}" ]; then
error_exit "${TAR_FILE} is not existed"
fi
if [ ! -d "${LOCAL_INSTALL_DIR}" ]; then
error_exit "local install directory is not existed"
fi
if [ -d "${LOCAL_TARGET_DIR}" ]; then
rm -rf "${LOCAL_TARGET_DIR}" || error_exit "rm old xllm directory failed"
fi
tar -xzf "${TAR_FILE}" -C "${LOCAL_INSTALL_DIR}" || error_exit "extract failed"
}
main() {
cd_bin_dir
read_version
check_files
create_dirs
copy_headers
copy_so
package_tar
cleanup_temp
extract_to_local
echo -e "install file: \033[33m${TAR_FILE}\033[0m"
echo -e "install path: \033[33m/usr/local/${TEMP_DIR}\033[0m"
}
main

View File

@@ -0,0 +1,254 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <folly/executors/CPUThreadPoolExecutor.h>
#include "core/common/instance_name.h"
#include "core/distributed_runtime/llm_master.h"
#include "core/framework/request/request_output.h"
#include "core/framework/request/request_params.h"
#include "core/util/uuid.h"
#include "types.h"
namespace xllm {
struct LLMCore {
// List of loaded model identifiers
std::vector<std::string> model_ids;
// Master controller for LLM runtime management
std::unique_ptr<LLMMaster> master;
// Thread pool for asynchronous task execution
std::unique_ptr<folly::CPUThreadPoolExecutor> executor;
};
namespace detail {
namespace {
thread_local ShortUUID short_uuid;
std::string generate_request_id() {
return "xllm-" + InstanceName::name()->get_name_hash() + "-" +
short_uuid.random();
}
} // namespace
enum class InterfaceType { COMPLETIONS, CHAT_COMPLETIONS };
RequestParams transfer_request_params(
const XLLM_RequestParams& request_params) {
RequestParams xllm_request_params;
xllm_request_params.echo = request_params.echo;
xllm_request_params.offline = request_params.offline;
xllm_request_params.logprobs = request_params.logprobs;
xllm_request_params.best_of = request_params.best_of;
xllm_request_params.slo_ms = request_params.slo_ms;
xllm_request_params.top_k = request_params.top_k;
xllm_request_params.top_p = request_params.top_p;
xllm_request_params.ignore_eos = request_params.ignore_eos;
xllm_request_params.skip_special_tokens = request_params.skip_special_tokens;
xllm_request_params.n = request_params.n;
xllm_request_params.max_tokens = request_params.max_tokens;
xllm_request_params.frequency_penalty = request_params.frequency_penalty;
xllm_request_params.presence_penalty = request_params.presence_penalty;
xllm_request_params.repetition_penalty = request_params.repetition_penalty;
xllm_request_params.stop = request_params.stop;
xllm_request_params.stop_token_ids = request_params.stop_token_ids;
xllm_request_params.beam_width = request_params.beam_width;
xllm_request_params.num_return_sequences =
request_params.num_return_sequences;
xllm_request_params.top_logprobs = request_params.top_logprobs;
return xllm_request_params;
}
XLLM_Response build_success_response(const RequestOutput& output,
const InterfaceType& if_type,
const std::string& request_id,
int64_t created_time,
const std::string& model) {
XLLM_Response response;
response.status_code = XLLM_StatusCode::kSuccess;
response.id = request_id;
response.created = created_time;
response.model = model;
if (if_type == InterfaceType::COMPLETIONS) {
response.object = "text_completion";
} else if (if_type == InterfaceType::CHAT_COMPLETIONS) {
response.object = "chat.completion";
}
response.choices.reserve(output.outputs.size());
for (const auto& output : output.outputs) {
XLLM_Choice choice;
choice.index = output.index;
if (output.logprobs.has_value()) {
std::vector<XLLM_LogProb> xllm_logprobs;
xllm_logprobs.reserve(output.logprobs.value().size());
for (const auto& logprob : output.logprobs.value()) {
XLLM_LogProb xllm_logprob;
xllm_logprob.token = logprob.token;
xllm_logprob.token_id = logprob.token_id;
xllm_logprob.logprob = logprob.logprob;
if (logprob.top_logprobs.has_value()) {
xllm_logprob.top_logprobs.reserve(
logprob.top_logprobs.value().size());
for (const auto& top_logprob : logprob.top_logprobs.value()) {
XLLM_LogProbData xllm_logprob_data;
xllm_logprob_data.token = top_logprob.token;
xllm_logprob_data.token_id = top_logprob.token_id;
xllm_logprob_data.logprob = top_logprob.logprob;
xllm_logprob.top_logprobs.emplace_back(xllm_logprob_data);
}
}
xllm_logprobs.emplace_back(xllm_logprob);
}
choice.logprobs = xllm_logprobs;
}
if (if_type == InterfaceType::COMPLETIONS) {
choice.text = output.text;
} else if (if_type == InterfaceType::CHAT_COMPLETIONS) {
XLLM_ChatMessage chat_message;
chat_message.role = "assistant";
chat_message.content = output.text;
choice.message = chat_message;
}
if (output.finish_reason.has_value()) {
choice.finish_reason = output.finish_reason.value();
}
response.choices.emplace_back(choice);
}
if (output.usage.has_value()) {
const auto& usage = output.usage.value();
response.usage.prompt_tokens = usage.num_prompt_tokens;
response.usage.completion_tokens = usage.num_generated_tokens;
response.usage.total_tokens = usage.num_total_tokens;
}
return response;
}
XLLM_Response build_error_response(const std::string& request_id,
XLLM_StatusCode status_code,
const std::string& error_info) {
XLLM_Response response;
response.status_code = status_code;
response.error_info = error_info;
response.id = request_id.empty() ? "unknown_request" : request_id;
LOG(ERROR) << "Request [" << response.id << "] error: " << error_info
<< " (code: " << static_cast<int>(response.status_code) << ")";
return response;
}
template <typename InputType>
XLLM_Response handle_inference_request(LLMCore* llm_core,
const std::string& model_id,
const InputType& input,
uint32_t timeout_ms,
const XLLM_RequestParams& request_params,
InterfaceType interface_type) {
if (!llm_core) {
return build_error_response(
"", XLLM_StatusCode::kNotInitialized, "LLM is not initialized");
}
auto it = std::find(
llm_core->model_ids.begin(), llm_core->model_ids.end(), model_id);
if (it == llm_core->model_ids.end()) {
return build_error_response("",
XLLM_StatusCode::kModelNotFound,
"Specified model ID not loaded: " + model_id);
}
RequestParams xllm_request_params = transfer_request_params(request_params);
std::string request_id = xllm_request_params.request_id.empty()
? generate_request_id()
: xllm_request_params.request_id;
xllm_request_params.request_id = request_id;
int64_t created_time = absl::ToUnixSeconds(absl::Now());
try {
auto promise_ptr = std::make_shared<folly::Promise<XLLM_Response>>();
auto weak_promise =
std::weak_ptr<folly::Promise<XLLM_Response>>(promise_ptr);
auto future = promise_ptr->getSemiFuture();
llm_core->master->handle_request(
input,
std::nullopt,
xllm_request_params,
std::nullopt,
[model_id,
request_id,
created_time,
interface_type,
weak_promise,
timeout_ms](const RequestOutput& req_output) -> bool {
auto promise_ptr = weak_promise.lock();
if (!promise_ptr) {
return false;
}
try {
XLLM_Response response = build_success_response(
req_output, interface_type, request_id, created_time, model_id);
promise_ptr->setValue(std::move(response));
} catch (const folly::PromiseAlreadySatisfied& e) {
return false;
}
return true;
});
return std::move(future)
.via(llm_core->executor.get())
.within(std::chrono::milliseconds(timeout_ms))
.thenTry([](folly::Try<XLLM_Response>&& result) {
if (result.hasValue()) {
return std::move(result).value();
} else {
result.throwUnlessValue();
return XLLM_Response{};
}
})
.get();
} catch (const folly::FutureTimeout& e) {
return build_error_response(
request_id, XLLM_StatusCode::kTimeout, "Request timed out");
} catch (const std::exception& e) {
return build_error_response(
request_id,
XLLM_StatusCode::kInternalError,
"Failed to handle request: " + std::string(e.what()));
}
}
} // namespace detail
} // namespace xllm

View File

@@ -0,0 +1,186 @@
/* Copyright 2025 The xLLM 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 "llm.h"
#include <folly/Unit.h>
#include <folly/experimental/coro/Timeout.h>
#include <folly/futures/Future.h>
#include <glog/logging.h>
#include <pthread.h>
#include <atomic>
#include <exception>
#include "core/common/global_flags.h"
#include "internal.h"
namespace xllm {
namespace {
static std::atomic<bool> g_glog_inited = false;
static pthread_mutex_t g_log_init_mutex = PTHREAD_MUTEX_INITIALIZER;
void InitGlog(const std::string& log_dir) {
pthread_mutex_lock(&g_log_init_mutex);
if (!g_glog_inited) {
google::InitGoogleLogging("xllm");
google::SetLogDestination(google::INFO,
(log_dir + "/xllm.log.INFO.").c_str());
google::SetLogDestination(google::WARNING,
(log_dir + "/xllm.log.WARNING.").c_str());
google::SetLogDestination(google::ERROR,
(log_dir + "/xllm.log.ERROR.").c_str());
google::SetStderrLogging(google::FATAL);
g_glog_inited = true;
}
pthread_mutex_unlock(&g_log_init_mutex);
}
} // namespace
LLM::LLM() = default;
LLM::~LLM() {
if (nullptr != llm_core_) {
delete llm_core_;
llm_core_ = nullptr;
}
}
bool LLM::Initialize(const std::string& model_path,
const std::string& devices,
const XLLM_InitLLMOptions& init_options) {
if (!init_options.log_dir.empty()) {
InitGlog(init_options.log_dir);
}
if (!std::filesystem::exists(model_path)) {
LOG(ERROR) << "model path[" << model_path << "] does not exist";
return false;
}
try {
Options options;
options.model_path(model_path)
.task_type(init_options.task)
.devices(devices)
.draft_model_path(init_options.draft_model)
.draft_devices(init_options.draft_devices)
.backend("llm")
.block_size(init_options.block_size)
.max_cache_size(init_options.max_cache_size)
.max_memory_utilization(init_options.max_memory_utilization)
.enable_prefix_cache(init_options.enable_prefix_cache)
.max_tokens_per_batch(init_options.max_tokens_per_batch)
.max_seqs_per_batch(init_options.max_seqs_per_batch)
.max_tokens_per_chunk_for_prefill(
init_options.max_tokens_per_chunk_for_prefill)
.num_speculative_tokens(init_options.num_speculative_tokens)
.num_request_handling_threads(init_options.num_request_handling_threads)
.communication_backend(init_options.communication_backend)
.rank_tablefile(init_options.rank_tablefile)
.expert_parallel_degree(init_options.expert_parallel_degree)
.enable_chunked_prefill(init_options.enable_chunked_prefill)
.enable_prefill_sp(init_options.enable_prefill_sp)
.master_node_addr(init_options.master_node_addr)
.device_ip(init_options.device_ip)
.transfer_listen_port(init_options.transfer_listen_port)
.nnodes(init_options.nnodes)
.node_rank(init_options.node_rank)
.dp_size(init_options.dp_size)
.ep_size(init_options.ep_size)
.instance_name(init_options.instance_name)
.enable_disagg_pd(init_options.enable_disagg_pd)
.enable_schedule_overlap(init_options.enable_schedule_overlap)
.enable_pd_ooc(init_options.enable_pd_ooc)
.kv_cache_transfer_mode(init_options.kv_cache_transfer_mode)
.disable_ttft_profiling(init_options.disable_ttft_profiling)
.enable_forward_interruption(init_options.enable_forward_interruption)
.enable_shm(init_options.enable_shm)
.input_shm_size(init_options.input_shm_size)
.output_shm_size(init_options.output_shm_size)
.is_local(init_options.is_local)
.server_idx(init_options.server_idx);
#if !defined(USE_NPU) && !defined(USE_CUDA)
FLAGS_enable_block_copy_kernel = false;
#endif
llm_core_ = new LLMCore();
llm_core_->master = std::make_unique<LLMMaster>(options);
llm_core_->master->run();
size_t cpu_cores = std::thread::hardware_concurrency();
size_t thread_num = std::clamp((cpu_cores == 0) ? 8 : cpu_cores / 2,
static_cast<size_t>(4),
static_cast<size_t>(16));
llm_core_->executor =
std::make_unique<folly::CPUThreadPoolExecutor>(thread_num);
std::filesystem::path model_path_fs =
std::filesystem::path(model_path).lexically_normal();
std::string model_id;
if (model_path_fs.has_filename()) {
model_id = model_path_fs.filename().string();
} else if (!model_path_fs.empty()) {
model_id = model_path_fs.string();
} else {
model_id = "default";
}
llm_core_->model_ids.emplace_back(model_id);
return true;
} catch (const std::exception& e) {
LOG(ERROR) << "LLM initialization failed: " << e.what();
if (nullptr != llm_core_) {
delete llm_core_;
llm_core_ = nullptr;
}
return false;
}
}
XLLM_Response LLM::Completions(const std::string& model_id,
const std::string& prompt,
uint32_t timeout_ms,
const XLLM_RequestParams& request_params) {
return detail::handle_inference_request(llm_core_,
model_id,
prompt,
timeout_ms,
request_params,
detail::InterfaceType::COMPLETIONS);
}
XLLM_Response LLM::ChatCompletions(
const std::string& model_id,
const std::vector<XLLM_ChatMessage>& messages,
uint32_t timeout_ms,
const XLLM_RequestParams& request_params) {
std::vector<Message> internal_messages;
internal_messages.reserve(messages.size());
for (const auto& msg : messages) {
internal_messages.emplace_back(msg.role, msg.content);
}
return detail::handle_inference_request(
llm_core_,
model_id,
internal_messages,
timeout_ms,
request_params,
detail::InterfaceType::CHAT_COMPLETIONS);
}
} // namespace xllm

View File

@@ -0,0 +1,94 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <atomic>
#include <optional>
#include <string>
#include "macros.h"
#include "types.h"
namespace xllm {
// Forward declaration
struct LLMCore;
// A wrapper for loading, initializing, and text generation functions of large
// language models
class XLLM_CAPI_EXPORT LLM {
public:
LLM();
virtual ~LLM();
LLM(const LLM&) = delete;
LLM& operator=(const LLM&) = delete;
LLM(LLM&&) noexcept = delete;
LLM& operator=(LLM&&) noexcept = delete;
/**
* @brief Initialize the model: Load model files and configure runtime
* environment
* @param model_path Path to model files
* @param devices Device configuration (format: "npu:1" for specific NPU,
* "auto" for auto-selection)
* @param init_options Advanced initialization options, Provided default
* configuration
* @return bool true if initialization succeeds; false if fails
* @note Must be called before Completions/ChatCompletions, and only needs to
* be called once
*/
bool Initialize(const std::string& model_path,
const std::string& devices,
const XLLM_InitLLMOptions& init_options);
/**
* @brief Generate completions for the given prompt
* @param model_id ID of the loaded model
* @param prompt Input prompt text
* @param timeout_ms Timeout in milliseconds
* @param request_params Request parameters (temperature, max tokens, etc.)
* @return XLLM_Response Response containing generated text and
* metadata
*/
XLLM_Response Completions(const std::string& model_id,
const std::string& prompt,
uint32_t timeout_ms,
const XLLM_RequestParams& request_params);
/**
* @brief Generates chat completions based on a sequence of conversation
* messages
* @param model_id ID of the loaded model
* @param messages A list of XLLM_ChatMessage objects representing the
* conversation history, each message contains a role (user/assistant/system)
* and content text
* @param timeout_ms Timeout in milliseconds
* @param request_params Request parameters (temperature, max tokens, etc.)
* @return XLLM_Response Response containing generated text and
* metadata
*/
XLLM_Response ChatCompletions(const std::string& model_id,
const std::vector<XLLM_ChatMessage>& messages,
uint32_t timeout_ms,
const XLLM_RequestParams& request_params);
private:
// Opaque pointer to internal LLM core implementation
LLMCore* llm_core_ = nullptr;
};
} // namespace xllm

View File

@@ -0,0 +1,27 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
namespace xllm {
#ifdef XLLM_CAPI_WEAK
#define XLLM_CAPI_EXPORT \
__attribute__((visibility("default"))) __attribute((weak))
#else
#define XLLM_CAPI_EXPORT __attribute__((visibility("default")))
#endif // XLLM_CAPI_WEAK
} // namespace xllm

View File

@@ -0,0 +1,318 @@
/* Copyright 2025 The xLLM 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.
==============================================================================*/
#pragma once
#include <cstdint>
#include <functional>
#include <optional>
#include <string>
#include <vector>
#include "macros.h"
namespace xllm {
struct XLLM_CAPI_EXPORT XLLM_ChatMessage {
// The role of the messages author. One of "system", "user", "assistant".
std::string role;
// The content of the message. null for assistant messages with function
// calls.
std::string content;
};
struct XLLM_CAPI_EXPORT XLLM_InitLLMOptions {
bool enable_chunked_prefill = false;
// Whether to enable prefill-only sequence parallel
bool enable_prefill_sp = false;
bool enable_prefix_cache = false;
// Whether to enable disaggregated prefill and decode execution
bool enable_disagg_pd = false;
// Whether to enable online-offline co-location in disaggregated PD mode
bool enable_pd_ooc = false;
// Whether to enable schedule overlap
bool enable_schedule_overlap = false;
// Whether to disable TTFT profiling
bool disable_ttft_profiling = false;
// Whether to enable forward interruption
bool enable_forward_interruption = false;
// Whether to enable shared memory for executing model
bool enable_shm = false;
// Input shared memory size (MB)
uint64_t input_shm_size = 1024;
// Output shared memory size (MB)
uint64_t output_shm_size = 128;
bool is_local = true;
// The KVCacheTranfer listen port
int transfer_listen_port = 26000;
// The number of multi-nodes
int nnodes = 1;
// The node rank
int node_rank = 0;
// Data parallel size for MLA attention
int dp_size = 1;
// Expert parallel size for MoE model
int ep_size = 1;
// Number of slots per kv cache block. Default is 128
int block_size = 128;
// Max gpu memory size for kv cache. Default is 0, which means cache size is
// caculated by available memory
int max_cache_size = 0;
// Max number of tokens per batch
int max_tokens_per_batch = 20480;
// Max number of sequences per batch
int max_seqs_per_batch = 256;
// Max number of token per chunk in prefill stage
int max_tokens_per_chunk_for_prefill = -1;
// Number of speculative tokens
int num_speculative_tokens = 0;
// Number of threads for handling input requests
int num_request_handling_threads = 4;
// Expert parallel degree
int expert_parallel_degree = 0;
// The fraction of GPU memory to be used for model inference, including model
// weights and kv cache
float max_memory_utilization = 0.9;
// The task to use the model for(e.g. generate, embed)
std::string task = "generate";
// NPU communication backend.(e.g. lccl, hccl). When enable dp, use hccl
std::string communication_backend = "lccl";
// ATB HCCL rank table file
std::string rank_tablefile = "";
// The role of instance(e.g. DEFAULT, PREFILL, DECODE, MIX)
std::string instance_role = "DEFAULT";
std::string device_ip = "";
// The master address for multi-node distributed serving(e.g. 10.18.1.1:9999)
std::string master_node_addr = "127.0.0.1:18899";
std::string instance_name = "";
// The mode of kv cache transfer(e.g. PUSH, PULL)
std::string kv_cache_transfer_mode = "PUSH";
std::string log_dir;
// draft hf model path to the model file
std::optional<std::string> draft_model = std::nullopt;
// Devices to run the draft model on, e.g. npu:0, npu:0,npu:1
std::optional<std::string> draft_devices = std::nullopt;
// Index ID for internal server ID, which must be set different values
// if the model supports multiple version or there are multiple models.
int64_t server_idx = 0;
};
struct XLLM_CAPI_EXPORT XLLM_RequestParams {
// Whether to include the original prompt in the response. default = false
bool echo = false;
// whether is a offline request. default = false.
bool offline = false;
// Whether to return the log probabilities of the tokens. default = false.
bool logprobs = false;
// Whether to ignore the end of sequence token. default = false.
bool ignore_eos = false;
// Whether to skip special tokens in the output text. default = true.
bool skip_special_tokens = true;
// Number of completions to return for each prompt. default = 1
uint32_t n = 1;
// Number of tokens to generate
// the prompt token count + max_tokens can't exceed the model's max context
// length.
uint32_t max_tokens = 5120;
// Number of sequences to generate for each prompt and select n best among.
std::optional<uint32_t> best_of;
int32_t ttlt_slo_ms;
int32_t ttft_slo_ms;
int32_t tpot_slo_ms;
int32_t beam_width = 0;
int32_t num_return_sequences = 0;
// Number of top log probabilities to return. default = 0.
int64_t top_logprobs = 0;
// top_k sampling cutoff. default = -1 to disable.
int64_t top_k = -1;
// top_p sampling cutoff, between [0.0, 1.0]. default = 1.0
float top_p = 1.0;
// frequency penalty to reduce the likelihood of generating the same word
// multiple times. values between [0.0, 2.0]. 0.0 means no penalty. default =
// 0.0 Positive values penalize new tokens based on their existing frequency
// in the text.
float frequency_penalty = 0.0;
// presence penalty to reduce the likelihood of generating words already in
// the prompt. values between [-2.0, 2.0]. Positive values penalize new tokens
// based on their existing in the prompt. default = 0.0
float presence_penalty = 0.0;
// Repetition penalty to penalize new tokens based on their occurence in the
// text. values > 1.0 encourage the model to use new tokens, while values
// < 1.0 encourage the model to repeat tokens. default = 1.0
float repetition_penalty = 1.0;
// Temperature of the sampling, between [0, 2]. default = 0.0
// higher value will make the ouput more random.
float temperature = 0.0;
// The list of strings to stop generating further tokens.
// the output will contain the stop string.
std::optional<std::vector<std::string>> stop;
// The list of token ids to stop generating further tokens.
std::optional<std::vector<int32_t>> stop_token_ids;
};
enum XLLM_CAPI_EXPORT XLLM_StatusCode {
kSuccess = 0, // Request succeeded
kNotInitialized = 1, // LLM instance not initialized
kModelNotFound = 2, // Specified model ID not loaded
kTimeout = 3, // Request timed out
kInvalidRequest = 4, // Invalid input parameters
kInternalError = 5, // Internal system error
};
struct XLLM_CAPI_EXPORT XLLM_Usage {
// The number of tokens in the prompt.
int32_t prompt_tokens;
// The number of tokens in the generated completion.
int32_t completion_tokens;
// The total number of tokens used in the request (prompt + completion).
int32_t total_tokens;
};
struct XLLM_CAPI_EXPORT XLLM_LogProbData {
// Token
std::string token;
// Token id.
int32_t token_id;
// Log probability of the token.
float logprob;
};
struct XLLM_CAPI_EXPORT XLLM_LogProb {
// Token
std::string token;
// Token id.
int32_t token_id;
// Log probability of the token.
float logprob;
// Log probability of top tokens.
std::vector<XLLM_LogProbData> top_logprobs;
};
struct XLLM_CAPI_EXPORT XLLM_Choice {
// The index of the generated completion
uint32_t index;
// The generated text for completions inference
std::optional<std::string> text;
// The generated item for rec inference
std::optional<uint64_t> item_id;
// The generated message for chatcompletions inference
std::optional<XLLM_ChatMessage> message;
// The log probabilities of output tokens.
std::optional<std::vector<XLLM_LogProb>> logprobs;
// The reason of the model stoped generating tokens.
// "stop" - the model hit a natural stop point or a provided stop sequence.
// "length" - the maximum number of tokens specified in the request was
// reached. "function_call" - the model called a function.
std::string finish_reason;
};
struct XLLM_CAPI_EXPORT XLLM_Response {
// Return code indicating request status (0 = success, non-zero = error)
XLLM_StatusCode status_code = XLLM_StatusCode::kSuccess;
// Optional error details (populated if status_code != kSuccess)
std::string error_info;
// Unique id for the completion request
std::string id;
// The object type, which is always "text_completion".
std::string object;
// The unix timestamp (in seconds) of when the completion was created.
int64_t created;
// The model used for the completion
std::string model;
// List of generated completion choices for the input prompt
std::vector<XLLM_Choice> choices;
// Usage statistics for the completion request.
XLLM_Usage usage;
};
} // namespace xllm

View File

@@ -0,0 +1 @@
"""Compiler-side utilities for xLLM build and AOT flows."""

View File

@@ -0,0 +1 @@
"""TileLang AOT compiler support for xLLM."""

View File

@@ -0,0 +1,13 @@
from __future__ import annotations
from .tilelang_ascend_install import (
PREPARE_ASCEND_COMMAND,
ensure_ascend_ready,
prepare_ascend,
)
__all__ = [
"PREPARE_ASCEND_COMMAND",
"ensure_ascend_ready",
"prepare_ascend",
]

View File

@@ -0,0 +1,75 @@
from __future__ import annotations
import argparse
from pathlib import Path
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Compile xLLM TileLang kernels and emit manifests."
)
parser.add_argument(
"--target",
required=True,
choices=["ascend", "cuda"],
help="Compilation target backend.",
)
parser.add_argument(
"--output-root",
required=True,
help="Output root for compiled TileLang artifacts.",
)
parser.add_argument(
"--device",
choices=["a2", "a3"],
default=None,
help="Ascend device type used to resolve build-time toolchain settings.",
)
parser.add_argument(
"--kernels",
nargs="*",
default=None,
help="Optional kernel names. Compile all registered kernels when omitted.",
)
parser.add_argument(
"--force",
action="store_true",
help="Force recompilation even when cache is hit.",
)
return parser.parse_args(argv)
def main(argv: list[str] | None = None) -> None:
args = parse_args(argv)
output_root = Path(args.output_root).resolve()
output_root.mkdir(parents=True, exist_ok=True)
if args.target == "ascend":
from ..bootstrap import prepare_ascend
prepare_ascend()
from ..targets.ascend.build import build_kernels
manifests = build_kernels(
output_root=output_root,
kernel_names=args.kernels,
force=args.force,
device=args.device,
)
elif args.target == "cuda":
from ..targets.cuda.build import build_kernels
manifests = build_kernels(
output_root=output_root,
kernel_names=args.kernels,
force=args.force,
)
else:
raise ValueError(f"Unsupported target: {args.target}")
for manifest in manifests:
print(f"[INFO] built {manifest.target}:{manifest.kernel_name}")
print(f"[INFO] manifest: {Path(manifest.output_dir) / 'manifest.json'}")
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,29 @@
from __future__ import annotations
import argparse
from ..bootstrap import PREPARE_ASCEND_COMMAND, prepare_ascend
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Prepare third_party/tilelang-ascend for xLLM Ascend TileLang builds."
)
parser.add_argument(
"--force",
action="store_true",
help="Force rerunning install_ascend.sh even when cached artifacts look ready.",
)
return parser.parse_args(argv)
def main(argv: list[str] | None = None) -> None:
args = parse_args(argv)
tilelang_root = prepare_ascend(force=args.force)
print(f"[INFO] tilelang-ascend is ready under: {tilelang_root}")
print("[INFO] Next step: run your usual `python setup.py build ...` or `python setup.py test ...` command.")
print(f"[INFO] Re-run this step explicitly with: `{PREPARE_ASCEND_COMMAND}`")
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,50 @@
from __future__ import annotations
import hashlib
import json
from pathlib import Path
from typing import Any
from .manifest import KernelFamilyManifest
from .spec import KernelCompileSpec
from .toolchain import sha256_file
def compute_cache_key(
spec: KernelCompileSpec,
fingerprint: dict[str, Any],
dependency_files: list[str | Path],
) -> str:
payload = {
"spec": spec.cache_key_material(),
"fingerprint": fingerprint,
"dependencies": {
str(Path(path).resolve()): sha256_file(path) for path in dependency_files
},
}
encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode(
"utf-8"
)
return hashlib.sha256(encoded).hexdigest()
def is_cache_hit(
manifest_path: str | Path, variant_key: str, expected_cache_key: str
) -> bool:
path = Path(manifest_path)
if not path.is_file():
return False
try:
manifest = KernelFamilyManifest.read(path)
except Exception:
return False
variant = manifest.get_variant(variant_key)
if variant is None:
return False
if variant.cache_key != expected_cache_key:
return False
return Path(variant.generated_source).is_file() and Path(variant.compiled_binary).is_file()

View File

@@ -0,0 +1,106 @@
from __future__ import annotations
import json
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import Any
from .spec import DispatchField
@dataclass
class KernelAbiParameter:
cpp_type: str
name: str
@dataclass
class KernelAbi:
return_type: str
parameters: list[KernelAbiParameter] = field(default_factory=list)
def to_json_dict(self) -> dict[str, Any]:
return asdict(self)
@dataclass
class KernelVariantManifest:
variant_key: str
specialization: dict[str, Any]
generated_source: str
compiled_binary: str
entry_symbol: str
cache_key: str
dispatch_values: dict[str, Any] = field(default_factory=dict)
toolchain_options: dict[str, Any] = field(default_factory=dict)
fingerprint: dict[str, Any] = field(default_factory=dict)
compile_definitions: list[str] = field(default_factory=list)
def to_json_dict(self) -> dict[str, Any]:
return asdict(self)
@dataclass
class KernelFamilyManifest:
target: str
kernel_name: str
output_dir: str
variants_inc: str
registry_inc: str = ""
dispatch_schema: list[DispatchField] = field(default_factory=list)
kernel_abi: KernelAbi | None = None
variants: list[KernelVariantManifest] = field(default_factory=list)
schema_version: int = 2
def to_json_dict(self) -> dict[str, Any]:
data = asdict(self)
data["dispatch_schema"] = [asdict(field) for field in self.dispatch_schema]
data["kernel_abi"] = (
None if self.kernel_abi is None else self.kernel_abi.to_json_dict()
)
data["variants"] = [variant.to_json_dict() for variant in self.variants]
return data
@property
def manifest_path(self) -> Path:
return Path(self.output_dir) / "manifest.json"
def write(self, path: str | Path) -> None:
output = Path(path)
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(
json.dumps(self.to_json_dict(), indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
@classmethod
def read(cls, path: str | Path) -> "KernelFamilyManifest":
data = json.loads(Path(path).read_text(encoding="utf-8"))
dispatch_schema = [
DispatchField(**field) for field in data.pop("dispatch_schema", [])
]
kernel_abi_data = data.pop("kernel_abi", None)
kernel_abi = None
if kernel_abi_data is not None:
kernel_abi = KernelAbi(
return_type=kernel_abi_data["return_type"],
parameters=[
KernelAbiParameter(**param)
for param in kernel_abi_data.get("parameters", [])
],
)
variants = [
KernelVariantManifest(**variant) for variant in data.pop("variants", [])
]
return cls(
dispatch_schema=dispatch_schema,
kernel_abi=kernel_abi,
variants=variants,
**data,
)
def get_variant(self, variant_key: str) -> KernelVariantManifest | None:
for variant in self.variants:
if variant.variant_key == variant_key:
return variant
return None

View File

@@ -0,0 +1,268 @@
from __future__ import annotations
import re
from dataclasses import dataclass, field
from typing import Any
_REGISTER_KERNEL_ATTR = "__xllm_tilelang_registered_kernel__"
_ENTRY_SYMBOL_CONTEXT_KEY = "entry_symbol"
_C_IDENTIFIER_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
_SAFE_VARIANT_KEY_PATTERN = re.compile(r"^[A-Za-z0-9_]+$")
_SUPPORTED_DISPATCH_FIELD_KINDS = frozenset({"int32", "dtype"})
_SPECIALIZATION_CONFIG_KEYS = frozenset(
{
"variant_key",
"specialization",
"compile_definitions",
"kernel_name",
"target",
"entry_name",
"source_entry_symbol",
}
)
@dataclass(frozen=True)
class DispatchField:
name: str
kind: str
def validate(self) -> None:
if not _C_IDENTIFIER_PATTERN.match(self.name):
raise ValueError(
"DispatchField.name must be a valid C/C++ identifier"
)
if self.kind not in _SUPPORTED_DISPATCH_FIELD_KINDS:
supported = ", ".join(sorted(_SUPPORTED_DISPATCH_FIELD_KINDS))
raise ValueError(
f"DispatchField.kind must be one of: {supported}"
)
@dataclass(frozen=True)
class KernelCompileSpec:
target: str
kernel_name: str
module_name: str
variant_key: str
specialization: dict[str, Any] = field(default_factory=dict)
dispatch_values: dict[str, Any] = field(default_factory=dict)
entry_name: str | None = None
source_entry_symbol: str = "call"
def cache_key_material(self) -> dict[str, Any]:
return {
"target": self.target,
"kernel_name": self.kernel_name,
"module_name": self.module_name,
"variant_key": self.variant_key,
"specialization": self.specialization,
"dispatch_values": self.dispatch_values,
"entry_name": self.entry_name,
"source_entry_symbol": self.source_entry_symbol,
}
@dataclass(frozen=True)
class KernelSpec:
variant_key: str
specialization: dict[str, Any] = field(default_factory=dict)
compile_definitions: dict[str, str] = field(default_factory=dict)
kernel_name: str | None = None
target: str = "ascend"
entry_name: str | None = None
source_entry_symbol: str = "call"
def validate(self) -> None:
if not self.variant_key:
raise ValueError("KernelSpec.variant_key must not be empty")
if not _SAFE_VARIANT_KEY_PATTERN.match(self.variant_key):
raise ValueError(
"KernelSpec.variant_key must contain only letters, digits, or "
"underscore"
)
if not self.specialization:
raise ValueError("KernelSpec.specialization must not be empty")
if self.kernel_name is not None and not _C_IDENTIFIER_PATTERN.match(
self.kernel_name
):
raise ValueError(
"KernelSpec.kernel_name must be a valid C/C++ identifier"
)
if self.entry_name is not None and not _C_IDENTIFIER_PATTERN.match(
self.entry_name
):
raise ValueError("KernelSpec.entry_name must be a valid C/C++ identifier")
context_keys = set(self.specialization)
context_keys.add(_ENTRY_SYMBOL_CONTEXT_KEY)
uses_entry_symbol = False
for macro_name, context_key in self.compile_definitions.items():
if context_key not in context_keys:
raise KeyError(
"KernelSpec.compile_definitions references unknown context "
f"key {context_key!r} for macro {macro_name!r}"
)
uses_entry_symbol = (
uses_entry_symbol or context_key == _ENTRY_SYMBOL_CONTEXT_KEY
)
if uses_entry_symbol and not self.entry_name:
raise ValueError(
"KernelSpec.entry_name is required when "
"compile_definitions references 'entry_symbol'"
)
def to_compile_spec(
self, *, module_name: str, dispatch_schema: list[DispatchField]
) -> KernelCompileSpec:
dispatch_values = {
field.name: self.specialization[field.name] for field in dispatch_schema
}
return KernelCompileSpec(
target=self.target,
kernel_name=self.kernel_name or module_name,
module_name=module_name,
variant_key=self.variant_key,
specialization=dict(self.specialization),
dispatch_values=dispatch_values,
entry_name=self.entry_name,
source_entry_symbol=self.source_entry_symbol,
)
def render_compile_definitions(self, *, entry_symbol: str) -> list[str]:
self.validate()
context = dict(self.specialization)
context[_ENTRY_SYMBOL_CONTEXT_KEY] = entry_symbol
definitions: list[str] = []
for macro_name, context_key in self.compile_definitions.items():
if context_key not in context:
raise KeyError(
"compile_definitions references unknown context key "
f"{context_key!r} for macro {macro_name!r}"
)
definitions.append(f"{macro_name}={context[context_key]}")
return definitions
class TilelangKernel:
"""Marker base class for TileLang kernel generator classes."""
TARGET = "ascend"
KERNEL_NAME: str | None = None
ENTRY_NAME: str | None = None
COMPILE_DEFINITIONS: dict[str, str] = {}
SOURCE_ENTRY_SYMBOL = "call"
DISPATCH_SCHEMA: list[DispatchField] = []
SPECIALIZATIONS: list[dict[str, Any] | KernelSpec] = []
@classmethod
def dispatch_schema(cls) -> list[DispatchField]:
if not cls.DISPATCH_SCHEMA:
raise NotImplementedError(
f"{cls.__name__} must define non-empty DISPATCH_SCHEMA"
)
normalized: list[DispatchField] = []
seen_names: set[str] = set()
for index, field in enumerate(cls.DISPATCH_SCHEMA):
if not isinstance(field, DispatchField):
raise TypeError(
f"{cls.__name__}.DISPATCH_SCHEMA[{index}] must be DispatchField, "
f"got {type(field).__name__}"
)
field.validate()
if field.name in seen_names:
raise ValueError(
f"{cls.__name__}.DISPATCH_SCHEMA contains duplicate field "
f"{field.name!r}"
)
seen_names.add(field.name)
normalized.append(field)
return normalized
@classmethod
def specs(cls) -> list[KernelSpec]:
if not cls.SPECIALIZATIONS:
raise NotImplementedError(
f"{cls.__name__} must define non-empty SPECIALIZATIONS or override "
"specs()"
)
return [
cls._specialization_to_spec(specialization)
for specialization in cls.SPECIALIZATIONS
]
@classmethod
def _specialization_to_spec(
cls, specialization: dict[str, Any] | KernelSpec
) -> KernelSpec:
if isinstance(specialization, KernelSpec):
return specialization
if not isinstance(specialization, dict):
raise TypeError(
f"{cls.__name__}.SPECIALIZATIONS entries must be dict or KernelSpec, "
f"got {type(specialization).__name__}"
)
specialization_data = dict(specialization)
specialization_fields = specialization_data.pop("specialization", None)
if specialization_fields is None:
specialization_fields = {
key: value
for key, value in specialization_data.items()
if key not in _SPECIALIZATION_CONFIG_KEYS
}
for key in specialization_fields:
specialization_data.pop(key)
elif not isinstance(specialization_fields, dict):
raise TypeError(
f"{cls.__name__}.SPECIALIZATIONS specialization must be dict, got "
f"{type(specialization_fields).__name__}"
)
compile_definitions = dict(cls.COMPILE_DEFINITIONS)
compile_definitions.update(
dict(specialization_data.pop("compile_definitions", {}))
)
variant_key = specialization_data.pop("variant_key", None)
if not isinstance(variant_key, str) or not variant_key:
raise ValueError(
f"{cls.__name__}.SPECIALIZATIONS entries must define non-empty "
"'variant_key'"
)
spec = KernelSpec(
variant_key=variant_key,
specialization=dict(specialization_fields),
compile_definitions=compile_definitions,
kernel_name=specialization_data.pop("kernel_name", cls.KERNEL_NAME),
target=specialization_data.pop("target", cls.TARGET),
entry_name=specialization_data.pop("entry_name", cls.ENTRY_NAME),
source_entry_symbol=specialization_data.pop(
"source_entry_symbol", cls.SOURCE_ENTRY_SYMBOL
),
)
if specialization_data:
unknown_keys = ", ".join(sorted(specialization_data))
raise KeyError(
f"{cls.__name__}.SPECIALIZATIONS contains unsupported config keys: "
f"{unknown_keys}"
)
return spec
def register_kernel(cls: type[TilelangKernel]) -> type[TilelangKernel]:
setattr(cls, _REGISTER_KERNEL_ATTR, True)
return cls
def is_registered_kernel_class(obj: object) -> bool:
return isinstance(obj, type) and bool(obj.__dict__.get(_REGISTER_KERNEL_ATTR, False))

View File

@@ -0,0 +1,88 @@
from __future__ import annotations
import hashlib
import os
import shutil
import subprocess
import sys
from pathlib import Path
from typing import Sequence
def repo_root() -> Path:
return Path(__file__).resolve().parents[4]
def default_tilelang_root() -> Path:
return repo_root() / "third_party" / "tilelang-ascend"
def resolve_tilelang_root() -> Path:
value = os.environ.get("TL_ROOT", "").strip()
if value:
return Path(value).resolve()
return default_tilelang_root().resolve()
def require_env(name: str) -> str:
value = os.environ.get(name, "").strip()
if not value:
raise RuntimeError(f"Required environment variable is not set: {name}")
return value
def prepend_pythonpath(env: dict[str, str], path: str) -> None:
current = env.get("PYTHONPATH", "")
items = [item for item in current.split(os.pathsep) if item]
items = [item for item in items if item != path]
items.insert(0, path)
env["PYTHONPATH"] = os.pathsep.join(items)
def prepare_tilelang_import(tilelang_root: str | Path | None = None) -> Path:
tl_root = (
Path(tilelang_root).resolve() if tilelang_root is not None else resolve_tilelang_root()
)
os.environ["TL_ROOT"] = str(tl_root)
prepend_pythonpath(os.environ, str(tl_root))
tl_root_str = str(tl_root)
# Keep TL_ROOT at sys.path front to avoid resolving the sibling
# package xllm/compiler/tilelang as top-level `tilelang`.
sys.path = [p for p in sys.path if p != tl_root_str]
sys.path.insert(0, tl_root_str)
os.environ.setdefault("ACL_OP_INIT_MODE", "1")
return tl_root
def run_checked(
cmd: Sequence[str],
*,
cwd: str | Path | None = None,
env: dict[str, str] | None = None,
) -> None:
subprocess.check_call(list(cmd), cwd=cwd, env=env)
def sha256_file(path: str | Path) -> str:
data = Path(path).read_bytes()
return hashlib.sha256(data).hexdigest()
def git_head(path: str | Path) -> str:
repo_path = str(Path(path).resolve())
result = subprocess.run(
["git", "-c", f"safe.directory={repo_path}", "-C", repo_path, "rev-parse", "HEAD"],
text=True,
capture_output=True,
check=False,
)
if result.returncode != 0:
return ""
return result.stdout.strip()
def find_required_executable(name: str) -> str:
executable = shutil.which(name)
if not executable:
raise RuntimeError(f"Required executable was not found in PATH: {name}")
return executable

View File

@@ -0,0 +1,23 @@
diff --git a/requirements-build.txt b/requirements-build.txt
--- a/requirements-build.txt
+++ b/requirements-build.txt
@@ -1,6 +1,5 @@
# Should be mirrored in pyproject.toml
build
-cmake>=3.26
packaging
setuptools>=61
torch
diff --git a/install_ascend.sh b/install_ascend.sh
--- a/install_ascend.sh
+++ b/install_ascend.sh
@@ -140,8 +140,8 @@ echo "Building TileLang with make..."
# Other wise, make will use all available cores
# and it may cause the system to be unresponsive
CORES=$(nproc)
MAKE_JOBS=$(( CORES * 50 / 100 ))
-make -j${MAKE_JOBS}
+make -j
if [ $? -ne 0 ]; then
echo "Error: TileLang build failed."

View File

@@ -0,0 +1,76 @@
from __future__ import annotations
import re
from ...common.manifest import KernelAbi, KernelAbiParameter
def rename_entry_symbol(source: str, source_entry_symbol: str, entry_symbol: str) -> str:
pattern = rf"\b{re.escape(source_entry_symbol)}\b"
return re.sub(pattern, entry_symbol, source)
def rename_variant_internal_symbols(source: str, variant_key: str) -> str:
symbol_names: set[str] = set()
symbol_names.update(
re.findall(
r'extern\s+"C"\s+__global__\s+__aicore__\s+void\s+([A-Za-z_][A-Za-z0-9_]*)\s*\(',
source,
)
)
symbol_names.update(
re.findall(r"\bvoid\s+([A-Za-z_][A-Za-z0-9_]*_tiling)\s*\(", source)
)
renamed_source = source
for symbol_name in sorted(symbol_names, key=len, reverse=True):
renamed_source = re.sub(
rf"\b{re.escape(symbol_name)}\b",
f"{symbol_name}__{variant_key}",
renamed_source,
)
return renamed_source
def normalize_cpp_type(cpp_type: str) -> str:
normalized = re.sub(r"\s+", " ", cpp_type).strip()
normalized = re.sub(r"\s*([*&]+)\s*", r"\1", normalized)
return normalized
def parse_kernel_abi(source: str, entry_symbol: str) -> KernelAbi:
pattern = re.compile(
rf'extern\s+"C"\s+'
rf"(?P<return_type>[^(){{}};]+?)\s+"
rf"{re.escape(entry_symbol)}\s*\("
r"(?P<params>[^)]*)\)\s*\{",
re.MULTILINE,
)
match = pattern.search(source)
if match is None:
raise ValueError(
f"Failed to parse exported entry ABI for symbol {entry_symbol!r}"
)
return_type = normalize_cpp_type(match.group("return_type"))
params_text = match.group("params").strip()
parameters: list[KernelAbiParameter] = []
if params_text and params_text != "void":
for param in (part.strip() for part in params_text.split(",")):
parsed = re.match(
r"(?P<type>.+?[\*&]?)\s*(?P<name>[A-Za-z_][A-Za-z0-9_]*)$",
param,
)
if parsed is None:
raise ValueError(
"Failed to parse kernel ABI parameter "
f"{param!r} for symbol {entry_symbol!r}"
)
parameters.append(
KernelAbiParameter(
cpp_type=normalize_cpp_type(parsed.group("type")),
name=parsed.group("name"),
)
)
return KernelAbi(return_type=return_type, parameters=parameters)

View File

@@ -0,0 +1,53 @@
import importlib
import os
import pkgutil
import re
from dataclasses import dataclass
from pathlib import Path
from ...common.manifest import KernelFamilyManifest
from ...common.toolchain import find_required_executable
from .kernel_family_builder import build_kernel_family as _build_kernel_family
from .kernel_registry import RegisteredKernelFamily, get_default_families
from .toolchain import resolve_build_context
def build_kernel_family(
family: RegisteredKernelFamily,
output_root: str | Path,
force: bool = False,
device: str | None = None,
) -> KernelFamilyManifest:
context = resolve_build_context(
device=device,
bisheng_executable=find_required_executable("bisheng"),
)
return _build_kernel_family(
family,
output_root=output_root,
context=context,
force=force,
)
def build_kernels(
output_root: str | Path,
kernel_names: list[str] | None = None,
force: bool = False,
device: str | None = None,
) -> list[KernelFamilyManifest]:
context = resolve_build_context(
device=device,
bisheng_executable=find_required_executable("bisheng"),
)
manifests = []
for family in get_default_families(kernel_names):
manifests.append(
_build_kernel_family(
family,
output_root=output_root,
context=context,
force=force,
)
)
return manifests

View File

@@ -0,0 +1,368 @@
from __future__ import annotations
from concurrent.futures import ProcessPoolExecutor, as_completed
from dataclasses import dataclass
import importlib
import os
from pathlib import Path
from ...common.cache import compute_cache_key, is_cache_hit
from ...common.manifest import KernelAbi, KernelFamilyManifest, KernelVariantManifest
from ...common.spec import DispatchField, KernelCompileSpec, KernelSpec, TilelangKernel
from ...common.toolchain import repo_root, run_checked
from . import abi_entry, kernel_registry, toolchain
from .kernel_registry import RegisteredKernelFamily
from .kernels import utils as kernel_utils
from .kernels.utils import render_family_registry_inc, render_family_variants_inc
from .toolchain import AscendBuildContext, TILELANG_BISHENG_COMMON_FLAGS
@dataclass(frozen=True)
class _VariantBuildPlan:
compile_spec: KernelCompileSpec
kernel_spec: KernelSpec
generated_source: Path
compiled_binary: Path
cache_key: str
@dataclass(frozen=True)
class _VariantBuildResult:
manifest: KernelVariantManifest
kernel_abi: KernelAbi
@dataclass(frozen=True)
class _VariantWorkerArgs:
kernel_cls_module: str
kernel_cls_name: str
plan: _VariantBuildPlan
entry_symbol: str
bisheng_executable: str
bisheng_arch: str
include_dirs: tuple[str, ...]
toolchain_options: dict
fingerprint: dict
compile_cwd: str
def _variant_entry_symbol(spec: KernelCompileSpec) -> str:
kernel_entry_name = spec.entry_name or spec.kernel_name
return f"{kernel_entry_name}__{spec.variant_key}_call"
def _run_variant_worker(args: _VariantWorkerArgs) -> _VariantBuildResult:
mod = importlib.import_module(args.kernel_cls_module)
kernel_cls = getattr(mod, args.kernel_cls_name)
plan = args.plan
compile_spec = plan.compile_spec
kernel_spec = plan.kernel_spec
source = kernel_cls.generate_source(**compile_spec.specialization)
rendered_source = abi_entry.rename_variant_internal_symbols(
abi_entry.rename_entry_symbol(
source, compile_spec.source_entry_symbol, args.entry_symbol
),
compile_spec.variant_key,
)
kernel_abi = abi_entry.parse_kernel_abi(rendered_source, args.entry_symbol)
plan.generated_source.write_text(rendered_source, encoding="utf-8")
compile_cmd = [
args.bisheng_executable,
f"--npu-arch={args.bisheng_arch}",
*TILELANG_BISHENG_COMMON_FLAGS,
f"-Dg_tilingKey=g_tilingKey__{compile_spec.variant_key}",
*[f"-I{d}" for d in args.include_dirs],
str(plan.generated_source),
"-c",
"-o",
str(plan.compiled_binary),
]
run_checked(compile_cmd, cwd=args.compile_cwd)
manifest = KernelVariantManifest(
variant_key=compile_spec.variant_key,
specialization=dict(compile_spec.specialization),
dispatch_values=dict(compile_spec.dispatch_values),
generated_source=str(plan.generated_source),
compiled_binary=str(plan.compiled_binary),
entry_symbol=args.entry_symbol,
cache_key=plan.cache_key,
toolchain_options=dict(args.toolchain_options),
fingerprint=dict(args.fingerprint),
compile_definitions=kernel_spec.render_compile_definitions(
entry_symbol=args.entry_symbol
),
)
return _VariantBuildResult(manifest=manifest, kernel_abi=kernel_abi)
def _read_family_manifest(path: Path) -> KernelFamilyManifest | None:
if not path.is_file():
return None
try:
return KernelFamilyManifest.read(path)
except Exception:
return None
def _render_variants_inc(
kernel_name: str,
kernel_cls: type[TilelangKernel],
dispatch_schema: list[DispatchField],
variants: list[KernelVariantManifest],
) -> str:
renderer = getattr(kernel_cls, "render_variants_inc", None)
if renderer is None:
return render_family_variants_inc(
kernel_name=kernel_name,
dispatch_schema=dispatch_schema,
variants=variants,
)
if not callable(renderer):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' defines "
"non-callable render_variants_inc"
)
rendered = renderer(variants, dispatch_schema)
if not isinstance(rendered, str):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' "
"render_variants_inc(...) must return str"
)
return rendered
def _render_registry_inc(
kernel_name: str,
kernel_cls: type[TilelangKernel],
dispatch_schema: list[DispatchField],
kernel_abi: KernelAbi,
variants: list[KernelVariantManifest],
) -> str:
renderer = getattr(kernel_cls, "render_registry_inc", None)
if renderer is None:
return render_family_registry_inc(
kernel_name=kernel_name,
dispatch_schema=dispatch_schema,
kernel_abi=kernel_abi,
variants=variants,
)
if not callable(renderer):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' defines "
"non-callable render_registry_inc"
)
rendered = renderer(variants, dispatch_schema, kernel_abi)
if not isinstance(rendered, str):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' "
"render_registry_inc(...) must return str"
)
return rendered
def _build_dependency_files(family: RegisteredKernelFamily) -> list[Path]:
# Keep cache invalidation aligned with the split builder implementation.
files = [
Path(family.module.__file__).resolve(),
Path(__file__).resolve(),
Path(toolchain.__file__).resolve(),
Path(kernel_registry.__file__).resolve(),
Path(abi_entry.__file__).resolve(),
Path(kernel_utils.__file__).resolve(),
Path(__file__).resolve().with_name("build.py"),
]
deduped: list[Path] = []
seen: set[Path] = set()
for path in files:
if path in seen:
continue
seen.add(path)
deduped.append(path)
return deduped
def build_kernel_family(
family: RegisteredKernelFamily,
output_root: str | Path,
context: AscendBuildContext,
force: bool = False,
) -> KernelFamilyManifest:
family_output_dir = Path(output_root) / "targets" / "ascend" / family.kernel_name
family_output_dir.mkdir(parents=True, exist_ok=True)
manifest_path = family_output_dir / "manifest.json"
existing_manifest = _read_family_manifest(manifest_path)
dependency_files = _build_dependency_files(family)
variant_manifest_by_key: dict[str, KernelVariantManifest] = {}
uncached_plans: list[_VariantBuildPlan] = []
family_kernel_abi: KernelAbi | None = None
for compile_spec, kernel_spec in family.spec_pairs:
if compile_spec.target != "ascend":
raise ValueError(
f"Unsupported target for Ascend build.py: {compile_spec.target}"
)
variant_output_dir = family_output_dir / compile_spec.variant_key
variant_output_dir.mkdir(parents=True, exist_ok=True)
generated_source = (
variant_output_dir
/ f"{compile_spec.kernel_name}_{compile_spec.variant_key}_kernel.cpp"
)
compiled_binary = (
variant_output_dir
/ f"{compile_spec.kernel_name}_{compile_spec.variant_key}_kernel.o"
)
cache_key = compute_cache_key(
compile_spec,
context.fingerprint,
dependency_files,
)
cached_variant = (
existing_manifest.get_variant(compile_spec.variant_key)
if existing_manifest is not None
else None
)
if (
not force
and cached_variant is not None
and Path(cached_variant.generated_source).is_file()
and Path(cached_variant.compiled_binary).is_file()
and is_cache_hit(manifest_path, compile_spec.variant_key, cache_key)
):
cached_source = Path(cached_variant.generated_source).read_text(
encoding="utf-8"
)
kernel_abi = abi_entry.parse_kernel_abi(
cached_source, cached_variant.entry_symbol
)
if family_kernel_abi is None:
family_kernel_abi = kernel_abi
elif kernel_abi != family_kernel_abi:
raise ValueError(
"All variants in a TileLang kernel must share the same exported "
f"C ABI. Mismatch found in variant {compile_spec.variant_key!r}."
)
variant_manifest_by_key[compile_spec.variant_key] = KernelVariantManifest(
variant_key=compile_spec.variant_key,
specialization=dict(compile_spec.specialization),
dispatch_values=dict(compile_spec.dispatch_values),
generated_source=cached_variant.generated_source,
compiled_binary=cached_variant.compiled_binary,
entry_symbol=cached_variant.entry_symbol,
cache_key=cached_variant.cache_key,
toolchain_options=dict(context.toolchain_options),
fingerprint=dict(context.fingerprint),
compile_definitions=kernel_spec.render_compile_definitions(
entry_symbol=cached_variant.entry_symbol
),
)
continue
uncached_plans.append(
_VariantBuildPlan(
compile_spec=compile_spec,
kernel_spec=kernel_spec,
generated_source=generated_source,
compiled_binary=compiled_binary,
cache_key=cache_key,
)
)
compile_cwd = str(repo_root())
kernel_cls_module = family.kernel_cls.__module__
kernel_cls_name = family.kernel_cls.__name__
worker_args_list: list[_VariantWorkerArgs] = []
for plan in uncached_plans:
worker_args_list.append(
_VariantWorkerArgs(
kernel_cls_module=kernel_cls_module,
kernel_cls_name=kernel_cls_name,
plan=plan,
entry_symbol=_variant_entry_symbol(plan.compile_spec),
bisheng_executable=context.bisheng_executable,
bisheng_arch=context.bisheng_arch,
include_dirs=tuple(str(d) for d in context.include_dirs),
toolchain_options=dict(context.toolchain_options),
fingerprint=dict(context.fingerprint),
compile_cwd=compile_cwd,
)
)
if worker_args_list:
max_workers = max(1, os.cpu_count() or 1)
with ProcessPoolExecutor(max_workers=max_workers) as executor:
future_to_args = {
executor.submit(_run_variant_worker, args): args
for args in worker_args_list
}
for future in as_completed(future_to_args):
args = future_to_args[future]
try:
result = future.result()
except Exception as exc:
raise RuntimeError(
"Ascend variant build failed for variant "
f"{args.plan.compile_spec.variant_key!r}"
) from exc
if family_kernel_abi is None:
family_kernel_abi = result.kernel_abi
elif result.kernel_abi != family_kernel_abi:
raise ValueError(
"All variants in a TileLang kernel must share the same exported "
"C ABI. Mismatch found in variant "
f"{result.manifest.variant_key!r}."
)
variant_manifest_by_key[result.manifest.variant_key] = result.manifest
variant_manifests: list[KernelVariantManifest] = [
variant_manifest_by_key[compile_spec.variant_key]
for compile_spec, _ in family.spec_pairs
]
if family_kernel_abi is None:
raise ValueError(
f"TileLang kernel {family.kernel_name!r} produced no exported kernel ABI"
)
variants_inc_path = family_output_dir / "variants.inc"
variants_inc_path.write_text(
_render_variants_inc(
family.kernel_name,
family.kernel_cls,
family.dispatch_schema,
variant_manifests,
),
encoding="utf-8",
)
registry_inc_path = family_output_dir / "registry.inc"
registry_inc_path.write_text(
_render_registry_inc(
family.kernel_name,
family.kernel_cls,
family.dispatch_schema,
family_kernel_abi,
variant_manifests,
),
encoding="utf-8",
)
manifest = KernelFamilyManifest(
target="ascend",
kernel_name=family.kernel_name,
output_dir=str(family_output_dir),
variants_inc=str(variants_inc_path),
registry_inc=str(registry_inc_path),
dispatch_schema=list(family.dispatch_schema),
kernel_abi=family_kernel_abi,
variants=variant_manifests,
)
manifest.write(manifest_path)
return manifest

View File

@@ -0,0 +1,212 @@
from __future__ import annotations
import importlib
import pkgutil
from dataclasses import dataclass
from pathlib import Path
from types import ModuleType
from ...common.spec import (
DispatchField,
KernelCompileSpec,
KernelSpec,
TilelangKernel,
is_registered_kernel_class,
)
from ...common.toolchain import prepare_tilelang_import
@dataclass(frozen=True)
class RegisteredKernelFamily:
module: ModuleType
kernel_cls: type[TilelangKernel]
module_name: str
kernel_name: str
dispatch_schema: list[DispatchField]
spec_pairs: list[tuple[KernelCompileSpec, KernelSpec]]
def _load_kernel_module(module_name: str) -> ModuleType:
prepare_tilelang_import()
return importlib.import_module(f"{__package__}.kernels.{module_name}")
def _kernels_dir() -> Path:
return Path(__file__).resolve().parent / "kernels"
def _iter_kernel_module_names() -> list[str]:
return sorted(
module.name
for module in pkgutil.iter_modules([str(_kernels_dir())])
if not module.name.startswith("_")
)
def _resolve_registered_kernel_class(
module_name: str,
) -> tuple[ModuleType, type[TilelangKernel] | None]:
module = _load_kernel_module(module_name)
kernel_classes = [
obj
for obj in vars(module).values()
if isinstance(obj, type)
and obj.__module__ == module.__name__
and is_registered_kernel_class(obj)
]
if not kernel_classes:
return module, None
if len(kernel_classes) > 1:
kernel_names = ", ".join(sorted(cls.__name__ for cls in kernel_classes))
raise TypeError(
f"TileLang kernel module {module_name!r} must define at most one "
f"@register_kernel class, found: {kernel_names}"
)
kernel_cls = kernel_classes[0]
if not issubclass(kernel_cls, TilelangKernel):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' must inherit "
"TilelangKernel"
)
return module, kernel_cls
def _load_registered_kernel_family(
module_name: str,
) -> RegisteredKernelFamily | None:
module, kernel_cls = _resolve_registered_kernel_class(module_name)
if kernel_cls is None:
return None
generate_source = kernel_cls.__dict__.get("generate_source")
if not isinstance(generate_source, (staticmethod, classmethod)):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' must define "
"callable generate_source(...) as @staticmethod or @classmethod"
)
resolved_generate_source = getattr(kernel_cls, "generate_source", None)
if not callable(resolved_generate_source):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' must define "
"callable generate_source(...)"
)
resolved_specs = getattr(kernel_cls, "specs", None)
if not callable(resolved_specs):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' must define "
"callable specs() -> list[KernelSpec]"
)
resolved_dispatch_schema = getattr(kernel_cls, "dispatch_schema", None)
if not callable(resolved_dispatch_schema):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' must define "
"callable dispatch_schema() -> list[DispatchField]"
)
try:
kernel_specs = resolved_specs()
except NotImplementedError as exc:
raise TypeError(str(exc)) from exc
if not isinstance(kernel_specs, list) or not kernel_specs:
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' must return a "
"non-empty list[KernelSpec] from specs()"
)
try:
dispatch_schema = resolved_dispatch_schema()
except NotImplementedError as exc:
raise TypeError(str(exc)) from exc
if not isinstance(dispatch_schema, list) or not dispatch_schema:
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' must return a "
"non-empty list[DispatchField] from dispatch_schema()"
)
for index, field in enumerate(dispatch_schema):
if not isinstance(field, DispatchField):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' "
f"dispatch_schema()[{index}] must be DispatchField"
)
family_kernel_name: str | None = None
seen_variant_keys: set[str] = set()
spec_pairs: list[tuple[KernelCompileSpec, KernelSpec]] = []
for index, kernel_spec in enumerate(kernel_specs):
if not isinstance(kernel_spec, KernelSpec):
raise TypeError(
f"registered kernel class '{kernel_cls.__name__}' specs()[{index}] "
"must be KernelSpec"
)
kernel_spec.validate()
missing_dispatch_fields = [
field.name
for field in dispatch_schema
if field.name not in kernel_spec.specialization
]
if missing_dispatch_fields:
raise ValueError(
f"registered kernel class '{kernel_cls.__name__}' specs()[{index}] "
"is missing DISPATCH_SCHEMA fields: "
f"{', '.join(missing_dispatch_fields)}"
)
compile_spec = kernel_spec.to_compile_spec(
module_name=module_name,
dispatch_schema=dispatch_schema,
)
if family_kernel_name is None:
family_kernel_name = compile_spec.kernel_name
elif compile_spec.kernel_name != family_kernel_name:
raise ValueError(
f"registered kernel class '{kernel_cls.__name__}' must return "
"KernelSpec entries with the same kernel_name"
)
if compile_spec.variant_key in seen_variant_keys:
raise ValueError(
f"registered kernel class '{kernel_cls.__name__}' has duplicate "
f"variant_key {compile_spec.variant_key!r}"
)
seen_variant_keys.add(compile_spec.variant_key)
spec_pairs.append((compile_spec, kernel_spec))
assert family_kernel_name is not None
return RegisteredKernelFamily(
module=module,
kernel_cls=kernel_cls,
module_name=module_name,
kernel_name=family_kernel_name,
dispatch_schema=dispatch_schema,
spec_pairs=spec_pairs,
)
def registered_families() -> dict[str, RegisteredKernelFamily]:
families: dict[str, RegisteredKernelFamily] = {}
for module_name in _iter_kernel_module_names():
family = _load_registered_kernel_family(module_name)
if family is None:
continue
if family.kernel_name in families:
raise ValueError(
"Duplicate Ascend TileLang kernel_name registered: "
f"{family.kernel_name}"
)
families[family.kernel_name] = family
return families
def get_default_families(
kernel_names: list[str] | None = None,
) -> list[RegisteredKernelFamily]:
families = registered_families()
if kernel_names is None:
return list(families.values())
missing = [name for name in kernel_names if name not in families]
if missing:
raise ValueError(f"Unknown Ascend TileLang kernels: {', '.join(missing)}")
return [families[name] for name in kernel_names]

View File

@@ -0,0 +1,673 @@
#!/usr/bin/env python3
import argparse
from pathlib import Path
import tilelang
import tilelang.language as T
from compiler.tilelang.targets.ascend.kernels.utils import (
DEFAULT_ASCEND_PASS_CONFIGS,
detect_vec_core_num,
mte2_wait_mte3,
mte3_notify_mte2,
)
from compiler.tilelang.common.spec import (
DispatchField,
TilelangKernel,
register_kernel,
)
DEFAULT_NUM_HEADS = 32
DEFAULT_DTYPE = "bf16"
DEFAULT_MAX_BATCH = 262144
DEFAULT_MAX_HEADS = 128
REF_CHECK_NUM_BATCHES = 16
REF_CHECK_NUM_HEADS = (1, 16, 32, 48, 64, 128)
VEC_NUM = 2
VECTOR_BYTES_PER_ITER = 256
SUPPORTED_NUM_HEADS = (4, 6, 8, 12, 16, 24, 32, 48, 64, 128)
MAX_VEC_CORE_NUM = detect_vec_core_num()
BATCH_SIZE_SPECIALIZATIONS = tuple(range(2, 49, 2))
# Dedicated MTE3->MTE2 event for chunk-to-chunk UB reuse. The numeric
# id has no semantic meaning; it only needs to be paired and not collide
# with other MTE3_MTE2 syncs in this kernel.
CHUNK_OUTPUT_STORED_EVENT = 7
def select_launch_block_num(*, num_batches: int, vec_core_num: int) -> int:
"""Pick launch block_num by current batch size."""
if num_batches <= 0:
raise ValueError(f"num_batches({num_batches}) must be > 0")
if vec_core_num <= 0:
raise ValueError(f"vec_core_num({vec_core_num}) must be > 0")
return min(num_batches, vec_core_num)
def _dtype_size_in_bytes(dtype: str) -> int:
sizes = {
"float16": 2,
"bfloat16": 2,
"float32": 4,
}
if dtype not in sizes:
raise ValueError(f"Unsupported dtype for vector alignment: {dtype}")
return sizes[dtype]
def _align_count_to_vector_bytes(count: int, dtype: str) -> int:
elem_bytes = _dtype_size_in_bytes(dtype)
elems_per_iter = VECTOR_BYTES_PER_ITER // elem_bytes
return ((count + elems_per_iter - 1) // elems_per_iter) * elems_per_iter
DATACOPY_MIN_BYTES = 32
def _can_use_bulk_dma(num_heads: int, rows_per_iter: int) -> bool:
"""Check if multi-row DMA is viable (GM and UB strides match)."""
f32_vec_elems = VECTOR_BYTES_PER_ITER // 4
if (rows_per_iter * num_heads) % f32_vec_elems != 0:
return False
bf16_block_elems = DATACOPY_MIN_BYTES // 2
if num_heads % bf16_block_elems != 0:
return False
return True
UB_BUDGET_BYTES = 64 * 1024
def _compute_rows_per_iter(num_heads: int, ub_dim: int | None = None) -> int:
"""Compute how many rows to process per loop iteration based on UB budget."""
if ub_dim is None:
ub_dim = _align_count_to_vector_bytes(num_heads, "float32")
cmp_mask_bytes = (ub_dim + 7) // 8
shared_bytes = 1 * ub_dim * 4
per_row_bytes = (
2 * ub_dim * 4
+ 2 * ub_dim * 2
+ 5 * ub_dim * 4
+ 1 * ub_dim * 1
+ 1 * cmp_mask_bytes
)
max_rows = (UB_BUDGET_BYTES - shared_bytes) // per_row_bytes
if max_rows >= 32:
return 32
if max_rows >= 16:
return 16
if max_rows >= 8:
return 8
if max_rows >= 4:
return 4
if max_rows >= 2:
return 2
return 1
def build_fused_gdn_gating_kernel(
*,
batch_size: int,
compile_max_batch: int,
num_heads: int,
):
if num_heads not in SUPPORTED_NUM_HEADS:
raise ValueError(
"fused_gdn_gating only supports num_heads in "
f"{SUPPORTED_NUM_HEADS}, got {num_heads}"
)
if batch_size <= 0:
raise ValueError(f"batch_size({batch_size}) must be > 0")
if compile_max_batch <= 0:
raise ValueError(
f"compile_max_batch({compile_max_batch}) must be > 0"
)
if batch_size > compile_max_batch:
raise ValueError(
f"batch_size({batch_size}) must be <= compile_max_batch({compile_max_batch})"
)
vec_core_num = MAX_VEC_CORE_NUM
block_num = select_launch_block_num(
num_batches=batch_size, vec_core_num=vec_core_num
)
cubecore_block_num = block_num
task_num = block_num * VEC_NUM
acc_dtype = "float32"
input_dtype = "bfloat16"
mask_dtype = "uint8"
rows_conservative = _compute_rows_per_iter(num_heads)
use_bulk_dma = _can_use_bulk_dma(num_heads, rows_conservative)
if use_bulk_dma:
ub_tensor_dim = num_heads
rows_per_iter = _compute_rows_per_iter(num_heads, ub_dim=num_heads)
else:
ub_tensor_dim = _align_count_to_vector_bytes(num_heads, acc_dtype)
rows_per_iter = rows_conservative
if batch_size <= rows_per_iter:
rows_per_iter = max(1, batch_size // block_num)
use_bulk_dma = _can_use_bulk_dma(num_heads, rows_per_iter)
if use_bulk_dma:
ub_tensor_dim = num_heads
else:
ub_tensor_dim = _align_count_to_vector_bytes(num_heads, acc_dtype)
compare_select_mask_bytes = (ub_tensor_dim + 7) // 8
multi_count = rows_per_iter * ub_tensor_dim
multi_cmp_mask = rows_per_iter * compare_select_mask_bytes
@T.prim_func
def fused_gdn_gating_kernel(
A_log: T.Tensor((num_heads,), acc_dtype),
a: T.Tensor((compile_max_batch, num_heads), input_dtype),
b: T.Tensor((compile_max_batch, num_heads), input_dtype),
dt_bias: T.Tensor((num_heads,), acc_dtype),
g_out: T.Tensor((compile_max_batch, num_heads), acc_dtype),
beta_out: T.Tensor((compile_max_batch, num_heads), input_dtype),
num_batches: T.int32,
softplus_beta: T.float32,
softplus_threshold: T.float32,
):
with T.Kernel(cubecore_block_num, is_npu=True) as (cid, vid):
task_id = cid * VEC_NUM + vid
total_chunks = (num_batches + rows_per_iter - 1) // rows_per_iter
chunks_per_task = (total_chunks + task_num - 1) // task_num
chunk_start = task_id * chunks_per_task
chunks_left = T.if_then_else(
total_chunks > chunk_start,
total_chunks - chunk_start,
0,
)
num_chunks = T.if_then_else(
chunks_left < chunks_per_task,
chunks_left,
chunks_per_task,
)
with T.Scope("V"):
A_log_ub = T.alloc_shared((1, ub_tensor_dim), acc_dtype)
neg_exp_A_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), acc_dtype
)
dt_bias_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), acc_dtype
)
a_half_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), input_dtype
)
b_half_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), input_dtype
)
x_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), acc_dtype
)
beta_x_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), acc_dtype
)
softplus_abs_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), acc_dtype
)
softplus_tmp_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), acc_dtype
)
beta_fp32_ub = T.alloc_shared(
(rows_per_iter, ub_tensor_dim), acc_dtype
)
sigmoid_tmp_ub = T.alloc_ub(
(rows_per_iter, ub_tensor_dim), mask_dtype
)
softplus_cmp_mask_ub = T.alloc_ub(
(rows_per_iter, compare_select_mask_bytes), mask_dtype
)
# Preamble: load constants and replicate to R rows.
T.copy(A_log[0], A_log_ub[0, :num_heads])
T.tile.exp(A_log_ub, A_log_ub)
T.tile.mul(A_log_ub, A_log_ub, -1.0)
dt_bias_base_ub = T.alloc_shared(
(1, ub_tensor_dim), acc_dtype
)
T.copy(dt_bias[0], dt_bias_base_ub[0, :num_heads])
for r in T.serial(rows_per_iter):
T.copy(
A_log_ub[0, :ub_tensor_dim],
neg_exp_A_ub[r, :ub_tensor_dim],
)
T.copy(
dt_bias_base_ub[0, :ub_tensor_dim],
dt_bias_ub[r, :ub_tensor_dim],
)
if use_bulk_dma:
for chunk_idx in T.serial(num_chunks):
with T.If(chunk_idx > 0):
with T.Then():
mte2_wait_mte3(CHUNK_OUTPUT_STORED_EVENT)
base_row = (
(chunk_start + chunk_idx) * rows_per_iter
)
remaining = T.if_then_else(
num_batches > base_row,
num_batches - base_row,
0,
)
is_full_chunk = remaining >= rows_per_iter
with T.If(is_full_chunk):
with T.Then():
T.copy(a[base_row, 0], a_half_ub)
T.copy(b[base_row, 0], b_half_ub)
with T.Else():
for r in T.serial(remaining):
T.copy(
a[base_row + r, 0],
a_half_ub[r, :num_heads],
)
T.copy(
b[base_row + r, 0],
b_half_ub[r, :num_heads],
)
T.tile.cast(
x_ub, a_half_ub, "CAST_NONE", multi_count
)
T.tile.axpy(x_ub, dt_bias_ub, 1.0)
T.tile.mul(beta_x_ub, x_ub, softplus_beta)
T.tile.abs(softplus_abs_ub, beta_x_ub)
T.tile.mul(
softplus_tmp_ub, softplus_abs_ub, -1.0
)
T.tile.exp(beta_fp32_ub, softplus_tmp_ub)
T.tile.add(beta_fp32_ub, beta_fp32_ub, 1.0)
T.tile.ln(softplus_tmp_ub, beta_fp32_ub)
T.tile.compare(
softplus_cmp_mask_ub,
beta_x_ub,
softplus_threshold,
"GT",
)
T.tile.add(
beta_x_ub, beta_x_ub, softplus_abs_ub
)
T.tile.mul(
beta_x_ub, beta_x_ub, 0.5 / softplus_beta
)
T.tile.axpy(
beta_x_ub,
softplus_tmp_ub,
1.0 / softplus_beta,
)
T.tile.select(
beta_x_ub,
softplus_cmp_mask_ub,
x_ub,
beta_x_ub,
"VSEL_TENSOR_TENSOR_MODE",
)
T.tile.cast(
x_ub, b_half_ub, "CAST_NONE", multi_count
)
T.tile.sigmoid(
beta_fp32_ub, x_ub, sigmoid_tmp_ub
)
T.tile.mul(x_ub, neg_exp_A_ub, beta_x_ub)
T.tile.cast(
b_half_ub,
beta_fp32_ub,
"CAST_RINT",
multi_count,
)
with T.If(is_full_chunk):
with T.Then():
T.copy(x_ub, g_out[base_row, 0])
T.copy(b_half_ub, beta_out[base_row, 0])
with T.Else():
for r in T.serial(remaining):
T.copy(
x_ub[r, :num_heads],
g_out[base_row + r, :],
)
T.copy(
b_half_ub[r, :num_heads],
beta_out[base_row + r, :],
)
with T.If(chunk_idx < num_chunks - 1):
with T.Then():
mte3_notify_mte2(CHUNK_OUTPUT_STORED_EVENT)
else:
for chunk_idx in T.serial(num_chunks):
with T.If(chunk_idx > 0):
with T.Then():
mte2_wait_mte3(CHUNK_OUTPUT_STORED_EVENT)
base_row = (
(chunk_start + chunk_idx) * rows_per_iter
)
remaining = T.if_then_else(
num_batches > base_row,
num_batches - base_row,
0,
)
valid_rows = T.if_then_else(
remaining >= rows_per_iter,
rows_per_iter,
remaining,
)
for r in T.serial(valid_rows):
T.copy(
a[base_row + r, 0],
a_half_ub[r, :num_heads],
)
T.copy(
b[base_row + r, 0],
b_half_ub[r, :num_heads],
)
T.tile.cast(
x_ub, a_half_ub, "CAST_NONE", multi_count
)
T.tile.axpy(x_ub, dt_bias_ub, 1.0)
T.tile.mul(beta_x_ub, x_ub, softplus_beta)
T.tile.abs(softplus_abs_ub, beta_x_ub)
T.tile.mul(
softplus_tmp_ub, softplus_abs_ub, -1.0
)
T.tile.exp(beta_fp32_ub, softplus_tmp_ub)
T.tile.add(beta_fp32_ub, beta_fp32_ub, 1.0)
T.tile.ln(softplus_tmp_ub, beta_fp32_ub)
T.tile.compare(
softplus_cmp_mask_ub,
beta_x_ub,
softplus_threshold,
"GT",
)
T.tile.add(
beta_x_ub, beta_x_ub, softplus_abs_ub
)
T.tile.mul(
beta_x_ub, beta_x_ub, 0.5 / softplus_beta
)
T.tile.axpy(
beta_x_ub,
softplus_tmp_ub,
1.0 / softplus_beta,
)
T.tile.select(
beta_x_ub,
softplus_cmp_mask_ub,
x_ub,
beta_x_ub,
"VSEL_TENSOR_TENSOR_MODE",
)
T.tile.cast(
x_ub, b_half_ub, "CAST_NONE", multi_count
)
T.tile.sigmoid(
beta_fp32_ub, x_ub, sigmoid_tmp_ub
)
T.tile.mul(x_ub, neg_exp_A_ub, beta_x_ub)
T.tile.cast(
b_half_ub,
beta_fp32_ub,
"CAST_RINT",
multi_count,
)
for r in T.serial(valid_rows):
T.copy(
x_ub[r, :num_heads],
g_out[base_row + r, :],
)
T.copy(
b_half_ub[r, :num_heads],
beta_out[base_row + r, :],
)
with T.If(chunk_idx < num_chunks - 1):
with T.Then():
mte3_notify_mte2(CHUNK_OUTPUT_STORED_EVENT)
return fused_gdn_gating_kernel
@tilelang.jit(pass_configs=DEFAULT_ASCEND_PASS_CONFIGS)
def fused_gdn_gating_kernel_jit(
num_batches: int,
compile_max_batch: int,
num_heads: int,
):
return build_fused_gdn_gating_kernel(
batch_size=num_batches,
compile_max_batch=compile_max_batch,
num_heads=num_heads,
)
@register_kernel
class FusedGdnGatingKernel(TilelangKernel):
DISPATCH_SCHEMA = [
DispatchField("batch_size", "int32"),
DispatchField("num_heads", "int32"),
DispatchField("dtype", "dtype"),
]
SPECIALIZATIONS = [
{
"variant_key": f"bs{batch_size}_nh{num_heads}_bf16",
"batch_size": batch_size,
"num_heads": num_heads,
"dtype": DEFAULT_DTYPE,
}
for num_heads in SUPPORTED_NUM_HEADS
for batch_size in BATCH_SIZE_SPECIALIZATIONS
]
@staticmethod
def generate_source(batch_size: int, num_heads: int, dtype: str) -> str:
if dtype != DEFAULT_DTYPE:
raise ValueError(
f"fused_gdn_gating only supports dtype={DEFAULT_DTYPE}, got {dtype}"
)
if num_heads not in SUPPORTED_NUM_HEADS:
raise ValueError(
"fused_gdn_gating only supports num_heads in "
f"{SUPPORTED_NUM_HEADS}, got {num_heads}"
)
if batch_size not in BATCH_SIZE_SPECIALIZATIONS:
raise ValueError(
"fused_gdn_gating only supports batch_size in "
f"{BATCH_SIZE_SPECIALIZATIONS}, got {batch_size}"
)
tilelang.disable_cache()
tilelang_kernel = build_fused_gdn_gating_kernel(
batch_size=batch_size,
compile_max_batch=DEFAULT_MAX_BATCH,
num_heads=num_heads,
)
with tilelang.tvm.transform.PassContext(
opt_level=3, config=DEFAULT_ASCEND_PASS_CONFIGS
):
kernel = tilelang.engine.lower(tilelang_kernel)
return kernel.kernel_source
def _torch_fused_gdn_gating(
A_log: "torch.Tensor",
a: "torch.Tensor",
b: "torch.Tensor",
dt_bias: "torch.Tensor",
softplus_beta: float,
softplus_threshold: float,
) -> tuple["torch.Tensor", "torch.Tensor"]:
import torch
import torch.nn.functional as F
softplus_out = F.softplus(
a.to(torch.float32) + dt_bias,
beta=softplus_beta,
threshold=softplus_threshold,
)
g_ref = -A_log.exp() * softplus_out
beta_ref = torch.sigmoid(b.to(torch.float32)).to(torch.bfloat16)
return g_ref, beta_ref
def _run_ref_check(
*,
num_batches: int,
num_heads: int,
compile_max_batch: int,
softplus_beta: float,
softplus_threshold: float,
) -> None:
import torch
if not hasattr(torch, "npu") or not torch.npu.is_available():
print("[WARN] Skip fused_gdn_gating reference check: NPU is not available")
return
if num_batches <= 0:
raise ValueError(f"num_batches({num_batches}) must be > 0")
if num_batches > compile_max_batch:
raise ValueError(
f"num_batches({num_batches}) must be <= compile_max_batch({compile_max_batch})"
)
torch.manual_seed(42)
device = torch.device("npu")
A_log = torch.randn((num_heads,), device=device, dtype=torch.float32)
a = torch.randn((num_batches, num_heads), device=device, dtype=torch.bfloat16)
b = torch.randn((num_batches, num_heads), device=device, dtype=torch.bfloat16)
dt_bias = torch.randn((num_heads,), device=device, dtype=torch.float32)
g_out = torch.empty((num_batches, num_heads), device=device, dtype=torch.float32)
beta_out = torch.empty(
(num_batches, num_heads), device=device, dtype=torch.bfloat16
)
kernel = fused_gdn_gating_kernel_jit(
num_batches=num_batches,
compile_max_batch=num_batches,
num_heads=num_heads,
)
kernel(
A_log,
a,
b,
dt_bias,
g_out,
beta_out,
num_batches,
softplus_beta,
softplus_threshold,
)
torch.npu.synchronize()
g_ref, beta_ref = _torch_fused_gdn_gating(
A_log=A_log,
a=a,
b=b,
dt_bias=dt_bias,
softplus_beta=softplus_beta,
softplus_threshold=softplus_threshold,
)
torch.testing.assert_close(g_out, g_ref, rtol=1e-3, atol=1e-3)
torch.testing.assert_close(
beta_out.to(torch.float32),
beta_ref.to(torch.float32),
rtol=1e-2,
atol=1e-2,
)
print(f"[INFO] fused_gdn_gating output matches torch reference for num_heads={num_heads}")
def _run_ref_suite(
*,
num_batches: int,
compile_max_batch: int,
softplus_beta: float,
softplus_threshold: float,
ref_num_heads_list: list[int],
) -> None:
for num_heads in ref_num_heads_list:
_run_ref_check(
num_batches=num_batches,
num_heads=num_heads,
compile_max_batch=compile_max_batch,
softplus_beta=softplus_beta,
softplus_threshold=softplus_threshold,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Generate TileLang AscendC source for fused_gdn_gating AOT kernel."
)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument(
"--batch-size",
type=int,
default=max(BATCH_SIZE_SPECIALIZATIONS),
help=(
"Batch-size specialization used for source generation. "
f"Supported values: {BATCH_SIZE_SPECIALIZATIONS}"
),
)
parser.add_argument("--num-heads", type=int, default=DEFAULT_NUM_HEADS)
parser.add_argument("--dtype", type=str, default=DEFAULT_DTYPE)
parser.add_argument(
"--skip-ref-check",
action="store_true",
help="Skip runtime torch-reference check.",
)
parser.add_argument(
"--ref-num-batches",
type=int,
default=REF_CHECK_NUM_BATCHES,
help="Batch size used by the optional torch-reference check.",
)
parser.add_argument(
"--softplus-beta",
type=float,
default=1.0,
help="Softplus beta used by the optional torch-reference check.",
)
parser.add_argument(
"--softplus-threshold",
type=float,
default=20.0,
help="Softplus threshold used by the optional torch-reference check.",
)
parser.add_argument(
"--ref-num-heads-list",
type=int,
nargs="+",
default=list(REF_CHECK_NUM_HEADS),
help="Head counts covered by the optional bf16 torch-reference test suite.",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
source = FusedGdnGatingKernel.generate_source(
batch_size=args.batch_size,
num_heads=args.num_heads,
dtype=args.dtype,
)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(source, encoding="utf-8")
if not args.skip_ref_check:
_run_ref_suite(
num_batches=args.ref_num_batches,
compile_max_batch=DEFAULT_MAX_BATCH,
softplus_beta=args.softplus_beta,
softplus_threshold=args.softplus_threshold,
ref_num_heads_list=args.ref_num_heads_list,
)
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,304 @@
#!/usr/bin/env python3
import argparse
from pathlib import Path
import tilelang
import tilelang.language as T
from .utils import (
DEFAULT_ASCEND_PASS_CONFIGS,
detect_vec_core_num,
)
from ....common.spec import DispatchField, TilelangKernel, register_kernel
DEFAULT_HEAD_DIM = 576
DEFAULT_ROPE_DIM = 64
DEFAULT_DTYPE = "bf16"
SECONDARY_HEAD_DIM = 128
SECONDARY_ROPE_DIM = 128
VEC_NUM = 2
FIXED_UB_BUFFER_BYTES = 64 * 1024
REF_CHECK_NUM_TOKENS = 16
# AOT kernel tensor signatures still require static first-dim bounds.
# Keep a sufficiently large compile-time upper bound so runtime rows
# (`num_tokens * num_heads`) used by wrapper tests stay in-range.
MIN_COMPILE_NUM_TOKENS = 65536
# Per-row bytes in UB for this kernel:
# x_half(2) + x(4) + sin_half(2) + sin(4) + cos_half(2) + cos(4)
# + x_rotate(4) + out(4) + mask(4) = 30 bytes per rope element.
UB_BYTES_PER_ROW_PER_ROPE_ELEM = 30
def _derive_max_rows_num_in_ub(rope_dim: int, ub_buffer_bytes: int) -> int:
if ub_buffer_bytes <= 0:
raise ValueError(f"ub_buffer_bytes({ub_buffer_bytes}) must be > 0")
if rope_dim <= 0:
raise ValueError(f"rope_dim({rope_dim}) must be > 0")
bytes_per_row = UB_BYTES_PER_ROW_PER_ROPE_ELEM * rope_dim
max_rows = ub_buffer_bytes // bytes_per_row
if max_rows <= 0:
raise ValueError(
"UB budget is too small for current rope_dim: "
f"ub_buffer_bytes={ub_buffer_bytes}, rope_dim={rope_dim}"
)
return max_rows
def build_rope_kernel(
head_dim: int,
rope_dim: int,
vec_core_num: int,
ub_buffer_bytes: int,
):
if rope_dim % 2 != 0:
raise ValueError(f"rope_dim({rope_dim}) must be even")
if rope_dim > head_dim:
raise ValueError(f"rope_dim({rope_dim}) must be <= head_dim({head_dim})")
if vec_core_num <= 0:
raise ValueError(f"vec_core_num({vec_core_num}) must be > 0")
if vec_core_num % VEC_NUM != 0:
raise ValueError(
f"vec_core_num({vec_core_num}) must be divisible by VEC_NUM({VEC_NUM})"
)
task_num = vec_core_num
m_num = vec_core_num // VEC_NUM
max_rows_num_in_ub = _derive_max_rows_num_in_ub(
rope_dim=rope_dim,
ub_buffer_bytes=ub_buffer_bytes,
)
# Current AOT path fixes launch block_num at compile time, so runtime input
# shape only changes per-task workload splitting. The tensor signature still
# needs a static upper bound for the first dimension.
compile_num_tokens = max(task_num * max_rows_num_in_ub, MIN_COMPILE_NUM_TOKENS)
compile_flatten_width = compile_num_tokens * head_dim
acc_dtype = "float32"
mask_dtype = "uint32"
@T.prim_func
def rope_in_place_kernel(
x_in: T.Tensor((1, compile_flatten_width), "bfloat16"),
sin: T.Tensor((compile_num_tokens, rope_dim), "bfloat16"),
cos: T.Tensor((compile_num_tokens, rope_dim), "bfloat16"),
x_out: T.Tensor((1, compile_flatten_width), "bfloat16"),
num_tokens: T.int32,
x_stride: T.int32,
):
with T.Kernel(m_num, is_npu=True) as (cid, vid):
task_id = cid * VEC_NUM + vid
block_m = (num_tokens + task_num - 1) // task_num
row_start = task_id * block_m
rows_left = T.if_then_else(
num_tokens > row_start, num_tokens - row_start, 0
)
num_rows_per_vec = T.if_then_else(
rows_left < block_m,
rows_left,
block_m,
)
with T.Scope("V"):
mask_ub = T.alloc_ub([1, rope_dim], mask_dtype)
for j in T.serial(rope_dim // 2):
mask_ub[0, 2 * j] = 4 * (2 * j + 1)
mask_ub[0, 2 * j + 1] = 4 * (2 * j)
sin_mask_ub = T.alloc_ub((rope_dim,), acc_dtype)
T.tile.fill(sin_mask_ub, 1.0)
for i in T.serial(rope_dim):
if i % 2 == 0:
sin_mask_ub[i] = -1.0
x_half_ub = T.alloc_shared([1, rope_dim], "bfloat16")
x_ub = T.alloc_shared([1, rope_dim], acc_dtype)
sin_half_ub = T.alloc_shared([1, rope_dim], "bfloat16")
sin_ub = T.alloc_shared([1, rope_dim], acc_dtype)
cos_half_ub = T.alloc_shared([1, rope_dim], "bfloat16")
cos_ub = T.alloc_shared([1, rope_dim], acc_dtype)
x_rotate_ub = T.alloc_shared([1, rope_dim], acc_dtype)
out_ub = T.alloc_shared([1, rope_dim], acc_dtype)
for row_local in T.serial(num_rows_per_vec):
row = row_start + row_local
row_offset = row * x_stride
T.copy(x_in[0, row_offset], x_half_ub[0, :])
T.copy(sin[row, :], sin_half_ub[0, :])
T.copy(cos[row, :], cos_half_ub[0, :])
T.tile.cast(x_ub, x_half_ub, "CAST_NONE", rope_dim)
T.tile.cast(sin_ub, sin_half_ub, "CAST_NONE", rope_dim)
T.tile.cast(cos_ub, cos_half_ub, "CAST_NONE", rope_dim)
T.tile.mul(sin_ub[0, :], sin_ub[0, :], sin_mask_ub)
T.tile.gather(x_rotate_ub, x_ub, mask_ub, 0)
T.tile.mul(x_ub, x_ub, cos_ub)
T.tile.mul(x_rotate_ub, x_rotate_ub, sin_ub)
T.tile.add(out_ub, x_ub, x_rotate_ub)
T.tile.cast(x_half_ub, out_ub, "CAST_RINT", rope_dim)
T.copy(x_half_ub[0, :], x_out[0, row_offset])
return rope_in_place_kernel
@tilelang.jit(pass_configs=DEFAULT_ASCEND_PASS_CONFIGS)
def rope_in_place_kernel_jit(
head_dim: int,
rope_dim: int,
vec_core_num: int,
ub_buffer_bytes: int,
):
return build_rope_kernel(
head_dim=head_dim,
rope_dim=rope_dim,
vec_core_num=vec_core_num,
ub_buffer_bytes=ub_buffer_bytes,
)
@register_kernel
class RopeKernel(TilelangKernel):
DISPATCH_SCHEMA = [
DispatchField("head_dim", "int32"),
DispatchField("rope_dim", "int32"),
DispatchField("dtype", "dtype"),
]
SPECIALIZATIONS = [
{
"variant_key": "hd128_rd128_bf16",
"head_dim": SECONDARY_HEAD_DIM,
"rope_dim": SECONDARY_ROPE_DIM,
"dtype": DEFAULT_DTYPE,
},
{
"variant_key": "hd576_rd64_bf16",
"head_dim": DEFAULT_HEAD_DIM,
"rope_dim": DEFAULT_ROPE_DIM,
"dtype": DEFAULT_DTYPE,
},
]
@staticmethod
def generate_source(head_dim: int, rope_dim: int, dtype: str) -> str:
if dtype != DEFAULT_DTYPE:
raise ValueError(
f"RoPE TileLang kernel only supports dtype={DEFAULT_DTYPE}, got {dtype}"
)
tilelang.disable_cache()
vec_core_num = detect_vec_core_num()
ub_buffer_bytes = FIXED_UB_BUFFER_BYTES
tilelang_kernel = build_rope_kernel(
head_dim=head_dim,
rope_dim=rope_dim,
vec_core_num=vec_core_num,
ub_buffer_bytes=ub_buffer_bytes,
)
with tilelang.tvm.transform.PassContext(
opt_level=3, config=DEFAULT_ASCEND_PASS_CONFIGS
):
kernel = tilelang.engine.lower(tilelang_kernel)
return kernel.kernel_source
def _torch_rope_ref_rows(
x: "torch.Tensor",
sin: "torch.Tensor",
cos: "torch.Tensor",
dim_start: int,
) -> "torch.Tensor":
import torch
x_fp32 = x.to(torch.float32)
sin_fp32 = sin.to(torch.float32)
cos_fp32 = cos.to(torch.float32)
rope_dim = sin_fp32.shape[1]
x_part = x_fp32[:, dim_start : dim_start + rope_dim]
x_reshape = x_part.reshape(x_part.shape[0], -1, 2)
x0 = x_reshape[:, :, 0]
x1 = x_reshape[:, :, 1]
x_rot = torch.stack([-x1, x0], dim=-1).reshape_as(x_part)
out = x.clone()
out[:, dim_start : dim_start + rope_dim] = (
x_part * cos_fp32 + x_rot * sin_fp32
).to(torch.bfloat16)
return out
def _run_ref_check(
num_tokens: int,
head_dim: int,
rope_dim: int,
vec_core_num: int,
ub_buffer_bytes: int,
) -> None:
import torch
if not hasattr(torch, "npu") or not torch.npu.is_available():
print("[WARN] Skip RoPE reference check: NPU is not available")
return
torch.manual_seed(42)
device = torch.device("npu")
x_in = torch.randn((num_tokens, head_dim), device=device, dtype=torch.bfloat16)
sin = torch.randn((num_tokens, rope_dim), device=device, dtype=torch.bfloat16)
cos = torch.randn((num_tokens, rope_dim), device=device, dtype=torch.bfloat16)
x_out = x_in.clone()
x_in_flat = x_in.view(1, -1)
x_out_flat = x_out.view(1, -1)
kernel = rope_in_place_kernel_jit(
head_dim=head_dim,
rope_dim=rope_dim,
vec_core_num=vec_core_num,
ub_buffer_bytes=ub_buffer_bytes,
)
kernel(x_in_flat, sin, cos, x_out_flat, num_tokens, head_dim)
torch.npu.synchronize()
x_ref = _torch_rope_ref_rows(x_in, sin, cos, 0)
torch.testing.assert_close(x_out, x_ref, rtol=1e-3, atol=1e-3)
print("[INFO] RoPE output matches torch reference")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Generate TileLang AscendC source for RoPE AOT kernel."
)
parser.add_argument("--output", required=True, help="Output AscendC .cpp file")
parser.add_argument("--head-dim", type=int, default=DEFAULT_HEAD_DIM)
parser.add_argument("--rope-dim", type=int, default=DEFAULT_ROPE_DIM)
parser.add_argument("--dtype", default=DEFAULT_DTYPE)
parser.add_argument(
"--skip-ref-check",
action="store_true",
help="Skip runtime torch-reference check.",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output).resolve()
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(
RopeKernel.generate_source(
head_dim=args.head_dim,
rope_dim=args.rope_dim,
dtype=args.dtype,
),
encoding="utf-8",
)
if not args.skip_ref_check:
_run_ref_check(
num_tokens=REF_CHECK_NUM_TOKENS,
head_dim=args.head_dim,
rope_dim=args.rope_dim,
vec_core_num=detect_vec_core_num(),
ub_buffer_bytes=FIXED_UB_BUFFER_BYTES,
)
if __name__ == "__main__":
main()

View File

@@ -0,0 +1,317 @@
import tilelang
import tilelang.language as T
from typing import Any
from ....common.manifest import KernelAbi, KernelVariantManifest
from ....common.spec import DispatchField
DEFAULT_ASCEND_PASS_CONFIGS = {
# Use raw pass-config strings to avoid hard dependency on
# tilelang.PassConfigKey export timing/version.
"tl.ascend_auto_sync": True,
"tl.ascend_memory_planning": True,
"tl.ascend_auto_cross_core_sync": True,
"tl.ascend_auto_cv_combine": True,
}
DEFAULT_ASCEND_BISHENG_ARCH = "dav-2201"
ASCEND_VEC_CORE_NUM_PROPERTY_KEYS = (
"vector_core_num",
"aiv_core_num",
"vec_core_num",
)
def detect_vec_core_num(default_vec_core_num: int = 48) -> int:
try:
import torch
if hasattr(torch, "npu") and torch.npu.is_available():
props = torch.npu.get_device_properties(torch.npu.current_device())
for key in ASCEND_VEC_CORE_NUM_PROPERTY_KEYS:
value = getattr(props, key, None)
if isinstance(value, int) and value > 0:
return value
except Exception:
pass
return default_vec_core_num
def _snake_to_pascal(name: str) -> str:
parts = [part for part in name.split("_") if part]
return "".join(part[:1].upper() + part[1:] for part in parts)
def _dispatch_field_suffix(name: str) -> str:
parts = [part for part in name.split("_") if part]
mapped = {"dtype": "DType"}
return "".join(mapped.get(part, part[:1].upper() + part[1:]) for part in parts)
def _dtype_enum_suffix(dtype_name: str) -> str:
common_suffixes = {
"bf16": "BF16",
"fp16": "Float16",
"fp32": "Float32",
"float16": "Float16",
"float32": "Float32",
"int8": "Int8",
"int32": "Int32",
"uint8": "UInt8",
}
if dtype_name in common_suffixes:
return common_suffixes[dtype_name]
return _snake_to_pascal(dtype_name)
def _dispatch_field_cpp_type(field: DispatchField) -> str:
field_types = {
"int32": "int32_t",
"dtype": "TilelangDType",
}
return field_types[field.kind]
def _render_dispatch_value_literal(*, field: DispatchField, value: Any) -> str:
if field.kind == "int32":
if not isinstance(value, int) or isinstance(value, bool):
raise TypeError(
f"Unsupported int32 dispatch value for {field.name!r}: {value!r}"
)
return str(value)
if field.kind == "dtype":
if not isinstance(value, str):
raise TypeError(
f"Unsupported dtype dispatch value for {field.name!r}: {value!r}"
)
dtype_suffix = _dtype_enum_suffix(value)
return f"TilelangDType::k{dtype_suffix}"
raise TypeError(
f"Unsupported dispatch field kind {field.kind!r} for {field.name!r}"
)
def _validate_dispatch_values(
*, kernel_name: str, dispatch_schema: list[DispatchField], variant: KernelVariantManifest
) -> None:
schema_names = [field.name for field in dispatch_schema]
missing_keys = [name for name in schema_names if name not in variant.dispatch_values]
extra_keys = [name for name in variant.dispatch_values if name not in schema_names]
if missing_keys or extra_keys:
raise ValueError(
f"TileLang kernel family {kernel_name!r} variant "
f"{variant.variant_key!r} has inconsistent dispatch values: "
f"missing={missing_keys}, extra={extra_keys}"
)
def render_family_variants_inc(
*,
kernel_name: str,
dispatch_schema: list[DispatchField],
variants: list[KernelVariantManifest],
) -> str:
if not variants:
return ""
if not dispatch_schema:
raise ValueError(
f"TileLang kernel family {kernel_name!r} has empty dispatch schema"
)
macro_name = f"XLLM_TL_{kernel_name.upper()}_VARIANT"
lines: list[str] = []
for variant in variants:
_validate_dispatch_values(
kernel_name=kernel_name,
dispatch_schema=dispatch_schema,
variant=variant,
)
variant_args = [
_render_dispatch_value_literal(
field=field,
value=variant.dispatch_values[field.name],
)
for field in dispatch_schema
]
variant_args.append(f"\"{variant.variant_key}\"")
variant_args.append(variant.entry_symbol)
lines.append(f"{macro_name}({', '.join(variant_args)})")
return "\n".join(lines) + "\n"
def render_family_registry_inc(
*,
kernel_name: str,
dispatch_schema: list[DispatchField],
kernel_abi: KernelAbi,
variants: list[KernelVariantManifest],
) -> str:
if not variants:
return ""
if not dispatch_schema:
raise ValueError(
f"TileLang kernel family {kernel_name!r} has empty dispatch schema"
)
family_prefix = _snake_to_pascal(kernel_name)
specialization_type = f"{family_prefix}Specialization"
kernel_fn_type = f"{family_prefix}KernelFn"
registry_name = f"k{family_prefix}Registry"
entry_type = f"KernelEntry<{specialization_type}, {kernel_fn_type}>"
field_wrapper_types = {
field.name: f"{family_prefix}{_dispatch_field_suffix(field.name)}"
for field in dispatch_schema
}
symbol_declarations: list[str] = []
registry_entries: list[str] = []
for variant in variants:
_validate_dispatch_values(
kernel_name=kernel_name,
dispatch_schema=dispatch_schema,
variant=variant,
)
specialization_args = ", ".join(
f"{field_wrapper_types[field.name]}{{"
f"{_render_dispatch_value_literal(field=field, value=variant.dispatch_values[field.name])}"
f"}}"
for field in dispatch_schema
)
symbol_declarations.append(
f'extern "C" function_type_t<{kernel_fn_type}> {variant.entry_symbol};'
)
registry_entries.append(
f' {entry_type}{{make_{kernel_name}_specialization({specialization_args}), '
f'"{variant.variant_key}", &{variant.entry_symbol}}},'
)
struct_fields = [
f" {_dispatch_field_cpp_type(field)} {field.name};" for field in dispatch_schema
]
equality_terms = [f"lhs.{field.name} == rhs.{field.name}" for field in dispatch_schema]
builder_params = [
f"{field_wrapper_types[field.name]} {field.name}" for field in dispatch_schema
]
builder_values = ", ".join(f"{field.name}.value" for field in dispatch_schema)
function_params = ", ".join(
f"{parameter.cpp_type} {parameter.name}" for parameter in kernel_abi.parameters
)
lines = [
f"struct {specialization_type} {{",
*struct_fields,
"};",
"",
f"constexpr bool operator==(const {specialization_type}& lhs,",
f" const {specialization_type}& rhs) {{",
" return " + " && ".join(equality_terms) + ";",
"}",
"",
]
for field in dispatch_schema:
lines.extend(
[
f"struct {field_wrapper_types[field.name]} {{",
f" {_dispatch_field_cpp_type(field)} value;",
"};",
"",
]
)
lines.extend(
[
f"constexpr {specialization_type} make_{kernel_name}_specialization(",
" " + ", ".join(builder_params) + ") {",
f" return {specialization_type}{{{builder_values}}};",
"}",
"",
f"using {kernel_fn_type} = {kernel_abi.return_type} (*)({function_params});",
"",
*symbol_declarations,
"",
f"constexpr std::array<{entry_type}, {len(variants)}> {registry_name}{{{{",
*registry_entries,
"}};",
"",
f"inline const {entry_type}* find_{kernel_name}_kernel_entry(",
f" const {specialization_type}& specialization) {{",
f" return find_kernel_entry({registry_name}, specialization);",
"}",
"",
f"inline std::string available_{kernel_name}_variant_keys() {{",
f" return available_variant_keys({registry_name});",
"}",
]
)
return "\n".join(lines) + "\n"
# ---------------------------------------------------------------------------
# Pipeline sync macros for TileLang Ascend kernels with
# tl.ascend_auto_sync=False. These helpers keep kernel implementations small
# and centralize flag-direction naming.
# ---------------------------------------------------------------------------
@T.macro
def mte2_notify_v(event_id: T.int32):
T.set_flag("mte2", "v", event_id)
@T.macro
def v_wait_mte2(event_id: T.int32):
T.wait_flag("mte2", "v", event_id)
@T.macro
def v_notify_mte2(event_id: T.int32):
T.set_flag("v", "mte2", event_id)
@T.macro
def mte2_wait_v(event_id: T.int32):
T.wait_flag("v", "mte2", event_id)
@T.macro
def v_notify_mte3(event_id: T.int32):
T.set_flag("v", "mte3", event_id)
@T.macro
def mte3_wait_v(event_id: T.int32):
T.wait_flag("v", "mte3", event_id)
@T.macro
def mte3_notify_v(event_id: T.int32):
T.set_flag("mte3", "v", event_id)
@T.macro
def v_wait_mte3(event_id: T.int32):
T.wait_flag("mte3", "v", event_id)
@T.macro
def mte3_notify_mte2(event_id: T.int32):
T.set_flag("mte3", "mte2", event_id)
@T.macro
def mte2_wait_mte3(event_id: T.int32):
T.wait_flag("mte3", "mte2", event_id)
@T.macro
def mte2_notify_mte3(event_id: T.int32):
T.set_flag("mte2", "mte3", event_id)
@T.macro
def mte3_wait_mte2(event_id: T.int32):
T.wait_flag("mte2", "mte3", event_id)

View File

@@ -0,0 +1,134 @@
from __future__ import annotations
import os
from dataclasses import dataclass
from pathlib import Path
from ...common.toolchain import git_head, require_env
from .kernels.utils import DEFAULT_ASCEND_BISHENG_ARCH
TILELANG_BISHENG_COMMON_FLAGS = [
"-O2",
"-std=c++17",
"-xasc",
"-fPIC",
"-Wno-macro-redefined",
"-Wno-ignored-attributes",
"-Wno-non-c-typedef-for-linkage",
"-DBACKEND_HYBM",
]
ASCEND_DEVICE_TO_BISHENG_ARCH = {
"a2": DEFAULT_ASCEND_BISHENG_ARCH,
"a3": DEFAULT_ASCEND_BISHENG_ARCH,
}
@dataclass(frozen=True)
class AscendBuildContext:
device: str | None
bisheng_arch: str
bisheng_executable: str
toolchain_options: dict[str, str]
fingerprint: dict[str, str]
include_dirs: list[str]
def normalize_ascend_device(device: str | None) -> str | None:
if device is None:
return None
normalized = device.strip().lower()
if not normalized:
return None
if normalized not in ASCEND_DEVICE_TO_BISHENG_ARCH:
supported = ", ".join(sorted(ASCEND_DEVICE_TO_BISHENG_ARCH))
raise ValueError(
f"Unsupported Ascend TileLang device {device!r}. Expected one of: "
f"{supported}"
)
return normalized
def resolve_bisheng_arch(device: str | None) -> tuple[str | None, str]:
normalized_device = normalize_ascend_device(device)
if normalized_device is None:
print(
"[WARN] TileLang Ascend build did not receive --device. Falling back "
f"to default bisheng_arch={DEFAULT_ASCEND_BISHENG_ARCH}. Prefer "
"running via xLLM main build path or pass --device a2|a3 explicitly."
)
return None, DEFAULT_ASCEND_BISHENG_ARCH
return normalized_device, ASCEND_DEVICE_TO_BISHENG_ARCH[normalized_device]
def build_toolchain_options(device: str | None, bisheng_arch: str) -> dict[str, str]:
toolchain_options = {"bisheng_arch": bisheng_arch}
if device is not None:
toolchain_options["device"] = device
return toolchain_options
def resolve_npu_home_path() -> str:
for env_name in ("NPU_HOME_PATH", "NPU_TOOLKIT_HOME"):
value = os.environ.get(env_name, "").strip()
if value:
return value
for candidate in (
"/usr/local/Ascend/ascend-toolkit/latest",
"/usr/local/Ascend/ascend-toolkit",
):
if Path(candidate).exists():
return candidate
raise RuntimeError(
"Required NPU toolkit root is not set. Expected NPU_HOME_PATH or "
"NPU_TOOLKIT_HOME, or a standard install path under "
"/usr/local/Ascend/ascend-toolkit."
)
def bisheng_include_dirs() -> list[str]:
tl_root = require_env("TL_ROOT")
npu_home_path = resolve_npu_home_path()
return [
f"{npu_home_path}/include",
f"{npu_home_path}/include/experiment/runtime",
f"{npu_home_path}/include/experiment/msprof",
f"{npu_home_path}/compiler/tikcpp",
f"{npu_home_path}/compiler/tikcpp/tikcfw",
f"{npu_home_path}/compiler/tikcpp/tikcfw/impl",
f"{npu_home_path}/compiler/tikcpp/tikcfw/interface",
f"{tl_root}/3rdparty/catlass/include",
f"{tl_root}/3rdparty/shmem/include",
f"{tl_root}/3rdparty/shmem/src/device",
f"{tl_root}/src",
]
def build_fingerprint(bisheng_executable: str, bisheng_arch: str) -> dict[str, str]:
tl_root = require_env("TL_ROOT")
npu_home_path = resolve_npu_home_path()
return {
"target": "ascend",
"tl_root": tl_root,
"tilelang_git_head": git_head(tl_root),
"npu_home_path": npu_home_path,
"bisheng_executable": bisheng_executable,
"bisheng_arch": bisheng_arch,
}
def resolve_build_context(device: str | None, bisheng_executable: str) -> AscendBuildContext:
normalized_device, bisheng_arch = resolve_bisheng_arch(device)
fingerprint = build_fingerprint(bisheng_executable, bisheng_arch)
if normalized_device is not None:
fingerprint["device"] = normalized_device
return AscendBuildContext(
device=normalized_device,
bisheng_arch=bisheng_arch,
bisheng_executable=bisheng_executable,
toolchain_options=build_toolchain_options(normalized_device, bisheng_arch),
fingerprint=fingerprint,
include_dirs=bisheng_include_dirs(),
)

View File

@@ -0,0 +1,18 @@
from __future__ import annotations
from pathlib import Path
from ...common.manifest import KernelFamilyManifest
def build_kernels(
output_root: str | Path,
kernel_names: list[str] | None = None,
force: bool = False,
) -> list[KernelFamilyManifest]:
if kernel_names:
raise NotImplementedError(
"CUDA TileLang AOT build pipeline is scaffolded but no kernels are "
"registered yet."
)
return []

View File

@@ -0,0 +1,378 @@
from __future__ import annotations
import os
import shlex
import subprocess
from dataclasses import dataclass
from pathlib import Path
from scripts.build_support.env import set_npu_envs
from .common.toolchain import (
default_tilelang_root,
git_head,
prepare_tilelang_import,
repo_root,
resolve_tilelang_root,
)
PREPARE_ASCEND_COMMAND = "python xllm/compiler/tilelang_launcher.py prepare-ascend"
@dataclass(frozen=True)
class TilelangPrepareState:
tilelang_root: Path
cann_set_env: Path
current_head: str
cached_head: str | None
artifacts_ready: bool
import_ok: bool
import_detail: str
def _ready_error(message: str) -> RuntimeError:
return RuntimeError(f"{message}\nRun `{PREPARE_ASCEND_COMMAND}` first.")
def _find_cann_set_env() -> Path | None:
candidates: list[Path] = []
npu_home_path = os.environ.get("NPU_HOME_PATH", "").strip()
if npu_home_path:
toolkit_root = Path(npu_home_path).resolve()
candidates.append(toolkit_root / "set_env.sh")
candidates.append(toolkit_root.parent / "set_env.sh")
candidates.extend(
[
Path("/usr/local/Ascend/ascend-toolkit/set_env.sh"),
Path("/usr/local/Ascend/ascend-toolkit/latest/set_env.sh"),
]
)
for script in candidates:
if script.is_file():
return script.resolve()
return None
def resolve_cann_set_env() -> Path:
cann_set_env = _find_cann_set_env()
if cann_set_env is not None:
return cann_set_env
set_npu_envs()
cann_set_env = _find_cann_set_env()
if cann_set_env is not None:
return cann_set_env
raise RuntimeError(
"[ERROR] Cannot find CANN set_env.sh. Expected a path like "
"/usr/local/Ascend/ascend-toolkit/set_env.sh."
)
def ensure_tilelang_submodules(tilelang_root: str | Path) -> Path:
tl_root = Path(tilelang_root).resolve()
required_markers = {
"3rdparty/catlass/CMakeLists.txt": tl_root / "3rdparty" / "catlass" / "CMakeLists.txt",
"3rdparty/composable_kernel/CMakeLists.txt": (
tl_root / "3rdparty" / "composable_kernel" / "CMakeLists.txt"
),
"3rdparty/cutlass/CMakeLists.txt": tl_root / "3rdparty" / "cutlass" / "CMakeLists.txt",
"3rdparty/pto-isa/CMakeLists.txt": tl_root / "3rdparty" / "pto-isa" / "CMakeLists.txt",
"3rdparty/shmem/CMakeLists.txt": tl_root / "3rdparty" / "shmem" / "CMakeLists.txt",
"3rdparty/tvm/CMakeLists.txt": tl_root / "3rdparty" / "tvm" / "CMakeLists.txt",
}
missing = [name for name, path in required_markers.items() if not path.is_file()]
if missing:
if (tl_root / ".git").exists():
repair_hint = (
"Run "
f"`git -C {shlex.quote(str(tl_root))} submodule update --init --recursive` "
"first."
)
else:
bundled_root = default_tilelang_root().resolve()
bundled_repair_cmd = (
f"git -C {shlex.quote(str(repo_root()))} "
"submodule update --init --recursive third_party/tilelang-ascend"
)
if tl_root == bundled_root:
repair_hint = (
"Sync the bundled TileLang checkout from the xLLM repo root: "
f"`{bundled_repair_cmd}`."
)
else:
repair_hint = (
f"`TL_ROOT={tl_root}` is not a git checkout. "
"Point TL_ROOT at a fully initialized tilelang-ascend clone, "
f"or sync the bundled checkout with `{bundled_repair_cmd}`."
)
raise RuntimeError(
"[ERROR] tilelang-ascend nested dependencies are incomplete: "
f"missing {', '.join(missing)}. "
f"{repair_hint}"
)
return tl_root
def tilelang_git_head_cache_path(tilelang_root: str | Path) -> Path:
return Path(tilelang_root).resolve() / "build" / ".xllm_tilelang_git_head_cached"
def read_tilelang_git_head_cached(tilelang_root: str | Path) -> str | None:
cache_path = tilelang_git_head_cache_path(tilelang_root)
if not cache_path.is_file():
return None
value = cache_path.read_text(encoding="utf-8").strip()
return value or None
def write_tilelang_git_head_cached(tilelang_root: str | Path, head: str) -> None:
cache_path = tilelang_git_head_cache_path(tilelang_root)
cache_path.parent.mkdir(parents=True, exist_ok=True)
cache_path.write_text(head + "\n", encoding="utf-8")
def tilelang_artifacts_ready(tilelang_root: str | Path) -> bool:
tl_root = Path(tilelang_root).resolve()
required = [
tl_root / "build" / "libtilelang_module.so",
tl_root / "build" / "libtilelang.so",
tl_root / "build" / "tvm" / "libtvm.so",
]
return all(path.exists() for path in required)
def verify_tilelang_import(tilelang_root: str | Path) -> tuple[bool, str]:
tl_root = prepare_tilelang_import(tilelang_root)
env = os.environ.copy()
env["TL_ROOT"] = str(tl_root)
pythonpath = env.get("PYTHONPATH", "")
pythonpath_items = [item for item in pythonpath.split(os.pathsep) if item]
if str(tl_root) not in pythonpath_items:
pythonpath_items.insert(0, str(tl_root))
env["PYTHONPATH"] = os.pathsep.join(pythonpath_items)
cmd = [
"bash",
"-lc",
"python - <<'PY'\n"
"import tilelang\n"
"print(getattr(tilelang, '__file__', '<unknown>'))\n"
"PY",
]
result = subprocess.run(
cmd,
text=True,
capture_output=True,
check=False,
env=env,
)
if result.returncode != 0:
detail = (result.stderr or result.stdout).strip()
return False, detail
return True, result.stdout.strip()
def _tilelang_patch_dir() -> Path:
return Path(__file__).resolve().parent / "patches" / "tilelang_ascend"
def _tilelang_install_patch_path() -> Path:
patch_path = _tilelang_patch_dir() / "0001-install-ascend.patch"
if not patch_path.is_file():
raise RuntimeError(f"[ERROR] Missing TileLang patch file: {patch_path}")
return patch_path
def _git_apply_base_cmd(repo_root: Path) -> list[str]:
return [
"git",
"-c",
f"safe.directory={repo_root}",
"-C",
str(repo_root),
"apply",
"--whitespace=nowarn",
]
def _check_git_patch_state(repo_root: Path, patch_path: Path) -> str:
apply_check = subprocess.run(
_git_apply_base_cmd(repo_root) + ["--check", str(patch_path)],
text=True,
capture_output=True,
check=False,
)
if apply_check.returncode == 0:
return "unapplied"
reverse_check = subprocess.run(
_git_apply_base_cmd(repo_root) + ["--reverse", "--check", str(patch_path)],
text=True,
capture_output=True,
check=False,
)
if reverse_check.returncode == 0:
return "applied"
apply_detail = (apply_check.stderr or apply_check.stdout).strip()
reverse_detail = (reverse_check.stderr or reverse_check.stdout).strip()
raise RuntimeError(
"[ERROR] Failed to match TileLang patch "
f"{patch_path.name} against {repo_root}.\n"
f"apply --check: {apply_detail or '<no output>'}\n"
f"reverse --check: {reverse_detail or '<no output>'}"
)
def _apply_git_patch(repo_root: Path, patch_path: Path, message: str) -> None:
if _check_git_patch_state(repo_root, patch_path) == "applied":
return
subprocess.check_call(_git_apply_base_cmd(repo_root) + [str(patch_path)])
print(message)
def _restore_git_patch(repo_root: Path, patch_path: Path) -> None:
if _check_git_patch_state(repo_root, patch_path) != "applied":
return
subprocess.check_call(_git_apply_base_cmd(repo_root) + ["--reverse", str(patch_path)])
def _patch_tilelang_install_tree(tilelang_root: str | Path) -> None:
tl_root = Path(tilelang_root).resolve()
required_files = (
tl_root / "install_ascend.sh",
tl_root / "requirements-build.txt",
)
missing = [str(path.name) for path in required_files if not path.is_file()]
if missing:
raise RuntimeError(
"[ERROR] Missing tilelang install files: " + ", ".join(missing)
)
_apply_git_patch(
tl_root,
_tilelang_install_patch_path(),
"[INFO] Applied tilelang install patch",
)
def _restore_tilelang_install_tree(tilelang_root: str | Path) -> None:
tl_root = Path(tilelang_root).resolve()
_restore_git_patch(tl_root, _tilelang_install_patch_path())
def _run_tilelang_install(tilelang_root: str | Path, cann_set_env: str | Path) -> None:
tl_root = ensure_tilelang_submodules(tilelang_root)
_patch_tilelang_install_tree(tl_root)
cmd = (
f"source {shlex.quote(str(cann_set_env))} && "
"bash install_ascend.sh && "
"source set_env.sh"
)
env = os.environ.copy()
git_config_count = int(env.get("GIT_CONFIG_COUNT", "0") or "0")
env[f"GIT_CONFIG_KEY_{git_config_count}"] = "safe.directory"
env[f"GIT_CONFIG_VALUE_{git_config_count}"] = str(tl_root)
env["GIT_CONFIG_COUNT"] = str(git_config_count + 1)
try:
subprocess.check_call(
["bash", "-lc", cmd],
cwd=str(tl_root),
env=env,
)
finally:
_restore_tilelang_install_tree(tl_root)
def collect_prepare_state() -> TilelangPrepareState:
tilelang_root = ensure_tilelang_submodules(resolve_tilelang_root())
prepare_tilelang_import(tilelang_root)
return TilelangPrepareState(
tilelang_root=tilelang_root,
cann_set_env=resolve_cann_set_env(),
current_head=git_head(tilelang_root),
cached_head=read_tilelang_git_head_cached(tilelang_root),
artifacts_ready=tilelang_artifacts_ready(tilelang_root),
import_ok=False,
import_detail="",
)
def refresh_prepare_state_import(state: TilelangPrepareState) -> TilelangPrepareState:
import_ok, import_detail = verify_tilelang_import(state.tilelang_root)
return TilelangPrepareState(
tilelang_root=state.tilelang_root,
cann_set_env=state.cann_set_env,
current_head=state.current_head,
cached_head=state.cached_head,
artifacts_ready=tilelang_artifacts_ready(state.tilelang_root),
import_ok=import_ok,
import_detail=import_detail,
)
def prepare_state() -> TilelangPrepareState:
return refresh_prepare_state_import(collect_prepare_state())
def install_reasons(state: TilelangPrepareState, *, force: bool) -> list[str]:
reasons: list[str] = []
if force:
reasons.append("forced")
if state.cached_head is None:
reasons.append("HEAD cache missing")
elif state.current_head != state.cached_head:
reasons.append("HEAD changed")
if not state.artifacts_ready:
reasons.append("artifacts missing")
if not state.import_ok:
reasons.append("tilelang import failed")
return list(dict.fromkeys(reasons))
def ensure_ascend_ready() -> Path:
set_npu_envs()
state = prepare_state()
if not state.artifacts_ready:
raise _ready_error(
"[ERROR] tilelang-ascend artifacts are missing under "
f"{state.tilelang_root / 'build'}."
)
if not state.import_ok:
raise _ready_error(
"[ERROR] Failed to import tilelang after configuring TL_ROOT="
f"{state.tilelang_root}: {state.import_detail}"
)
return state.tilelang_root
def prepare_ascend(*, force: bool = False) -> Path:
set_npu_envs()
state = prepare_state()
reasons = install_reasons(state, force=force)
if reasons:
print("[INFO] Preparing tilelang-ascend: " + "; ".join(reasons))
_run_tilelang_install(state.tilelang_root, state.cann_set_env)
prepare_tilelang_import(state.tilelang_root)
write_tilelang_git_head_cached(state.tilelang_root, state.current_head)
state = prepare_state()
if not state.artifacts_ready:
raise RuntimeError(
"[ERROR] tilelang-ascend artifacts are still missing after prepare."
)
if not state.import_ok:
raise RuntimeError(
"[ERROR] tilelang import still failed after prepare: "
f"{state.import_detail}"
)
print(f"[INFO] tilelang import success: {state.import_detail}")
return state.tilelang_root

View File

@@ -0,0 +1,56 @@
from __future__ import annotations
import argparse
import sys
from pathlib import Path
def _bootstrap_import_paths() -> None:
compiler_dir = Path(__file__).resolve().parent
package_root = compiler_dir.parent
repo_root = package_root.parent
for path in (repo_root, package_root):
path_str = str(path)
if path_str not in sys.path:
sys.path.insert(0, path_str)
def _build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Source-tree launcher for xLLM TileLang prepare/compile flows."
)
subparsers = parser.add_subparsers(dest="command")
subparsers.required = True
subparsers.add_parser(
"prepare-ascend",
add_help=False,
help="Prepare third_party/tilelang-ascend for Ascend TileLang builds.",
)
subparsers.add_parser(
"compile-kernels",
add_help=False,
help="Compile TileLang kernels and emit manifests.",
)
return parser
def main(argv: list[str] | None = None) -> None:
sys.setrecursionlimit(10000)
parser = _build_parser()
args, remainder = parser.parse_known_args(argv)
_bootstrap_import_paths()
if args.command == "prepare-ascend":
from compiler.tilelang.cli.prepare_ascend import main as entrypoint
elif args.command == "compile-kernels":
from compiler.tilelang.cli.compile_kernels import main as entrypoint
else: # pragma: no cover - argparse enforces choices
raise ValueError(f"Unsupported TileLang launcher command: {args.command}")
entrypoint(remainder)
if __name__ == "__main__":
main()

Some files were not shown because too many files have changed in this diff Show More