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:
133
upstream_ref/xllm/xllm/CMakeLists.txt
Normal file
133
upstream_ref/xllm/xllm/CMakeLists.txt
Normal 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()
|
||||
]])
|
||||
86
upstream_ref/xllm/xllm/__init__.py
Normal file
86
upstream_ref/xllm/xllm/__init__.py
Normal 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",
|
||||
]
|
||||
58
upstream_ref/xllm/xllm/api_service/CMakeLists.txt
Normal file
58
upstream_ref/xllm/xllm/api_service/CMakeLists.txt
Normal 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>
|
||||
)
|
||||
764
upstream_ref/xllm/xllm/api_service/anthropic_service_impl.cpp
Normal file
764
upstream_ref/xllm/xllm/api_service/anthropic_service_impl.cpp
Normal 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
|
||||
45
upstream_ref/xllm/xllm/api_service/anthropic_service_impl.h
Normal file
45
upstream_ref/xllm/xllm/api_service/anthropic_service_impl.h
Normal 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
|
||||
1101
upstream_ref/xllm/xllm/api_service/api_service.cpp
Normal file
1101
upstream_ref/xllm/xllm/api_service/api_service.cpp
Normal file
File diff suppressed because it is too large
Load Diff
211
upstream_ref/xllm/xllm/api_service/api_service.h
Normal file
211
upstream_ref/xllm/xllm/api_service/api_service.h
Normal 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
|
||||
56
upstream_ref/xllm/xllm/api_service/api_service_impl.h
Normal file
56
upstream_ref/xllm/xllm/api_service/api_service_impl.h
Normal 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
|
||||
64
upstream_ref/xllm/xllm/api_service/call.cpp
Normal file
64
upstream_ref/xllm/xllm/api_service/call.cpp
Normal 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
|
||||
49
upstream_ref/xllm/xllm/api_service/call.h
Normal file
49
upstream_ref/xllm/xllm/api_service/call.h
Normal 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
|
||||
164
upstream_ref/xllm/xllm/api_service/chat_json_parser.cpp
Normal file
164
upstream_ref/xllm/xllm/api_service/chat_json_parser.cpp
Normal 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
|
||||
68
upstream_ref/xllm/xllm/api_service/chat_json_parser.h
Normal file
68
upstream_ref/xllm/xllm/api_service/chat_json_parser.h
Normal 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
|
||||
979
upstream_ref/xllm/xllm/api_service/chat_service_impl.cpp
Normal file
979
upstream_ref/xllm/xllm/api_service/chat_service_impl.cpp
Normal 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
|
||||
80
upstream_ref/xllm/xllm/api_service/chat_service_impl.h
Normal file
80
upstream_ref/xllm/xllm/api_service/chat_service_impl.h
Normal 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
|
||||
340
upstream_ref/xllm/xllm/api_service/completion_service_impl.cpp
Normal file
340
upstream_ref/xllm/xllm/api_service/completion_service_impl.cpp
Normal 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
|
||||
52
upstream_ref/xllm/xllm/api_service/completion_service_impl.h
Normal file
52
upstream_ref/xllm/xllm/api_service/completion_service_impl.h
Normal 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
|
||||
179
upstream_ref/xllm/xllm/api_service/embedding_output_builder.cpp
Normal file
179
upstream_ref/xllm/xllm/api_service/embedding_output_builder.cpp
Normal 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
|
||||
@@ -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
|
||||
187
upstream_ref/xllm/xllm/api_service/embedding_service_impl.cpp
Normal file
187
upstream_ref/xllm/xllm/api_service/embedding_service_impl.cpp
Normal 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
|
||||
58
upstream_ref/xllm/xllm/api_service/embedding_service_impl.h
Normal file
58
upstream_ref/xllm/xllm/api_service/embedding_service_impl.h
Normal 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
|
||||
@@ -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
|
||||
@@ -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
|
||||
86
upstream_ref/xllm/xllm/api_service/mm_service_utils.h
Normal file
86
upstream_ref/xllm/xllm/api_service/mm_service_utils.h
Normal 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
|
||||
62
upstream_ref/xllm/xllm/api_service/models_service_impl.cpp
Normal file
62
upstream_ref/xllm/xllm/api_service/models_service_impl.cpp
Normal 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
|
||||
43
upstream_ref/xllm/xllm/api_service/models_service_impl.h
Normal file
43
upstream_ref/xllm/xllm/api_service/models_service_impl.h
Normal 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
|
||||
125
upstream_ref/xllm/xllm/api_service/non_stream_call.h
Normal file
125
upstream_ref/xllm/xllm/api_service/non_stream_call.h
Normal 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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
159
upstream_ref/xllm/xllm/api_service/rerank_service_impl.cpp
Normal file
159
upstream_ref/xllm/xllm/api_service/rerank_service_impl.cpp
Normal 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
|
||||
100
upstream_ref/xllm/xllm/api_service/rerank_service_impl.h
Normal file
100
upstream_ref/xllm/xllm/api_service/rerank_service_impl.h
Normal 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
|
||||
471
upstream_ref/xllm/xllm/api_service/sample_service_impl.cpp
Normal file
471
upstream_ref/xllm/xllm/api_service/sample_service_impl.cpp
Normal 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,
|
||||
¶ms.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
|
||||
74
upstream_ref/xllm/xllm/api_service/sample_service_impl.h
Normal file
74
upstream_ref/xllm/xllm/api_service/sample_service_impl.h
Normal 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
|
||||
124
upstream_ref/xllm/xllm/api_service/service_impl_factory.cpp
Normal file
124
upstream_ref/xllm/xllm/api_service/service_impl_factory.cpp
Normal 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
|
||||
37
upstream_ref/xllm/xllm/api_service/service_impl_factory.h
Normal file
37
upstream_ref/xllm/xllm/api_service/service_impl_factory.h
Normal 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
|
||||
49
upstream_ref/xllm/xllm/api_service/serving_mode.h
Normal file
49
upstream_ref/xllm/xllm/api_service/serving_mode.h
Normal 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
|
||||
223
upstream_ref/xllm/xllm/api_service/stream_call.h
Normal file
223
upstream_ref/xllm/xllm/api_service/stream_call.h
Normal 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
|
||||
87
upstream_ref/xllm/xllm/api_service/stream_output_parser.cpp
Normal file
87
upstream_ref/xllm/xllm/api_service/stream_output_parser.cpp
Normal 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
|
||||
69
upstream_ref/xllm/xllm/api_service/stream_output_parser.h
Normal file
69
upstream_ref/xllm/xllm/api_service/stream_output_parser.h
Normal 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
|
||||
157
upstream_ref/xllm/xllm/api_service/utils.h
Normal file
157
upstream_ref/xllm/xllm/api_service/utils.h
Normal 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
|
||||
54
upstream_ref/xllm/xllm/api_service/xllm_metrics.h
Normal file
54
upstream_ref/xllm/xllm/api_service/xllm_metrics.h
Normal 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
|
||||
104
upstream_ref/xllm/xllm/c_api/README.md
Normal file
104
upstream_ref/xllm/xllm/c_api/README.md
Normal 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
|
||||
```
|
||||
161
upstream_ref/xllm/xllm/c_api/default.h
Normal file
161
upstream_ref/xllm/xllm/c_api/default.h
Normal 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
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
751
upstream_ref/xllm/xllm/c_api/internal/helper.cpp
Normal file
751
upstream_ref/xllm/xllm/c_api/internal/helper.cpp
Normal 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
|
||||
172
upstream_ref/xllm/xllm/c_api/internal/helper.h
Normal file
172
upstream_ref/xllm/xllm/c_api/internal/helper.h
Normal 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
|
||||
221
upstream_ref/xllm/xllm/c_api/internal/llm.cpp
Normal file
221
upstream_ref/xllm/xllm/c_api/internal/llm.cpp
Normal 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);
|
||||
}
|
||||
432
upstream_ref/xllm/xllm/c_api/internal/rec.cpp
Normal file
432
upstream_ref/xllm/xllm/c_api/internal/rec.cpp
Normal 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);
|
||||
}
|
||||
227
upstream_ref/xllm/xllm/c_api/llm.h
Normal file
227
upstream_ref/xllm/xllm/c_api/llm.h
Normal 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
|
||||
339
upstream_ref/xllm/xllm/c_api/rec.h
Normal file
339
upstream_ref/xllm/xllm/c_api/rec.h
Normal 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
|
||||
280
upstream_ref/xllm/xllm/c_api/test/CMakeLists.txt
Normal file
280
upstream_ref/xllm/xllm/c_api/test/CMakeLists.txt
Normal 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"
|
||||
)
|
||||
109
upstream_ref/xllm/xllm/c_api/test/README.md
Normal file
109
upstream_ref/xllm/xllm/c_api/test/README.md
Normal 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。
|
||||
647
upstream_ref/xllm/xllm/c_api/test/utils.cpp
Normal file
647
upstream_ref/xllm/xllm/c_api/test/utils.cpp
Normal 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
|
||||
146
upstream_ref/xllm/xllm/c_api/test/utils.h
Normal file
146
upstream_ref/xllm/xllm/c_api/test/utils.h
Normal 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_
|
||||
320
upstream_ref/xllm/xllm/c_api/test/xllm_test.cpp
Normal file
320
upstream_ref/xllm/xllm/c_api/test/xllm_test.cpp
Normal 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(¶ms);
|
||||
} else {
|
||||
xllm_rec_request_params_default(¶ms);
|
||||
}
|
||||
if (request->params().ByteSizeLong() > 0) {
|
||||
PbToXllmRequestParams(request->params(), ¶ms);
|
||||
}
|
||||
|
||||
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,
|
||||
¶ms);
|
||||
} 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,
|
||||
¶ms);
|
||||
} else if (fn == "xllm_rec_text_completions") {
|
||||
raw = xllm_rec_text_completions(rec_handler_,
|
||||
model_id,
|
||||
request->prompt().c_str(),
|
||||
timeout_ms,
|
||||
¶ms);
|
||||
} 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,
|
||||
¶ms);
|
||||
} 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,
|
||||
¶ms);
|
||||
} 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,
|
||||
¶ms);
|
||||
} 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;
|
||||
}
|
||||
24
upstream_ref/xllm/xllm/c_api/test/xllm_test.flags
Normal file
24
upstream_ref/xllm/xllm/c_api/test/xllm_test.flags
Normal 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)
|
||||
224
upstream_ref/xllm/xllm/c_api/test/xllm_test.proto
Normal file
224
upstream_ref/xllm/xllm/c_api/test/xllm_test.proto
Normal 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);
|
||||
}
|
||||
118
upstream_ref/xllm/xllm/c_api/tools/install.sh
Executable file
118
upstream_ref/xllm/xllm/c_api/tools/install.sh
Executable 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
|
||||
627
upstream_ref/xllm/xllm/c_api/types.h
Normal file
627
upstream_ref/xllm/xllm/c_api/types.h
Normal 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
|
||||
44
upstream_ref/xllm/xllm/cc_api/README.md
Normal file
44
upstream_ref/xllm/xllm/cc_api/README.md
Normal 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
|
||||
```
|
||||
97
upstream_ref/xllm/xllm/cc_api/examples/CMakeLists.txt
Normal file
97
upstream_ref/xllm/xllm/cc_api/examples/CMakeLists.txt
Normal 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"
|
||||
)
|
||||
@@ -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;
|
||||
}
|
||||
97
upstream_ref/xllm/xllm/cc_api/examples/service_request.cpp
Normal file
97
upstream_ref/xllm/xllm/cc_api/examples/service_request.cpp
Normal 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
|
||||
28
upstream_ref/xllm/xllm/cc_api/examples/service_request.h
Normal file
28
upstream_ref/xllm/xllm/cc_api/examples/service_request.h
Normal 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
|
||||
@@ -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;
|
||||
}
|
||||
14
upstream_ref/xllm/xllm/cc_api/examples/start-llm-instance.sh
Executable file
14
upstream_ref/xllm/xllm/cc_api/examples/start-llm-instance.sh
Executable 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
|
||||
119
upstream_ref/xllm/xllm/cc_api/install.sh
Executable file
119
upstream_ref/xllm/xllm/cc_api/install.sh
Executable 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
|
||||
254
upstream_ref/xllm/xllm/cc_api/internal.h
Normal file
254
upstream_ref/xllm/xllm/cc_api/internal.h
Normal 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
|
||||
186
upstream_ref/xllm/xllm/cc_api/llm.cpp
Normal file
186
upstream_ref/xllm/xllm/cc_api/llm.cpp
Normal 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
|
||||
94
upstream_ref/xllm/xllm/cc_api/llm.h
Normal file
94
upstream_ref/xllm/xllm/cc_api/llm.h
Normal 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
|
||||
27
upstream_ref/xllm/xllm/cc_api/macros.h
Normal file
27
upstream_ref/xllm/xllm/cc_api/macros.h
Normal 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
|
||||
318
upstream_ref/xllm/xllm/cc_api/types.h
Normal file
318
upstream_ref/xllm/xllm/cc_api/types.h
Normal 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
|
||||
1
upstream_ref/xllm/xllm/compiler/__init__.py
Normal file
1
upstream_ref/xllm/xllm/compiler/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""Compiler-side utilities for xLLM build and AOT flows."""
|
||||
1
upstream_ref/xllm/xllm/compiler/tilelang/__init__.py
Normal file
1
upstream_ref/xllm/xllm/compiler/tilelang/__init__.py
Normal file
@@ -0,0 +1 @@
|
||||
"""TileLang AOT compiler support for xLLM."""
|
||||
13
upstream_ref/xllm/xllm/compiler/tilelang/bootstrap.py
Normal file
13
upstream_ref/xllm/xllm/compiler/tilelang/bootstrap.py
Normal 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",
|
||||
]
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
50
upstream_ref/xllm/xllm/compiler/tilelang/common/cache.py
Normal file
50
upstream_ref/xllm/xllm/compiler/tilelang/common/cache.py
Normal 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()
|
||||
106
upstream_ref/xllm/xllm/compiler/tilelang/common/manifest.py
Normal file
106
upstream_ref/xllm/xllm/compiler/tilelang/common/manifest.py
Normal 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
|
||||
268
upstream_ref/xllm/xllm/compiler/tilelang/common/spec.py
Normal file
268
upstream_ref/xllm/xllm/compiler/tilelang/common/spec.py
Normal 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))
|
||||
88
upstream_ref/xllm/xllm/compiler/tilelang/common/toolchain.py
Normal file
88
upstream_ref/xllm/xllm/compiler/tilelang/common/toolchain.py
Normal 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
|
||||
@@ -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."
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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]
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
@@ -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(),
|
||||
)
|
||||
@@ -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 []
|
||||
@@ -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
|
||||
56
upstream_ref/xllm/xllm/compiler/tilelang_launcher.py
Normal file
56
upstream_ref/xllm/xllm/compiler/tilelang_launcher.py
Normal 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
Reference in New Issue
Block a user