Use llama_chat_apply_template for generic model support; add --chat-template flag
This commit is contained in:
parent
5180d7b508
commit
7e9a014f4c
35
main.cpp
35
main.cpp
@ -110,6 +110,7 @@ static void print_usage(const char * prog) {
|
|||||||
" --temp <f> sampling temperature (0 = greedy, default: 0)\n"
|
" --temp <f> sampling temperature (0 = greedy, default: 0)\n"
|
||||||
" -s <seed> random seed\n"
|
" -s <seed> random seed\n"
|
||||||
" --json output in JSON format\n"
|
" --json output in JSON format\n"
|
||||||
|
" --chat-template <name> chat template (default: auto from model)\n"
|
||||||
" -h show this help\n"
|
" -h show this help\n"
|
||||||
"\n"
|
"\n"
|
||||||
"If <image_path> is omitted, reads image from clipboard.\n"
|
"If <image_path> is omitted, reads image from clipboard.\n"
|
||||||
@ -126,6 +127,7 @@ int main(int argc, char ** argv) {
|
|||||||
std::string mmproj_path;
|
std::string mmproj_path;
|
||||||
std::string image_path;
|
std::string image_path;
|
||||||
std::string prompt = "OCR";
|
std::string prompt = "OCR";
|
||||||
|
std::string chat_template;
|
||||||
int n_threads = (int)std::thread::hardware_concurrency();
|
int n_threads = (int)std::thread::hardware_concurrency();
|
||||||
int n_gpu_layers = -1;
|
int n_gpu_layers = -1;
|
||||||
int n_ctx = 8192;
|
int n_ctx = 8192;
|
||||||
@ -148,6 +150,7 @@ int main(int argc, char ** argv) {
|
|||||||
{"mm", required_argument, nullptr, 2},
|
{"mm", required_argument, nullptr, 2},
|
||||||
{"ngl", required_argument, nullptr, 3},
|
{"ngl", required_argument, nullptr, 3},
|
||||||
{"model", required_argument, nullptr, 'm'},
|
{"model", required_argument, nullptr, 'm'},
|
||||||
|
{"chat-template", required_argument, nullptr, 4},
|
||||||
{"help", no_argument, nullptr, 'h'},
|
{"help", no_argument, nullptr, 'h'},
|
||||||
{nullptr, 0, nullptr, 0}
|
{nullptr, 0, nullptr, 0}
|
||||||
};
|
};
|
||||||
@ -159,6 +162,7 @@ int main(int argc, char ** argv) {
|
|||||||
case 1: json_output = true; break;
|
case 1: json_output = true; break;
|
||||||
case 2: mmproj_path = optarg; break;
|
case 2: mmproj_path = optarg; break;
|
||||||
case 3: n_gpu_layers = std::stoi(optarg); break;
|
case 3: n_gpu_layers = std::stoi(optarg); break;
|
||||||
|
case 4: chat_template = optarg; break;
|
||||||
case 'm': model_path = optarg; break;
|
case 'm': model_path = optarg; break;
|
||||||
case 'p': prompt = optarg; break;
|
case 'p': prompt = optarg; break;
|
||||||
case 't': n_threads = std::stoi(optarg); break;
|
case 't': n_threads = std::stoi(optarg); break;
|
||||||
@ -254,18 +258,35 @@ int main(int argc, char ** argv) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Construct prompt using GLM-OCR chat template format
|
// Build user message with image marker placeholder
|
||||||
// The model expects: [gMASK]<sop><|user|>\n<__media__>PROMPT\n<|assistant|>\n
|
std::string user_content = std::string(mtmd_default_marker()) + prompt;
|
||||||
std::string full_prompt = "[gMASK]<sop><|user|>\n";
|
|
||||||
full_prompt += mtmd_default_marker();
|
|
||||||
full_prompt += prompt;
|
|
||||||
full_prompt += "\n<|assistant|>\n";
|
|
||||||
|
|
||||||
|
// Apply chat template for any model
|
||||||
|
std::string full_prompt;
|
||||||
|
const char * tmpl = chat_template.empty()
|
||||||
|
? llama_model_chat_template(model, nullptr)
|
||||||
|
: chat_template.c_str();
|
||||||
|
if (tmpl) {
|
||||||
|
llama_chat_message msg[1] = {{"user", user_content.c_str()}};
|
||||||
|
char buf[8192];
|
||||||
|
int32_t n = llama_chat_apply_template(tmpl, msg, 1, true, buf, sizeof(buf));
|
||||||
|
if (n > 0) {
|
||||||
|
full_prompt = std::string(buf, n);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: use chatml format for models without recognizable template
|
||||||
|
if (full_prompt.empty()) {
|
||||||
|
full_prompt = "<|im_start|>user\n" + user_content + "<|im_end|>\n<|im_start|>assistant\n";
|
||||||
|
}
|
||||||
|
|
||||||
|
// Template output already includes role markers and special tokens (e.g. [gMASK]<sop>)
|
||||||
|
// So we don't add BOS separately — add_special=false
|
||||||
mtmd_input_chunks * chunks = mtmd_input_chunks_init();
|
mtmd_input_chunks * chunks = mtmd_input_chunks_init();
|
||||||
|
|
||||||
mtmd_input_text in_text;
|
mtmd_input_text in_text;
|
||||||
in_text.text = full_prompt.c_str();
|
in_text.text = full_prompt.c_str();
|
||||||
in_text.add_special = false; // template already includes [gMASK]/<sop>
|
in_text.add_special = false;
|
||||||
in_text.parse_special = true;
|
in_text.parse_special = true;
|
||||||
|
|
||||||
const mtmd_bitmap * bitmaps[1] = { bmp };
|
const mtmd_bitmap * bitmaps[1] = { bmp };
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user