add vision support

This commit is contained in:
Xuan Son Nguyen
2026-10-01 20:35:43 +02:00
parent 9efdb2f5d7
commit afabc13b56
10 changed files with 175 additions and 13 deletions
+89 -5
View File
@@ -164,6 +164,68 @@ std::vector<server_decision_question> server_decision_context::parse_questions(c
return questions;
}
//
// images
//
static const size_t DECISION_MAX_IMAGES = 8;
static void decision_load_image(const json & url, std::vector<raw_buffer> & files) {
if (!url.is_string() || !string_starts_with(url.get<std::string>(), "data:image/")) {
throw std::invalid_argument("images must be data URLs (data:image/...;base64,...)");
}
if (files.size() >= DECISION_MAX_IMAGES) {
throw std::invalid_argument(string_format("too many images, the maximum is %zu", DECISION_MAX_IMAGES));
}
handle_media(files, url.get<std::string>(), "");
}
json server_decision_context::parse_state(const json & body, std::vector<raw_buffer> & files) const {
if (body.contains("images") && !body.at("images").is_null()) {
if (!body.at("images").is_array()) {
throw std::invalid_argument("\"images\" must be an array");
}
for (const auto & url : body.at("images")) {
decision_load_image(url, files);
}
}
const json & state = body.at("state");
const bool is_wrapped = state.is_object() && state.contains("messages");
const json & messages = is_wrapped ? state.at("messages") : state;
if (!messages.is_array()) {
return state;
}
// chat messages: take the image parts out of the content
json messages_out = json::array();
for (const auto & msg : messages) {
if (!msg.is_object() || !msg.contains("content") || !msg.at("content").is_array()) {
messages_out.push_back(msg);
continue;
}
json content = json::array();
for (const auto & part : msg.at("content")) {
if (part.is_object() && json_value(part, "type", std::string()) == "image_url" && part.contains("image_url")) {
const json & image_url = part.at("image_url");
decision_load_image(image_url.is_object() && image_url.contains("url") ? image_url.at("url") : image_url, files);
} else {
content.push_back(part);
}
}
json msg_out = msg;
msg_out["content"] = content;
messages_out.push_back(msg_out);
}
if (!is_wrapped) {
return messages_out;
}
json state_out = state;
state_out["messages"] = messages_out;
return state_out;
}
//
// prompt
//
@@ -192,7 +254,7 @@ static json decision_replace_text(const json & val, const std::string & search,
return val;
}
std::string server_decision_context::render(const json & state, const server_decision_question & question) const {
std::string server_decision_context::render(const json & state, const server_decision_question & question, size_t n_images) const {
json options = json::array();
for (const auto & opt : question.options) {
options.push_back(json{
@@ -214,6 +276,16 @@ std::string server_decision_context::render(const json & state, const server_dec
inp = decision_replace_text(inp, text_marker, " ");
}
// the template puts one media marker per image
json images = json::array();
if (n_images > 0) {
inp = decision_replace_text(inp, get_media_marker(), " ");
for (size_t i = 0; i < n_images; i++) {
images.push_back(get_media_marker());
}
}
inp["images"] = images;
jinja::context ctx(tmpl->source());
jinja::global_from_json(ctx, inp, false);
jinja::runtime runtime(ctx);
@@ -221,15 +293,27 @@ std::string server_decision_context::render(const json & state, const server_dec
return jinja::runtime::gather_string_parts(results)->as_string().str();
}
void server_decision_context::fill_task(const json & state, const server_decision_question & question, server_task & task) const {
llama_tokens tokens = common_tokenize(vocab, render(state, question), false, true);
void server_decision_context::fill_task(
const json & state,
const server_decision_question & question,
const std::vector<raw_buffer> & files,
mtmd_context * mctx,
const mtmd_helper_init_opt & init_opt,
server_task & task) const {
const std::string prompt = render(state, question, files.size());
if (type == COMMON_DECISION_TYPE_OPENJEV) {
task.decision.labels.assign(labels.begin(), labels.begin() + question.options.size());
} else {
fill_task_laya(tokens, question, task);
if (!files.empty()) {
task.tokens = process_mtmd_prompt(mctx, prompt, files, init_opt);
return;
}
}
llama_tokens tokens = common_tokenize(vocab, prompt, false, true);
if (type == COMMON_DECISION_TYPE_LAYA) {
fill_task_laya(tokens, question, task);
}
task.tokens = server_tokens(tokens, false);
}