Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions xllm/api_service/embedding_service_impl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,13 @@ limitations under the License.

#include <glog/logging.h>

#include <algorithm>
#include <string>

#include "common/instance_name.h"
#include "distributed_runtime/llm_master.h"
#include "embedding_output_builder.h"
#include "framework/config/model_config.h"
#include "framework/request/request_params.h"
#include "mm_service_utils.h"
#include "util/utils.h"
Expand Down Expand Up @@ -159,6 +161,21 @@ void MMEmbeddingServiceImpl::process_async_impl(
req_messages, messages, call, master_->get_image_limit())) {
return;
}

// mm_embed encodes multimodal inputs, so a text-only request is invalid.
if (::xllm::ModelConfig::get_instance().task() == "mm_embed") {
const bool has_multimodal_content =
std::any_of(messages.begin(), messages.end(), [](Message& msg) {
return msg.has_mm_content();
});
if (!has_multimodal_content) {
call->finish_with_error(
StatusCode::INVALID_ARGUMENT,
"mm_embed request must contain multimodal content.");
return;
}
}

auto request_id = request_params.request_id;

auto payload = call->take_request_payload();
Expand Down
12 changes: 12 additions & 0 deletions xllm/core/common/message.h
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,18 @@ struct Message {
return count;
}

bool has_mm_content() const {
if (std::holds_alternative<std::string>(content)) {
return false;
}
for (const auto& item : std::get<MMContentVec>(content)) {
if (item.type != "text") {
return true;
}
}
return false;
}

std::string role;
Content content;

Expand Down
20 changes: 14 additions & 6 deletions xllm/core/distributed_runtime/worker_service.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -175,7 +175,7 @@ void WorkerService::step(ForwardInput& fwd_input,
torch::Tensor& top_tokens,
torch::Tensor& top_logprobs,
torch::Tensor& embeddings,
std::vector<torch::Tensor>& mm_embeddings,
std::vector<std::vector<torch::Tensor>>& mm_embeddings,
std::vector<torch::Tensor>& dit_images,
std::vector<std::string>& dit_text_output,
torch::Tensor& expert_load_data,
Expand Down Expand Up @@ -216,9 +216,14 @@ void WorkerService::step(ForwardInput& fwd_input,

mm_embeddings.clear();
mm_embeddings.reserve(sample_output.mm_embeddings.size());
for (auto mm_embedding : sample_output.mm_embeddings) {
mm_embeddings.emplace_back(
safe_to(mm_embedding, torch::kCPU, /*non_blocking=*/true));
for (const auto& seq_mm_embeddings : sample_output.mm_embeddings) {
std::vector<torch::Tensor> seq_out;
seq_out.reserve(seq_mm_embeddings.size());
for (const auto& mm_embedding : seq_mm_embeddings) {
seq_out.emplace_back(
safe_to(mm_embedding, torch::kCPU, /*non_blocking=*/true));
}
mm_embeddings.emplace_back(std::move(seq_out));
}

dit_images.clear();
Expand Down Expand Up @@ -318,7 +323,7 @@ void WorkerService::create_polling_shm_thread(
torch::Tensor top_tokens;
torch::Tensor top_logprobs;
torch::Tensor embeddings;
std::vector<torch::Tensor> mm_embeddings;
std::vector<std::vector<torch::Tensor>> mm_embeddings;
std::vector<torch::Tensor> dit_images;
std::vector<std::string> dit_text_output;
torch::Tensor expert_load_data;
Expand Down Expand Up @@ -747,7 +752,7 @@ void WorkerService::ExecuteModel(::google::protobuf::RpcController* controller,
torch::Tensor top_tokens;
torch::Tensor top_logprobs;
torch::Tensor embeddings;
std::vector<torch::Tensor> mm_embeddings;
std::vector<std::vector<torch::Tensor>> mm_embeddings;
std::vector<torch::Tensor> dit_images;
std::vector<std::string> dit_text_output;
torch::Tensor expert_load_data;
Expand Down Expand Up @@ -777,6 +782,7 @@ void WorkerService::ExecuteModel(::google::protobuf::RpcController* controller,
top_tokens,
top_logprobs,
embeddings,
mm_embeddings,
expert_load_data,
prepared_layer_id,
src_seq_idxes,
Expand Down Expand Up @@ -901,11 +907,13 @@ void WorkerService::GetLastStepResult(
if (next_tokens.defined() || !dit_images.empty() ||
!dit_text_output.empty() ||
::xllm::EPLBConfig::get_instance().enable_eplb()) {
const std::vector<std::vector<torch::Tensor>> mm_embeddings;
forward_output_to_proto(next_tokens,
logprobs,
top_tokens,
top_logprobs,
embeddings,
mm_embeddings,
expert_load_data,
prepared_layer_id,
src_seq_idxes,
Expand Down
2 changes: 1 addition & 1 deletion xllm/core/distributed_runtime/worker_service.h
Original file line number Diff line number Diff line change
Expand Up @@ -157,7 +157,7 @@ class WorkerService : public proto::DistributeWorker {
torch::Tensor& top_tokens,
torch::Tensor& top_logprobs,
torch::Tensor& embeddings,
std::vector<torch::Tensor>& mm_embeddings,
std::vector<std::vector<torch::Tensor>>& mm_embeddings,
std::vector<torch::Tensor>& dit_images,
std::vector<std::string>& dit_text_output,
torch::Tensor& expert_load_data,
Expand Down
39 changes: 8 additions & 31 deletions xllm/core/framework/batch/batch.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -473,43 +473,20 @@ void Batch::refresh_onerec_prefill_output_targets() {

void Batch::process_sample_output(const RawForwardOutput& raw_output,
bool replace_fake_token) {
if (raw_output.mm_embeddings.size() > 0) {
// mm embed task
int64_t mm_embedding_idx = 0;
const auto sequences = get_sequences();
for (auto* seq : sequences) {
int64_t mm_item_count = seq->mm_data().size();
if (mm_item_count <= 0) {
continue;
}
std::vector<torch::Tensor> seq_mm_embeddings;
// if we want to return the full embeding of images and prompts,
// the output is a single embedding tensor, else it would be a vector of
// image embeddings
int64_t output_tensor_size =
::xllm::ModelConfig::get_instance().enable_return_mm_full_embeddings()
? 1
: mm_item_count;
seq_mm_embeddings.reserve(output_tensor_size);
for (int64_t i = mm_embedding_idx;
i < mm_embedding_idx + output_tensor_size;
++i) {
CHECK_LT(i, raw_output.mm_embeddings.size());
seq_mm_embeddings.push_back(raw_output.mm_embeddings[i]);
}
seq->update_mm_embeddings(seq_mm_embeddings);
// we only support complete mm embedding in one iteration now
CHECK(seq->finished());
mm_embedding_idx += output_tensor_size;
}
}

for (size_t output_idx = 0; output_idx < output_targets_.size();
++output_idx) {
const auto& target = output_targets_[output_idx];
auto* seq = target.sequence;
CHECK(seq != nullptr);

if (output_idx < raw_output.outputs.size()) {
const auto& seq_mm_embeddings =
raw_output.outputs[output_idx].mm_embeddings;
if (!seq_mm_embeddings.empty()) {
seq->update_mm_embeddings(seq_mm_embeddings);
}
}

if (!target.from_sample_slot) {
if (seq->finished()) {
continue;
Expand Down
3 changes: 1 addition & 2 deletions xllm/core/framework/sampling/sampling_params.h
Original file line number Diff line number Diff line change
Expand Up @@ -190,8 +190,7 @@ struct SampleOutput {
// directly without re-selecting. Only set on the CP target prefill path.
torch::Tensor selected_embeddings;

// each element is a FloatTensor
std::vector<torch::Tensor> mm_embeddings;
std::vector<std::vector<torch::Tensor>> mm_embeddings;
};

} // namespace xllm
8 changes: 3 additions & 5 deletions xllm/core/runtime/embed_vlm_worker_impl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -93,17 +93,15 @@ std::optional<ForwardOutput> EmbedVLMWorkerImpl::step(
sample_output.mm_embeddings.reserve(q_seq_len_vec.size());
int32_t token_start_idx = 0;
for (auto seq_len : q_seq_len_vec) {
auto image_embed =
auto seq_embed =
embeddings.slice(0, token_start_idx, token_start_idx + seq_len);
sample_output.mm_embeddings.emplace_back(image_embed);
sample_output.mm_embeddings.push_back({seq_embed});
token_start_idx += seq_len;
}
output.sample_output = sample_output;
} else {
sample_output.embeddings = embeddings;
output.sample_output = sample_output;
output.embedding = embeddings;
}
output.sample_output = sample_output;
}
ret = device_.synchronize_default_stream();
return output;
Expand Down
4 changes: 2 additions & 2 deletions xllm/core/runtime/forward_params.h
Original file line number Diff line number Diff line change
Expand Up @@ -933,6 +933,8 @@ struct ForwardOutput {

struct RawSampleOutput {
std::vector<RawToken> tokens; // num tokens
// multimodal embedding output for this sequence
std::vector<torch::Tensor> mm_embeddings;
};

struct RawForwardOutput {
Expand All @@ -946,8 +948,6 @@ struct RawForwardOutput {

// batch-level beam output for Rec multi-round mode
std::vector<int32_t> beam_sequence_group; // flattened 2D
// multimodal embedding output
std::vector<torch::Tensor> mm_embeddings;
// dit output data
DiTForwardOutput dit_forward_output;
};
Expand Down
19 changes: 11 additions & 8 deletions xllm/core/runtime/forward_shared_memory_manager.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2521,6 +2521,7 @@ size_t calculate_raw_sample_output_size(const RawSampleOutput& sample) {
for (const auto& token : sample.tokens) {
size += calculate_raw_token_size(token);
}
size += get_vector_tensor_size(sample.mm_embeddings);
return size;
}

Expand All @@ -2537,8 +2538,6 @@ size_t calculate_raw_forward_output_size(const RawForwardOutput& output) {
size += get_vector_size(output.out_tokens);
size += get_vector_size(output.out_logprobs);
size += type_size<int32_t>; // prepared_layer_id
// mm_embedding_data
size += get_vector_tensor_size(output.mm_embeddings);
const bool has_dit_forward_output =
!output.dit_forward_output.tensors.empty();
size += type_size<bool>;
Expand Down Expand Up @@ -2567,6 +2566,7 @@ void write_raw_sample_output(char*& buffer, const RawSampleOutput& sample) {
for (const auto& token : sample.tokens) {
write_raw_token(buffer, token);
}
write_vector_tensor(buffer, sample.mm_embeddings);
}

void read_raw_token(const char*& buffer, RawToken& token) {
Expand Down Expand Up @@ -2594,6 +2594,7 @@ void read_raw_sample_output(const char*& buffer, RawSampleOutput& sample) {
for (auto& token : sample.tokens) {
read_raw_token(buffer, token);
}
read_vector_tensor(buffer, sample.mm_embeddings);
}

void deserialize_raw_forward_output(const char* buffer,
Expand All @@ -2612,8 +2613,6 @@ void deserialize_raw_forward_output(const char* buffer,

read_data(buffer, output.prepared_layer_id);

read_vector_tensor(buffer, output.mm_embeddings);

bool has_dit_forward_output = false;
read_data(buffer, has_dit_forward_output);
if (has_dit_forward_output) {
Expand All @@ -2635,7 +2634,6 @@ void serialize_raw_forward_output(const RawForwardOutput& output,

write_data(buffer, output.prepared_layer_id);

write_vector_tensor(buffer, output.mm_embeddings);
const bool has_dit_forward_output =
!output.dit_forward_output.tensors.empty();
write_data(buffer, has_dit_forward_output);
Expand Down Expand Up @@ -2850,7 +2848,7 @@ void convert_tensor_to_raw_output(
const torch::Tensor& top_tokens,
const torch::Tensor& top_logprobs,
const torch::Tensor& embeddings,
const std::vector<torch::Tensor>& mm_embeddings,
const std::vector<std::vector<torch::Tensor>>& mm_embeddings,
const std::vector<torch::Tensor>& dit_images,
const std::vector<std::string>& dit_text_output,
const torch::Tensor& expert_load_data,
Expand Down Expand Up @@ -2894,9 +2892,11 @@ void convert_tensor_to_raw_output(
if (embeddings.defined() && embeddings.numel() > 0) {
num_seqs = std::max(num_seqs, static_cast<int32_t>(embeddings.size(0)));
}
if (!mm_embeddings.empty()) {
num_seqs = std::max(num_seqs, static_cast<int32_t>(mm_embeddings.size()));
}

raw_output.outputs.reserve(num_seqs);
raw_output.mm_embeddings = mm_embeddings;
raw_output.dit_forward_output.tensors = dit_images;
raw_output.dit_forward_output.text_output = dit_text_output;
for (int32_t output_idx = 0; output_idx < num_seqs; ++output_idx) {
Expand Down Expand Up @@ -2970,6 +2970,9 @@ void convert_tensor_to_raw_output(

raw_sample_output.tokens.push_back(std::move(raw_token));
}
if (output_idx < static_cast<int32_t>(mm_embeddings.size())) {
raw_sample_output.mm_embeddings = mm_embeddings[output_idx];
}
raw_output.outputs.push_back(std::move(raw_sample_output));
}
}
Expand Down Expand Up @@ -3180,7 +3183,7 @@ bool ForwardSharedMemoryManager::raw_output_write(
const torch::Tensor& top_tokens,
const torch::Tensor& top_logprobs,
const torch::Tensor& embeddings,
const std::vector<torch::Tensor>& mm_embeddings,
const std::vector<std::vector<torch::Tensor>>& mm_embeddings,
const std::vector<torch::Tensor>& dit_images,
const std::vector<std::string>& dit_text_output,
const torch::Tensor& expert_load_data,
Expand Down
27 changes: 14 additions & 13 deletions xllm/core/runtime/forward_shared_memory_manager.h
Original file line number Diff line number Diff line change
Expand Up @@ -105,19 +105,20 @@ class ForwardSharedMemoryManager : public SharedMemoryManager {

bool input_write(const ForwardInput& input);
void input_read(ForwardInput& input, const torch::Device& device);
bool raw_output_write(const torch::Tensor& next_tokens,
const torch::Tensor& logprobs,
const torch::Tensor& top_tokens,
const torch::Tensor& top_logprobs,
const torch::Tensor& embeddings,
const std::vector<torch::Tensor>& mm_embeddings,
const std::vector<torch::Tensor>& dit_images,
const std::vector<std::string>& dit_text_output,
const torch::Tensor& expert_load_data,
int32_t prepared_layer_id,
const torch::Tensor& src_seq_idxes,
const torch::Tensor& out_tokens,
const torch::Tensor& out_logprobs);
bool raw_output_write(
const torch::Tensor& next_tokens,
const torch::Tensor& logprobs,
const torch::Tensor& top_tokens,
const torch::Tensor& top_logprobs,
const torch::Tensor& embeddings,
const std::vector<std::vector<torch::Tensor>>& mm_embeddings,
const std::vector<torch::Tensor>& dit_images,
const std::vector<std::string>& dit_text_output,
const torch::Tensor& expert_load_data,
int32_t prepared_layer_id,
const torch::Tensor& src_seq_idxes,
const torch::Tensor& out_tokens,
const torch::Tensor& out_logprobs);
void raw_output_read(RawForwardOutput& outputs);

void clear();
Expand Down
19 changes: 18 additions & 1 deletion xllm/core/runtime/mm_embed_vlm_worker_impl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,24 @@ std::optional<ForwardOutput> MMEmbedVLMWorkerImpl::step(

ForwardOutput output;
SampleOutput sample_output;
sample_output.mm_embeddings = mm_embeddings;
// Group the flattened per-image embeddings by sequence, using each sequence's
// image count from the host-side mm_data. Outer = sequence, inner = images.
const auto& mm_data_vec = input.input_params.multimodal.mm_data.mm_data_vec();
sample_output.mm_embeddings.reserve(mm_data_vec.size());
size_t image_idx = 0;
for (const auto& seq_mm_data : mm_data_vec) {
const size_t seq_image_count = seq_mm_data.size();
std::vector<torch::Tensor> seq_mm_embeddings;
seq_mm_embeddings.reserve(seq_image_count);
for (size_t i = 0; i < seq_image_count; ++i) {
CHECK_LT(image_idx, mm_embeddings.size());
seq_mm_embeddings.push_back(mm_embeddings[image_idx++]);
}
sample_output.mm_embeddings.push_back(std::move(seq_mm_embeddings));
}
CHECK_EQ(image_idx, mm_embeddings.size())
<< "mm_embedding count mismatch: grouped " << image_idx << " but got "
<< mm_embeddings.size();
output.sample_output = sample_output;

return output;
Expand Down
Loading
Loading